from __future__ import annotations import argparse, json, os, time, types from pathlib import Path import torch MODEL_ID=os.environ.get('FLUX2_MODEL_ID','black-forest-labs/FLUX.2-klein-4B') OUT=Path(os.environ.get('CF_ONNX_OUT','/tmp/cf-flux2-onnx')) class Flux2Export(torch.nn.Module): def __init__(self, model): super().__init__(); self.model=model def forward(self, hidden_states, encoder_hidden_states, timestep, img_ids, txt_ids): return self.model(hidden_states=hidden_states, encoder_hidden_states=encoder_hidden_states, timestep=timestep, img_ids=img_ids, txt_ids=txt_ids, guidance=None, joint_attention_kwargs=None, return_dict=False)[0] def capture_inputs(size:int): from diffusers import Flux2KleinPipeline pipe=Flux2KleinPipeline.from_pretrained(MODEL_ID,torch_dtype=torch.bfloat16) # Keep the large text encoder/vae off GPU; transformer is moved only when needed. pipe.enable_model_cpu_offload() captured={} def probe(self, hidden_states, encoder_hidden_states=None, timestep=None, img_ids=None, txt_ids=None, guidance=None, joint_attention_kwargs=None, return_dict=True): for k,v in dict(hidden_states=hidden_states,encoder_hidden_states=encoder_hidden_states,timestep=timestep, img_ids=img_ids,txt_ids=txt_ids).items(): captured[k]=v.detach().cpu() raise RuntimeError('CF_CAPTURE_DONE') orig=pipe.transformer.forward; pipe.transformer.forward=types.MethodType(probe,pipe.transformer) try: pipe(prompt='a small friendly desktop robot companion, full body, neutral A-pose',height=size,width=size, num_inference_steps=1,guidance_scale=1.0,max_sequence_length=512) except RuntimeError as exc: if str(exc)!='CF_CAPTURE_DONE': raise pipe.transformer.forward=orig del pipe torch.cuda.empty_cache() return captured def main(): ap=argparse.ArgumentParser(); ap.add_argument('--size',type=int,choices=[512,768,1024],default=512); a=ap.parse_args() cap=capture_inputs(a.size) print('CAPTURED',{k:[list(v.shape),str(v.dtype)] for k,v in cap.items()},flush=True) from diffusers import Flux2Transformer2DModel # FP16 is the deployment format for the first TensorRT build. FP8 QDQ is a second-stage optimization. model=Flux2Transformer2DModel.from_pretrained(MODEL_ID,subfolder='transformer',torch_dtype=torch.float16).eval().cuda() inputs=tuple(cap[k].to(dtype=torch.float16 if k in {'hidden_states','encoder_hidden_states','timestep'} else cap[k].dtype) for k in ['hidden_states','encoder_hidden_states','timestep','img_ids','txt_ids']) inputs=tuple(x.cuda() for x in inputs) out=OUT/f'onnx/flux2-klein-4b/{a.size}/transformer.onnx'; out.parent.mkdir(parents=True,exist_ok=True) t=time.time() with torch.inference_mode(): torch.onnx.export(Flux2Export(model),inputs,str(out),opset_version=23,dynamo=True,external_data=True, input_names=['hidden_states','encoder_hidden_states','timestep','img_ids','txt_ids'],output_names=['sample'], optimize=True,dynamic_shapes=None) import onnx; onnx.checker.check_model(str(out)) meta={'size':a.size,'seconds':round(time.time()-t,3),'inputs':{k:{'shape':list(v.shape),'dtype':str(v.dtype)} for k,v in cap.items()},'model':MODEL_ID,'opset':23} (out.parent/'export_meta.json').write_text(json.dumps(meta,indent=2)); print(json.dumps(meta,indent=2)) if __name__=='__main__': main()