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