File size: 3,542 Bytes
8aca026
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
cc846b9
8aca026
 
cc846b9
8aca026
 
 
f17038f
8aca026
f17038f
1892c1c
f17038f
8aca026
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
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()