模型注册与初始化
注册到 MODEL_REGISTRY
vllm_norm 适配器通过 @register_model 装饰器将自身注册到框架的全局 MODEL_REGISTRY 中。当用户在命令行中指定 --model vllm_norm 时,框架会从注册表中查找并实例化该类。VLLM 类继承自抽象基类 LM(定义于 lm_eval/api/model.py),必须实现 loglikelihood()、loglikelihood_rolling() 和 generate_until() 三个核心方法。
@register_model("vllm_norm")
class VLLM(LM):
_DEFAULT_MAX_LENGTH = 2048
def __init__(self, url, max_tokens_to_sample: int = 2048, **kwargs):
super().__init__()
__init__() 的关键步骤
__init__() 方法执行了一系列初始化操作,可以分为以下几个阶段:
1. 加载 Tokenizer
从环境变量 MODEL_DIR 获取模型路径,使用 AutoTokenizer.from_pretrained() 加载分词器。如果默认加载失败(例如 fast tokenizer 不兼容),会自动降级为 use_fast=False 模式。
pretrained = os.environ.get(
"MODEL_DIR", "/data/minimax-dialogue/users/shennai/model/DeepSeek-V3-Base"
)
try:
self.tokenizer = AutoTokenizer.from_pretrained(pretrained, trust_remote_code=True)
except:
self.tokenizer = AutoTokenizer.from_pretrained(
pretrained, use_fast=False, trust_remote_code=True
)
2. 扫描推理服务器 IP
从 SERVER_IP_DIR 目录扫描所有 *.txt 文件以获取推理服务器 IP 地址。由于推理服务器可能尚未启动完成,这里实现了一个重试等待机制:每 5 秒重试一次,最多等待 5 分钟。如果超时仍无可用 IP,则抛出 RuntimeError 并附带详细的诊断信息。
self.urls = []
server_ip_dir = os.environ.get(
"SERVER_IP_DIR", "/minimax-dialogue/users/yize/data/server_ip_address"
)
# Wait for SERVER_IP_DIR to be populated (race with server startup).
# Retry every 5 seconds up to 5 minutes, then fail loudly.
_wait_interval_s = 5
_max_wait_s = 5 * 60
_waited_s = 0
while True:
urls: list[str] = []
for file in glob(os.path.join(server_ip_dir, "*.txt")):
try:
with open(file, "r") as f:
v = f.read().strip()
if v:
urls.append(v)
except Exception:
continue
if urls:
self.urls = urls
break
if _waited_s >= _max_wait_s:
raise RuntimeError(
"SERVER_IP_DIR has no non-empty '*.txt' ip files after waiting "
f"{_max_wait_s}s. SERVER_IP_DIR={server_ip_dir!r}. "
f"dir_listing={dir_listing!r}"
)
time.sleep(_wait_interval_s)
_waited_s += _wait_interval_s
每个推理服务器启动后会在 SERVER_IP_DIR 目录下创建一个 .txt 文件,文件内容为该服务器的 IP 地址。这种基于文件系统的服务发现模式非常适合分布式集群场景。
3. WAIT_AFTER_SERVER_START 额外等待
环境变量 WAIT_AFTER_SERVER_START 可选地配置一段额外等待时间。这解决了一个实际问题:服务器可能先写入 IP 文件,但模型加载尚未完成。等待结束后会重新扫描 IP,以获取在等待期间新启动的服务器。
wait_after_server_start = os.environ.get("WAIT_AFTER_SERVER_START", "0")
try:
wait_after_server_start_s = float(wait_after_server_start)
except Exception:
wait_after_server_start_s = 0.0
if wait_after_server_start_s > 0:
print(
f"[vllm_norm] WAIT_AFTER_SERVER_START={wait_after_server_start_s}s; sleeping...",
flush=True,
)
time.sleep(wait_after_server_start_s)
# Re-scan SERVER_IP_DIR after sleeping to pick up newly written ip files.
refreshed_urls: list[str] = []
for file in glob(os.path.join(server_ip_dir, "*.txt")):
try:
with open(file, "r") as f:
v = f.read().strip()
if v:
refreshed_urls.append(v)
except Exception:
continue
if refreshed_urls:
self.urls = refreshed_urls
4. 并发与模式配置
初始化过程的最后阶段完成了一系列运行时配置:
| 环境变量 | 作用 | 默认值 |
|---|---|---|
SERVING_PORT |
vLLM 服务端口(代码 fallback 默认 8000,但 run_evals_ft.sh 会设为 5002) |
8000(实际部署通常为 5002) |
NUM_PARRALLEL |
并发线程/协程数(thread_num) |
8 |
INCLUDE_PATH |
判断 pretrain/sft 模式(_is_pretrain) |
含 pretrain 路径 |
VLLM_ADD_THINKING_PREFIX |
是否添加 <think>\n 思考前缀 |
0(关闭) |
VLLM_SYSTEM_PROMPT |
自定义系统提示词 | 空 |
HTTP_CONNECTOR_LIMIT |
异步模式 HTTP 连接池总大小 | thread_num * 2 |
HTTP_CONNECTOR_LIMIT_PER_HOST |
每个 host 的连接数上限 | thread_num |
self.port = int(os.getenv("SERVING_PORT", 8000))
self.temperature = 0.01
self.max_tokens = 2000
self.top_p = 0.9
self._max_gen_toks = max_tokens_to_sample
self.thread_num = int(os.environ.get("NUM_PARRALLEL", "8"))
self._include_path = os.environ.get(
"INCLUDE_PATH",
"/output/lm-evaluation-harness/lm_eval/tasks/pretrain/",
)
self._is_pretrain = "pretrain" in self._include_path
self._add_assistant_thinking_prefix = (
int(os.environ.get("VLLM_ADD_THINKING_PREFIX", 0)) == 1
or int(os.environ.get("MUST_THINK", 0)) == 1
)
sys_prompt_raw = os.environ.get("VLLM_SYSTEM_PROMPT", '')
self._sys_prompt = sys_prompt_raw.encode().decode('unicode_escape') if sys_prompt_raw else ''
self._thinking_message = (
{"role": "assistant", "content": "<think>\n"}
if self._add_assistant_thinking_prefix else None
)
self._system_message = (
[{"role": "system", "content": self._sys_prompt}]
if self._sys_prompt else None
)
# HTTP 连接池配置优化
self.connector_limit = int(os.environ.get(
"HTTP_CONNECTOR_LIMIT", str(self.thread_num * 2)
))
self.connector_limit_per_host = int(os.environ.get(
"HTTP_CONNECTOR_LIMIT_PER_HOST", str(self.thread_num)
))
代码中有一条硬性断言:assert int(os.environ.get("ADD_BOD", 0)) == 0。BOD(Beginning Of Document)token 的添加应当在 HuggingFace 的 chat template 中完成,而非在评测框架中手动拼接。
请求构建
pretrain 模式 vs sft 模式
vllm_norm 支持两种推理模式,通过 _is_pretrain 标志区分:
- pretrain 模式:发送
POST /v1/completions请求,payload 包含原始文本prompt字段和logprobs=1,用于计算续写概率(loglikelihood 任务)。 - sft 模式:发送
POST /v1/chat/completions请求,payload 包含messages(OpenAI chat 格式),用于对话式生成任务。
_preprocess_payload() 方法
这是请求构建的核心方法,负责将原始 prompt 转换为 vLLM API 所需的 payload 格式。该方法在异步模式中被批量预处理调用,以减少每个请求的处理时间。
def _preprocess_payload(self, payload: dict) -> dict:
"""批量预处理 payload,减少每个请求的处理时间"""
prompt = payload.get("prompt", None)
payload["model"] = self._model_dir
if self._is_pretrain:
payload['logprobs'] = 1
return payload
if payload.get("prompt"):
payload.pop("prompt")
payload["return_entropy"] = self._return_entropy
if isinstance(prompt, list):
payload["messages"] = prompt
else:
prefix = self._prefix
suffix = self._suffix
if prefix and suffix and prompt.startswith(prefix) and prompt.endswith(suffix):
prompt = prompt[len(prefix) : -len(suffix)]
# 构建 messages
if BOS_TOKEN in prompt or EOS_TOKEN in prompt:
payload["messages"] = build_serving_from_chatml(prompt)
else:
payload["messages"] = [{"role": "user", "content": prompt}]
if self._add_assistant_thinking_prefix:
payload['continue_final_message'] = True
payload['add_generation_prompt'] = False
payload["messages"].append(self._thinking_message)
if self._system_message:
payload["messages"] = self._system_message + payload["messages"]
payload = specify_extra_kwargs(payload)
return payload
对于 sft 模式,prompt 的解析逻辑如下:
- 如果
prompt已经是list类型(messages 格式),直接使用。 - 如果是字符串,先尝试去除 SFT prefix/suffix 包装。
- 检查是否包含 ChatML 标记(
BOS_TOKEN/EOS_TOKEN),若包含则调用build_serving_from_chatml()解析。 - 否则,将纯文本封装为
[{"role": "user", "content": prompt}]。 - 如果启用了 thinking prefix,追加
{"role": "assistant", "content": "<think>\n"}并设置continue_final_message=True。 - 如果配置了 system prompt,将其插入 messages 列表的最前面。
具体示例:Pretrain vs SFT 请求 Payload 对比 示例
# ═══════════════════════════════════════════════════
# Pretrain 模式 — POST /v1/completions
# ═══════════════════════════════════════════════════
# 输入 prompt: "The capital of France is"
{
"prompt": "The capital of France is",
"model": "/data/models/Llama-2-7b",
"logprobs": 1,
"max_tokens": 1,
"stream": false,
"echo": true
}
# → 返回每个 token 的 logprob,用于计算 perplexity
# ═══════════════════════════════════════════════════
# SFT 模式 — POST /v1/chat/completions
# ═══════════════════════════════════════════════════
# 输入 prompt: "请问法国的首都是哪里?"
{
"model": "/data/models/ChatModel-7b",
"messages": [
{"role": "user", "content": "请问法国的首都是哪里?"}
],
"temperature": 0.01,
"top_p": 0.9,
"max_tokens": 2048,
"stream": false,
"stop": ["\n\n", "<|endoftext|>"],
"return_entropy": false
}
# ═══════════════════════════════════════════════════
# SFT 模式 + Thinking Prefix
# ═══════════════════════════════════════════════════
# VLLM_ADD_THINKING_PREFIX=1 时
{
"model": "/data/models/DeepSeek-R1",
"messages": [
{"role": "user", "content": "请问法国的首都是哪里?"},
{"role": "assistant", "content": "<think>\n"}
],
"continue_final_message": true,
"add_generation_prompt": false,
"temperature": 0.01,
"max_tokens": 2048,
"stream": false,
"stop": ["\n\n"]
}
build_serving_from_chatml()
该函数将自定义的 ChatML 格式(使用 ]~b] 和 [e~[ 作为 BOS/EOS 标记)解析为标准 OpenAI messages 列表。
BOS_TOKEN = ']~b]'
EOS_TOKEN = '[e~['
def build_serving_from_chatml(prompt):
system_data = []
settings = prompt.split(BOS_TOKEN)
for setting in settings:
if len(setting) == 0:
continue
if not setting.startswith('system'):
break
setting = setting[len('system '):]
fields = setting.split('=')
field_key = fields[0]
field_value = '='.join(fields[1:]).split('\n')
if field_key not in ('ai_setting', 'user_setting'):
continue
system_data.append({
'role': 'system',
field_key: field_value[0],
'text': '\n'.join(field_value[1:]).replace(EOS_TOKEN, '').strip()
})
dialogue_data = []
dialogues = prompt.split(BOS_TOKEN)[1:]
for i, dialogue in enumerate(dialogues):
if len(dialogue) == 0 or dialogue.startswith('system'):
continue
role = dialogue.split(' ')[0]
fields = dialogue.split('\n')
dialogue_data.append({
'role': role,
'name': fields[0].split('name=')[1],
'text': '\n'.join(fields[1:]).replace(EOS_TOKEN, '').strip(),
'trainable': False
})
if len(system_data) > 0:
conversations = [{'role': 'system', 'content': system_data[0]['text']}]
else:
conversations = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
for dd in dialogue_data:
conversations.append({
'role': 'user' if dd['role'] == 'user' else 'assistant',
'content': dd['text']
})
if conversations[-1]['role'] == 'assistant':
conversations = conversations[:-1]
return conversations
输入格式为 ]~b]system ai_setting=xxx\n系统提示内容[e~[\n]~b]user name=用户\n用户消息[e~[\n]~b]ai name=助手\n助手回复[e~[,输出为标准的 OpenAI messages 数组。注意:如果最后一条消息是 assistant 角色,会被自动移除(因为那是待生成的回复)。
specify_extra_kwargs()
该函数向 payload 中注入额外的控制参数:
def specify_extra_kwargs(payload):
timeout = os.environ.get("TIMEOUT", None)
if timeout:
payload["timeout"] = timeout
reasoning_effort = os.environ.get("REASONING_EFFORT", None)
if reasoning_effort:
payload["reasoning_effort"] = reasoning_effort
payload["chat_template_kwargs"] = {}
chat_template_file = os.environ.get("TOKENIZER_CHAT_TEMPLATE_FILE", None)
if chat_template_file:
with open(chat_template_file, "r") as f:
chat_template = f.read().strip()
payload["chat_template_kwargs"]["chat_template"] = chat_template
if os.environ.get("ENABLE_THINKING", False):
payload["chat_template_kwargs"]["enable_thinking"] = True
return payload
该函数支持通过环境变量注入:
TIMEOUT:请求超时时间REASONING_EFFORT:推理强度控制(用于支持 reasoning 模型)TOKENIZER_CHAT_TEMPLATE_FILE:从文件读取自定义 chat templateENABLE_THINKING:启用 thinking 模式(在 chat template 层面)
generate_until 实现
generate_until() 是框架中用于开放式文本生成(如 MMLU 选择题、代码生成等)的核心方法。它接收一组 Instance 请求,将它们按参数分组后批量发送给 vLLM 推理服务器。
主流程概览
Step 1: _collate() 排序
定义排序函数,按 token 长度降序排列请求。这样做有三个好处:
- 时间估计总是偏高而非偏低,更有利于进度规划
- 批次中第一个元素总是最长的,简化了 padding 逻辑
- OOM 错误会在靠前的位置发生,而非接近尾声时
def _collate(x):
# the negative sign on len(toks) sorts descending
toks = self.tok_encode(str(x[0]["text"]))
return -len(toks), x[0]["text"]
Step 2: Collator 分组
使用 Collator 工具类按 generation_kwargs 对请求进行分组。不同温度/采样参数的请求不能混在同一批次(例如 greedy sampling 和 temperature=0.8 不能混批),因此 Collator 的 grouping=True 确保了参数一致性。
re_ords = Collator([reg.args for reg in requests], _collate, grouping=True)
chunks = re_ords.get_batched(n=len(requests), batch_fn=None)
Step 3: 处理每组 chunk
对每组 chunk,代码执行以下步骤:
3a. 解析 VLM 视觉数据
首先尝试解包四元组 (contexts, all_gen_kwargs, doc_to_visual, docs)。如果成功,说明当前是视觉语言模型(VLM)任务,需要将图片转为 base64 编码。如果失败则降级为纯文本模式。
try:
# vlm任务
contexts, all_gen_kwargs, doc_to_visual, docs = zip(*chunk)
visuals = [
doc_to_visual[i](doc) for i, doc in enumerate(docs)
]
vision = True
except:
contexts, all_gen_kwargs = zip(*chunk)
vision = False
3b. 解析 until(stop tokens)
从 gen_kwargs 中提取 until(停止序列列表),并追加模型的 eos_token 作为额外的停止条件。同时解析 max_gen_toks(最大生成 token 数)。
gen_kwargs = all_gen_kwargs[0]
kwargs = copy.deepcopy(gen_kwargs)
if "until" in kwargs.keys():
until = kwargs.pop("until")
if isinstance(until, str):
until = [until]
elif not isinstance(until, list):
raise ValueError(
f"Expected `kwargs['until']` to be of type Union[str,list] but got {until}"
)
if not until:
until = [self.tokenizer.decode(self.eot_token_id)]
else:
until += [self.tokenizer.decode(self.eot_token_id)]
if "max_gen_toks" in kwargs.keys():
max_gen_toks = kwargs.pop("max_gen_toks")
else:
max_gen_toks = self.max_gen_toks
3c. 构建 payload list
为每个 context 构建一个独立的 HTTP 请求 payload,包含 prompt、temperature、top_p、max_tokens、stop 等参数。如果是 VLM 任务,还会添加 multi_modal_data 字段(包含 base64 编码的图片)。
pload = {
"prompt": "",
"temperature": kwargs.get("temperature", self.temperature),
"top_p": kwargs.get("top_p", self.top_p),
"max_tokens": max_gen_toks,
"stream": False,
"stop": until,
}
for i, context in enumerate(contexts):
if vision and DEFAULT_IMAGE_TOKEN not in context:
image_tokens = [DEFAULT_IMAGE_TOKEN] * len(visuals[i])
image_tokens = " ".join(image_tokens)
context = f"{image_tokens}\n{context}"
if vision:
text = sft_prefix + context + sft_suffix
else:
text = context["text"]
pload["prompt"] = text
if vision:
pload["multi_modal_data"] = dict(
image=[
self._image_to_base64(visuals[i][j])
for j in range(len(visuals[i]))
]
)
requests.append(copy.deepcopy(pload))
3d. 发送请求并收集结果
根据环境变量 USE_ASYNC 选择异步模式或线程池模式发送请求。最终通过 re_ords.get_original(res) 将结果恢复到原始顺序。
# perform batched generation
if os.environ.get("USE_ASYNC"):
cont, log_probs_list = self._model_generate_async(
requests=requests,
generate=True,
max_tokens=max_gen_toks,
stop=until,
**kwargs,
)
else:
cont, log_probs_list = self._model_generate(
requests=requests,
generate=True,
max_tokens=max_gen_toks,
stop=until,
**kwargs,
)
loglikelihood 实现
loglikelihood() 计算给定 context 下 continuation 的对数概率,是 pretrain 评测任务(如 perplexity 计算、多项选择题)的核心方法。
loglikelihood() 入口方法
该方法首先对每个 request 进行编码,生成 token 序列和对应的 mask(context 部分为 0,continuation 部分为 1)。
def loglikelihood(self, requests: List[Instance]) -> List[Tuple[float, bool]]:
new_reqs = []
for i, req in enumerate(requests):
context, continuation, gen_args = req.args
context = context["text"]
# 处理字符位置标记的格式
if isinstance(context, dict) and context.get("__char_spans__"):
full_text = context["full_text"]
target_char_spans = context["target_char_spans"]
tokens, mask = self._char_spans_to_token_mask(
full_text, target_char_spans
)
new_reqs.append(((context, continuation), tokens, mask))
continue
elif isinstance(context, list): # encoded texts with masks
text = context
mask = continuation
else:
if context == "":
context_enc, continuation_enc = (
[self.eot_token_id],
self.tok_encode(continuation),
)
else:
strip_context = gen_args.get("strip_context", True) if gen_args else True
context_enc, continuation_enc = self._encode_pair(
context, continuation, strip_context
)
text = context_enc + continuation_enc
mask = [0] * len(context_enc) + [1] * len(continuation_enc)
new_reqs.append(((context, continuation), text, mask))
return self._loglikelihood_tokens(new_reqs)
该方法支持三种输入格式:
- dict with __char_spans__:通过字符级别的起止位置标记目标区域,调用
_char_spans_to_token_mask()动态计算 token-level mask。 - encoded list with masks:context 已经是 token ID 列表,continuation 直接作为 mask。
- 标准文本格式:将 context 和 continuation 编码为 token 序列,context 部分 mask 为 0,continuation 部分 mask 为 1。
mask = [0, 0, 0, 1, 1, 1] 表示前 3 个 token 是 context(不计入 loglikelihood),后 3 个 token 是 continuation(需要计算 loglikelihood)。
具体示例:Token 编码与 Mask 生成全流程 示例
# 假设评测 HellaSwag 选择题(multiple_choice)
# context = "The man picked up the guitar and"
# continuation = " started strumming a tune"
# ═══════════════════════════════════════════════════
# Step 1: Tokenizer 编码
# ═══════════════════════════════════════════════════
context_enc = tokenizer.encode("The man picked up the guitar and")
# → [464, 582, 8347, 510, 278, 11847, 322] (7 个 token)
continuation_enc = tokenizer.encode(" started strumming a tune")
# → [4687, 851, 1822, 292, 263, 260, 2540] (7 个 token)
# 注意:需要用 _encode_pair() 去掉 continuation 开头被重复编码的 token
# ═══════════════════════════════════════════════════
# Step 2: 拼接 text 和生成 mask
# ═══════════════════════════════════════════════════
text = context_enc + continuation_enc
# → [464, 582, 8347, 510, 278, 11847, 322, 4687, 851, 1822, 292, 263, 260, 2540]
mask = [0]*len(context_enc) + [1]*len(continuation_enc)
# → [0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1]
# ├── context (不计入) ──┤├── continuation (计入) ──┤
# ═══════════════════════════════════════════════════
# Step 3: 发送请求并解析 logprobs
# ═══════════════════════════════════════════════════
# payload = {"prompt": "The man picked up the guitar and started strumming a tune",
# "max_tokens": 1, "logprobs": 1, "echo": true}
#
# vLLM 返回每个 token 的 logprob:
# token_logprobs = [None, -3.2, -1.5, -2.8, -0.3, -4.1, -1.2, -2.5, -3.8, -1.1, -2.3, -0.8, -3.5, -1.9]
# ↑ 第一个 token 无条件概率为 None
# Step 4: _parse_logprobs() 计算
# continuation_logprobs = sum(logprob * mask for ...)
# = (-2.5)*1 + (-3.8)*1 + (-1.1)*1 + (-2.3)*1 + (-0.8)*1 + (-3.5)*1 + (-1.9)*1
# = -15.9
#
# is_greedy: 检查每个 continuation position 的 top token 是否匹配
# → True(每个位置的最高概率 token 恰好就是 continuation token)
#
# 返回: (-15.9, True, "...")
_loglikelihood_tokens() 方法
这是实际发送请求并解析结果的方法。它构建 pretrain 模式的 payload:prompt=context+continuation、max_tokens=1、logprobs=1、echo=True。
def _loglikelihood_tokens(
self,
requests: List[Tuple[Tuple[str, str], List[int], List[int]]],
disable_tqdm: bool = False,
) -> List[Tuple[float, bool]]:
res = []
def _collate(x):
toks = x[1] + x[2]
return -len(toks), tuple(toks)
re_ord = Collator(requests, sort_fn=_collate)
chunks = re_ord.get_batched(n=len(requests), batch_fn=None)
for index, chunk in enumerate(chunks):
inputs = []
masks = []
ctxlens = []
for cache_key, text, mask in chunk:
ctxlen = mask.index(1)
try:
prompt = cache_key[0] + cache_key[1]
except:
prompt = cache_key[0]['full_text']
pload = {
"prompt": prompt,
"max_tokens": 1,
"stream": False,
"logprobs": 1,
"echo": True,
}
inputs.append(pload)
masks.append(mask)
ctxlens.append(ctxlen)
outputs, _ = self._model_generate(requests=inputs, generate=False)
for output, mask, ctxlen, (cache_key, text, _), pload in zip(
outputs, masks, ctxlens, chunk, inputs
):
answer = self._parse_logprobs(
tokens=text, outputs=output, ctxlen=ctxlen, masks=mask
)
res.append(answer)
return re_ord.get_original(res)
关键 payload 参数说明:
echo=True:让 vLLM 返回 prompt 本身每个 token 的 logprob(而不仅是生成的 token)max_tokens=1:我们不需要生成新内容,只需要 prompt 的 logprobslogprobs=1:返回每个位置 top-1 的 logprob
_parse_logprobs() 解析
该静态方法负责从 vLLM 返回的 logprobs 中提取三个关键信息:
@staticmethod
def _parse_logprobs(
tokens: List, outputs, ctxlen: int, masks: List
) -> Tuple[float, bool, str]:
# 过滤 None 值(第一个 token 无条件概率)
if len(outputs["logprobs"]["token_logprobs"]) != len(tokens):
outputs["logprobs"]["token_logprobs"] = [
logprob_dict for logprob_dict in outputs["logprobs"]["token_logprobs"]
if logprob_dict is not None
]
continuation_logprobs_dicts = convert_keys_to_int(outputs["logprobs"])
# 计算 continuation_logprobs:只累加 mask=1 部分的 logprob
continuation_logprobs = sum(
logprob_dict[1] * mask
for token, logprob_dict, mask in zip(
tokens[ctxlen:],
continuation_logprobs_dicts[ctxlen:],
masks[ctxlen:]
)
)
# 判断 is_greedy:top logprob token 是否匹配 continuation token
toplogprobs = outputs["logprobs"]["top_logprobs"]
is_greedy = True
for toplogprob, logprob_dict, mask in zip(
toplogprobs[ctxlen:], continuation_logprobs_dicts[ctxlen:], masks[ctxlen:]
):
if toplogprob and mask:
top_token = max(toplogprob, key=lambda x: toplogprob[x])
if top_token != logprob_dict[0]:
is_greedy = False
break
return (continuation_logprobs, is_greedy, json.dumps(info_dict_list))
返回值是一个三元组:
- continuation_logprobs(
float):continuation 所有 token 的对数概率之和 - is_greedy(
bool):每个 continuation token 是否都是该位置概率最高的 token(即 greedy decoding 是否会生成相同的 continuation) - info_dict_list(
str):JSON 格式的详细信息,包含每个 token 的 logprob 和 mask
并发与负载均衡
vllm_norm 提供了两种请求发送模式,分别适用于不同的部署场景。
架构概览
异步模式:_model_generate_async()
当设置 USE_ASYNC 环境变量时,使用 asyncio + aiohttp 发送并发请求。此模式要求只有一个 URL(assert len(self.urls) == 1),适用于单服务器高并发场景。
def _model_generate_async(
self,
requests: List[List[int]] = None,
generate: bool = False,
max_tokens: int = None,
stop: Optional[List[str]] = None,
**kwargs,
):
if generate:
sampling_kwargs = deepcopy(kwargs)
sampling_kwargs["max_tokens"] = max_tokens
if "max_new_tokens" in sampling_kwargs:
sampling_kwargs["max_tokens"] = sampling_kwargs["max_new_tokens"]
sampling_kwargs.pop("max_new_tokens")
sampling_kwargs["stop"] = stop
else:
sampling_kwargs = dict(temperature=0, prompt_logprobs=2, max_tokens=1)
res = [None] * len(requests)
log_probs_list = [None] * len(requests)
pbar = tqdm(total=len(requests), disable=(self.rank != 0))
assert len(self.urls) == 1, f"Only one URL is supported"
url = self.urls[0]
request_tuples = [(_idx, req) for _idx, req in enumerate(requests)]
asyncio.run(
self.__post_requests_to_server_async(url, request_tuples, pbar, res, log_probs_list)
)
return res, log_probs_list
异步请求的核心在 __post_requests_to_server_async() 中。它创建一个 aiohttp.ClientSession,配置了精细的 TCPConnector 参数:
timeout = aiohttp.ClientTimeout(total=None)
connector = aiohttp.TCPConnector(
limit=self.connector_limit, # 总连接池大小
limit_per_host=self.connector_limit_per_host, # 每个host的连接数
ttl_dns_cache=1e9, # DNS缓存时间(秒)
keepalive_timeout=300, # 连接保活时间(秒)
force_close=False, # 保持连接复用
enable_cleanup_closed=True, # 清理已关闭的连接
use_dns_cache=True, # 启用DNS缓存
)
semaphore = asyncio.Semaphore(self.thread_num)
async with aiohttp.ClientSession(timeout=timeout, connector=connector) as session:
tasks = [
async_post_func(url_completions, url_chat, session, semaphore, req, ...)
for req in preprocessed_tuples
]
await asyncio.gather(*tasks)
asyncio.Semaphore(self.thread_num) 限制了同时在 flight 中的请求数量。虽然所有 tasks 被一次性创建,但 semaphore 确保同时只有 thread_num 个请求在执行 HTTP 调用,避免压垮服务器。
线程池模式:_model_generate()
默认模式使用 ThreadPoolExecutor 发送请求。与异步模式不同,线程池模式支持多服务器负载均衡:为每个 URL 创建一个线程,线程内部再用子级 ThreadPoolExecutor 处理并发。
def _model_generate(
self,
requests: List[List[int]] = None,
generate: bool = False,
max_tokens: int = None,
stop: Optional[List[str]] = None,
**kwargs,
):
if generate:
sampling_kwargs = deepcopy(kwargs)
sampling_kwargs["max_tokens"] = max_tokens
sampling_kwargs["stop"] = stop
else:
sampling_kwargs = dict(temperature=0, prompt_logprobs=2, max_tokens=1)
request_queue = queue.Queue()
res = [None] * len(requests)
log_probs_list = [None] * len(requests)
for _idx, req in enumerate(requests):
request_queue.put((_idx, req))
# 每个 URL 一个线程,随机打乱避免热点
with ThreadPoolExecutor(max_workers=len(self.urls)) as executor:
futures = []
for url in random.sample(self.urls, len(self.urls)):
f = executor.submit(
self.__post_requests_to_server, url, request_queue, res, log_probs_list
)
futures.append(f)
for f in futures:
f.result()
return res, log_probs_list
注意 random.sample(self.urls, len(self.urls)) 的使用:它将 URL 列表随机打乱后分配给线程,实现了简单的随机轮询负载均衡。每个线程内部通过共享的 request_queue(线程安全的 queue.Queue)获取待处理请求,实现了工作窃取式的任务分发。
def __post_requests_to_server(self, url, request_queue, results, log_probs_list):
# 每个线程内部再创建 thread_num 个子线程
with ThreadPoolExecutor(max_workers=self.thread_num) as executor:
futures = []
for _ in range(self.thread_num):
f = executor.submit(
self.__post_one_request, url, request_queue, results, log_probs_list
)
futures.append(f)
for f in futures:
f.result()
错误处理与重试
两种模式都实现了相同的重试逻辑,最多重试 10000 次。错误处理分为三种情况:
| 错误类型 | 处理方式 | 等待时间 |
|---|---|---|
| "No available workers" | 服务器过载,长时间等待后重试 | 异步: 300s / 同步: 120s |
| Token 长度超限 | 动态减少 max_tokens 后重试;若输入本身超限则返回错误文本 |
立即重试 |
| 其他异常 | 打印 traceback 后短暂等待重试 | 0.5s |
error_message = response.json()["message"]
if (
"Please reduce the length of the messages or completion" in error_message
or "Please reduce the number of tokens in the input messages or the completion"
in error_message
):
numbers = re.findall(r'\d+', error_message)
if int(numbers[0]) > int(numbers[2]):
payload["max_tokens"] = int(numbers[0]) - int(numbers[2]) - 1
print(f"payload['max_tokens'] is set to: {payload['max_tokens']}")
else:
response = "输入token长度超过模型长度限制!!!!"
break
当 vLLM 返回"token 长度超限"错误时,代码会从错误信息中提取数字,计算出可用的 max_tokens 值(numbers[0] - numbers[2] - 1),然后自动降低生成长度重试。这避免了因为少量超长输入导致整个评测失败。
特殊处理
Thinking tokens 过滤
当模型启用了"思考"能力(如 DeepSeek-R1 风格的推理模型)时,模型输出可能包含 <think>...</think> 包裹的推理过程。评测时通常只需要最终回复,因此在 evaluator.py 中会过滤掉 thinking 内容。
相关的常量定义为:
THINKING_END_TOKEN = '</think>'THINKING_START_TOKEN = '<think>\n'
在 vllm_norm 中,thinking prefix 的添加通过 VLLM_ADD_THINKING_PREFIX 或 MUST_THINK 环境变量控制。当启用时,会向 messages 末尾追加一条 assistant 消息 {"role": "assistant", "content": "<think>\n"},并设置 continue_final_message=True + add_generation_prompt=False,让模型从思考开始生成。
对于 sft 模式的 chat/completions 接口,vLLM 还可能在返回中分离 reasoning_content 字段:
response_message = response_json["choices"][0]["message"]
if response_message.get("reasoning_content", None) is not None:
response.update({
"content": response_message["content"],
"reasoning_content": response_message.get("reasoning_content", None)
})
else:
response.update({"content": response_message["content"]})
if self._return_entropy:
entropy_list = response_json["choices"][0]['entropy']
entropy_list = entropy_list[-num_completion_tokens:]
response.update({
'entropy_list': [np.mean([e for e in entropy_list if isinstance(e, float)])]
})
BOD 处理
ADD_BOD 环境变量必须设为 0。BOD(Beginning Of Document)token 的添加应在 HuggingFace chat template 中处理,而非在评测框架中手动拼接。初始化时有一条断言确保这一点:
assert int(os.environ.get("ADD_BOD", 0)) == 0, \
"ADD_BOD must be 0 for vllm_norm, which should be supported in hf chat_template."
prefix/suffix 加载
通过 load_prefix() 和 load_suffix()(来自 lm_eval/tasks/sft/base/utils.py)加载 SFT 格式的前后缀。这些前后缀在评测 SFT 任务时被框架自动拼接到 prompt 的前后,但在发送给 vLLM 之前需要被去除(因为 vLLM 的 chat template 会自行处理格式化)。
prefix = self._prefix
suffix = self._suffix
if prefix and suffix and prompt.startswith(prefix) and prompt.endswith(suffix):
prompt = prompt[len(prefix) : -len(suffix)]
ChatML 格式解析细节
自定义 ChatML 使用了独特的 BOS/EOS 标记来避免与模型 tokenizer 的特殊 token 冲突:
| 标记 | 值 | 说明 |
|---|---|---|
BOS_TOKEN |
]~b] |
消息开始标记 |
EOS_TOKEN |
[e~[ |
消息结束标记 |
一条完整的 ChatML 消息格式如下:
]~b]system ai_setting=应事
系统提示内容[e~[
]~b]user name=用户
用户消息内容[e~[
]~b]ai name=应事
助手回复内容[e~[
build_serving_from_chatml() 会将其解析为标准 OpenAI messages 格式。解析过程分两个阶段:
- 提取 system 消息:扫描以
system开头的 block,提取ai_setting或user_setting的内容作为系统提示。 - 提取对话消息:扫描非 system 的 block,根据 role(
user/ai)映射为user/assistant角色。
VLM 图片处理
对于视觉语言模型(VLM)任务,_image_to_base64() 方法将 PIL Image 对象转换为 base64 编码字符串。图片数据通过 multi_modal_data 字段传递给 vLLM:
def _image_to_base64(self, img, format="jpeg"):
try:
buffer = BytesIO()
img.save(buffer, format=format)
img_str = base64.b64encode(buffer.getvalue()).decode("utf-8")
return img_str
except Exception as e:
print(f"转换过程出错: {str(e)}")
return None
_char_spans_to_token_mask()
当 loglikelihood 的输入使用字符级别的位置标记(__char_spans__ 格式)时,该方法将字符级别的 span 转换为 token 级别的 mask。它利用 tokenizer 的 return_offsets_mapping 功能获取每个 token 对应的字符偏移量,然后通过区间重叠判断生成 mask。
def _char_spans_to_token_mask(self, full_text, target_char_spans):
enc = self.tokenizer(
full_text, return_offsets_mapping=True, add_special_tokens=False
)
mask = [0] * len(enc["input_ids"])
offsets = enc["offset_mapping"]
for i, (start, end) in enumerate(offsets):
for t_start, t_end in target_char_spans:
# 区间重叠判断
if max(start, t_start) < min(end, t_end):
mask[i] = 1
break
return enc["input_ids"], mask
文件末尾的 convert_keys_to_int() 函数(L1158-1173)将 vLLM 返回的 logprobs 数据中的字符串 key 转换为整数。它实际返回的是 (token, token_logprob) 元组列表,供 _parse_logprobs() 使用。