| 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) |
| |
| 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 |
| |
| 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() |
|
|