From d7a37ef8841c45e13a5bea96eb9c49a278ad0bb5 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Wed, 30 Apr 2025 13:36:09 +0800 Subject: [PATCH] Update dist file download (#193) * Fix bug in api * Fix bug in api * Update dist file download * Update dist file download --- videox_fun/api/api.py | 12 ++--- videox_fun/api/api_multi_nodes.py | 89 ++++++++++++++++++++++++------- 2 files changed, 77 insertions(+), 24 deletions(-) diff --git a/videox_fun/api/api.py b/videox_fun/api/api.py index 1eb0b89..6434c1f 100755 --- a/videox_fun/api/api.py +++ b/videox_fun/api/api.py @@ -131,18 +131,18 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): if start_image is not None: if start_image.startswith('http'): start_image = save_url_image(start_image) - start_image = [Image.open(start_image)] + start_image = [Image.open(start_image).convert("RGB")] else: start_image = base64.b64decode(start_image) - start_image = [Image.open(BytesIO(start_image))] + start_image = [Image.open(BytesIO(start_image)).convert("RGB")] if end_image is not None: if end_image.startswith('http'): end_image = save_url_image(end_image) - end_image = [Image.open(end_image)] + end_image = [Image.open(end_image).convert("RGB")] else: end_image = base64.b64decode(end_image) - end_image = [Image.open(BytesIO(end_image))] + end_image = [Image.open(BytesIO(end_image)).convert("RGB")] if validation_video is not None: if validation_video.startswith('http'): @@ -165,10 +165,10 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): if ref_image is not None: if ref_image.startswith('http'): ref_image = save_url_image(ref_image) - ref_image = [Image.open(ref_image)] + ref_image = [Image.open(ref_image).convert("RGB")] else: ref_image = base64.b64decode(ref_image) - ref_image = [Image.open(BytesIO(ref_image))] + ref_image = [Image.open(BytesIO(ref_image)).convert("RGB")] try: save_sample_path, comment = controller.generate( diff --git a/videox_fun/api/api_multi_nodes.py b/videox_fun/api/api_multi_nodes.py index 8d54b61..9bf3785 100755 --- a/videox_fun/api/api_multi_nodes.py +++ b/videox_fun/api/api_multi_nodes.py @@ -1,16 +1,20 @@ # This file is modified from https://github.com/xdit-project/xDiT/blob/main/entrypoints/launch.py import base64 import gc +import hashlib +import io import os +import tempfile from io import BytesIO import gradio as gr +import requests import torch +import torch.distributed as dist from fastapi import FastAPI, HTTPException from PIL import Image -from .api import (encode_file_to_base64, save_base64_image, save_base64_video, - save_url_image, save_url_video) +from .api import download_from_url, encode_file_to_base64 try: import ray @@ -18,6 +22,56 @@ except: print("Ray is not installed. If you want to use multi gpus api. Please install it by running 'pip install ray'.") ray = None +def save_base64_video_dist(base64_string): + video_data = base64.b64decode(base64_string) + + md5_hash = hashlib.md5(video_data).hexdigest() + filename = f"{md5_hash}.mp4" + + temp_dir = tempfile.gettempdir() + file_path = os.path.join(temp_dir, filename) + + if dist.is_initialized(): + if dist.get_rank() == 0: + with open(file_path, 'wb') as video_file: + video_file.write(video_data) + dist.barrier() + else: + with open(file_path, 'wb') as video_file: + video_file.write(video_data) + return file_path + +def save_base64_image_dist(base64_string): + video_data = base64.b64decode(base64_string) + + md5_hash = hashlib.md5(video_data).hexdigest() + filename = f"{md5_hash}.jpg" + + temp_dir = tempfile.gettempdir() + file_path = os.path.join(temp_dir, filename) + + if dist.is_initialized(): + if dist.get_rank() == 0: + with open(file_path, 'wb') as video_file: + video_file.write(video_data) + dist.barrier() + else: + with open(file_path, 'wb') as video_file: + video_file.write(video_data) + return file_path + +def save_url_video_dist(url): + video_data = download_from_url(url) + if video_data: + return save_base64_video_dist(base64.b64encode(video_data)) + return None + +def save_url_image_dist(url): + image_data = download_from_url(url) + if image_data: + return save_base64_image_dist(base64.b64encode(image_data)) + return None + if ray is not None: @ray.remote(num_gpus=1) class MultiNodesGenerator: @@ -79,45 +133,45 @@ if ray is not None: if start_image is not None: if start_image.startswith('http'): - start_image = save_url_image(start_image) - start_image = [Image.open(start_image)] + start_image = save_url_image_dist(start_image) + start_image = [Image.open(start_image).convert("RGB")] else: start_image = base64.b64decode(start_image) - start_image = [Image.open(BytesIO(start_image))] + start_image = [Image.open(BytesIO(start_image)).convert("RGB")] if end_image is not None: if end_image.startswith('http'): - end_image = save_url_image(end_image) - end_image = [Image.open(end_image)] + end_image = save_url_image_dist(end_image) + end_image = [Image.open(end_image).convert("RGB")] else: end_image = base64.b64decode(end_image) - end_image = [Image.open(BytesIO(end_image))] + end_image = [Image.open(BytesIO(end_image)).convert("RGB")] if validation_video is not None: if validation_video.startswith('http'): - validation_video = save_url_video(validation_video) + validation_video = save_url_video_dist(validation_video) else: - validation_video = save_base64_video(validation_video) + validation_video = save_base64_video_dist(validation_video) if validation_video_mask is not None: if validation_video_mask.startswith('http'): - validation_video_mask = save_url_image(validation_video_mask) + validation_video_mask = save_url_image_dist(validation_video_mask) else: - validation_video_mask = save_base64_image(validation_video_mask) + validation_video_mask = save_base64_image_dist(validation_video_mask) if control_video is not None: if control_video.startswith('http'): - control_video = save_url_video(control_video) + control_video = save_url_video_dist(control_video) else: - control_video = save_base64_video(control_video) + control_video = save_base64_video_dist(control_video) if ref_image is not None: if ref_image.startswith('http'): - ref_image = save_url_image(ref_image) - ref_image = [Image.open(ref_image)] + ref_image = save_url_image_dist(ref_image) + ref_image = [Image.open(ref_image).convert("RGB")] else: ref_image = base64.b64decode(ref_image) - ref_image = [Image.open(BytesIO(ref_image))] + ref_image = [Image.open(BytesIO(ref_image)).convert("RGB")] try: save_sample_path, comment = self.controller.generate( @@ -163,7 +217,6 @@ if ray is not None: comment = f"Error. error information is {str(e)}" return {"message": comment, "save_sample_path": None, "base64_encoding": None} - import torch.distributed as dist if dist.get_rank() == 0: if save_sample_path != "": return {"message": comment, "save_sample_path": save_sample_path, "base64_encoding": encode_file_to_base64(save_sample_path)}