AI Logits处理器实践笔记
内容生成工具要求模型输出严格 JSON:status 只能是三个枚举值,level 必须是 1–10 的整数,还不能碰敏感词。光靠 prompt 写「请只输出合法 JSON」,跑一百次总有几条字段越界,json.loads 直接炸。
后来在解码阶段加 Logits 处理器,不合法 token 的概率直接压成零。这篇记从 HuggingFace 接入到自定义过滤器的做法。
问题背景
我在做一个内容生成工具,核心需求是让模型生成结构化数据。具体来说,需要输出 JSON 格式的字段,每个字段的值都有严格的限制:
status只能是active、inactive、pendinglevel必须是 1 到 10 的整数priority不能包含某些敏感词汇
一开始用的是常规的 prompt 方案:
system_prompt = """
你是一个数据生成助手。请输出符合以下格式的 JSON:
{
"status": "active|inactive|pending",
"level": 1-10 的整数,
"priority": 文本,不能包含敏感词
}
"""
然后现实是:模型有时候会输出 "status": "unknown",有时候 level 会变成 "11",priority 字段更是时不时冒出来不该出现的词汇。
prompt 层面的约束确实有效,但不完全可靠。你要的是"必须满足",而不是"最好满足"。
Logits 是什么?
先说人话版:logits 就是模型在某个位置对"下一个词是什么"的原始打分。
比如模型当前要生成下一个 token,它会计算整个词表中每个 token 的"可能性得分"。这个得分就叫 logits。通过 softmax 处理后,就变成了我们常说的概率分布。
打个比方:模型在生成句子时,它在某个点想了"下一步可以是什么",然后给了每个候选词一个分数。logits 就是这些分数,越高说明模型越倾向于选这个词。
关键点来了:这些分数是可以干预的。
如果你在采样之前,把某些 token 的 logits 调低或者调成负无穷,模型就会"认为"这些词不应该出现。这就是 logits 处理器的基本思路。
这个过滤过程大概是这个样子的:
基本实现思路
OpenAI 的 Python SDK 已经提供了 logprobs 参数,可以拿到模型的原始 logits 输出。但那是用来"看"的,不是用来"改"的。
要真正干预 logits,你需要在模型输出的 logits 上做操作,然后再进行采样。流程大概是这样的:
具体到代码层面,不同的框架实现方式不太一样。我用的是 Hugging Face Transformers,它提供了 LogitsProcessor 接口:
from transformers import LogitsProcessor
class CustomLogitsProcessor(LogitsProcessor):
def __init__(self, allowed_tokens):
self.allowed_tokens = allowed_tokens
def __call__(self, input_ids, scores):
# scores 就是 logits,形状是 (batch_size, vocab_size)
# 把不允许的 token 分数调成负无穷
mask = torch.ones_like(scores[0], dtype=torch.bool)
mask[self.allowed_tokens] = False
scores[0][mask] = float('-inf')
return scores
然后把它塞到生成配置里:
from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained("gpt2")
tokenizer = AutoTokenizer.from_pretrained("gpt2")
# 假设我们要限制某个位置只能输出特定词汇
allowed_status_tokens = [
tokenizer.convert_tokens_to_ids("active"),
tokenizer.convert_tokens_to_ids("inactive"),
tokenizer.convert_tokens_to_ids("pending"),
]
processor = CustomLogitsProcessor(allowed_status_tokens)
output = model.generate(
input_ids,
logits_processor=[processor],
max_new_tokens=10
)
这样模型在那个位置就只能输出 active、inactive 或 pending 中的一个,其他的都会被硬过滤掉。
第一次踩坑:分词的坑
刚实现完的时候,我以为这样就解决了问题。然后测试的时候发现:模型还是输出了 "status": "Active"(大写 A),然后 JSON 解析直接报错。
问题出在分词上。active 和 Active 在词表里可能是不同的 token,甚至可能被拆成多个子词。你只过滤了小写版本的 token,大写版本照样能输出。
解决方法有两个:
- 把所有可能的变体都加到 allowed_tokens 里
- 在 processor 里做 case-insensitive 处理
我选了第二种,因为更健壮:
class CaseInsensitiveLogitsProcessor(LogitsProcessor):
def __init__(self, allowed_words, tokenizer):
self.allowed_words = set(w.lower() for w in allowed_words)
self.tokenizer = tokenizer
def __call__(self, input_ids, scores):
vocab = self.tokenizer.get_vocab()
# 找出所有 allowed_words 对应的 token(包括大小写变体)
allowed_ids = set()
for word, idx in vocab.items():
if word.lower() in self.allowed_words:
allowed_ids.add(idx)
# 过滤
mask = torch.ones_like(scores[0], dtype=torch.bool)
mask[list(allowed_ids)] = False
scores[0][mask] = float('-inf')
return scores
这样不管模型想输出 active、Active 还是 ACTIVE,只要词根对就行。
第二次踩坑:上下文感知
接着来了新需求:level 字段的限制是 1 到 10,但这个限制只在 JSON 对应的字段位置生效。其他位置(比如描述文字里)应该可以正常出现数字。
直接用 processor 会把所有位置的 11-20 都过滤掉,这显然不对。
需要让 processor 知道"当前在哪个位置"。一个做法是解析当前已生成的内容,判断是不是在 level 字段。
用一个状态机来理解这个需求会更直观:
import re
class ContextAwareLogitsProcessor(LogitsProcessor):
def __init__(self, tokenizer):
self.tokenizer = tokenizer
self.number_tokens = {
i: tokenizer.convert_tokens_to_ids(str(i))
for i in range(1, 11)
}
def _is_in_level_field(self, input_ids):
# 把已生成的 token 转回文本
text = self.tokenizer.decode(input_ids[0])
# 判断是不是在 `"level": ` 后面
pattern = r'"level"\s*:\s*(\d*)$'
match = re.search(pattern, text[-50:]) # 只看最后 50 个字符
return match is not None
def __call__(self, input_ids, scores):
if self._is_in_level_field(input_ids):
# 只允许 1-10 的数字 token
mask = torch.ones_like(scores[0], dtype=torch.bool)
for num, token_id in self.number_tokens.items():
mask[token_id] = False
scores[0][mask] = float('-inf')
return scores
这个方法能工作,但有明显的性能问题:每次生成一个 token 都要重新 decode 一次,然后跑正则匹配。
优化方案是维护一个简单的状态机,记录当前解析到了 JSON 的哪个字段。这样效率会高很多,但代码复杂度也上去了。
第三次踩坑:多个处理器协同
到这一步,我有了三个处理器:
- 状态处理器:限制 status 字段只能用特定词
- 数字处理器:限制 level 字段只能用 1-10
- 敏感词处理器:过滤所有位置的关键词
然后把它们一起传给生成:
processors = [
StatusLogitsProcessor(tokenizer),
LevelLogitsProcessor(tokenizer),
SensitiveWordLogitsProcessor(tokenizer)
]
output = model.generate(input_ids, logits_processor=processors)
问题来了:这些处理器之间可能有冲突。比如敏感词列表里包含 active,那么状态处理器想让它输出 active,敏感词处理器却把它过滤掉了。
处理这个问题的思路是明确优先级。比如敏感词应该是最高优先级,状态处理器次之:
def __call__(self, input_ids, scores):
# 先执行自己的过滤
scores = super().__call__(input_ids, scores)
# 再调用下一个处理器,传递已经处理过的 scores
if self.next_processor:
scores = self.next_processor(input_ids, scores)
return scores
或者更直接的方案:把所有约束合并成一个 processor,内部统一处理优先级。
最终落地的方案
经过几轮迭代,最终的方案是:一个统一的处理器,内部维护字段状态和规则列表。
规则处理的大致流程是这样的:
class UnifiedLogitsProcessor(LogitsProcessor):
def __init__(self, tokenizer):
self.tokenizer = tokenizer
self.rules = {
'status': {
'type': 'enum',
'values': ['active', 'inactive', 'pending']
},
'level': {
'type': 'range',
'min': 1,
'max': 10
},
'_global': {
'type': 'exclude',
'words': ['敏感词1', '敏感词2']
}
}
def _get_current_field(self, input_ids):
text = self.tokenizer.decode(input_ids[0][-100:])
# 解析当前在哪个字段,返回 field_name 或 None
...
def __call__(self, input_ids, scores):
current_field = self._get_current_field(input_ids)
mask = torch.ones_like(scores[0], dtype=torch.bool)
# 全局过滤规则
global_rule = self.rules['_global']
for word in global_rule['words']:
token_ids = self._get_token_ids_for_word(word)
for tid in token_ids:
mask[tid] = False
# 字段级过滤规则
if current_field and current_field in self.rules:
rule = self.rules[current_field]
if rule['type'] == 'enum':
for value in rule['values']:
token_ids = self._get_token_ids_for_word(value)
for tid in token_ids:
mask[tid] = False # 注意这里是取反
elif rule['type'] == 'range':
for i in range(rule['min'], rule['max'] + 1):
tid = self.tokenizer.convert_tokens_to_ids(str(i))
mask[tid] = False
# 应用过滤
scores[0][mask] = float('-inf')
return scores
注意代码里有个小细节:mask 的逻辑是反的。mask[tid] = False 意味着"保留这个 token",mask[tid] = True 意味着"过滤掉这个 token"。
这个方案在性能和可维护性之间取了个平衡。解析当前字段的逻辑还是有点糙,但在实际场景中已经够用了。
实际效果
用这套方案跑了几个测试用例:
测试 1:基本字段约束
输入:生成一个包含 status 的 JSON
输出:
{
"status": "active",
"level": 7
}
结果:符合预期,status 字段没有出现不在允许列表里的词。
测试 2:边界情况
输入:强制让模型尝试输出 level=11
输出:模型自动跳过了 11,输出了 10
结果:硬约束生效,模型在遇到被过滤的 token 时会选择下一个合理的选项。
测试 3:敏感词过滤
输入:生成包含敏感词的描述
输出:描述里没有出现敏感词,模型自动用了同义词替换
结果:全局过滤规则正常工作。
但也发现了一些问题:当过滤规则太严格的时候,模型生成的流畅性会下降,有时候会陷入重复某个词的死循环。这说明 logits 处理器不是万能药,过度约束会损害生成质量。
还有其他路子吗?
除了 logits 处理器,还有几个思路:
后处理验证:先生成,再验证。不符合就重新生成。优点是简单,缺点是可能一直重复失败。
Structured Generation:像
guidance、outlines这样的库,专门做结构化生成。它们内部也是用 logits 过滤,但封装得更好用。微调模型:让模型学会只在特定位置输出特定内容。这个成本高,但长期效果好。
对于我的场景,logits 处理器已经够用了。但如果你的约束更复杂或者对生成质量要求更高,可以考虑 Structured Generation 库。
写在最后
logits 处理器本质上是给模型的"自由表达"加了一道闸。这道闸在你需要硬约束的时候很有用,但记住一个原则:能用 prompt 解决的问题,就不要动 logits。
prompt 更符合模型的"自然"行为,而 logits 过滤是强行干预。过度干预会让模型的行为变得僵硬,生成质量也会受影响。
就像和人沟通一样:你可以说清楚你的期望,也可以直接把某些选项划掉。前者更自然,后者更严格。选哪种,看你的需求有多硬。
折腾这一轮下来,最大的收获不是学会了 logits 过滤,而是重新理解了"约束"和"自由"之间的平衡。太自由的模型会失控,太严格的模型会失真。找到那个平衡点,才是做 AI 应用时需要持续优化的东西。
版权声明: 本文首发于 指尖魔法屋-AI Logits处理器实践笔记(https://blog.thinkmoon.cn/post/354-ai-logits-processor-raw-filter-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。