File size: 1,762 Bytes
67de8ac
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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()