import os
import sys
import time
import tensorrt as trt

onnx_path = "/workspace/purescale_4k_end2end.onnx"
engine_path = "/workspace/purescale_4k.engine"

print(f"=== [TensorRT Builder] Building {engine_path} from {onnx_path} ===")
t0 = time.time()

logger = trt.Logger(trt.Logger.INFO)
builder = trt.Builder(logger)
network = builder.create_network()
parser = trt.OnnxParser(network, logger)

with open(onnx_path, "rb") as f:
    if not parser.parse(f.read()):
        print("[-] Lỗi parse ONNX:")
        for error in range(parser.num_errors):
            print(parser.get_error(error))
        sys.exit(1)
        
print(f"✓ Parsed ONNX graph successfully! Inputs: {network.num_inputs}, Outputs: {network.num_outputs}")

config = builder.create_builder_config()

# Enable TF32 for fast Tensor Core acceleration
if hasattr(trt.BuilderFlag, "TF32"):
    config.set_flag(trt.BuilderFlag.TF32)

# Allocate workspace pool (4GB)
config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 4 * 1024 * 1024 * 1024)

print("[*] Compiling CUDA execution engine on RTX 4090...")
serialized_engine = builder.build_serialized_network(network, config)
if serialized_engine is None:
    print("[-] Build failed: serialized_engine is None!")
    sys.exit(1)

with open(engine_path, "wb") as f:
    f.write(serialized_engine)
        
dur = time.time() - t0
size_mb = os.path.getsize(engine_path) / (1024 * 1024)
print(f"✓ TensorRT Engine built successfully: {engine_path} ({size_mb:.1f} MB in {dur:.1f}s)")
