You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何解决TFT时间序列预测模型导出ONNX后输入参数缺失问题

解决TFT模型导出ONNX时输入参数缺失的问题

问题根源

你遇到的输入参数缺失问题,核心是构造的dummy input没有覆盖TFT模型forward方法的所有必填输入,同时包装模型的参数签名和TFT实际要求不匹配。pytorch-forecasting中的TFT模型forward方法需要的输入比你提取的多一个encoder_time_idx,且部分参数的传递逻辑需要对齐模型内部要求。

解决方案步骤

1. 确认TFT模型的完整输入参数

先打印模型的forward签名,明确所有必填项:

print(tft.forward)

TFT的标准forward输入参数为:encoder_cat, encoder_cont, encoder_time_idx, decoder_cat, decoder_cont, decoder_time_idx, encoder_target, decoder_target, encoder_lengths, decoder_lengths, groups, target_scale

2. 修正Dummy Input的构造

从验证数据加载器的batch中提取所有必填张量,补充遗漏的encoder_time_idx:

dummy_batch, _ = next(iter(val_dataloader))

# 提取所有必填输入张量
encoder_cat = dummy_batch["encoder_cat"]
encoder_cont = dummy_batch["encoder_cont"]
encoder_time_idx = dummy_batch["encoder_time_idx"]  # 补充遗漏的输入
decoder_cat = dummy_batch["decoder_cat"]
decoder_cont = dummy_batch["decoder_cont"]
decoder_time_idx = dummy_batch["decoder_time_idx"]
encoder_target = dummy_batch["encoder_target"]
decoder_target = dummy_batch["decoder_target"]
encoder_lengths = dummy_batch["encoder_lengths"]
decoder_lengths = dummy_batch["decoder_lengths"]
groups = dummy_batch["groups"]
target_scale = dummy_batch["target_scale"]

3. 调整包装模型的Forward逻辑

包装模型的forward参数要和TFT完全一致,避免字典传递可能的解析问题,直接按关键字参数传递:

class WrappedModel(torch.nn.Module):
    def __init__(self, model):
        super().__init__()
        self.model = model

    def forward(self, encoder_cat, encoder_cont, encoder_time_idx,
                decoder_cat, decoder_cont, decoder_time_idx,
                encoder_target, decoder_target, encoder_lengths,
                decoder_lengths, groups, target_scale):
        # 直接按TFT要求的参数传递,避免字典解析误差
        return self.model(
            encoder_cat=encoder_cat,
            encoder_cont=encoder_cont,
            encoder_time_idx=encoder_time_idx,
            decoder_cat=decoder_cat,
            decoder_cont=decoder_cont,
            decoder_time_idx=decoder_time_idx,
            encoder_target=encoder_target,
            decoder_target=decoder_target,
            encoder_lengths=encoder_lengths,
            decoder_lengths=decoder_lengths,
            groups=groups,
            target_scale=target_scale
        )

wrapped_model = WrappedModel(tft)

4. 修正ONNX导出配置

更新输入元组、input_names和dynamic_axes,确保所有输入都被正确声明:

# 构造完整的输入元组
dummy_input_tuple = (
    encoder_cat, encoder_cont, encoder_time_idx,
    decoder_cat, decoder_cont, decoder_time_idx,
    encoder_target, decoder_target, encoder_lengths,
    decoder_lengths, groups, target_scale
)

# 定义所有输入名称
input_names = [
    "encoder_cat", "encoder_cont", "encoder_time_idx",
    "decoder_cat", "decoder_cont", "decoder_time_idx",
    "encoder_target", "decoder_target", "encoder_lengths",
    "decoder_lengths", "groups", "target_scale"
]

# 更新动态轴配置
dynamic_axes = {
    "encoder_cat": {0: "batch_size", 1: "sequence_length"},
    "encoder_cont": {0: "batch_size", 1: "sequence_length"},
    "encoder_time_idx": {0: "batch_size", 1: "sequence_length"},  # 新增
    "decoder_cat": {0: "batch_size", 1: "sequence_length"},
    "decoder_cont": {0: "batch_size", 1: "sequence_length"},
    "decoder_time_idx": {0: "batch_size", 1: "sequence_length"},
    "encoder_target": {0: "batch_size", 1: "sequence_length"},
    "decoder_target": {0: "batch_size", 1: "sequence_length"},
    "encoder_lengths": {0: "batch_size"},
    "decoder_lengths": {0: "batch_size"},
    "groups": {0: "batch_size"},
    "target_scale": {0: "batch_size", 1: "num_features"},
    "output": {0: "batch_size", 1: "sequence_length"}
}

# 导出ONNX模型
torch.onnx.export(
    wrapped_model,
    dummy_input_tuple,
    onnx_model_path,
    input_names=input_names,
    output_names=["output"],
    dynamic_axes=dynamic_axes,
    opset_version=13,  # 升级opset版本,提升兼容性
    do_constant_folding=True,
    verbose=False
)

print(f"Model exported to {onnx_model_path}")

5. 验证ONNX模型加载

用ONNX Runtime验证加载是否正常:

import onnxruntime as ort

# 加载模型
sess = ort.InferenceSession(onnx_model_path)

# 检查输入输出名称
print("模型输入名称:", [inp.name for inp in sess.get_inputs()])
print("模型输出名称:", [out.name for out in sess.get_outputs()])

关键修改点说明

  • 补充了遗漏的encoder_time_idx输入,这是TFT模型forward的必填参数
  • 包装模型直接使用关键字参数传递,避免字典解析时可能的参数匹配错误
  • 升级opset版本到13,提升和ONNX Runtime的兼容性
  • 确保dynamic_axes覆盖所有输入张量,避免静态维度限制

内容的提问来源于stack exchange,提问作者Sudeeksha Vandrangi

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 04:07:33