如何在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
相关产品推荐
相关产品推荐

