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

如何在Hugging Face Transformers v4.33中实现代码补全的StoppingCriteria

实现自定义StoppingCriteria截断模型生成以精准统计推理时间

核心实现思路

要精准统计模型生成有效代码的推理时间,需在模型首次生成连续两个换行符(\n\n)时立即停止生成,避免无关内容的耗时统计。通过继承transformers.StoppingCriteria实现自定义截断逻辑,同时兼容单batch推理和num_return_sequences=k的场景,适配Codegen、Code LLAMA、WizardCoder等系列模型。

自定义StoppingCriteria实现

import torch
from transformers import StoppingCriteria, StoppingCriteriaList

class DoubleNewLineStoppingCriteria(StoppingCriteria):
    def __init__(self, tokenizer, prompt_lengths):
        self.tokenizer = tokenizer
        # 存储每个prompt的长度,跳过原prompt仅检查生成内容
        self.prompt_lengths = prompt_lengths

    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
        # 遍历batch内所有生成序列
        for idx in range(input_ids.shape[0]):
            # 提取当前序列的生成部分(跳过原prompt)
            generated_ids = input_ids[idx][self.prompt_lengths[idx]:]
            # 解码生成token为文本,跳过特殊token
            generated_text = self.tokenizer.decode(generated_ids, skip_special_tokens=True)
            # 检查是否出现连续两个换行
            if "\n\n" in generated_text:
                # 任一序列触发条件则全局停止生成
                return True
        return False

完整测试代码(适配HumanEval场景)

import time
from transformers import AutoTokenizer, AutoModelForCausalLM

# 模型配置(可替换为目标模型路径)
MODEL_NAME = "WizardLM/WizardCoder-15B-V1.0"
NUM_RETURN_SEQUENCES = 3  # 对应pass@k的k值
MAX_NEW_TOKENS = 512  # 兜底最大生成token数,防止无限生成

# 加载模型与tokenizer
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForCausalLM.from_pretrained(MODEL_NAME, torch_dtype=torch.bfloat16, device_map="auto")

# 示例HumanEval prompt(实际需遍历数据集)
humaneval_prompt = "def add(a, b):\n    \"\"\"Return the sum of a and b\"\"\"\n"

# 编码prompt并获取长度
inputs = tokenizer(humaneval_prompt, return_tensors="pt").to(model.device)
prompt_length = inputs.input_ids.shape[1]
# 批量生成时,每个样本对应NUM_RETURN_SEQUENCES个生成序列,因此长度列表需重复对应次数
prompt_lengths = [prompt_length] * NUM_RETURN_SEQUENCES

# 初始化自定义停止条件
stopping_criteria = StoppingCriteriaList([
    DoubleNewLineStoppingCriteria(tokenizer, prompt_lengths)
])

# 计时并执行生成
start_time = time.perf_counter()
outputs = model.generate(
    **inputs,
    max_new_tokens=MAX_NEW_TOKENS,
    num_return_sequences=NUM_RETURN_SEQUENCES,
    stopping_criteria=stopping_criteria,
    do_sample=True,  # pass@k场景通常需要采样,按需调整
    temperature=0.8,
    top_p=0.95
)
end_time = time.perf_counter()

# 计算时间指标
total_time = end_time - start_time
avg_time_per_sequence = total_time / NUM_RETURN_SEQUENCES

# 输出结果
print(f"总推理时间: {total_time:.2f}s")
print(f"单序列平均推理时间: {avg_time_per_sequence:.2f}s")

# 验证生成内容是否被正确截断
for idx, output in enumerate(outputs):
    generated_text = tokenizer.decode(output[prompt_length:], skip_special_tokens=True)
    print(f"\n生成序列 {idx+1}:")
    print(generated_text)
    print("-" * 50)

关键注意事项

  • prompt长度处理:必须传入每个样本的prompt长度,确保只检查模型生成的内容,避免误判prompt中已存在的\n\n。
  • 多序列兼容:自定义停止条件会遍历所有生成序列,任一序列触发\n\n即停止全局生成,适配pass@k的批量生成需求。
  • 兜底参数设置:max_new_tokens作为兜底限制,防止模型因未生成\n\n而无限生成。
  • 模型兼容性:该实现适配所有基于transformers的因果语言模型,无需针对Codegen、Code LLAMA、WizardCoder做特殊修改。

内容的提问来源于stack exchange,提问作者Boyuan Chen

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 05:25:17