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

将自定义模型部署到OpenSearch时遇RuntimeError: KeyError: token_type_ids

解决OpenSearch部署TorchScript模型时的token_type_ids KeyError问题

问题场景

将HuggingFace的sentence-transformers/LaBSE模型导出为TorchScript格式后,在OpenSearch中注册成功但部署失败,报错RuntimeError: KeyError: token_type_ids,即使导出时输入字典包含token_type_ids键仍触发错误。

错误日志

{
  "model_id": "Guem7I4BM4PGcXAkFYKV",
  "task_type": "DEPLOY_MODEL",
  "function_name": "TEXT_EMBEDDING",
  "state": "FAILED",
  "worker_node": [
    "d7zdcxaASDyFKGX7AhlG0g",
    "5fc7S9yJQQCWarLbAnVESg",
    "TPu-1KznRzCwj8tQAgHOMQ"
  ],
  "create_time": 1713367889874,
  "last_update_time": 1713368035088,
  "error": """                          {"d7zdcxaASDyFKGX7AhlG0g":"The following operation failed in the TorchScript interpreter.                          \nTraceback of TorchScript, serialized code (most recent call last):                          \n  File \"code/__torch__.py\", line 12, in forward                          \n    input_ids = inputs[\"input_ids\"]                          \n    attention_mask = inputs[\"attention_mask\"]                          \n    input = inputs[\"token_type_ids\"]                          \n            ~~~~~~~~~~~~~~~~~~~~~~~~ <--- HERE                          \n    _0 = (model).forward(input_ids, attention_mask, input, )                          \n    return {\"sentence_embedding\": _0}                          \n                          \nTraceback of TorchScript, original code (most recent call last):                          \n/Users/Library/Python/3.9/lib/python/site-packages/torch/jit/_trace.py(1074): trace_module                          \n/Users/Library/Python/3.9/lib/python/site-packages/torch/jit/_trace.py(806): trace                          \n/Users/opensearch/CustomModel/TEST/import_torch.py(43): export_to_torchscript                          \n/Users/opensearch/CustomModel/TEST/import_torch.py(52): <module>                          \nRuntimeError: KeyError: token_type_ids                          \n","5fc7S9yJQQCWarLbAnVESg":"The following operation failed in the TorchScript interpreter.                          \nTraceback of TorchScript, serialized code (most recent call last):                          \n  File \"code/__torch__.py\", line 12, in forward                          \n    input_ids = inputs[\"input_ids\"]                          \n    attention_mask = inputs[\"attention_mask\"]                          \n    input = inputs[\"token_type_ids\"]                          \n            ~~~~~~~~~~~~~~~~~~~~~~~~ <--- HERE                          \n    _0 = (model).forward(input_ids, attention_mask, input, )                          \n    return {\"sentence_embedding\": _0}                          \n                          \nTraceback of TorchScript, original code (most recent call last):                          \n/Users/Library/Python/3.9/lib/python/site-packages/torch/jit/_trace.py(1074): trace_module                          \n/Users/Library/Python/3.9/lib/python/site-packages/torch/jit/_trace.py(806): trace                          \n/Users/opensearch/CustomModel/TEST/import_torch.py(43): export_to_torchscript                          \n/Users/opensearch/CustomModel/TEST/import_torch.py(52): <module>                          \nRuntimeError: KeyError: token_type_ids                          \n","TPu-1KznRzCwj8tQAgHOMQ":"The following operation failed in the TorchScript interpreter.                          \nTraceback of TorchScript, serialized code (most recent call last):                          \n  File \"code/__torch__.py\", line 12, in forward                          \n    input_ids = inputs[\"input_ids\"]                          \n    attention_mask = inputs[\"attention_mask\"]                          \n    input = inputs[\"token_type_ids\"]                          \n            ~~~~~~~~~~~~~~~~~~~~~~~~ <--- HERE                          \n    _0 = (model).forward(input_ids, attention_mask, input, )                          \n    return {\"sentence_embedding\": _0}                          \n                          \nTraceback of TorchScript, original code (most recent call last):                          \n/Users/Library/Python/3.9/lib/python/site-packages/torch/jit/_trace.py(1074): trace_module                          \n/Users/Library/Python/3.9/lib/python/site-packages/torch/jit/_trace.py(806): trace                          \n/Users/opensearch/CustomModel/TEST/import_torch.py(43): export_to_torchscript                          \n/Users/opensearch/CustomModel/TEST/import_torch.py(52): <module>                          \nRuntimeError: KeyError: token_type_ids\n"}    """,
  "is_async": true
}

导出模型脚本

import torch
from transformers import AutoModel, AutoTokenizer, PreTrainedTokenizer
from transformers.utils import PaddingStrategy
from sentence_transformers import SentenceTransformer


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

    def forward(self, inputs: dict):
        with torch.no_grad():
            outputs = self.model(inputs)
        return {"sentence_embedding": outputs['sentence_embedding']}


def export_to_torchscript(model_name: str, is_sentence_transformer: bool, output_path: str, max_seq_length: int = 128):
    tokenizer: PreTrainedTokenizer = AutoTokenizer.from_pretrained(model_name)
    if is_sentence_transformer:
        model = SentenceTransformer(model_name, device="cpu")
    else:
        model = AutoModel.from_pretrained(model_name)
    model.eval()

    # Define example text
    text = "This is a test string"

    # Create inputs
    inputs = tokenizer(text, padding=PaddingStrategy.MAX_LENGTH, return_tensors="pt", max_length=max_seq_length)

    # Instantiate the wrapper class
    model_wrapper = TorchScriptWrapper(model)

    # Unpack HF batch encoding into a regular dict
    new_inputs = {}
    new_inputs["input_ids"] = inputs["input_ids"]
    new_inputs["attention_mask"] = inputs["attention_mask"]
    if inputs.get("token_type_ids", None) is not None:
        new_inputs["token_type_ids"] = inputs["token_type_ids"]

    # Trace the wrapper class
    traced_model = torch.jit.trace(model_wrapper, new_inputs, strict=False)

    # Save traced model to file
    traced_model.save(output_path)


if __name__ == "__main__":
    # Load pre-trained model and tokenizer
    model_name = "sentence-transformers/LaBSE"
    export_to_torchscript(model_name, True, output_path="torchscript_labse.pt")

问题原因

TorchScript的trace机制会基于示例输入生成固定执行逻辑,即使示例输入包含token_type_ids,但OpenSearch调用模型时可能未传入该参数,导致追踪生成的代码强制读取该键引发错误。此外,LaBSE模型本身并不依赖token_type_ids,单句嵌入场景无需该参数。

解决方案

方案1:修改Wrapper兼容缺失参数

调整TorchScriptWrapper的forward方法,仅在参数存在时传递:

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

    def forward(self, inputs: dict):
        with torch.no_grad():
            # 仅保留模型必需参数,条件添加token_type_ids
            model_inputs = {
                "input_ids": inputs["input_ids"],
                "attention_mask": inputs["attention_mask"]
            }
            if "token_type_ids" in inputs:
                model_inputs["token_type_ids"] = inputs["token_type_ids"]
            outputs = self.model(model_inputs)
        return {"sentence_embedding": outputs['sentence_embedding']}

方案2:改用Script模式导出

使用torch.jit.script替代trace,它会直接解析Python逻辑,更好处理条件分支:

# 替换原trace代码
traced_model = torch.jit.script(model_wrapper)

方案3:强制OpenSearch传入参数

在OpenSearch部署配置中确保分词器生成并传入token_type_ids,但灵活性较低,不推荐。

执行步骤

  1. 替换修改后的Wrapper类或导出方式
  2. 重新导出TorchScript模型
  3. 重新注册并部署到OpenSearch

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 21:34:57