Torch-TensorRT编译报错:bool类型不支持ONNX导出求解
解决Torch-TensorRT编译时"图降低过程中遇到未知类型bool,该类型不支持ONNX导出"的问题
错误原因
模型结构中存在ONNX导出不支持的bool类型张量操作或参数,Torch-TensorRT在转换流程中依赖ONNX导出,因此触发该报错。
具体解决方法
排查并修改模型中的bool类型逻辑
- 检查
ResNet2的代码,定位所有涉及torch.bool类型的张量运算、条件分支或可训练参数。 - 将bool张量转换为
torch.float32类型参与计算,若后续需要判断逻辑,仅在PyTorch前向传播的Python代码中转回bool,不要保留在计算图内。 - 替换计算图中的动态条件分支(如基于张量值的
if/else),改用torch.where等张量运算实现相同逻辑。
- 检查
切换Torch-TensorRT编译的IR模式
在compile函数中指定使用TorchScript IR,跳过ONNX导出步骤:import torch import torch_tensorrt torch.hub._validate_not_a_forked_repo=lambda a,b,c: True from src.models.resnet2 import ResNet2 # 加载模型并设置为评估模式 model = ResNet2(output_size = 2) model.load_state_dict(torch.load('Epoch_9_Valacc_0.911_9_.pth')) model.eval() # 新增评估模式设置 # 改用TorchScript IR编译 trt_model = torch_tensorrt.compile(model, inputs= [torch_tensorrt.Input((1, 3, 640, 640))], enabled_precisions= {torch.half}, debug=True, ir="torchscript" # 指定IR类型 )单独测试ONNX导出(可选)
如果必须使用ONNX路径,先单独执行ONNX导出命令定位具体问题:model.eval() torch.onnx.export(model, torch.randn(1,3,640,640), "test_model.onnx")根据ONNX导出的报错信息,针对性修改模型中的不兼容部分。
排查模型参数类型
检查是否存在bool类型的可训练参数,若有则转换为支持的类型:for name, param in model.named_parameters(): if param.dtype == torch.bool: print(f"发现bool类型参数: {name}") param.data = param.data.float() # 转换为float类型
内容的提问来源于stack exchange,提问作者rakeshKM
相关产品推荐
相关产品推荐

