dy2000/optimized-diffusers-code
0
1import argparse2from utils.llm_utils import LLMCodeOptimizer3from prompts import system_prompt, generate_prompt4from utils.pipeline_utils import determine_pipe_loading_memory5from utils.hardware_utils import (6 categorize_vram,7 categorize_ram,8 get_gpu_vram_gb,9 get_system_ram_gb,10 is_compile_friendly_gpu,11 is_fp8_friendly,12)13import torch14from pprint import pprint15 16 17def create_parser():18 parser = argparse.ArgumentParser()19 parser.add_argument(20 "--ckpt_id",21 type=str,22 default="black-forest-labs/FLUX.1-dev",23 help="Can be a repo id from the Hub or a local path where the checkpoint is stored.",24 )25 parser.add_argument(26 "--gemini_model",27 type=str,28 default="gemini-2.5-flash-preview-05-20",29 help="Gemini model to use. Choose from https://ai.google.dev/gemini-api/docs/models.",30 )31 parser.add_argument(32 "--variant",33 type=str,34 default=None,35 help="If the `ckpt_id` has variants, supply this flag to estimate compute. Example: 'fp16'.",36 )37 parser.add_argument(38 "--disable_bf16",39 action="store_true",40 help="When enabled the load memory is affected. Prefer not enabling this flag.",41 )42 parser.add_argument(43 "--enable_lossy",44 action="store_true",45 help="When enabled, the code will include snippets for enabling quantization.",46 )47 return parser48 49 50def main(args):51 if not torch.cuda.is_available():52 raise ValueError("Not supported for non-CUDA devices for now.")53 54 loading_mem_out = determine_pipe_loading_memory(args.ckpt_id, args.variant, args.disable_bf16)55 load_memory = loading_mem_out["total_loading_memory_gb"]56 ram_gb = get_system_ram_gb()57 ram_category = categorize_ram(ram_gb)58 if ram_gb is not None:59 print(f"\nSystem RAM: {ram_gb:.2f} GB")60 print(f"RAM Category: {ram_category}")61 else:62 print("\nCould not determine System RAM.")63 64 vram_gb = get_gpu_vram_gb()65 vram_category = categorize_vram(vram_gb)66 if vram_gb is not None:67 print(f"\nGPU VRAM: {vram_gb:.2f} GB")68 print(f"VRAM Category: {vram_category}")69 else:70 print("\nGPU VRAM check complete.")71 72 is_compile_friendly = is_compile_friendly_gpu()73 is_fp8_compatible = is_fp8_friendly()74 75 llm = LLMCodeOptimizer(model_name=args.gemini_model, system_prompt=system_prompt)76 current_generate_prompt = generate_prompt.format(77 ckpt_id=args.ckpt_id,78 pipeline_loading_memory=load_memory,79 available_system_ram=ram_gb,80 available_gpu_vram=vram_gb,81 enable_lossy_outputs=args.enable_lossy,82 is_fp8_supported=is_fp8_compatible,83 enable_torch_compile=is_compile_friendly,84 )85 pprint(f"{current_generate_prompt=}")86 print(llm(current_generate_prompt))87 88 89if __name__ == "__main__":90 parser = create_parser()91 args = parser.parse_args()92 main(args)93 