Compare commits
10 Commits
629f878c28
...
1902c81fdd
| Author | SHA1 | Date | |
|---|---|---|---|
| 1902c81fdd | |||
| f89bc60d59 | |||
| 810874ddb8 | |||
| c84151eef9 | |||
|
|
3d62430fd7 | ||
| 72aa7e690a | |||
| b5806731e0 | |||
| 47a4d9e72a | |||
| 3b8a567e9e | |||
| 50e3a05fb0 |
132
README.md
132
README.md
@@ -1,118 +1,46 @@
|
||||
# 天数智芯 天垓100 文本生成引擎(基于 vLLM 优化)
|
||||
|
||||
本项目是为**天数智芯-天垓100**加速卡深度优化的高性能文本生成推理引擎,基于开源 **vLLM** 框架进行架构级适配与增强,率先实现对 **Qwen3 系列**等最新大模型的高效支持。通过引入 **Prefix Caching**、PagedAttention 等先进优化技术,显著提升吞吐与响应速度,同时提供标准 **OpenAI 兼容 API 接口**,便于无缝集成现有应用生态。
|
||||
|
||||
## 支持模型
|
||||
|
||||
- **Qwen3**
|
||||
- **Llama3**
|
||||
- **DeepSeek-R1-Distill**
|
||||
- 其他兼容 vLLM 的 HuggingFace 模型(持续扩展中)
|
||||
|
||||
> 模型下载地址:[https://modelscope.cn/models/Qwen](https://modelscope.cn/models/Qwen)
|
||||
|
||||
---
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. 模型下载
|
||||
|
||||
从 ModelScope 下载所需模型(以 Qwen2.5-7B-Instruct 为例):
|
||||
|
||||
```bash
|
||||
modelscope download --model qwen/Qwen2.5-7B-Instruct README.md --local_dir /mnt/models/Qwen2.5-7B-Instruct
|
||||
```
|
||||
|
||||
> ⚠️ 请确保模型路径在后续 Docker 启动时正确挂载。
|
||||
|
||||
---
|
||||
|
||||
### 2. 拉取并构建 Docker 镜像
|
||||
|
||||
我们提供已预装天垓100驱动与vLLM优化版本的Docker镜像:
|
||||
# 天数智芯 天垓100 文本生成引擎(基于 vLLM 优化适配Qwen3.6-35B-A3B)
|
||||
|
||||
```
|
||||
# 本地构建
|
||||
docker build -t enginex-iluvatar-vllm:bi100 -f Dockerfile .
|
||||
docker build -t enginex-iluvatar-vllm:bi100-qwen3.6 -f Dockerfile .
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 3. 启动服务容器
|
||||
启动容器镜像
|
||||
|
||||
```bash
|
||||
docker run -it --rm -p 8000:80 \
|
||||
--name vllm-iluvatar \
|
||||
-v /mnt/models/Qwen2.5-7B-Instruct:/model:ro \
|
||||
--privileged \
|
||||
-e TENSOR_PARALLEL_SIZE=1 \
|
||||
-e PREFIX_CACHING=true \
|
||||
-e MAX_MODEL_LEN=10000 \
|
||||
enginex-iluvatar-vllm:bi100
|
||||
下载Qwen3.6-35B-A3B模型,并且需要将模型的config.json文件中architectures字段改成
|
||||
```json
|
||||
"architectures": [
|
||||
"Qwen3_5MoeForCausalLM"
|
||||
]
|
||||
```
|
||||
|
||||
> ✅ 参数说明:
|
||||
> - `PREFIX_CACHING=true`: 启用 Prefix Caching 优化,显著提升多请求共享前缀的推理效率
|
||||
> - `MAX_MODEL_LEN=10000`: 支持长上下文推理
|
||||
> - `--privileged`: 确保天垓100设备可见
|
||||
|
||||
---
|
||||
|
||||
## 4. 测试服务(使用 OpenAI 兼容接口)
|
||||
|
||||
服务启动后,可通过标准 OpenAI SDK 或 `curl` 进行测试。
|
||||
|
||||
### 示例:文本生成请求
|
||||
|
||||
```bash
|
||||
curl http://localhost:8000/v1/chat/completions \
|
||||
docker run -dit --network=host --ipc=host \
|
||||
-v /usr/src:/usr/src -v /lib/modules:/lib/modules -v /dev:/dev --privileged \
|
||||
-v /mnt/disk1/models/Qwen3.6-35B-A3B:/model:ro --entrypoint=python3 \
|
||||
-e CUDA_VISIBLE_DEVICES=4,5,6,7 -e VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 \
|
||||
enginex-iluvatar-vllm:bi100-qwen3.6 \
|
||||
-m vllm.entrypoints.openai.api_server \
|
||||
--model /model --port 1111 --served-model-name llm \
|
||||
--max-model-len 100000 --trust-remote-code -tp 4 --gpu-memory-utilization 0.95 \
|
||||
--max-num-seqs 1 --disable-log-requests --disable-frontend-multiprocessing \
|
||||
--max-num-batched-tokens 4096 --enable-chunked-prefill \
|
||||
--max-seq-len-to-capture 32768 --enable-auto-tool-choice \
|
||||
--tool-call-parser qwen3_coder --reasoning-parser qwen3
|
||||
```
|
||||
|
||||
请求
|
||||
```bash
|
||||
curl http://localhost:1111/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "qwen3-8b",
|
||||
"model": "llm",
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "请用中文介绍一下上海的特点。"}
|
||||
{"role": "user", "content": "Can you tell me the story of Snow White?"}
|
||||
],
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 512
|
||||
"max_tokens": 200,
|
||||
"temperature": 0.7
|
||||
}'
|
||||
```
|
||||
|
||||
### 使用 OpenAI Python SDK(需安装 `openai>=1.0`)
|
||||
|
||||
```python
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(base_url="http://localhost:8000/v1", api_key="none")
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="qwen3-8b",
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "请简要介绍杭州的特色文化。"}
|
||||
],
|
||||
max_tokens=512,
|
||||
temperature=0.7
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 测试结果对比(A100 vs 天垓100)
|
||||
|
||||
### 测试数据集
|
||||
|
||||
[chat_dataset_v0.json](chat_dataset_v0.json)
|
||||
|
||||
### 测试结果
|
||||
|
||||
在相同模型和输入条件下,测试平均输出速度(单位:字每秒),结果如下:
|
||||
|
||||
| 模型 | 天垓100 输出速度 | Nvidia A100 输出速度 |
|
||||
|--------|--------------------------|-------------------------------|
|
||||
| Qwen2.5-7B-Instruct | 36.8 | 112.4 |
|
||||
| Qwen2.5-1.5B-Instruct-AWQ | 72.4 | 100.8 |
|
||||
| Qwen/Qwen1.5-32B-Chat | 12.4 | 55.7 |
|
||||
|
||||
```
|
||||
@@ -1,42 +0,0 @@
|
||||
# 天数智芯 天垓100 文本生成引擎(基于 vLLM 优化适配Qwen3.6-27B)
|
||||
|
||||
```
|
||||
# 本地构建
|
||||
docker build -t enginex-iluvatar-vllm:bi100-qwen3.6 -f Dockerfile .
|
||||
```
|
||||
|
||||
|
||||
启动容器镜像
|
||||
|
||||
下载Qwen3.6-27B模型,并且需要将模型的config.json文件中architectures字段改成
|
||||
```json
|
||||
"architectures": [
|
||||
"Qwen3_5ForCausalLM"
|
||||
]
|
||||
```
|
||||
|
||||
```bash
|
||||
docker run -dit --network=host --ipc=host \
|
||||
-v /usr/src:/usr/src -v /lib/modules:/lib/modules -v /dev:/dev --privileged \
|
||||
--name vllm-iluvatar \
|
||||
-v /mnt/models/Qwen3.6-27B:/model:ro --entrypoint=python3 \
|
||||
enginex-iluvatar-vllm:bi100 \
|
||||
-m vllm.entrypoints.openai.api_server \
|
||||
--model /model --port 1111 --served-model-name llm \
|
||||
--max-model-len 10000 --enforce-eager --trust-remote-code -tp 4 --gpu-memory-utilization 0.95
|
||||
```
|
||||
|
||||
请求
|
||||
```bash
|
||||
curl http://localhost:1111/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "llm",
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Can you tell me the story of Snow White?"}
|
||||
],
|
||||
"max_tokens": 200,
|
||||
"temperature": 0.7
|
||||
}'
|
||||
```
|
||||
34
computility-run.yaml
Normal file
34
computility-run.yaml
Normal file
@@ -0,0 +1,34 @@
|
||||
concurrency: 1
|
||||
command:
|
||||
- python3
|
||||
- -m
|
||||
- vllm.entrypoints.openai.api_server
|
||||
- --model
|
||||
- /model
|
||||
- --served-model-name
|
||||
- llm
|
||||
- --max-model-len
|
||||
- '100000'
|
||||
- --gpu-memory-utilization
|
||||
- '0.9'
|
||||
- --trust-remote-code
|
||||
- -tp
|
||||
- '4'
|
||||
- --max-num-seqs
|
||||
- '1'
|
||||
- --disable-log-requests
|
||||
- --disable-frontend-multiprocessing
|
||||
- --max-num-batched-tokens
|
||||
- '8192'
|
||||
- --enable-chunked-prefill
|
||||
- --max-seq-len-to-capture
|
||||
- '32768'
|
||||
- --enable-auto-tool-choice
|
||||
- --tool-call-parser
|
||||
- qwen3_coder
|
||||
- --reasoning-parser
|
||||
- qwen3
|
||||
- --enable-prefix-caching
|
||||
env:
|
||||
- name: VLLM_ENGINE_ITERATION_TIMEOUT_S
|
||||
value: 3600
|
||||
@@ -113,6 +113,11 @@ class ConversationMessage(TypedDict, total=False):
|
||||
tool_calls: Optional[Iterable[ChatCompletionMessageToolCallParam]]
|
||||
"""The tool calls generated by the model, such as function calls."""
|
||||
|
||||
reasoning_content: Optional[str]
|
||||
"""Reasoning / thinking content for assistant messages.
|
||||
Passed directly to the chat template (Qwen3 reads message.reasoning_content
|
||||
natively) instead of being manually wrapped in <think>...</think>."""
|
||||
|
||||
|
||||
ModalityStr = Literal["image", "audio", "video"]
|
||||
_T = TypeVar("_T")
|
||||
@@ -480,15 +485,13 @@ def _parse_chat_message_content(
|
||||
if "tool_calls" in parsed_msg:
|
||||
result_msg["tool_calls"] = list(parsed_msg["tool_calls"])
|
||||
|
||||
# Prepend reasoning_content as <think>...</think> so the model
|
||||
# sees its own chain-of-thought in multi-turn conversations.
|
||||
reasoning = message.get("reasoning_content") # type: ignore[arg-type]
|
||||
# Pass reasoning content as a dedicated field so the chat template
|
||||
# can render it natively (Qwen3: message.reasoning_content branch).
|
||||
# Accept both "reasoning" (new vllm) and "reasoning_content" (ours).
|
||||
reasoning = (message.get("reasoning") # type: ignore[arg-type]
|
||||
or message.get("reasoning_content")) # type: ignore[arg-type]
|
||||
if reasoning and isinstance(reasoning, str):
|
||||
existing = result_msg.get("content") or ""
|
||||
result_msg["content"] = (
|
||||
f"<think>{reasoning}</think>\n\n{existing}"
|
||||
if existing else f"<think>{reasoning}</think>"
|
||||
)
|
||||
result_msg["reasoning_content"] = reasoning
|
||||
|
||||
elif role == "tool":
|
||||
parsed_msg = _ToolParser(message)
|
||||
|
||||
@@ -393,6 +393,20 @@ class PagedAttention:
|
||||
# --------------------------------------------------------------
|
||||
if ctx_len > 0:
|
||||
num_ctx_blocks = (ctx_len + block_size - 1) // block_size
|
||||
# Safety: if block_tables is too narrow this indicates a
|
||||
# prefix_cache_hit + chunked-prefill bug in model_runner.py
|
||||
# (Case 1 leaves prefix_cache_hit=True but block_table is
|
||||
# only computed_block_nums, not the full context blocks).
|
||||
# patch_model_runner.py fixes the root cause; this guard
|
||||
# prevents a zero-dim amax() crash if it still slips through.
|
||||
if num_ctx_blocks > block_tables.shape[1]:
|
||||
print(
|
||||
f"[paged_attn WARNING] seq {i}: num_ctx_blocks={num_ctx_blocks} "
|
||||
f"> block_tables.shape[1]={block_tables.shape[1]}, ctx_len={ctx_len}. "
|
||||
"Block table is undersized (prefix_cache_hit bug). "
|
||||
"Capping context to available blocks — attention may be incorrect.",
|
||||
file=sys.stderr, flush=True)
|
||||
num_ctx_blocks = block_tables.shape[1]
|
||||
for tile_blk in range(0, num_ctx_blocks, _BLOCKS_PER_TILE):
|
||||
blk_end = min(tile_blk + _BLOCKS_PER_TILE, num_ctx_blocks)
|
||||
blk_ids = block_tables[i, tile_blk:blk_end]
|
||||
|
||||
78
qwen3_6_scripts/patch_model_runner.py
Normal file
78
qwen3_6_scripts/patch_model_runner.py
Normal file
@@ -0,0 +1,78 @@
|
||||
"""
|
||||
Fix: prefix_cache_hit stays True for chunked-prefill chunk 2+ even when past cache.
|
||||
|
||||
Root cause:
|
||||
model_runner.py _compute_for_prefix_cache_hit has three cases:
|
||||
Case 1: prefix_cache_len <= context_len → "already past cache, do normal"
|
||||
Case 2: context_len < prefix_cache_len < seq_len → partial hit, correct
|
||||
Case 3: seq_len <= prefix_cache_len → full hit, reduce to 1 token
|
||||
|
||||
Case 1 does nothing (leaves prefix_cache_hit = True). Then in utils.py:
|
||||
if inter_data.prefix_cache_hit:
|
||||
block_table = computed_block_nums ← ONLY the original prefix blocks!
|
||||
|
||||
But context_len > prefix_cache_len means chunk 1 tokens (between prefix_cache_len
|
||||
and context_len) are ALSO in KV cache and need to be in block_table.
|
||||
block_table = computed_block_nums misses all chunk-1 blocks.
|
||||
|
||||
In _forward_prefix_pytorch:
|
||||
num_ctx_blocks = ceil(context_len / block_size) # e.g. 268
|
||||
block_tables.shape[1] = len(computed_block_nums) # e.g. 12 <-- too small!
|
||||
At tile_blk >= 12: blk_ids is empty → k_t shape [..., 0] → amax crash.
|
||||
|
||||
Fix:
|
||||
Set prefix_cache_hit = False for Case 1, so utils.py falls through to:
|
||||
elif chunked_prefill_enabled:
|
||||
block_table = block_tables[seq_id] ← full block table (prefix + chunk1)
|
||||
"""
|
||||
|
||||
import re
|
||||
import sys
|
||||
|
||||
CANDIDATE_PATHS = [
|
||||
"/usr/local/corex/lib64/python3/dist-packages/vllm/worker/model_runner.py",
|
||||
"/usr/local/corex/lib/python3/dist-packages/vllm/worker/model_runner.py",
|
||||
]
|
||||
|
||||
OLD_BLOCK = """\
|
||||
if prefix_cache_len <= context_len:
|
||||
# We already passed the cache hit region,
|
||||
# so do normal computation.
|
||||
pass"""
|
||||
|
||||
NEW_BLOCK = """\
|
||||
if prefix_cache_len <= context_len:
|
||||
# We already passed the cache hit region,
|
||||
# so do normal computation.
|
||||
# Must clear prefix_cache_hit so _add_seq_group uses the full
|
||||
# block_tables (prefix + previous-chunk blocks) instead of only
|
||||
# computed_block_nums (prefix only). Without this, block_tables
|
||||
# passed to _forward_prefix_pytorch is too narrow for context_len,
|
||||
# causing an empty blk_ids slice and a zero-dim amax() crash.
|
||||
inter_data.prefix_cache_hit = False"""
|
||||
|
||||
import os
|
||||
|
||||
patched = False
|
||||
for path in CANDIDATE_PATHS:
|
||||
if not os.path.exists(path):
|
||||
continue
|
||||
with open(path, "r") as f:
|
||||
src = f.read()
|
||||
if OLD_BLOCK not in src:
|
||||
if NEW_BLOCK in src:
|
||||
print(f"[patch_model_runner] already patched: {path}")
|
||||
patched = True
|
||||
break
|
||||
print(f"[patch_model_runner] WARNING: expected block not found in {path}, skipping")
|
||||
continue
|
||||
patched_src = src.replace(OLD_BLOCK, NEW_BLOCK, 1)
|
||||
with open(path, "w") as f:
|
||||
f.write(patched_src)
|
||||
print(f"[patch_model_runner] patched Case-1 prefix_cache_hit fix in: {path}")
|
||||
patched = True
|
||||
break
|
||||
|
||||
if not patched:
|
||||
print("[patch_model_runner] ERROR: could not find model_runner.py at any known path", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
@@ -8,8 +8,6 @@
|
||||
# are already correct for standard Triton 2.3.1 — do NOT overwrite them.
|
||||
# - DO NOT install BI-V150 corex Triton 2.1.0 (pkgs/triton): that causes
|
||||
# GPU hang on BI-V100 because the Triton CUDA PTX kernels are incompatible.
|
||||
#
|
||||
# Important Note: Qwen3.6-27B must apply TP=4,PP=2 combination in order to deploy using 8 GPUs
|
||||
|
||||
# Recommended server start command for TP=4 support 100K, need chunked prefill
|
||||
# CUDA_VISIBLE_DEVICES="4,5,6,7" VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 python3 -m vllm.entrypoints.openai.api_server \
|
||||
@@ -17,6 +15,14 @@
|
||||
# --max-model-len 100000 --enforce-eager --trust-remote-code -tp 4 --gpu-memory-utilization 0.95 \
|
||||
# --max-num-seqs 1 --disable-log-requests --disable-frontend-multiprocessing \
|
||||
# --max-num-batched-tokens 4096 --enable-chunked-prefill
|
||||
#
|
||||
# With prefix caching (GDN align-mode, requires chunked prefill):
|
||||
# CUDA_VISIBLE_DEVICES="4,5,6,7" VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 python3 -m vllm.entrypoints.openai.api_server \
|
||||
# --model /workspace/models/Qwen3.6-35B-A3B --port 1111 --served-model-name llm \
|
||||
# --max-model-len 150000 --trust-remote-code -tp 4 --gpu-memory-utilization 0.90 \
|
||||
# --max-num-seqs 1 --disable-log-requests --disable-frontend-multiprocessing \
|
||||
# --max-num-batched-tokens 8192 --enable-chunked-prefill --enable-prefix-caching \
|
||||
# --max-seq-len-to-capture 32768
|
||||
|
||||
# --- paged_attn.py: replace forward_prefix with pure-PyTorch fallback -------
|
||||
# The Triton context_attention_fwd kernel hangs BI-V100 GPUs permanently
|
||||
@@ -26,6 +32,15 @@
|
||||
# when context length is high
|
||||
cp ./paged_attn.py /usr/local/corex/lib/python3/dist-packages/vllm/attention/ops/paged_attn.py
|
||||
|
||||
# --- model_runner.py: fix prefix_cache_hit stays True in chunked-prefill chunk 2+ ---
|
||||
# Bug: _compute_for_prefix_cache_hit Case 1 (prefix_cache_len <= context_len)
|
||||
# leaves prefix_cache_hit=True. Then _add_seq_group uses block_table=computed_block_nums
|
||||
# (only the original prefix blocks), ignoring chunk-1 KV cache blocks.
|
||||
# _forward_prefix_pytorch then gets an undersized block_tables and crashes with
|
||||
# "amax(): Expected reduction dim -1 to have non-zero size" on the 2nd tile.
|
||||
# Fix: set prefix_cache_hit=False for Case 1 so the full block_tables is used.
|
||||
python3 ./patch_model_runner.py
|
||||
|
||||
# --- transformers: Qwen3_5 tokenizer / model files --------------------------
|
||||
pip install transformers==4.55.3 -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
cp -r ./qwen3_5 /usr/local/lib/python3.10/site-packages/transformers/models/
|
||||
@@ -42,8 +57,15 @@ python3 ./patch_vllm_qwen3_5.py
|
||||
# returns _cached_all_token_ids[-0:] == [0:] (the ENTIRE prompt+output list).
|
||||
# Each prefill chunk step adds prompt_len to previous_num_tokens, so a 10K
|
||||
# prompt processed in 3 chunks inflates completion_tokens by ~30K.
|
||||
# Also adds num_cached_tokens field to RequestMetrics for prefix-cache stats.
|
||||
cp ./sequence.py /usr/local/corex/lib/python3/dist-packages/vllm/sequence.py
|
||||
|
||||
# --- scheduler.py: record num_cached_tokens in RequestMetrics ----------------
|
||||
# Sets seq_group.metrics.num_cached_tokens = prefix_cache_len on first prefill
|
||||
# when --enable-prefix-caching is active, so serving_chat.py can report it in
|
||||
# usage.prompt_tokens_details.cached_tokens (OpenAI-compatible API response).
|
||||
cp ./scheduler.py /usr/local/corex/lib/python3/dist-packages/vllm/core/scheduler.py
|
||||
|
||||
# --- xformers: bypass cudnnFlashAttnForward (head_dim=256 > 128 limit) ------
|
||||
# Injects _run_sdpa_fallback (pure matmul+softmax) into xformers.py.
|
||||
# Required because head_dim=256 > 128 and ixformer flash attention either
|
||||
|
||||
@@ -99,11 +99,16 @@ class ModelList(OpenAIBaseModel):
|
||||
data: List[ModelCard] = Field(default_factory=list)
|
||||
|
||||
|
||||
class PromptTokensDetails(OpenAIBaseModel):
|
||||
cached_tokens: int = 0
|
||||
|
||||
|
||||
class UsageInfo(OpenAIBaseModel):
|
||||
prompt_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
completion_tokens: Optional[int] = 0
|
||||
reasoning_tokens: Optional[int] = None
|
||||
prompt_tokens_details: Optional[PromptTokensDetails] = None
|
||||
|
||||
|
||||
class RequestResponseMetadata(BaseModel):
|
||||
@@ -315,12 +320,20 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
prompt_logprobs = self.top_logprobs
|
||||
|
||||
guided_json_object = None
|
||||
if (self.response_format is not None
|
||||
and self.response_format.type == "json_object"):
|
||||
guided_json_object = True
|
||||
guided_json_from_schema = None
|
||||
if self.response_format is not None:
|
||||
if self.response_format.type == "json_object":
|
||||
guided_json_object = True
|
||||
elif (self.response_format.type == "json_schema"
|
||||
and self.response_format.json_schema is not None
|
||||
and self.response_format.json_schema.json_schema is not None):
|
||||
guided_json_from_schema = \
|
||||
self.response_format.json_schema.json_schema
|
||||
|
||||
guided_decoding = GuidedDecodingParams.from_optional(
|
||||
json=self._get_guided_json_from_tool() or self.guided_json,
|
||||
json=(self._get_guided_json_from_tool()
|
||||
or self.guided_json
|
||||
or guided_json_from_schema),
|
||||
regex=self.guided_regex,
|
||||
choice=self.guided_choice,
|
||||
grammar=self.guided_grammar,
|
||||
@@ -373,6 +386,35 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
||||
|
||||
return None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def normalize_messages(cls, data):
|
||||
"""Normalize incoming messages before pydantic union validation.
|
||||
|
||||
Real-world clients (e.g. from other providers) send assistant tool_call
|
||||
messages with content=null, which fails the strict Union type check.
|
||||
Replace null content with "" so validation passes.
|
||||
reasoning_content is intentionally kept — chat_utils.py wraps it as
|
||||
<think>...</think> for multi-turn reasoning history.
|
||||
"""
|
||||
messages = data.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
return data
|
||||
normalized = []
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
normalized.append(msg)
|
||||
continue
|
||||
if msg.get("content") is None:
|
||||
if msg.get("reasoning_content") is None:
|
||||
raise ValueError(
|
||||
"Each message must have at least one of 'content' or "
|
||||
"'reasoning_content'.")
|
||||
msg = {**msg, "content": ""}
|
||||
normalized.append(msg)
|
||||
data = {**data, "messages": normalized}
|
||||
return data
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def validate_stream_options(cls, data):
|
||||
@@ -609,12 +651,18 @@ class CompletionRequest(OpenAIBaseModel):
|
||||
echo_without_generation = self.echo and self.max_tokens == 0
|
||||
|
||||
guided_json_object = None
|
||||
if (self.response_format is not None
|
||||
and self.response_format.type == "json_object"):
|
||||
guided_json_object = True
|
||||
guided_json_from_schema = None
|
||||
if self.response_format is not None:
|
||||
if self.response_format.type == "json_object":
|
||||
guided_json_object = True
|
||||
elif (self.response_format.type == "json_schema"
|
||||
and self.response_format.json_schema is not None
|
||||
and self.response_format.json_schema.json_schema is not None):
|
||||
guided_json_from_schema = \
|
||||
self.response_format.json_schema.json_schema
|
||||
|
||||
guided_decoding = GuidedDecodingParams.from_optional(
|
||||
json=self.guided_json,
|
||||
json=self.guided_json or guided_json_from_schema,
|
||||
regex=self.guided_regex,
|
||||
choice=self.guided_choice,
|
||||
grammar=self.guided_grammar,
|
||||
|
||||
@@ -2,7 +2,8 @@
|
||||
# Pure-PyTorch DeltaNet (no fla / causal_conv1d dependency).
|
||||
# Text-only (no VL, no MTP).
|
||||
|
||||
from typing import Iterable, List, Optional, Tuple
|
||||
from collections import OrderedDict
|
||||
from typing import Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -420,9 +421,6 @@ class GatedDeltaNet(nn.Module):
|
||||
|
||||
else:
|
||||
# Decode: one token per sequence
|
||||
with open("/tmp/vllm_decode_debug.log", "a") as _f:
|
||||
_f.write(f"[deltanet decode] layer={self.layer_idx} num_seqs={hidden_states.shape[0]}\n")
|
||||
_f.flush()
|
||||
num_seqs = hidden_states.shape[0]
|
||||
weight_2d = self.conv1d_weight.squeeze(1)
|
||||
|
||||
@@ -452,17 +450,47 @@ class GatedDeltaNet(nn.Module):
|
||||
q = q.repeat_interleave(self.head_expand_ratio, dim=2)
|
||||
k = k.repeat_interleave(self.head_expand_ratio, dim=2)
|
||||
|
||||
core_out, last_state = _torch_recurrent_gated_delta_rule(
|
||||
q, k, v, g, beta,
|
||||
initial_state=temporal_state,
|
||||
output_final_state=True,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
# Inlined decode recurrent step (seq_len=1).
|
||||
# Replaces _torch_recurrent_gated_delta_rule to avoid 5 transpose+
|
||||
# contiguous+float32 copies, core_out allocation, and Python loop.
|
||||
# Uses bmm/baddbmm_ to eliminate 3 large (B,H,k,v) intermediate tensors.
|
||||
# temporal_state: (B, H_v, k_dim, v_dim) float32 — updated in-place.
|
||||
orig_dtype = q.dtype
|
||||
_scale = self.head_k_dim ** -0.5
|
||||
|
||||
q_t = _l2norm(q.squeeze(1)).float() * _scale # (B, H_v, k_dim)
|
||||
k_t = _l2norm(k.squeeze(1)).float() # (B, H_v, k_dim)
|
||||
v_t = v.squeeze(1).float() # (B, H_v, v_dim)
|
||||
g_t = g.squeeze(1).float().exp_() # (B, H_v)
|
||||
bt = beta.squeeze(1).float() # (B, H_v)
|
||||
|
||||
# Decay state in-place: (B, H_v, k_dim, v_dim) *= scalar per head
|
||||
temporal_state.mul_(g_t[:, :, None, None])
|
||||
|
||||
# Reshape to batched-matmul layout: (B*H_v, k_dim, v_dim)
|
||||
ts_flat = temporal_state.view(-1, self.head_k_dim, self.head_v_dim)
|
||||
BH = ts_flat.shape[0]
|
||||
|
||||
# kv_mem = k_t @ temporal_state shape: (B*H_v, 1, k_dim) @ (B*H_v, k_dim, v_dim)
|
||||
kv_mem = torch.bmm(
|
||||
k_t.view(BH, 1, self.head_k_dim), ts_flat
|
||||
).view(num_seqs, local_num_v, self.head_v_dim) # (B, H_v, v_dim)
|
||||
|
||||
delta = (v_t - kv_mem) * bt[:, :, None] # (B, H_v, v_dim)
|
||||
|
||||
# State update: temporal_state += outer(k_t, delta) fused, no intermediate
|
||||
ts_flat.baddbmm_(
|
||||
k_t.view(BH, self.head_k_dim, 1),
|
||||
delta.view(BH, 1, self.head_v_dim),
|
||||
)
|
||||
if last_state is not None:
|
||||
temporal_state.copy_(last_state)
|
||||
|
||||
# Output: core_out = q_t @ updated temporal_state
|
||||
core_out = torch.bmm(
|
||||
q_t.view(BH, 1, self.head_k_dim), ts_flat
|
||||
).view(num_seqs, local_num_v, self.head_v_dim).to(orig_dtype)
|
||||
# core_out: (B, H_v, v_dim) = (num_seqs, local_num_v, head_v_dim) already
|
||||
|
||||
z = z_all.reshape(num_seqs, local_num_v, self.head_v_dim)
|
||||
core_out = core_out.reshape(num_seqs, local_num_v, self.head_v_dim)
|
||||
normed = self.norm(
|
||||
core_out.reshape(-1, self.head_v_dim),
|
||||
z.reshape(-1, self.head_v_dim))
|
||||
@@ -733,28 +761,52 @@ class Qwen3_5MoeSparseBlock(nn.Module):
|
||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||
topk_weights = topk_weights.to(hidden_states.dtype)
|
||||
|
||||
out = torch.zeros_like(hidden_states)
|
||||
w13 = self.experts.w13_weight # (E, 2*I, H)
|
||||
w2 = self.experts.w2_weight # (E, H, I)
|
||||
|
||||
for eid in range(self.num_experts):
|
||||
# Tokens routed to this expert
|
||||
mask = (topk_ids == eid) # (T, top_k) bool
|
||||
tok_ids, topk_pos = mask.nonzero(as_tuple=True)
|
||||
if tok_ids.numel() == 0:
|
||||
continue
|
||||
T = hidden_states.shape[0]
|
||||
if T == 1:
|
||||
# Fast path: single token (decode).
|
||||
# Batched GEMM: replace top_k separate F.linear calls with 2 fused ops.
|
||||
# gate_up: 1 large GEMM (1,H) × (K*2*I,H)^T → (1, K*2*I)
|
||||
# down: 1 bmm (K,H,I) @ (K,I,1) → (K,H)
|
||||
# Total: 3 kernel launches vs previous 16 (top_k*2).
|
||||
eids = topk_ids[0] # (K,)
|
||||
ws = topk_weights[0].to(hidden_states.dtype) # (K,)
|
||||
w13_sel = w13[eids] # (K, 2*I, H)
|
||||
w2_sel = w2[eids] # (K, H, I)
|
||||
|
||||
tokens = hidden_states[tok_ids] # (n, H)
|
||||
# gate + up projection (ColumnParallel shard)
|
||||
gate_up = F.linear(tokens, w13[eid]) # (n, 2*I)
|
||||
gate, up = gate_up.chunk(2, dim=-1)
|
||||
act = F.silu(gate) * up # (n, I)
|
||||
# down projection (RowParallel shard) — result is partial
|
||||
# F.linear(x, W) = x @ W.T; w2[eid]: (H, I) → x @ W.T = (n,H) ✓
|
||||
expert_out = F.linear(act, w2[eid]) # (n, H)
|
||||
H = hidden_states.shape[-1]
|
||||
|
||||
weights = topk_weights[tok_ids, topk_pos].unsqueeze(-1)
|
||||
out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype))
|
||||
gate_up = F.linear(
|
||||
hidden_states,
|
||||
w13_sel.reshape(-1, H), # (K*2*I, H) — contiguous after indexing
|
||||
) # (1, K*2*I)
|
||||
gate_up = gate_up.view(self.top_k, -1) # (K, 2*I)
|
||||
gate, up = gate_up.chunk(2, dim=-1) # (K, I) each
|
||||
act = F.silu(gate) * up # (K, I)
|
||||
|
||||
# bmm: (K,H,I) @ (K,I,1) → (K,H,1) → (K,H)
|
||||
expert_out = torch.bmm(w2_sel, act.unsqueeze(-1)).squeeze(-1) # (K, H)
|
||||
|
||||
out = (expert_out * ws.unsqueeze(-1)).sum(0, keepdim=True).to(
|
||||
hidden_states.dtype) # (1, H)
|
||||
else:
|
||||
# General path (prefill / multi-seq): loop over unique active experts.
|
||||
# At most T*top_k unique experts, always <= num_experts.
|
||||
out = torch.zeros_like(hidden_states)
|
||||
unique_eids = topk_ids.view(-1).unique().tolist()
|
||||
for eid in unique_eids:
|
||||
eid = int(eid)
|
||||
mask = (topk_ids == eid) # (T, top_k)
|
||||
tok_ids, topk_pos = mask.nonzero(as_tuple=True)
|
||||
tokens = hidden_states[tok_ids] # (n, H)
|
||||
gate_up = F.linear(tokens, w13[eid]) # (n, 2*I)
|
||||
gate, up = gate_up.chunk(2, dim=-1)
|
||||
act = F.silu(gate) * up # (n, I)
|
||||
expert_out = F.linear(act, w2[eid]) # (n, H)
|
||||
weights = topk_weights[tok_ids, topk_pos].unsqueeze(-1)
|
||||
out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype))
|
||||
|
||||
return out # partial, all-reduce done in forward()
|
||||
|
||||
@@ -982,6 +1034,15 @@ class Qwen3_5ForCausalLM(nn.Module, HasInnerState, SupportsLoRA):
|
||||
# Lazy initialised in first forward call
|
||||
self.mamba_cache: Optional[MambaCacheManager] = None
|
||||
|
||||
# GDN prefix state cache (align mode): stores (conv_states, temporal_states) snapshots
|
||||
# at KV-block boundaries so that prefix-cache-hit requests can restore correct GDN state.
|
||||
# Key: tuple of physical block IDs covering the cached prefix
|
||||
# Value: (conv_states_cpu, temporal_states_cpu) each of shape (num_gdn_layers, ...)
|
||||
self._gdn_prefix_cache: OrderedDict = OrderedDict()
|
||||
self._gdn_prefix_cache_max: int = 16 # ~16 × 16 MB ≈ 256 MB CPU RAM
|
||||
self._block_size: int = (cache_config.block_size
|
||||
if cache_config is not None else 16)
|
||||
|
||||
def _get_mamba_cache_shape(self):
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
# Each sequence's state is stored in float32
|
||||
@@ -1018,9 +1079,69 @@ class Qwen3_5ForCausalLM(nn.Module, HasInnerState, SupportsLoRA):
|
||||
# temporal_states: (num_linear_layers, batch, local_num_v, k_dim, v_dim)
|
||||
conv_states, temporal_states = mamba_tensors
|
||||
|
||||
# ── GDN prefix-cache align mode: inject saved state on prefix hit ─────
|
||||
# Conditions: prefill pass, batch=1, context_len > 0 (prefix cached or
|
||||
# previous chunk already processed), block_tables available.
|
||||
# We always attempt a lookup: for subsequent chunked-prefill chunks the
|
||||
# key matches our own saved state (same data already in slot → no-op).
|
||||
# For a true cross-request prefix hit the key matches a previous request.
|
||||
_is_single_seq_prefill = (
|
||||
attn_metadata is not None
|
||||
and attn_metadata.num_prefill_tokens > 0
|
||||
and conv_states.shape[1] == 1 # batch == 1
|
||||
and getattr(attn_metadata, 'context_lens_tensor', None) is not None
|
||||
and getattr(attn_metadata, 'block_tables', None) is not None
|
||||
and attn_metadata.block_tables.numel() > 0
|
||||
)
|
||||
if _is_single_seq_prefill:
|
||||
context_len = int(attn_metadata.context_lens_tensor[0].item())
|
||||
if context_len > 0:
|
||||
num_prefix_blocks = context_len // self._block_size
|
||||
if (num_prefix_blocks > 0
|
||||
and attn_metadata.block_tables.shape[1] >= num_prefix_blocks):
|
||||
lookup_key = tuple(
|
||||
attn_metadata.block_tables[0, :num_prefix_blocks]
|
||||
.cpu().tolist())
|
||||
if lookup_key in self._gdn_prefix_cache:
|
||||
saved_conv, saved_temporal = self._gdn_prefix_cache[lookup_key]
|
||||
conv_states[:, 0].copy_(
|
||||
saved_conv.to(conv_states.device), non_blocking=True)
|
||||
temporal_states[:, 0].copy_(
|
||||
saved_temporal.to(temporal_states.device), non_blocking=True)
|
||||
self._gdn_prefix_cache.move_to_end(lookup_key)
|
||||
logger.debug("GDN prefix cache hit: prefix_len=%d blocks=%d",
|
||||
context_len, num_prefix_blocks)
|
||||
# ── End inject ──────────────────────────────────────────────────────────
|
||||
|
||||
hidden_states = self.model(
|
||||
input_ids, positions, kv_caches, attn_metadata,
|
||||
conv_states, temporal_states)
|
||||
|
||||
# ── GDN prefix-cache align mode: save state after this prefill chunk ───
|
||||
# Save state keyed by ALL complete KV blocks processed so far.
|
||||
# Next requests reusing this prefix will restore from here.
|
||||
if _is_single_seq_prefill:
|
||||
context_len = int(attn_metadata.context_lens_tensor[0].item())
|
||||
query_len = attn_metadata.num_prefill_tokens
|
||||
total_processed = context_len + query_len
|
||||
num_complete_blocks = total_processed // self._block_size
|
||||
if (num_complete_blocks > 0
|
||||
and attn_metadata.block_tables.shape[1] >= num_complete_blocks):
|
||||
save_key = tuple(
|
||||
attn_metadata.block_tables[0, :num_complete_blocks]
|
||||
.cpu().tolist())
|
||||
# Move to end (LRU: most recent = last) and update value
|
||||
if save_key in self._gdn_prefix_cache:
|
||||
self._gdn_prefix_cache.move_to_end(save_key)
|
||||
self._gdn_prefix_cache[save_key] = (
|
||||
conv_states[:, 0].cpu().clone(),
|
||||
temporal_states[:, 0].cpu().clone(),
|
||||
)
|
||||
# Evict oldest entries beyond max
|
||||
while len(self._gdn_prefix_cache) > self._gdn_prefix_cache_max:
|
||||
self._gdn_prefix_cache.popitem(last=False)
|
||||
# ── End save ────────────────────────────────────────────────────────────
|
||||
|
||||
return hidden_states
|
||||
|
||||
def compute_logits(
|
||||
@@ -1201,6 +1322,35 @@ class Qwen3_5MoeForCausalLM(Qwen3_5ForCausalLM):
|
||||
weight_loader(param, loaded_weight)
|
||||
continue
|
||||
|
||||
# --- Individual expert weights (FT checkpoint: experts.{i}.{proj}.weight) ---
|
||||
# Standard transformers fine-tuning saves each expert separately instead of
|
||||
# the pre-merged (num_experts, ...) tensors in the original checkpoint.
|
||||
if ".mlp.experts." in name:
|
||||
parts = name.split(".mlp.experts.", 1)
|
||||
expert_rest = parts[1] # e.g. "0.gate_proj.weight"
|
||||
dot_pos = expert_rest.find(".")
|
||||
if dot_pos > 0 and expert_rest[:dot_pos].isdigit():
|
||||
eid = int(expert_rest[:dot_pos])
|
||||
proj_raw = expert_rest[dot_pos + 1:]
|
||||
proj = proj_raw[:-7] if proj_raw.endswith(".weight") else proj_raw
|
||||
prefix = parts[0] # e.g. "model.layers.0"
|
||||
if proj == "gate_proj":
|
||||
w13_name = f"{prefix}.mlp.experts.w13_weight"
|
||||
if w13_name in params_dict:
|
||||
param = params_dict[w13_name]
|
||||
param.weight_loader(param, loaded_weight, "w1_weight", "w1", eid)
|
||||
elif proj == "up_proj":
|
||||
w13_name = f"{prefix}.mlp.experts.w13_weight"
|
||||
if w13_name in params_dict:
|
||||
param = params_dict[w13_name]
|
||||
param.weight_loader(param, loaded_weight, "w3_weight", "w3", eid)
|
||||
elif proj == "down_proj":
|
||||
w2_name = f"{prefix}.mlp.experts.w2_weight"
|
||||
if w2_name in params_dict:
|
||||
param = params_dict[w2_name]
|
||||
param.weight_loader(param, loaded_weight, "w2_weight", "w2", eid)
|
||||
continue
|
||||
|
||||
# --- Stacked / standard weights ---
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
|
||||
1656
qwen3_6_scripts/scheduler.py
Normal file
1656
qwen3_6_scripts/scheduler.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -119,6 +119,7 @@ class RequestMetrics:
|
||||
scheduler_time: Optional[float] = None
|
||||
model_forward_time: Optional[float] = None
|
||||
model_execute_time: Optional[float] = None
|
||||
num_cached_tokens: Optional[int] = None
|
||||
|
||||
|
||||
class SequenceDataDelta(
|
||||
|
||||
@@ -25,7 +25,7 @@ from vllm.entrypoints.openai.protocol import (
|
||||
ChatCompletionResponseChoice, ChatCompletionResponseStreamChoice,
|
||||
ChatCompletionStreamResponse, ChatMessage, DeltaFunctionCall, DeltaMessage,
|
||||
DeltaToolCall, ErrorResponse, FunctionCall, RequestResponseMetadata,
|
||||
ToolCall, UsageInfo)
|
||||
PromptTokensDetails, ToolCall, UsageInfo)
|
||||
from vllm.entrypoints.openai.serving_engine import (BaseModelPath,
|
||||
LoRAModulePath,
|
||||
OpenAIServing,
|
||||
@@ -179,6 +179,16 @@ class OpenAIServingChat(OpenAIServing):
|
||||
logger.exception("Error in loading multi-modal data")
|
||||
return self.create_error_response(str(e))
|
||||
|
||||
# n > max_num_seqs deadlock guard: scheduler uses break (not continue)
|
||||
# when can_schedule(num_new_seqs=n) fails, so an n that exceeds
|
||||
# max_num_seqs permanently blocks the entire waiting queue with no error.
|
||||
_sched_cfg = await self.engine_client.get_scheduler_config()
|
||||
_max_seqs = _sched_cfg.max_num_seqs
|
||||
if request.n is not None and request.n > _max_seqs:
|
||||
return self.create_error_response(
|
||||
f"n={request.n} exceeds max_num_seqs={_max_seqs}. "
|
||||
f"Use n<={_max_seqs} or omit n.")
|
||||
|
||||
# validation for OpenAI tools
|
||||
# tool_choice = "required" is not supported
|
||||
if request.tool_choice == "required":
|
||||
@@ -284,12 +294,12 @@ class OpenAIServingChat(OpenAIServing):
|
||||
if request.stream:
|
||||
return self.chat_completion_stream_generator(
|
||||
request, result_generator, request_id, conversation, tokenizer,
|
||||
request_metadata)
|
||||
request_metadata, raw_request=raw_request)
|
||||
|
||||
try:
|
||||
return await self.chat_completion_full_generator(
|
||||
request, result_generator, request_id, conversation, tokenizer,
|
||||
request_metadata)
|
||||
request_metadata, raw_request=raw_request)
|
||||
except ValueError as e:
|
||||
# TODO: Use a vllm-specific Validation Error
|
||||
return self.create_error_response(str(e))
|
||||
@@ -307,6 +317,7 @@ class OpenAIServingChat(OpenAIServing):
|
||||
conversation: List[ConversationMessage],
|
||||
tokenizer: AnyTokenizer,
|
||||
request_metadata: RequestResponseMetadata,
|
||||
raw_request: Optional[Request] = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
model_name = self.base_model_paths[0].name
|
||||
created_time = int(time.time())
|
||||
@@ -318,6 +329,7 @@ class OpenAIServingChat(OpenAIServing):
|
||||
previous_num_tokens = [0] * num_choices
|
||||
finish_reason_sent = [False] * num_choices
|
||||
num_prompt_tokens = 0
|
||||
num_cached_tokens: Optional[int] = None
|
||||
|
||||
if isinstance(request.tool_choice, ChatCompletionNamedToolChoiceParam):
|
||||
tool_choice_function_name = request.tool_choice.function.name
|
||||
@@ -379,12 +391,37 @@ class OpenAIServingChat(OpenAIServing):
|
||||
yield "data: [DONE]\n\n"
|
||||
return
|
||||
|
||||
# Background task: poll is_disconnected() every 300 ms and abort the
|
||||
# engine request as soon as the client goes away. This catches the
|
||||
# case where the HTTP layer (Starlette/uvicorn) does not actively read
|
||||
# the receive channel during streaming, so is_disconnected() in
|
||||
# iterate_with_cancellation never fires during fast decode.
|
||||
_disconnect_watcher: Optional[asyncio.Task] = None
|
||||
if raw_request is not None:
|
||||
async def _watch_disconnect() -> None:
|
||||
try:
|
||||
while True:
|
||||
if await raw_request.is_disconnected():
|
||||
logger.info(
|
||||
"Client disconnected (decode watcher), "
|
||||
"aborting request %s", request_id)
|
||||
await self.engine_client.abort(request_id)
|
||||
return
|
||||
await asyncio.sleep(0.3)
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
_disconnect_watcher = asyncio.ensure_future(_watch_disconnect())
|
||||
|
||||
try:
|
||||
async for res in result_generator:
|
||||
if res.prompt_token_ids is not None:
|
||||
num_prompt_tokens = len(res.prompt_token_ids)
|
||||
if res.encoder_prompt_token_ids is not None:
|
||||
num_prompt_tokens += len(res.encoder_prompt_token_ids)
|
||||
if (num_cached_tokens is None
|
||||
and res.metrics is not None
|
||||
and res.metrics.num_cached_tokens is not None):
|
||||
num_cached_tokens = res.metrics.num_cached_tokens
|
||||
|
||||
# We need to do it here, because if there are exceptions in
|
||||
# the result_generator, it needs to be sent as the FIRST
|
||||
@@ -560,9 +597,14 @@ class OpenAIServingChat(OpenAIServing):
|
||||
# if the message delta is None (e.g. because it was a
|
||||
# "control token" for tool calls or the parser otherwise
|
||||
# wasn't ready to send a token, then
|
||||
# get the next token without streaming a chunk
|
||||
# get the next token without streaming a chunk.
|
||||
# However, if this is the finish token we must NOT skip —
|
||||
# the finish block updates reasoning_token_counts, sets
|
||||
# finish_reason_sent, and flushes the final usage chunk.
|
||||
if delta_message is None:
|
||||
continue
|
||||
if output.finish_reason is None:
|
||||
continue
|
||||
delta_message = DeltaMessage()
|
||||
|
||||
if output.finish_reason is None:
|
||||
# Send token-by-token response for each request.n
|
||||
@@ -686,6 +728,9 @@ class OpenAIServingChat(OpenAIServing):
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=num_prompt_tokens + completion_tokens,
|
||||
reasoning_tokens=total_reasoning,
|
||||
prompt_tokens_details=(
|
||||
PromptTokensDetails(cached_tokens=num_cached_tokens)
|
||||
if num_cached_tokens is not None else None),
|
||||
)
|
||||
|
||||
final_usage_chunk = ChatCompletionStreamResponse(
|
||||
@@ -708,11 +753,27 @@ class OpenAIServingChat(OpenAIServing):
|
||||
total_tokens=num_prompt_tokens + num_completion_tokens,
|
||||
reasoning_tokens=total_reasoning)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
# Client disconnected via CancelledError path; abort engine request.
|
||||
await self.engine_client.abort(request_id)
|
||||
return
|
||||
except ValueError as e:
|
||||
# TODO: Use a vllm-specific Validation Error
|
||||
logger.error("error in chat completion stream generator: %s", e)
|
||||
data = self.create_streaming_error_response(str(e))
|
||||
yield f"data: {data}\n\n"
|
||||
finally:
|
||||
# Stop the disconnect watcher (it may already be done if it fired).
|
||||
if _disconnect_watcher is not None and not _disconnect_watcher.done():
|
||||
_disconnect_watcher.cancel()
|
||||
try:
|
||||
await _disconnect_watcher
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
# Covers GeneratorExit when Starlette calls aclose() on disconnect
|
||||
# during decode (tokens arrive fast so CancelledError path is not
|
||||
# always triggered). abort() is a no-op for already-finished requests.
|
||||
await self.engine_client.abort(request_id)
|
||||
# Send the final done message after all response.n are finished
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
@@ -724,17 +785,47 @@ class OpenAIServingChat(OpenAIServing):
|
||||
conversation: List[ConversationMessage],
|
||||
tokenizer: AnyTokenizer,
|
||||
request_metadata: RequestResponseMetadata,
|
||||
raw_request: Optional[Request] = None,
|
||||
) -> Union[ErrorResponse, ChatCompletionResponse]:
|
||||
|
||||
model_name = self.base_model_paths[0].name
|
||||
created_time = int(time.time())
|
||||
final_res: Optional[RequestOutput] = None
|
||||
|
||||
# Background watcher: same logic as the streaming path — polls
|
||||
# is_disconnected() every 300 ms so that a client disconnect during
|
||||
# non-streaming decode is caught even when uvicorn isn't actively
|
||||
# reading the receive channel.
|
||||
_disconnect_watcher: Optional[asyncio.Task] = None
|
||||
if raw_request is not None:
|
||||
async def _watch_disconnect() -> None:
|
||||
try:
|
||||
while True:
|
||||
if await raw_request.is_disconnected():
|
||||
logger.info(
|
||||
"Client disconnected (non-stream watcher), "
|
||||
"aborting request %s", request_id)
|
||||
await self.engine_client.abort(request_id)
|
||||
return
|
||||
await asyncio.sleep(0.3)
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
_disconnect_watcher = asyncio.ensure_future(_watch_disconnect())
|
||||
|
||||
try:
|
||||
async for res in result_generator:
|
||||
final_res = res
|
||||
except asyncio.CancelledError:
|
||||
await self.engine_client.abort(request_id)
|
||||
return self.create_error_response("Client disconnected")
|
||||
finally:
|
||||
if _disconnect_watcher is not None and not _disconnect_watcher.done():
|
||||
_disconnect_watcher.cancel()
|
||||
try:
|
||||
await _disconnect_watcher
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
await self.engine_client.abort(request_id)
|
||||
|
||||
assert final_res is not None
|
||||
|
||||
@@ -876,11 +967,16 @@ class OpenAIServingChat(OpenAIServing):
|
||||
total_reasoning_tokens = sum(
|
||||
rp.count_reasoning_tokens(list(output.token_ids))
|
||||
for output in final_res.outputs)
|
||||
num_cached_tokens = (final_res.metrics.num_cached_tokens
|
||||
if final_res.metrics is not None else None)
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=num_prompt_tokens,
|
||||
completion_tokens=num_generated_tokens,
|
||||
total_tokens=num_prompt_tokens + num_generated_tokens,
|
||||
reasoning_tokens=total_reasoning_tokens,
|
||||
prompt_tokens_details=(
|
||||
PromptTokensDetails(cached_tokens=num_cached_tokens)
|
||||
if num_cached_tokens is not None else None),
|
||||
)
|
||||
|
||||
request_metadata.final_usage_info = usage
|
||||
|
||||
Reference in New Issue
Block a user