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