通过AWS SageMaker使用Facebook M2M-100模型时如何指定forced_bos_token_id
解决方案
你遇到的参数失效问题有两个核心原因:
- SageMaker HuggingFace 推理容器规定,所有生成控制参数必须放在
parameters字段下,不能与inputs平级 forced_bos_token_id需要传入对应语言的数字ID而非字符串代码,同时M2M100要求分词前必须指定源语言,默认推理容器没有适配这个逻辑,需要自定义推理脚本处理参数传递。
完整实现步骤
步骤1:编写自定义推理脚本
在本地创建code目录,在目录下新建inference.py,内容如下:
from transformers import M2M100ForConditionalGeneration, M2M100Tokenizer def model_fn(model_dir): # 加载模型和分词器 model = M2M100ForConditionalGeneration.from_pretrained(model_dir) tokenizer = M2M100Tokenizer.from_pretrained(model_dir) return {"model": model, "tokenizer": tokenizer} def predict_fn(data, model_and_tokenizer): model = model_and_tokenizer["model"] tokenizer = model_and_tokenizer["tokenizer"] # 获取请求参数 inputs = data.pop("inputs", data) src_lang = data.pop("src_lang", "en") tgt_lang = data.pop("tgt_lang", "zh") # 设置源语言 tokenizer.src_lang = src_lang encoded_input = tokenizer(inputs, return_tensors="pt").to(model.device) # 生成翻译结果 generated_tokens = model.generate( **encoded_input, forced_bos_token_id=tokenizer.get_lang_id(tgt_lang), **data ) # 解码返回 return tokenizer.batch_decode(generated_tokens, skip_special_tokens=True)
步骤2:修改模型部署代码
更新原来的部署代码,指定自定义脚本路径:
from sagemaker.huggingface import HuggingFaceModel import sagemaker role = sagemaker.get_execution_role() hub = { 'HF_MODEL_ID':'facebook/m2m100_1.2B', 'HF_TASK':'text2text-generation' } huggingface_model = HuggingFaceModel( transformers_version='4.17.0', # 建议升级到更高版本适配M2M100 pytorch_version='1.10.2', py_version='py38', env=hub, role=role, source_dir="./code", # 自定义脚本所在目录 entry_point="inference.py" # 自定义脚本文件名 ) # 部署端点 predictor = huggingface_model.deploy( initial_instance_count=1, instance_type='ml.g4dn.xlarge' # 建议用GPU实例,1.2B参数CPU推理速度极慢 )
步骤3:调用时指定源语言和目标语言
部署完成后,按以下格式调用即可:
# 示例:中文翻译为法语 result = predictor.predict({ "inputs": "生活就像一盒巧克力。", "src_lang": "zh", "tgt_lang": "fr" }) print(result) # 输出:["La vie est comme une boîte de chocolat."]
临时替代方案(无需重新部署)
如果暂时不想修改部署逻辑,也可以提前在本地拿到对应语言的数字ID,按以下格式调用(仅当你的部署版本默认设置了正确的源语言时生效,不推荐长期使用):
# 示例:源语言默认是英文,翻译为法语 result = predictor.predict({ "inputs": "The answer to the universe is", "parameters": { "forced_bos_token_id": 135 # 法语fr对应的数字ID,可提前通过tokenizer.get_lang_id("fr")获取 } })
内容的提问来源于stack exchange,提问作者rudolfovic
相关产品推荐
相关产品推荐

