Companion-Forge-L4-ONNX / scripts /export_flux2_transformer.py
patdev's picture
Use streaming ONNX checker for multi-GB FLUX.2 export
1892c1c verified
Raw
History Blame Contribute Delete
3.54 kB
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()