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