Compare commits
9 Commits
85f3240c98
...
abd3d5640a
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
abd3d5640a | ||
|
|
4daa30a267 | ||
|
|
c077f7fd40 | ||
|
|
e832687893 | ||
|
|
d44ec4d8db | ||
|
|
4698cd5687 | ||
|
|
53d816b154 | ||
|
|
c59b529315 | ||
|
|
87cc24b819 |
78
PRD.md
78
PRD.md
@@ -57,3 +57,81 @@ prefix_prefill.py, logits_processor.py, mamba_cache.py, arg_utils.py
|
||||
4. ✅ d03 tool_call thinking耗尽 → 自动禁用thinking
|
||||
5. ✅ 内存碎片OOM → PYTORCH_CUDA_ALLOC_CONF
|
||||
6. ✅ 模型层代码破坏CoreX → patch_ops.sh只部署serving层
|
||||
|
||||
## CCCL tuning_select_if.cuh → serving_chat.py 映射
|
||||
|
||||
### 设计思想翻译
|
||||
CCCL三级分发:compute_capability → sm_tuning → benchmark参数
|
||||
我们三级分发:请求类型 → 处理路径 → Sub168实测参数
|
||||
|
||||
### 参数对应关系
|
||||
| CCCL概念 | 我们的对应 |
|
||||
|---------|-----------|
|
||||
| compute_capability (SM80/90/100) | 请求类型 (tool_call/reasoning/basic) |
|
||||
| input_size (1/2/4/8 bytes) | 请求复杂度 (simple/multimodal/multi-turn) |
|
||||
| flagged/unflagged | has_tools/no_tools |
|
||||
| keep_rejects/discard | enable_thinking/disable_thinking |
|
||||
| threads_per_block | max_tokens cap |
|
||||
| items_per_thread | default_max_tokens计算 |
|
||||
| delay_constructor | token budget 分配策略 |
|
||||
| benchmark注释 (4个加速比) | Sub168日志实测数据 |
|
||||
|
||||
### Sub168 benchmark数据(=我们的tuning表)
|
||||
| 请求类型 | 时间 | token数 | TPS |
|
||||
|---------|------|---------|-----|
|
||||
| d01 basic | 8.49s | 139 | 16.4 |
|
||||
| d03 tool_call | 2.12s | ~34 | ~16 |
|
||||
| d04 reasoning | 17.78s | 1192 | 67 |
|
||||
| d07 reasoning+content | 61.11s | 4451 | 72.8 |
|
||||
| replay avg | - | - | 11.86 |
|
||||
|
||||
## CCCL agent_rle.cuh → streaming SSE 映射
|
||||
|
||||
| CCCL agent_rle | serving_chat.py |
|
||||
|---------------|----------------|
|
||||
| streaming_context.num_uniques() | reasoning_token_counts[i] |
|
||||
| streaming_context.base_offset() | previous_num_tokens[i] |
|
||||
| BlockDiscontinuity (值变化检测) | reasoning_end_arr[i] (</think>检测) |
|
||||
| per-partition isolated state | per-choice state arrays |
|
||||
| ScatterDirect (压缩输出) | delta_message分发 |
|
||||
|
||||
## CCCL adjacent_difference → streaming delta 映射
|
||||
|
||||
| CCCL adjacent_diff | serving_chat.py |
|
||||
|-------------------|----------------|
|
||||
| SubtractLeftCopy | delta_text = output.text (保留原始+输出差值) |
|
||||
| previous element | previous_texts[i] |
|
||||
| current = prev + delta | current_text = previous_text + delta_text |
|
||||
| update prev = current | previous_texts[i] = current_text |
|
||||
|
||||
## CCCL tuning_batched_topk.cuh → 采样策略映射
|
||||
|
||||
### 设计思想
|
||||
6级worker_policy按tile size递减排列。运行时选最小够用的配置。
|
||||
multi_worker_policy用于超大segment的协作处理。
|
||||
|
||||
### 映射
|
||||
| CCCL batched_topk | 我们的对应 |
|
||||
|-------------------|-----------|
|
||||
| worker_policy.items_per_thread (2-64) | max_tokens cap (2048/8192) |
|
||||
| segment size → policy selection | 请求类型 → cap选择 |
|
||||
| epilogue_policy (收尾阶段) | finish_reason处理 |
|
||||
| multi_worker_policy | n>1多choice并行 |
|
||||
|
||||
### 不改sampling params的原因
|
||||
t2_temperature系列测试明确验证temperature传递。
|
||||
覆盖默认值会导致测试失败。当前策略正确。
|
||||
|
||||
## CCCL block_reduce_warp_reductions → DeltaNet chunk_size
|
||||
|
||||
### 设计思想
|
||||
Sequential path dominance → reduce per-iteration work.
|
||||
CCCL: thread count固定时,减少items_per_thread让每个thread做更少work。
|
||||
我们: Python loop iterations固定(=chunk_size),减少chunk_size从32→16。
|
||||
|
||||
## CCCL execution/exception.cuh → api_server.py _select_error_policy
|
||||
|
||||
### 设计思想
|
||||
Device code: exception_ptr永远false,不假装能恢复,直接fail fast。
|
||||
Host code: 用标准exception。
|
||||
我们: _select_error_policy按exception类型分发——OOM→503, dead→503, validation→400。
|
||||
|
||||
@@ -18,9 +18,10 @@
|
||||
# - d01: 95.87s, d03_tool_call: FAIL in 49s
|
||||
#
|
||||
# CONCLUSION: Sub168 succeeds by using BASE IMAGE native model code.
|
||||
# Our custom qwen3_5.py/model_runner/etc BREAKS CoreX acceleration.
|
||||
# qwen3_5.py MUST be deployed — base image registry references it but
|
||||
# the module file is missing (causes ModuleNotFoundError on startup).
|
||||
#
|
||||
# DO NOT deploy: qwen3_5.py, model_runner.py, _custom_ops.py,
|
||||
# DO NOT deploy: model_runner.py, _custom_ops.py,
|
||||
# sampler.py, scheduler.py, sequence.py, xformers.py, paged_attn.py,
|
||||
# prefix_prefill.py, logits_processor.py, mamba_cache.py, arg_utils.py
|
||||
# ==========================================================================
|
||||
@@ -62,7 +63,15 @@ else
|
||||
echo "[patch_ops] WARNING: transformers/models not found"
|
||||
fi
|
||||
|
||||
# 2. Registry — only if base image doesn't already have Qwen3_5
|
||||
# 2. Model module — qwen3_5.py MUST exist for registry to import.
|
||||
# Base image registry lists Qwen3_5ForCausalLM/Qwen3_5MoeForCausalLM
|
||||
# but the actual module file may be missing (causes ModuleNotFoundError
|
||||
# on startup: "No module named 'vllm.model_executor.models.qwen3_5'").
|
||||
# Deploy our qwen3_5.py so the module can be imported.
|
||||
cp ./qwen3_5.py "$VLLM/model_executor/models/qwen3_5.py" 2>/dev/null && \
|
||||
echo "[patch_ops] qwen3_5.py deployed (model module)" || true
|
||||
|
||||
# 2b. Registry — only if base image doesn't already have Qwen3_5
|
||||
if grep -q "Qwen3_5ForCausalLM" "$VLLM/model_executor/models/registry.py" 2>/dev/null; then
|
||||
echo "[patch_ops] registry already has Qwen3_5 — NOT overwriting"
|
||||
else
|
||||
@@ -99,6 +108,7 @@ for P in /usr/local/corex/lib/python3/dist-packages/vllm \
|
||||
done
|
||||
if [ -n "$VLLM2" ]; then
|
||||
echo "[patch_ops] Second vllm at: $VLLM2"
|
||||
cp ./qwen3_5.py "$VLLM2/model_executor/models/qwen3_5.py" 2>/dev/null || true
|
||||
if ! grep -q "Qwen3_5ForCausalLM" "$VLLM2/model_executor/models/registry.py" 2>/dev/null; then
|
||||
cp ./registry.py "$VLLM2/model_executor/models/registry.py" 2>/dev/null || true
|
||||
fi
|
||||
@@ -113,5 +123,6 @@ if [ -n "$VLLM2" ]; then
|
||||
cp ./chat_utils.py "$VLLM2/entrypoints/chat_utils.py" 2>/dev/null || true
|
||||
fi
|
||||
|
||||
echo "[patch_ops] DONE — serving-only patches, CoreX native model PRESERVED"
|
||||
echo "[patch_ops] NOT deployed (base image native): qwen3_5.py, model_runner.py, _custom_ops.py, sampler.py, scheduler.py, sequence.py, xformers.py, paged_attn.py, prefix_prefill.py, logits_processor.py, mamba_cache.py, arg_utils.py"
|
||||
echo "[patch_ops] DONE — serving layer + qwen3_5.py model module deployed"
|
||||
echo "[patch_ops] Deployed: qwen3_5.py (model module, required for registry import)"
|
||||
echo "[patch_ops] NOT deployed (base image native): model_runner.py, _custom_ops.py, sampler.py, scheduler.py, sequence.py, xformers.py, paged_attn.py, prefix_prefill.py, logits_processor.py, mamba_cache.py, arg_utils.py"
|
||||
|
||||
@@ -42,123 +42,6 @@ from vllm.model_executor.models.interfaces import HasInnerState, SupportsLoRA
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Hardware-aware policy dispatch (translated from CCCL cc_dispatch.cuh)
|
||||
#
|
||||
# cc_dispatch.cuh's design:
|
||||
# 1. Runtime: detect device compute capability
|
||||
# 2. policy_selector(cc) → returns kernel config (threads, items, algorithm)
|
||||
# 3. lowest_cc_resolver: merge CCs with identical policies → fewer instantiations
|
||||
# 4. dispatch_compute_cap: bridge runtime detection → compile-time specialization
|
||||
#
|
||||
# Translation to Python/PyTorch:
|
||||
# 1. Runtime: detect BI-V100 capabilities (SMEM, cuSOLVER, MoE kernels)
|
||||
# 2. _hw_policy → returns DeltaNet chunk_size, MoE strategy, solve method
|
||||
# 3. Capabilities detected once at module load, cached globally
|
||||
# 4. All kernel code reads from _hw_policy instead of hardcoded constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class _HardwarePolicy:
|
||||
"""CCCL cc_dispatch equivalent: detect hardware once, select policies."""
|
||||
|
||||
def __init__(self):
|
||||
self._detected = False
|
||||
# Defaults (safe for any hardware)
|
||||
self.deltanet_chunk_size = 64
|
||||
self.deltanet_prefill_chunk = 4096
|
||||
self.solve_triangular_available = False
|
||||
self.moe_native_topk = False
|
||||
self.moe_native_align = False
|
||||
self.moe_native_invoke = False
|
||||
self.smem_bytes = 49152 # 48KB default for BI-V100
|
||||
|
||||
def detect(self, device: torch.device = None):
|
||||
"""Run once to probe hardware capabilities. CCCL: policy_selector(cc)."""
|
||||
if self._detected:
|
||||
return
|
||||
self._detected = True
|
||||
|
||||
if device is None:
|
||||
if not torch.cuda.is_available():
|
||||
return
|
||||
device = torch.device("cuda:0")
|
||||
|
||||
# Probe SMEM (CCCL: compute_capability → SMEM size)
|
||||
try:
|
||||
idx = device.index if device.index is not None else 0
|
||||
props = torch.cuda.get_device_properties(idx)
|
||||
self.smem_bytes = props.total_memory # not SMEM, but available
|
||||
# BI-V100: 48KB confirmed via ixsmi
|
||||
self.smem_bytes = 49152
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Probe cuSOLVER/cuBLAS trsm (CCCL: check if kernel exists for this CC)
|
||||
try:
|
||||
test_A = torch.eye(4, device=device, dtype=torch.float32)
|
||||
test_b = torch.ones(4, 2, device=device, dtype=torch.float32)
|
||||
torch.linalg.solve_triangular(test_A, test_b, upper=False)
|
||||
self.solve_triangular_available = True
|
||||
except RuntimeError:
|
||||
self.solve_triangular_available = False
|
||||
|
||||
# tuning_transform_tile.cuh pick_tile_size translation:
|
||||
# Derive DeltaNet chunk_size from hardware params, not hardcode.
|
||||
#
|
||||
# CCCL formula:
|
||||
# items_for_vec = ceil(vector_bytes / min_elem_size)
|
||||
# items_for_latency = target_bytes_in_flight / (occupancy × threads × bytes_per_iter)
|
||||
# tile_size = max(items_for_vec, items_for_latency), rounded to power of 2
|
||||
#
|
||||
# For DeltaNet: chunk_size controls the (C×C) matrix in _forward_sub_lower.
|
||||
# Memory per chunk ≈ 2 × C² × sizeof(float32) × batch × heads (decay_mask + A matrix)
|
||||
# On BI-V100 with 48KB SMEM (not directly usable from PyTorch but indicates
|
||||
# hardware tier), and ~16GB GPU memory for KV cache + model:
|
||||
#
|
||||
# solve_triangular path: one cuBLAS call per chunk, larger = fewer calls
|
||||
# Python loop path: C iterations per chunk, smaller = fewer iterations
|
||||
if self.solve_triangular_available:
|
||||
# Like CCCL max_items_per_thread=32 with threads=128 → tile=4096:
|
||||
# larger chunk = amortize kernel launch overhead
|
||||
self.deltanet_chunk_size = 64
|
||||
else:
|
||||
# Like CCCL reducing items for MUFU-heavy small-elem ops:
|
||||
# smaller chunk = fewer Python loop iterations (C iterations)
|
||||
# 32 iterations vs 64 = 2× fewer kernel launches in the loop
|
||||
self.deltanet_chunk_size = 32
|
||||
|
||||
# Prefill sub-chunk: controls peak memory per DeltaNet forward call.
|
||||
# CCCL target = cc_to_min_bytes_in_flight(cc): BI-V100 ≈ lower tier.
|
||||
# Qwen3.5 DeltaNet state: (B, heads, k_dim, v_dim) ≈ (1,6,64,64)×4B = 96KB/layer
|
||||
# With _DNN_CHUNK=4096 tokens: working memory ≈ 4096×hidden×4B ≈ 60MB
|
||||
# With _DNN_CHUNK=2048: ≈ 30MB — leaves more room for KV cache
|
||||
# BI-V100 at 0.95 GPU util with 256K context needs memory headroom
|
||||
self.deltanet_prefill_chunk = 4096
|
||||
|
||||
# Probe MoE native kernels (CCCL: check op availability per CC)
|
||||
try:
|
||||
import ixformer.functions as ixf_F
|
||||
self.moe_native_topk = hasattr(ixf_F, 'vllm_moe_topk_softmax')
|
||||
self.moe_native_align = hasattr(ixf_F, 'vllm_moe_align_block_size')
|
||||
self.moe_native_invoke = hasattr(ixf_F, 'vllm_invoke_fused_moe_kernel')
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Log detected policy (CCCL: policy is logged/printed for debugging)
|
||||
logger.info(
|
||||
"HardwarePolicy detected: chunk=%d solve_tri=%s "
|
||||
"moe_native=[topk=%s align=%s invoke=%s]",
|
||||
self.deltanet_chunk_size,
|
||||
self.solve_triangular_available,
|
||||
self.moe_native_topk,
|
||||
self.moe_native_align,
|
||||
self.moe_native_invoke)
|
||||
|
||||
|
||||
# Global singleton (CCCL: policies are constexpr globals)
|
||||
_hw_policy = _HardwarePolicy()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pure-PyTorch DeltaNet kernels (fallbacks from transformers 5.2.0)
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -191,16 +74,12 @@ def _torch_chunk_gated_delta_rule(
|
||||
value: torch.Tensor, # (batch, seq, num_heads, head_v_dim)
|
||||
g: torch.Tensor, # (batch, seq, num_heads)
|
||||
beta: torch.Tensor, # (batch, seq, num_heads)
|
||||
chunk_size: int = 0, # 0 = use _hw_policy.deltanet_chunk_size
|
||||
chunk_size: int = 64,
|
||||
initial_state: Optional[torch.Tensor] = None,
|
||||
output_final_state: bool = False,
|
||||
use_qk_l2norm_in_kernel: bool = False,
|
||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
initial_dtype = query.dtype
|
||||
# cc_dispatch: resolve chunk_size from hardware policy
|
||||
if chunk_size <= 0:
|
||||
_hw_policy.detect(query.device)
|
||||
chunk_size = _hw_policy.deltanet_chunk_size
|
||||
if use_qk_l2norm_in_kernel:
|
||||
query = _l2norm(query)
|
||||
key = _l2norm(key)
|
||||
@@ -232,67 +111,16 @@ def _torch_chunk_gated_delta_rule(
|
||||
torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device),
|
||||
diagonal=0)
|
||||
|
||||
# CCCL overflow_cast_t pattern: clamp BEFORE accumulation, not after.
|
||||
# Without pre-clamp, cumsum of large g values produces huge numbers
|
||||
# that downstream exp() and matmul amplify into NaN.
|
||||
# BI-V100 docker logs show 99.98-100% NaN rate in every GatedDeltaNet layer.
|
||||
#
|
||||
# Pre-clamp: limit each g element so cumsum over chunk_size stays bounded.
|
||||
# With chunk_size=64 and per-element clamp ±0.3, cumsum range ≈ ±19.2.
|
||||
# Post-clamp to ±12 keeps exp(g_diff) ≤ exp(24) ≈ 2.6e10 — safe for
|
||||
# float32 matmul accumulation (k_dim=64 → max product ~1.7e12, within float32).
|
||||
g = g.clamp(-0.5, 0.5)
|
||||
g = g.cumsum(dim=-1)
|
||||
g = g.clamp(-12.0, 12.0)
|
||||
decay_mask = ((g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float()).tril()
|
||||
|
||||
# Lower-triangular solve WITHOUT libcusolver (not available on BI-V100).
|
||||
#
|
||||
# Computes (I - A)^{-1} @ RHS where A is strictly lower-triangular.
|
||||
# A = (k_beta @ key^T) * decay_mask, masked to lower triangle.
|
||||
#
|
||||
# Forward substitution: x[0] = rhs[0]; x[i] = rhs[i] + A[i,:i] @ x[:i]
|
||||
# Vectorized as batched matmul over chunk rows — no Python loop per row.
|
||||
# Uses torch.triangular_solve (LAPACK-based, works without cuSOLVER)
|
||||
# as primary path, with manual row-loop as fallback.
|
||||
A = ((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(mask_upper, 0)
|
||||
|
||||
# For solve: (I-A) @ X = RHS → X = (I-A)^{-1} @ RHS
|
||||
# Since (I-A) is lower-triangular with 1s on diagonal, and A is strictly
|
||||
# lower-triangular, we can use a row-by-row forward substitution.
|
||||
# This avoids cuSOLVER entirely — only needs basic matmul and indexing.
|
||||
|
||||
def _forward_sub_lower(A_lower, rhs):
|
||||
"""Solve (I - A_lower) @ X = RHS.
|
||||
|
||||
cc_dispatch pattern: _hw_policy.solve_triangular_available was probed
|
||||
once at startup. No per-call try/except overhead.
|
||||
"""
|
||||
C = rhs.shape[-2]
|
||||
if _hw_policy.solve_triangular_available:
|
||||
eye = torch.eye(C, dtype=A_lower.dtype, device=A_lower.device)
|
||||
IminusA = eye - A_lower
|
||||
return torch.linalg.solve_triangular(
|
||||
IminusA, rhs, upper=False, unitriangular=True)
|
||||
else:
|
||||
# Python forward substitution fallback with numerical stability.
|
||||
# CCCL overflow_cast pattern: clamp intermediate results per row
|
||||
# to prevent the A @ x accumulation from amplifying small errors
|
||||
# into NaN. Without this, BI-V100 shows 100% NaN in every DeltaNet layer.
|
||||
x = torch.zeros_like(rhs)
|
||||
x[..., 0, :] = rhs[..., 0, :].clamp(-1e4, 1e4)
|
||||
for i in range(1, C):
|
||||
correction = (A_lower[..., i, :i].unsqueeze(-2) @ x[..., :i, :]).squeeze(-2)
|
||||
x[..., i, :] = (rhs[..., i, :] + correction).clamp(-1e4, 1e4)
|
||||
return x
|
||||
|
||||
value = _forward_sub_lower(A, v_beta)
|
||||
|
||||
# Clamp g.exp() to prevent k_cumdecay from having extreme values
|
||||
# that would amplify in the forward substitution loop.
|
||||
k_cumdecay = _forward_sub_lower(A, k_beta * g.exp().clamp(-1e4, 1e4).unsqueeze(-1))
|
||||
|
||||
del A # free memory
|
||||
attn = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(mask_upper, 0)
|
||||
for i in range(1, chunk_size):
|
||||
row = attn[..., i, :i].clone()
|
||||
sub = attn[..., :i, :i].clone()
|
||||
attn[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
|
||||
attn = attn + torch.eye(chunk_size, dtype=attn.dtype, device=attn.device)
|
||||
value = attn @ v_beta
|
||||
k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1))
|
||||
|
||||
last_state = (
|
||||
torch.zeros(batch, num_heads, k_dim, v_dim, dtype=value.dtype, device=value.device)
|
||||
@@ -304,37 +132,18 @@ def _torch_chunk_gated_delta_rule(
|
||||
torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device),
|
||||
diagonal=1)
|
||||
|
||||
# CCCL block_scan.cuh BLOCK_SCAN_RAKING_MEMOIZE strategy:
|
||||
# Precompute all per-chunk exp values outside the loop, eliminating
|
||||
# redundant exp() inside the sequential cross-chunk scan.
|
||||
# RAKING_MEMOIZE: "preserve upsweep segment values in registers while
|
||||
# performing warp-synchronous scan, allowing downsweep not to re-read."
|
||||
num_chunks = total_len // chunk_size
|
||||
# g shape: (batch, heads, num_chunks, chunk_size)
|
||||
# g_exp_full[i] = exp(g[:,:,i,:]) for attn_inter computation
|
||||
g_exp_full = g.exp() # (batch, heads, num_chunks, chunk_size)
|
||||
# g_last_exp[i] = exp(g[:,:,i,-1]) for state decay
|
||||
g_last_exp = g_exp_full[:, :, :, -1] # (batch, heads, num_chunks)
|
||||
# g_diff_exp[i] = exp(g[:,:,i,-1] - g[:,:,i,:]) for k_i weighting
|
||||
g_diff_exp = (g[:, :, :, -1:] - g).exp() # (batch, heads, num_chunks, chunk_size)
|
||||
|
||||
for i in range(num_chunks):
|
||||
for i in range(total_len // chunk_size):
|
||||
q_i, k_i, v_i = query[:, :, i], key[:, :, i], value[:, :, i]
|
||||
attn_i = (q_i @ k_i.transpose(-1, -2) * decay_mask[:, :, i]).masked_fill_(mask_upper2, 0)
|
||||
v_prime = k_cumdecay[:, :, i] @ last_state
|
||||
v_new = v_i - v_prime
|
||||
# Use precomputed exp (MEMOIZE: no redundant exp in loop body)
|
||||
attn_inter = (q_i * g_exp_full[:, :, i, :, None]) @ last_state
|
||||
attn_inter = (q_i * g[:, :, i, :, None].exp()) @ last_state
|
||||
core_out[:, :, i] = attn_inter + attn_i @ v_new
|
||||
last_state = (
|
||||
last_state * g_last_exp[:, :, i, None, None]
|
||||
+ (k_i * g_diff_exp[:, :, i, :, None])
|
||||
last_state * g[:, :, i, -1, None, None].exp()
|
||||
+ (k_i * (g[:, :, i, -1, None] - g[:, :, i]).exp()[..., None])
|
||||
.transpose(-1, -2) @ v_new
|
||||
)
|
||||
# CCCL numerical guard: clamp state to prevent cross-chunk accumulation
|
||||
# from amplifying into NaN. State elements represent k_dim × v_dim
|
||||
# attention memory; values beyond ±1e4 indicate numerical runaway.
|
||||
last_state = last_state.clamp(-1e4, 1e4)
|
||||
|
||||
if not output_final_state:
|
||||
last_state = None
|
||||
@@ -559,15 +368,7 @@ class GatedDeltaNet(nn.Module):
|
||||
v = v.reshape(1, seq_len, local_num_v, self.head_v_dim)
|
||||
|
||||
beta = b_all[s:e].sigmoid().unsqueeze(0) # (1, seq_len, local_num_v)
|
||||
# CCCL overflow_cast pattern: clamp before exp to prevent
|
||||
# overflow → NaN cascade. Tightened to [-5,5] because:
|
||||
# A_log.exp() range [0.007, 148.4] — moderate decay rates.
|
||||
# Multiplied by softplus(a + dt_bias) ≈ [0.7, 10] → g ≈ [-1484, -0.005]
|
||||
# Per-element g then gets clamped to [-0.5, 0.5] in chunk_gated_delta_rule.
|
||||
# The tighter clamp here prevents A_log outliers from creating
|
||||
# extreme g values before the chunk-level clamp catches them.
|
||||
_A_safe = self.A_log.float().clamp(-5.0, 5.0)
|
||||
g = (-_A_safe.exp()
|
||||
g = (-self.A_log.float().exp()
|
||||
* F.softplus(a_all[s:e].float() + self.dt_bias)
|
||||
).unsqueeze(0) # (1, seq_len, local_num_v)
|
||||
|
||||
@@ -580,8 +381,7 @@ class GatedDeltaNet(nn.Module):
|
||||
# Full 18K: tensors [1,6,282,64,64]=220 MB each → ~990 MB/call.
|
||||
# With _DNN_CHUNK=4096: [1,6,64,64,64]=6 MB each → ~137 MB/call.
|
||||
# State is chained via initial_state / output_final_state.
|
||||
_hw_policy.detect(hidden_states.device)
|
||||
_DNN_CHUNK = _hw_policy.deltanet_prefill_chunk
|
||||
_DNN_CHUNK = 4096
|
||||
cur_state = temporal_state[si:si + 1].clone()
|
||||
core_out_parts = []
|
||||
for sc_start in range(0, seq_len, _DNN_CHUNK):
|
||||
@@ -613,20 +413,9 @@ class GatedDeltaNet(nn.Module):
|
||||
outputs.append(out)
|
||||
|
||||
result = torch.cat(outputs, dim=0)
|
||||
# thrust::all_of early termination pattern (bench/all_of/basic.cu):
|
||||
# Check a small sample first — if no NaN in sample, skip full scan.
|
||||
# MismatchAt=0.01 insight: NaN usually appears early or everywhere.
|
||||
# Sample first 64 elements + last 64 — covers most failure modes.
|
||||
_n = result.numel()
|
||||
_sample_ok = True
|
||||
if _n > 128:
|
||||
_s = result.view(-1)
|
||||
_sample_ok = not (torch.isnan(_s[:64]).any() or torch.isnan(_s[-64:]).any())
|
||||
if not _sample_ok or (_n <= 128 and torch.isnan(result).any()):
|
||||
# Full scan only when sample detected NaN
|
||||
nan_frac = torch.isnan(result).float().mean().item()
|
||||
if torch.isnan(result).any():
|
||||
logger.warning("NaN in prefill GatedDeltaNet layer %d (frac=%.4f), replacing with zeros",
|
||||
self.layer_idx, nan_frac)
|
||||
self.layer_idx, torch.isnan(result).float().mean().item())
|
||||
result = torch.nan_to_num(result, nan=0.0)
|
||||
return result
|
||||
|
||||
@@ -654,9 +443,7 @@ class GatedDeltaNet(nn.Module):
|
||||
v = v.reshape(num_seqs, 1, local_num_v, self.head_v_dim)
|
||||
|
||||
beta = b_all.sigmoid().unsqueeze(1) # (num_seqs, 1, local_num_v)
|
||||
# CCCL overflow_cast pattern: tightened to [-5,5] matching prefill path
|
||||
_A_safe = self.A_log.float().clamp(-5.0, 5.0)
|
||||
g = (-_A_safe.exp()
|
||||
g = (-self.A_log.float().exp()
|
||||
* F.softplus(a_all.float() + self.dt_bias)
|
||||
).unsqueeze(1) # (num_seqs, 1, local_num_v)
|
||||
|
||||
@@ -674,7 +461,7 @@ class GatedDeltaNet(nn.Module):
|
||||
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().clamp_(-12.0, 12.0).exp_() # (B, H_v) overflow_cast tightened
|
||||
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
|
||||
@@ -696,8 +483,6 @@ class GatedDeltaNet(nn.Module):
|
||||
k_t.view(BH, self.head_k_dim, 1),
|
||||
delta.view(BH, 1, self.head_v_dim),
|
||||
)
|
||||
# CCCL numerical guard: clamp decode state (same as prefill cross-chunk)
|
||||
ts_flat.clamp_(-1e4, 1e4)
|
||||
|
||||
# Output: core_out = q_t @ updated temporal_state
|
||||
core_out = torch.bmm(
|
||||
@@ -711,18 +496,9 @@ class GatedDeltaNet(nn.Module):
|
||||
z.reshape(-1, self.head_v_dim))
|
||||
normed = normed.reshape(num_seqs, -1)
|
||||
out, _ = self.out_proj(normed)
|
||||
# thrust::all_of early termination: sample check before full scan
|
||||
_n = out.numel()
|
||||
_has_nan = False
|
||||
if _n > 128:
|
||||
_s = out.view(-1)
|
||||
_has_nan = torch.isnan(_s[:64]).any() or torch.isnan(_s[-64:]).any()
|
||||
else:
|
||||
_has_nan = torch.isnan(out).any().item()
|
||||
if _has_nan:
|
||||
nan_frac = torch.isnan(out).float().mean().item()
|
||||
if torch.isnan(out).any():
|
||||
logger.warning("NaN in decode GatedDeltaNet layer %d (frac=%.4f), replacing with zeros",
|
||||
self.layer_idx, nan_frac)
|
||||
self.layer_idx, torch.isnan(out).float().mean().item())
|
||||
out = torch.nan_to_num(out, nan=0.0)
|
||||
return out
|
||||
|
||||
@@ -915,16 +691,10 @@ class Qwen3_5MLP(nn.Module):
|
||||
class Qwen3_5MoeSparseBlock(nn.Module):
|
||||
"""Replaces Qwen3_5MLP for qwen3_5_moe_text layers.
|
||||
|
||||
FusedMoE stores expert weights and provides native ixformer forward kernel.
|
||||
Forward tries the native fused kernel first (one CUDA launch for all experts),
|
||||
falling back to _pure_pytorch_experts if the native kernel fails on BI-V100.
|
||||
|
||||
CCCL architecture insight (dispatch_reduce_by_key.cuh):
|
||||
The native fused_moe_kernel implements the same pattern as CCCL's
|
||||
DeviceReduceByKey — sort tokens by expert_id, pad to block boundary
|
||||
(moe_align_block_size), then one kernel processes all expert-token pairs
|
||||
with block-level parallelism. This is the architecturally correct approach
|
||||
vs the fallback's Python for-loop over experts.
|
||||
FusedMoE is used ONLY for weight storage and loading (create_weights /
|
||||
weight_loader are pure PyTorch). Its forward kernel is bypassed because
|
||||
ixformer on BI-V100 lacks vllm_moe_topk_softmax / vllm_invoke_fused_moe_kernel.
|
||||
Routing and expert computation use a pure-PyTorch loop instead.
|
||||
|
||||
Shared expert uses RowParallelLinear(reduce_results=False) so both paths
|
||||
produce partial (pre-all-reduce) outputs that are combined before a single
|
||||
@@ -972,14 +742,6 @@ class Qwen3_5MoeSparseBlock(nn.Module):
|
||||
self.shared_expert_gate = ReplicatedLinear(
|
||||
hidden_size, 1, bias=False, quant_config=quant_config)
|
||||
|
||||
# sync_handler.cuh: register resources at init, initialize once.
|
||||
# Pre-declare MoE strategy here (resolved on first forward when device
|
||||
# is known). _use_native_moe is set to None = "not yet decided".
|
||||
# This avoids hasattr() checks in the forward hot path.
|
||||
self._use_native_moe: Optional[bool] = None
|
||||
self._moe_out_buf: Optional[torch.Tensor] = None
|
||||
self._moe_out_buf_key: Optional[tuple] = None
|
||||
|
||||
def _pure_pytorch_experts(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -1030,132 +792,27 @@ class Qwen3_5MoeSparseBlock(nn.Module):
|
||||
out = (expert_out * ws.unsqueeze(-1)).sum(0, keepdim=True).to(
|
||||
hidden_states.dtype) # (1, H)
|
||||
else:
|
||||
# General path (prefill / multi-seq): CCCL histogram sort+reduce pattern.
|
||||
#
|
||||
# CCCL insight (thrust/examples/histogram.cu sparse_histogram):
|
||||
# sort data → reduce_by_key over contiguous segments.
|
||||
# Applied to MoE: sort (token, expert) pairs by expert_id so all tokens
|
||||
# routed to the same expert are contiguous, then process each expert's
|
||||
# batch with a single F.linear call.
|
||||
#
|
||||
# Previous code: for-loop over unique experts, each with F.linear.
|
||||
# With 256 experts × top_k=8 ≈ up to 256 active experts → 512 F.linear calls.
|
||||
# New code: sort + segment → same number of F.linear calls but with
|
||||
# contiguous token batches (better GPU occupancy) + no Python dict lookup.
|
||||
#
|
||||
# Further optimization: group experts by similar token count and pad
|
||||
# to enable batched GEMM across expert groups (CCCL segmented_reduce pattern).
|
||||
# TODO: implement when we have benchmark data showing this path is hot.
|
||||
|
||||
# smem_resource_raw.cuh: reuse buffer across calls.
|
||||
# CCCL manages SMEM as multi-stage ping-pong: same memory, different
|
||||
# stages. We do the same: keep a class-level buffer, resize only if
|
||||
# shape changes, zero in-place instead of allocating.
|
||||
_buf_key = (T, hidden_states.shape[-1])
|
||||
if not hasattr(self, '_moe_out_buf') or self._moe_out_buf_key != _buf_key:
|
||||
self._moe_out_buf = torch.zeros_like(hidden_states)
|
||||
self._moe_out_buf_key = _buf_key
|
||||
else:
|
||||
self._moe_out_buf.zero_()
|
||||
out = self._moe_out_buf
|
||||
|
||||
# Flatten all (token, expert) assignments: (T*top_k,) pairs
|
||||
flat_eids = topk_ids.view(-1) # (T*K,)
|
||||
flat_tok_ids = torch.arange(T, device=hidden_states.device).unsqueeze(1) \
|
||||
.expand(-1, self.top_k).reshape(-1) # (T*K,)
|
||||
flat_topk_pos = torch.arange(self.top_k, device=hidden_states.device) \
|
||||
.unsqueeze(0).expand(T, -1).reshape(-1) # (T*K,)
|
||||
|
||||
# Sort by expert_id — CCCL histogram pattern: sort brings equal keys together
|
||||
sort_idx = flat_eids.argsort(stable=True)
|
||||
sorted_eids = flat_eids[sort_idx]
|
||||
sorted_tok_ids = flat_tok_ids[sort_idx]
|
||||
sorted_topk_pos = flat_topk_pos[sort_idx]
|
||||
|
||||
# thrust/examples/mode.cu complete pipeline translation:
|
||||
# sort → unique_count → reduce_by_key(data, constant_iterator<1>) → max_element
|
||||
# torch.unique_consecutive = sort's reduce_by_key in one fused call.
|
||||
# Returns (unique_keys, inverse, counts) — mode.cu builds the same from
|
||||
# sort + reduce_by_key(data, constant_iterator<1>, keys, counts).
|
||||
# Replaces: changes detection → nonzero → concat → 3 separate GPU ops.
|
||||
seg_eids, _inv, seg_counts = torch.unique_consecutive(
|
||||
sorted_eids, return_inverse=True, return_counts=True)
|
||||
seg_ends = seg_counts.cumsum(0)
|
||||
seg_starts = torch.cat([
|
||||
torch.zeros(1, dtype=seg_ends.dtype, device=seg_ends.device),
|
||||
seg_ends[:-1]])
|
||||
|
||||
# Process each expert segment
|
||||
seg_starts_cpu = seg_starts.tolist()
|
||||
seg_ends_cpu = seg_ends.tolist()
|
||||
seg_eids_cpu = seg_eids.tolist()
|
||||
for seg_i in range(len(seg_starts_cpu)):
|
||||
s, e = seg_starts_cpu[seg_i], seg_ends_cpu[seg_i]
|
||||
eid = seg_eids_cpu[seg_i]
|
||||
tok_ids_seg = sorted_tok_ids[s:e]
|
||||
topk_pos_seg = sorted_topk_pos[s:e]
|
||||
|
||||
# dispatch_copy_mdspan.cuh: check if data is exhaustive (contiguous).
|
||||
# If token IDs form a contiguous range, use slice (zero-copy)
|
||||
# instead of fancy indexing (allocates new tensor).
|
||||
n_seg = e - s
|
||||
first_tok = int(tok_ids_seg[0])
|
||||
if n_seg > 1 and int(tok_ids_seg[-1]) == first_tok + n_seg - 1:
|
||||
# Fast path: contiguous slice (no copy)
|
||||
tokens = hidden_states[first_tok:first_tok + n_seg]
|
||||
else:
|
||||
# Slow path: gather by index
|
||||
tokens = hidden_states[tok_ids_seg]
|
||||
# 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_seg, topk_pos_seg].unsqueeze(-1)
|
||||
out.index_add_(0, tok_ids_seg, (expert_out * weights).to(out.dtype))
|
||||
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()
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
router_logits, _ = self.gate(hidden_states)
|
||||
|
||||
# Try native FusedMoE path first (ixformer kernel).
|
||||
# CCCL dispatch_reduce_by_key.cuh insight: the native fused kernel does
|
||||
# sort-by-expert + block-aligned GEMM in one launch — architecturally
|
||||
# identical to CCCL's AgentReduceByKey::ConsumeRange.
|
||||
# One fused kernel vs our _pure_pytorch_experts' 256× F.linear calls.
|
||||
#
|
||||
# _custom_ops.py confirms ixformer HAS these ops:
|
||||
# ixf_F.vllm_moe_topk_softmax
|
||||
# ixf_F.vllm_moe_align_block_size
|
||||
# ixf_F.vllm_invoke_fused_moe_kernel
|
||||
# The original comment "ixformer lacks MoE kernels" may have been
|
||||
# wrong or outdated. Try native first, catch and fallback if it fails.
|
||||
# cc_dispatch + sync_handler: strategy resolved on first call,
|
||||
# pre-registered field checked as None (no hasattr overhead).
|
||||
if self._use_native_moe is None:
|
||||
_hw_policy.detect(hidden_states.device)
|
||||
# Only try native if at least align+invoke are available
|
||||
# (topk_softmax has PyTorch fallback in _custom_ops.py)
|
||||
self._use_native_moe = (
|
||||
_hw_policy.moe_native_align and _hw_policy.moe_native_invoke)
|
||||
if not self._use_native_moe:
|
||||
logger.info(
|
||||
"HardwarePolicy: MoE native kernels unavailable "
|
||||
"(align=%s invoke=%s), using PyTorch experts.",
|
||||
_hw_policy.moe_native_align, _hw_policy.moe_native_invoke)
|
||||
|
||||
if self._use_native_moe:
|
||||
try:
|
||||
routed_out = self.experts(hidden_states, router_logits)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"FusedMoE native kernel failed (%s: %s), "
|
||||
"falling back to pure PyTorch experts permanently.",
|
||||
type(e).__name__, e)
|
||||
self._use_native_moe = False
|
||||
routed_out = self._pure_pytorch_experts(hidden_states, router_logits)
|
||||
else:
|
||||
routed_out = self._pure_pytorch_experts(hidden_states, router_logits)
|
||||
routed_out = self._pure_pytorch_experts(hidden_states, router_logits)
|
||||
|
||||
gate_up, _ = self.shared_expert_gate_up(hidden_states)
|
||||
shared_out = self.act_fn(gate_up)
|
||||
|
||||
1369
qwen3_6_scripts/qwen3_5_base_original.py
Normal file
1369
qwen3_6_scripts/qwen3_5_base_original.py
Normal file
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user