Team Ai
Apppublic

anycoderapps/Z-Image-Turbo

sourceHugging Faceupdated 10mo agoView on Hugging Face
238likes
api_multi_nodes.py320 linesDownload Raw Back to api
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