AI Logits处理器实践笔记

内容生成工具要求模型输出严格 JSON:status 只能是三个枚举值,level 必须是 1–10 的整数,还不能碰敏感词。光靠 prompt 写「请只输出合法 JSON」,跑一百次总有几条字段越界,json.loads 直接炸。

后来在解码阶段加 Logits 处理器,不合法 token 的概率直接压成零。这篇记从 HuggingFace 接入到自定义过滤器的做法。

问题背景

我在做一个内容生成工具,核心需求是让模型生成结构化数据。具体来说,需要输出 JSON 格式的字段,每个字段的值都有严格的限制:

  • status 只能是 activeinactivepending
  • level 必须是 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 处理器的基本思路。

这个过滤过程大概是这个样子的:

graph TB subgraph 原始输出 A[词表<br/>active: 2.3<br/>inactive: 1.8<br/>unknown: 1.5<br/>pending: 2.1] end subgraph Logits处理器 B[识别目标token<br/>active, inactive, pending] C[过滤其他token<br/>unknown → -∞] end subgraph 过滤后 D[active: 2.3<br/>inactive: 1.8<br/>unknown: -∞<br/>pending: 2.1] end A --> B B --> C C --> D

基本实现思路

OpenAI 的 Python SDK 已经提供了 logprobs 参数,可以拿到模型的原始 logits 输出。但那是用来"看"的,不是用来"改"的。

要真正干预 logits,你需要在模型输出的 logits 上做操作,然后再进行采样。流程大概是这样的:

graph LR A[模型原始输出] --> B[Logits 处理器] B --> C[过滤后 logits] C --> D[采样] D --> E[生成下一个 token]

具体到代码层面,不同的框架实现方式不太一样。我用的是 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
)

这样模型在那个位置就只能输出 activeinactivepending 中的一个,其他的都会被硬过滤掉。

第一次踩坑:分词的坑

刚实现完的时候,我以为这样就解决了问题。然后测试的时候发现:模型还是输出了 "status": "Active"(大写 A),然后 JSON 解析直接报错。

问题出在分词上。activeActive 在词表里可能是不同的 token,甚至可能被拆成多个子词。你只过滤了小写版本的 token,大写版本照样能输出。

解决方法有两个:

  1. 把所有可能的变体都加到 allowed_tokens 里
  2. 在 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

这样不管模型想输出 activeActive 还是 ACTIVE,只要词根对就行。

第二次踩坑:上下文感知

接着来了新需求:level 字段的限制是 1 到 10,但这个限制只在 JSON 对应的字段位置生效。其他位置(比如描述文字里)应该可以正常出现数字。

直接用 processor 会把所有位置的 11-20 都过滤掉,这显然不对。

需要让 processor 知道"当前在哪个位置"。一个做法是解析当前已生成的内容,判断是不是在 level 字段。

用一个状态机来理解这个需求会更直观:

stateDiagram-v2 [*] --> JSON开始 JSON开始 --> status字段: 遇到 "status": JSON开始 --> level字段: 遇到 "level": JSON开始 --> priority字段: 遇到 "priority:" JSON开始 --> JSON结束: 遇到 "}" status字段 --> 字段间: 遇到 "," level字段 --> 字段间: 遇到 "," priority字段 --> 字段间: 遇到 "," 字段间 --> status字段: 遇到 "status": 字段间 --> level字段: 遇到 "level": 字段间 --> priority字段: 遇到 "priority:" 字段间 --> JSON结束: 遇到 "}" JSON结束 --> [*] note right of level字段 在这个状态时 只允许输出 1-10 end note
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 的哪个字段。这样效率会高很多,但代码复杂度也上去了。

第三次踩坑:多个处理器协同

到这一步,我有了三个处理器:

  1. 状态处理器:限制 status 字段只能用特定词
  2. 数字处理器:限制 level 字段只能用 1-10
  3. 敏感词处理器:过滤所有位置的关键词

然后把它们一起传给生成:

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,内部统一处理优先级。

最终落地的方案

经过几轮迭代,最终的方案是:一个统一的处理器,内部维护字段状态和规则列表。

规则处理的大致流程是这样的:

graph TD A[收到新的 input_ids] --> B[解析当前字段状态] B --> C[初始化 mask] C --> D{有全局规则?} D -->|是| E[应用全局排除规则] D -->|否| F{有字段级规则?} E --> F F -->|是| G{规则类型?} F -->|否| H[应用 mask 过滤] G -->|enum| I[只保留允许的 token] G -->|range| J[只保留范围内的 token] G -->|exclude| K[排除指定 token] I --> H J --> H K --> H H --> L[返回过滤后的 scores]
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 处理器,还有几个思路:

  1. 后处理验证:先生成,再验证。不符合就重新生成。优点是简单,缺点是可能一直重复失败。

  2. Structured Generation:像 guidanceoutlines 这样的库,专门做结构化生成。它们内部也是用 logits 过滤,但封装得更好用。

  3. 微调模型:让模型学会只在特定位置输出特定内容。这个成本高,但长期效果好。

对于我的场景,logits 处理器已经够用了。但如果你的约束更复杂或者对生成质量要求更高,可以考虑 Structured Generation 库。

写在最后

logits 处理器本质上是给模型的"自由表达"加了一道闸。这道闸在你需要硬约束的时候很有用,但记住一个原则:能用 prompt 解决的问题,就不要动 logits。

prompt 更符合模型的"自然"行为,而 logits 过滤是强行干预。过度干预会让模型的行为变得僵硬,生成质量也会受影响。

就像和人沟通一样:你可以说清楚你的期望,也可以直接把某些选项划掉。前者更自然,后者更严格。选哪种,看你的需求有多硬。

折腾这一轮下来,最大的收获不是学会了 logits 过滤,而是重新理解了"约束"和"自由"之间的平衡。太自由的模型会失控,太严格的模型会失真。找到那个平衡点,才是做 AI 应用时需要持续优化的东西。

版权声明: 本文首发于 指尖魔法屋-AI Logits处理器实践笔记https://blog.thinkmoon.cn/post/354-ai-logits-processor-raw-filter-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!