如何限制Huggingface BERT encoder-decoder模型解码器的生成词汇
限制Huggingface BERT encoder-decoder模型解码器词汇范围的方法
有三种常用的实现方案,可根据你的使用场景选择:
1. 生成时通过prefix_allowed_tokens_fn参数动态限制
这是无需修改模型结构、上手最快的方案,适合临时限制词汇的场景:
- 首先将你指定的小型词汇表转换为模型tokenizer对应的id列表,同时将生成必需的特殊token(开始符、结束符、填充符等)加入列表
- 自定义允许token返回函数,在调用
generate方法时传入该函数即可,解码器会仅从你指定的token中选取得分最高的结果 - 示例代码如下:
# 预处理允许的token id列表 allowed_token_ids = [tokenizer.convert_tokens_to_ids(tok) for tok in your_custom_vocab] # 补充必要特殊token allowed_token_ids += [tokenizer.cls_token_id, tokenizer.sep_token_id, tokenizer.pad_token_id] # 定义允许token返回函数 def allowed_tokens_fn(batch_id, input_ids): return allowed_token_ids # 生成时传入参数 outputs = model.generate( input_ids=input_ids, prefix_allowed_tokens_fn=allowed_tokens_fn, max_new_tokens=64, # 其余生成参数 )
2. 固定修改模型输出层
如果长期固定使用该小型词汇表,推荐直接修改模型结构,提升推理和训练效率:
- 先构建小型词汇表的id映射,将原生词汇表中你需要的token映射为新的连续id
- 提取原生lm_head层中对应允许token的权重,替换模型的lm_head层为输出维度匹配小型词汇表大小的新层
- 后续微调、推理都直接使用修改后的模型,无需额外加过滤逻辑,运行速度比第一种方案更快
3. 生成后过滤(不推荐)
仅适合极小批量临时测试场景:生成完整文本后逐token校验,将不在允许列表的token替换为指定备选token,该方案效率低,且容易出现语义不通的问题。
注意事项
无论使用哪种方案,都需要将生成必需的特殊token加入允许词汇列表,否则会出现生成无法终止、输出乱码的异常。
内容的提问来源于stack exchange,提问作者Joseph Harvey
相关产品推荐
相关产品推荐

