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