from diffusers import DiffusionPipeline, TorchAoConfig from diffusers.quantizers import PipelineQuantizationConfig from utils.fa3_processor import FlashFluxAttnProcessor3_0 import torch # quantize the Flux transformer with FP8 pipe = DiffusionPipeline.from_pretrained( "black-forest-labs/FLUX.1-dev", torch_dtype=torch.bfloat16, quantization_config=PipelineQuantizationConfig( quant_mapping={"transformer": TorchAoConfig("float8dq_e4m3_row")} ) ).to("cuda") # use Flash-attention 3 pipe.transformer.set_attn_processor(FlashFluxAttnProcessor3_0()) # use torch.compile() pipe.transformer.compile(fullgraph=True, mode="max-autotune") # perform inference pipe_kwargs = { "prompt": "A cat holding a sign that says hello world", "height": 1024, "width": 1024, "guidance_scale": 3.5, "num_inference_steps": 28, "max_sequence_length": 512, } # first time will be slower, subsequent runs will be faster image = pipe(**pipe_kwargs).images[0]