CoolFace
Apppublic

dy2000/optimized-diffusers-code

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
e2e_example.py93 linesDownload Raw Back to root
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