将自定义模型部署到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,但灵活性较低,不推荐。
执行步骤
- 替换修改后的Wrapper类或导出方式
- 重新导出TorchScript模型
- 重新注册并部署到OpenSearch
内容的提问来源于stack exchange,提问作者SegTree
相关产品推荐
相关产品推荐

