如何阻止Transformer生成函数输出特定词汇?(T5模型场景)
解决T5模型生成时阻止特定词汇的方法
当然可以实现,下面是两种常用的方案:
方法一:使用bad_words_ids参数直接阻止目标词汇
transformers库的generate方法自带bad_words_ids参数,能让模型避免生成指定的词汇序列。具体操作如下:
- 先把要阻止的词汇转换成对应的token ID列表:
stopwords = ["park", "offers", "offer"] # 注意包含衍生形式,比如示例里的"offers" bad_words_ids = [tokenizer.encode(word, add_special_tokens=False) for word in stopwords]
- 在调用
generate时传入该参数:
sequence_ids = model.generate(input_ids, bad_words_ids=bad_words_ids) sequences = tokenizer.batch_decode(sequence_ids) print(sequences)
这个参数会匹配连续的token序列,只要模型生成的token序列和bad_words_ids里的某一项完全一致,就会被阻止。
方法二:通过修改logits强制降低目标词汇的生成概率
如果需要更精细的控制,可以自定义Logits处理器,把目标词汇对应的token的logit设为负无穷,让模型几乎不会选择这些token:
- 定义自定义Logits处理器:
from transformers import LogitsProcessor class StopWordsProcessor(LogitsProcessor): def __init__(self, stop_token_ids): self.stop_token_ids = stop_token_ids def __call__(self, input_ids, scores): # 将目标token的logit设为负无穷,彻底排除生成可能 for token_id in self.stop_token_ids: scores[:, token_id] = -float("inf") return scores
- 获取要阻止的词汇对应的所有token ID(包括子词拆分后的token):
stopwords = ["park", "offer"] stop_token_ids = [] for word in stopwords: tokens = tokenizer.encode(word, add_special_tokens=False) stop_token_ids.extend(tokens) # 去重避免重复处理 stop_token_ids = list(set(stop_token_ids))
- 调用
generate时传入处理器:
processor = StopWordsProcessor(stop_token_ids) sequence_ids = model.generate(input_ids, logits_processor=[processor]) sequences = tokenizer.batch_decode(sequence_ids) print(sequences)
需要注意:第二种方法针对单个token,如果目标词汇被拆分成多个子词,要把所有子词的token ID都加入阻止列表;而第一种方法更适合阻止完整的词汇或短语。
内容的提问来源于stack exchange,提问作者littleworth
相关产品推荐
相关产品推荐

