Companion-Forge-L4-ONNX / scripts /build_flux2_trt.py
patdev's picture
Add static FLUX.2 Klein TensorRT L4 builder
67de8ac verified
Raw
History Blame Contribute Delete
1.76 kB
from __future__ import annotations
import argparse,json,time
from pathlib import Path
import tensorrt as trt
LOGGER=trt.Logger(trt.Logger.INFO)
def main():
ap=argparse.ArgumentParser();ap.add_argument('--root',default='.');ap.add_argument('--size',type=int,default=512);a=ap.parse_args()
root=Path(a.root);src=root/f'onnx/flux2-klein-4b/{a.size}/transformer.onnx';dst=root/f'engines/l4-sm89/flux2-klein-4b-{a.size}.plan'
builder=trt.Builder(LOGGER)
flags=0
if hasattr(trt.NetworkDefinitionCreationFlag,'STRONGLY_TYPED'):
flags |= 1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)
network=builder.create_network(flags);parser=trt.OnnxParser(network,LOGGER)
# parse_from_file is required so TensorRT can resolve transformer.onnx.data beside the graph.
if not parser.parse_from_file(str(src)):
raise RuntimeError('\n'.join(str(parser.get_error(i)) for i in range(parser.num_errors)))
config=builder.create_builder_config();config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE,3*(1<<30))
if hasattr(config,'builder_optimization_level'):config.builder_optimization_level=3
t=time.time();blob=builder.build_serialized_network(network,config)
if blob is None:raise RuntimeError('TensorRT build returned None')
dst.parent.mkdir(parents=True,exist_ok=True);dst.write_bytes(blob)
import torch
meta={'size':a.size,'seconds':round(time.time()-t,3),'bytes':dst.stat().st_size,'gpu':torch.cuda.get_device_name(0),'cc':list(torch.cuda.get_device_capability()),'cuda':torch.version.cuda,'tensorrt':trt.__version__,'precision':'fp16-static'}
(root/f'flux2_trt_{a.size}_meta.json').write_text(json.dumps(meta,indent=2));print(json.dumps(meta,indent=2))
if __name__=='__main__':main()