anycoderapps/Z-Image-Turbo
238
1# This file is modified from https://github.com/xdit-project/xDiT/blob/main/entrypoints/launch.py2import base643import gc4import hashlib5import io6import os7import tempfile8from io import BytesIO9 10import gradio as gr11import requests12import torch13import torch.distributed as dist14from fastapi import FastAPI, HTTPException15from PIL import Image16 17from .api import download_from_url, encode_file_to_base6418 19try:20 import ray21except:22 print("Ray is not installed. If you want to use multi gpus api. Please install it by running 'pip install ray'.")23 ray = None24 25def save_base64_video_dist(base64_string):26 video_data = base64.b64decode(base64_string)27 28 md5_hash = hashlib.md5(video_data).hexdigest()29 filename = f"{md5_hash}.mp4" 30 31 temp_dir = tempfile.gettempdir()32 file_path = os.path.join(temp_dir, filename)33 34 if dist.is_initialized():35 if dist.get_rank() == 0:36 with open(file_path, 'wb') as video_file:37 video_file.write(video_data)38 dist.barrier()39 else:40 with open(file_path, 'wb') as video_file:41 video_file.write(video_data)42 return file_path43 44def save_base64_image_dist(base64_string):45 video_data = base64.b64decode(base64_string)46 47 md5_hash = hashlib.md5(video_data).hexdigest()48 filename = f"{md5_hash}.jpg" 49 50 temp_dir = tempfile.gettempdir()51 file_path = os.path.join(temp_dir, filename)52 53 if dist.is_initialized():54 if dist.get_rank() == 0:55 with open(file_path, 'wb') as video_file:56 video_file.write(video_data)57 dist.barrier()58 else:59 with open(file_path, 'wb') as video_file:60 video_file.write(video_data)61 return file_path62 63def save_url_video_dist(url):64 video_data = download_from_url(url)65 if video_data:66 return save_base64_video_dist(base64.b64encode(video_data))67 return None68 69def save_url_image_dist(url):70 image_data = download_from_url(url)71 if image_data:72 return save_base64_image_dist(base64.b64encode(image_data))73 return None74 75if ray is not None:76 @ray.remote(num_gpus=1)77 class MultiNodesGenerator:78 def __init__(79 self, rank: int, world_size: int, Controller,80 GPU_memory_mode, scheduler_dict, model_name=None, model_type="Inpaint", 81 config_path=None, ulysses_degree=1, ring_degree=1,82 fsdp_dit=False, fsdp_text_encoder=False, compile_dit=False, 83 weight_dtype=None, savedir_sample=None,84 ):85 # Set PyTorch distributed environment variables86 os.environ["RANK"] = str(rank)87 os.environ["WORLD_SIZE"] = str(world_size)88 os.environ["MASTER_ADDR"] = "127.0.0.1"89 os.environ["MASTER_PORT"] = "29500"90 91 self.rank = rank92 self.controller = Controller(93 GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path, 94 ulysses_degree=ulysses_degree, ring_degree=ring_degree, 95 fsdp_dit=fsdp_dit, fsdp_text_encoder=fsdp_text_encoder, compile_dit=compile_dit, 96 weight_dtype=weight_dtype, savedir_sample=savedir_sample,97 )98 99 def generate(self, datas):100 try:101 base_model_path = datas.get('base_model_path', 'none')102 base_model_2_path = datas.get('base_model_2_path', 'none')103 lora_model_path = datas.get('lora_model_path', 'none')104 lora_model_2_path = datas.get('lora_model_2_path', 'none')105 lora_alpha_slider = datas.get('lora_alpha_slider', 0.55)106 prompt_textbox = datas.get('prompt_textbox', None)107 negative_prompt_textbox = datas.get('negative_prompt_textbox', 'The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. ')108 sampler_dropdown = datas.get('sampler_dropdown', 'Euler')109 sample_step_slider = datas.get('sample_step_slider', 30)110 resize_method = datas.get('resize_method', "Generate by")111 width_slider = datas.get('width_slider', 672)112 height_slider = datas.get('height_slider', 384)113 base_resolution = datas.get('base_resolution', 512)114 is_image = datas.get('is_image', False)115 generation_method = datas.get('generation_method', False)116 length_slider = datas.get('length_slider', 49)117 overlap_video_length = datas.get('overlap_video_length', 4)118 partial_video_length = datas.get('partial_video_length', 72)119 cfg_scale_slider = datas.get('cfg_scale_slider', 6)120 start_image = datas.get('start_image', None)121 end_image = datas.get('end_image', None)122 validation_video = datas.get('validation_video', None)123 validation_video_mask = datas.get('validation_video_mask', None)124 control_video = datas.get('control_video', None)125 denoise_strength = datas.get('denoise_strength', 0.70)126 seed_textbox = datas.get("seed_textbox", 43)127 128 ref_image = datas.get('ref_image', None)129 enable_teacache = datas.get('enable_teacache', True)130 teacache_threshold = datas.get('teacache_threshold', 0.10)131 num_skip_start_steps = datas.get('num_skip_start_steps', 1)132 teacache_offload = datas.get('teacache_offload', False)133 cfg_skip_ratio = datas.get('cfg_skip_ratio', 0)134 enable_riflex = datas.get('enable_riflex', False)135 riflex_k = datas.get('riflex_k', 6)136 fps = datas.get('fps', None)137 138 generation_method = "Image Generation" if is_image else generation_method139 140 if start_image is not None:141 if start_image.startswith('http'):142 start_image = save_url_image_dist(start_image)143 start_image = [Image.open(start_image).convert("RGB")]144 else:145 start_image = base64.b64decode(start_image)146 start_image = [Image.open(BytesIO(start_image)).convert("RGB")]147 148 if end_image is not None:149 if end_image.startswith('http'):150 end_image = save_url_image_dist(end_image)151 end_image = [Image.open(end_image).convert("RGB")]152 else:153 end_image = base64.b64decode(end_image)154 end_image = [Image.open(BytesIO(end_image)).convert("RGB")]155 156 if validation_video is not None:157 if validation_video.startswith('http'):158 validation_video = save_url_video_dist(validation_video)159 else:160 validation_video = save_base64_video_dist(validation_video)161 162 if validation_video_mask is not None:163 if validation_video_mask.startswith('http'):164 validation_video_mask = save_url_image_dist(validation_video_mask)165 else:166 validation_video_mask = save_base64_image_dist(validation_video_mask)167 168 if control_video is not None:169 if control_video.startswith('http'):170 control_video = save_url_video_dist(control_video)171 else:172 control_video = save_base64_video_dist(control_video)173 174 if ref_image is not None:175 if ref_image.startswith('http'):176 ref_image = save_url_image_dist(ref_image)177 ref_image = [Image.open(ref_image).convert("RGB")]178 else:179 ref_image = base64.b64decode(ref_image)180 ref_image = [Image.open(BytesIO(ref_image)).convert("RGB")]181 182 try:183 save_sample_path, comment = self.controller.generate(184 "",185 base_model_path,186 lora_model_path, 187 lora_alpha_slider,188 prompt_textbox, 189 negative_prompt_textbox, 190 sampler_dropdown, 191 sample_step_slider, 192 resize_method,193 width_slider, 194 height_slider, 195 base_resolution,196 generation_method,197 length_slider, 198 overlap_video_length, 199 partial_video_length, 200 cfg_scale_slider, 201 start_image,202 end_image,203 validation_video,204 validation_video_mask, 205 control_video, 206 denoise_strength,207 seed_textbox,208 ref_image = ref_image,209 enable_teacache = enable_teacache, 210 teacache_threshold = teacache_threshold, 211 num_skip_start_steps = num_skip_start_steps, 212 teacache_offload = teacache_offload, 213 cfg_skip_ratio = cfg_skip_ratio,214 enable_riflex = enable_riflex, 215 riflex_k = riflex_k, 216 base_model_2_dropdown = base_model_2_path,217 lora_model_2_dropdown = lora_model_2_path,218 fps = fps,219 is_api = True,220 )221 except Exception as e:222 gc.collect()223 torch.cuda.empty_cache()224 torch.cuda.ipc_collect()225 save_sample_path = ""226 comment = f"Error. error information is {str(e)}"227 if dist.is_initialized():228 if dist.get_rank() == 0:229 return {"message": comment, "save_sample_path": None, "base64_encoding": None}230 else:231 return None232 else:233 return {"message": comment, "save_sample_path": None, "base64_encoding": None}234 235 236 if dist.is_initialized():237 if dist.get_rank() == 0:238 if save_sample_path != "":239 return {"message": comment, "save_sample_path": save_sample_path, "base64_encoding": encode_file_to_base64(save_sample_path)}240 else:241 return {"message": comment, "save_sample_path": None, "base64_encoding": None}242 else:243 return None244 else:245 if save_sample_path != "":246 return {"message": comment, "save_sample_path": save_sample_path, "base64_encoding": encode_file_to_base64(save_sample_path)}247 else:248 return {"message": comment, "save_sample_path": None, "base64_encoding": None}249 250 except Exception as e:251 print(f"Error generating: {str(e)}")252 comment = f"Error generating: {str(e)}"253 if dist.is_initialized():254 if dist.get_rank() == 0:255 return {"message": comment, "save_sample_path": None, "base64_encoding": None}256 else:257 return None258 else:259 return {"message": comment, "save_sample_path": None, "base64_encoding": None}260 261 class MultiNodesEngine:262 def __init__(263 self, 264 world_size, 265 Controller,266 GPU_memory_mode, 267 scheduler_dict, 268 model_name, 269 model_type, 270 config_path,271 ulysses_degree=1, 272 ring_degree=1, 273 fsdp_dit=False,274 fsdp_text_encoder=False,275 compile_dit=False,276 weight_dtype=torch.bfloat16,277 savedir_sample="samples"278 ):279 # Ensure Ray is initialized280 if not ray.is_initialized():281 ray.init()282 283 num_workers = world_size284 self.workers = [285 MultiNodesGenerator.remote(286 rank, world_size, Controller, 287 GPU_memory_mode, scheduler_dict, model_name=model_name, model_type=model_type, config_path=config_path, 288 ulysses_degree=ulysses_degree, ring_degree=ring_degree, 289 fsdp_dit=fsdp_dit, fsdp_text_encoder=fsdp_text_encoder, compile_dit=compile_dit, 290 weight_dtype=weight_dtype, savedir_sample=savedir_sample,291 )292 for rank in range(num_workers)293 ]294 print("Update workers done")295 296 async def generate(self, data):297 results = ray.get([298 worker.generate.remote(data)299 for worker in self.workers300 ])301 302 return next(path for path in results if path is not None) 303 304 def multi_nodes_infer_forward_api(_: gr.Blocks, app: FastAPI, engine):305 306 @app.post("/videox_fun/infer_forward")307 async def _multi_nodes_infer_forward_api(308 datas: dict,309 ):310 try:311 result = await engine.generate(datas)312 return result313 except Exception as e:314 if isinstance(e, HTTPException):315 raise e316 raise HTTPException(status_code=500, detail=str(e))317else:318 MultiNodesEngine = None319 MultiNodesGenerator = None320 multi_nodes_infer_forward_api = None