Compare commits

...

10 Commits

Author SHA1 Message Date
1902c81fdd fix issue of loading weight 2026-06-30 09:55:13 +08:00
f89bc60d59 fix multiple issues 2026-06-26 17:23:55 +08:00
810874ddb8 enable prefix caching 2026-06-26 13:27:52 +08:00
c84151eef9 fix issues 2026-06-26 12:55:02 +08:00
liwei02
3d62430fd7 调整配置参数 2026-06-25 17:36:43 +08:00
72aa7e690a Add README and start commands 2026-06-23 17:17:22 +08:00
b5806731e0 some op overhead optimization 2026-06-19 11:19:39 +08:00
47a4d9e72a fix no reasoning token issue 2026-06-18 12:21:05 +08:00
3b8a567e9e fix serving issues when requesting real data 2026-06-12 17:57:23 +08:00
50e3a05fb0 fix incorrect MoE step to ensure decoding speed 2026-06-12 11:44:50 +08:00
12 changed files with 2184 additions and 196 deletions

132
README.md
View File

@@ -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 |
```

View File

@@ -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
View 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

View File

@@ -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)

View File

@@ -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]

View 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)

View File

@@ -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

View File

@@ -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,

View File

@@ -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

File diff suppressed because it is too large Load Diff

View File

@@ -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(

View File

@@ -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