File size: 4,176 Bytes
2eb3593
 
 
 
 
 
 
 
 
 
 
 
 
ebdc7e7
2eb3593
 
7acb541
 
 
 
 
 
 
 
 
 
ebdc7e7
3254de2
ebdc7e7
 
 
7acb541
 
2eb3593
7acb541
2eb3593
 
 
 
 
 
 
 
 
 
 
 
7acb541
 
 
 
 
 
 
 
 
2eb3593
 
 
 
 
 
 
 
 
 
 
7acb541
2eb3593
 
 
 
ebdc7e7
 
2eb3593
ebdc7e7
 
 
 
 
2eb3593
ebdc7e7
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
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
from __future__ import annotations
import argparse, json, time
from pathlib import Path
import tensorrt as trt
import torch

LOGGER = trt.Logger(trt.Logger.INFO)

COMPONENTS = {
    "ss-flow-solo": ("onnx/anigen/ss-flow-solo/model.onnx", "engines/l4-sm89/anigen-ss-flow-solo.plan"),
    "ss-decoder": ("onnx/anigen/ss-decoder/model.onnx", "engines/l4-sm89/anigen-ss-decoder.plan"),
    "slat-flow-core": ("onnx/anigen/slat-flow-core/model.onnx", "engines/l4-sm89/anigen-slat-flow-core.plan"),
    "slat-decoder-core": ("onnx/anigen/slat-decoder-core/model.onnx", "engines/l4-sm89/anigen-slat-decoder-core.plan"),
    "skin-decoder": ("onnx/anigen/skin-decoder/model.onnx", "engines/l4-sm89/anigen-skin-decoder.plan"),
}

PROFILES = {
    "slat-flow-core": {
        "geo": ((1,128,1024),(1,4096,1024),(1,16384,1024)),
        "skin": ((1,128,512),(1,4096,512),(1,16384,512)),
        "skl": ((1,16,512),(1,1024,512),(1,8192,512)),
        "mod_geo": ((1,1024),(1,1024),(1,1024)),
        "mod_skin": ((1,512),(1,512),(1,512)),
        "mod_skl": ((1,512),(1,512),(1,512)),
        "cond": ((1,1374,1024),(1,1374,1024),(1,1374,1024)),
    },
    "skin-decoder": {
        "vertex_features": ((1,128,4),(1,25000,4),(1,300000,4)),
        "joint_features": ((1,2,4),(1,24,4),(1,64,4)),
        "parents": ((1,2),(1,24),(1,64)),
    },
}


def build(name: str, src: Path, dst: Path, workspace_gib: int, opt_level: int):
    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, workspace_gib * (1 << 30))
    if hasattr(config, "builder_optimization_level"):
        config.builder_optimization_level = opt_level
    if name in PROFILES:
        profile = builder.create_optimization_profile()
        by_name = PROFILES[name]
        for i in range(network.num_inputs):
            tensor = network.get_input(i)
            if tensor.name in by_name:
                mn, op, mx = by_name[tensor.name]
                profile.set_shape(tensor.name, mn, op, mx)
        config.add_optimization_profile(profile)
    started = 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)
    return {
        "seconds": round(time.time() - started, 3),
        "bytes": dst.stat().st_size,
        "inputs": {network.get_input(i).name: list(network.get_input(i).shape) for i in range(network.num_inputs)},
        "outputs": {network.get_output(i).name: list(network.get_output(i).shape) for i in range(network.num_outputs)},
        "profile": PROFILES.get(name),
    }


def main():
    ap = argparse.ArgumentParser(); ap.add_argument("--root", default="/tmp/cf-anigen-onnx"); ap.add_argument("--component", choices=list(COMPONENTS)+["all"], default="all"); ap.add_argument("--workspace-gib", type=int, default=6); ap.add_argument("--opt-level", type=int, default=4); a=ap.parse_args()
    root=Path(a.root); names=list(COMPONENTS) if a.component=='all' else [a.component]; result={}
    for name in names:
        src_rel,dst_rel=COMPONENTS[name]; src,dst=root/src_rel,root/dst_rel
        if not src.exists(): print(f"SKIP {name}: {src} missing",flush=True); continue
        print(f"BUILD {name}: {src} -> {dst}",flush=True); result[name]=build(name,src,dst,a.workspace_gib,a.opt_level)
    result['environment']={'gpu':torch.cuda.get_device_name(0),'cc':list(torch.cuda.get_device_capability(0)),'cuda':torch.version.cuda,'tensorrt':trt.__version__}
    meta=root/f'trt_build_anigen_{a.component}_meta.json'; meta.write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2),flush=True)

if __name__=='__main__': main()