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

如何阻止Transformer生成函数输出特定词汇?(T5模型场景)

解决T5模型生成时阻止特定词汇的方法

当然可以实现,下面是两种常用的方案:

方法一:使用bad_words_ids参数直接阻止目标词汇

transformers库的generate方法自带bad_words_ids参数,能让模型避免生成指定的词汇序列。具体操作如下:

  1. 先把要阻止的词汇转换成对应的token ID列表:
stopwords = ["park", "offers", "offer"]  # 注意包含衍生形式,比如示例里的"offers"
bad_words_ids = [tokenizer.encode(word, add_special_tokens=False) for word in stopwords]
  1. 在调用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:

  1. 定义自定义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
  1. 获取要阻止的词汇对应的所有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))
  1. 调用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 16:52:39