Update dist file download (#193)
* Fix bug in api * Fix bug in api * Update dist file download * Update dist file download
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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)}
|
||||
|
||||
Reference in New Issue
Block a user