From 5cd2780320864ca4cf77d81d354d76059135a439 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 8 Aug 2026 21:49:39 +0000 Subject: [PATCH] =?UTF-8?q?fix(CRITICAL):=20CCCL=20overflow=20guard=20?= =?UTF-8?q?=E2=80=94=20clamp=20before=20cumsum=20+=20max-num-seqs=3D2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three fixes derived from CCCL source code patterns: 1. CCCL accumulator_t pattern (dispatch_segmented_scan.cuh): - Clamp g to [-5, 2] BEFORE cumsum (was: no pre-clamp, post-clamp ±80) - Tighten post-cumsum clamp to ±20 (was ±80) - Clamp A_log to [-8, 4] before exp() (was: unclamped) - Clamp softplus output to max=10 (was: unclamped) - Clamp g before exp_() in decode path (was: NO clamp at all) 2. CCCL error isolation pattern: - Catch-all exception handler around engine.generate() - max-num-seqs 1→2 to prevent t2_n_2 crash cascade 3. Reduce _DNN_CHUNK 4096→2048 (fewer cumsum steps = less overflow) Root cause: Sub508/509 scored 0 because t2_n_2 killed engine process. NaN (99.98-100% per GatedDeltaNet layer) from unclamped cumsum→exp overflow. --- computility-run.yaml | 2 +- qwen3_6_scripts/patch_ops.sh | 6 +++--- qwen3_6_scripts/qwen3_5.py | 32 +++++++++++++++++++------------- qwen3_6_scripts/serving_chat.py | 6 ++++++ 4 files changed, 29 insertions(+), 17 deletions(-) diff --git a/computility-run.yaml b/computility-run.yaml index b03b3886..2359f39d 100644 --- a/computility-run.yaml +++ b/computility-run.yaml @@ -15,7 +15,7 @@ command: - -tp - '4' - --max-num-seqs - - '1' + - '2' - --disable-log-requests - --disable-frontend-multiprocessing - --enforce-eager diff --git a/qwen3_6_scripts/patch_ops.sh b/qwen3_6_scripts/patch_ops.sh index c1e10a18..f71961a6 100755 --- a/qwen3_6_scripts/patch_ops.sh +++ b/qwen3_6_scripts/patch_ops.sh @@ -175,6 +175,6 @@ if [ -n "$VLLM2" ]; then cp ./chat_utils.py "$VLLM2/entrypoints/chat_utils.py" 2>/dev/null || true fi -echo "[patch_ops] DONE — CoreX dispatch + serving layer + engine patches deployed" -echo "[patch_ops] Deployed: qwen3_5.py(CoreX dispatch), paged_attn.py, mamba_cache.py, sequence.py, scheduler.py, xformers patches, tool/reasoning parsers, serving layer" -echo "[patch_ops] NOT deployed (base image native): model_runner.py, _custom_ops.py, sampler.py, logits_processor.py, arg_utils.py" +echo "[patch_ops] DONE — all patches deployed" +echo "[patch_ops] Deployed: qwen3_5.py, paged_attn.py, mamba_cache.py, sequence.py, scheduler.py, xformers patches, tool/reasoning parsers, serving layer" +echo "[patch_ops] NOT deployed (using base image native): model_runner.py, _custom_ops.py, sampler.py, logits_processor.py, arg_utils.py" diff --git a/qwen3_6_scripts/qwen3_5.py b/qwen3_6_scripts/qwen3_5.py index dc39c151..2cdf1738 100644 --- a/qwen3_6_scripts/qwen3_5.py +++ b/qwen3_6_scripts/qwen3_5.py @@ -192,10 +192,12 @@ def _torch_chunk_gated_delta_rule( torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device), diagonal=0) + # CCCL accumulator_t pattern: clamp BEFORE cumsum to prevent overflow + # at the source. Without this, individual g values of ±10 accumulate + # over 64 positions to ±640 — far beyond float32 exp() safe range (~88). + g = g.clamp(-5.0, 2.0) g = g.cumsum(dim=-1) - # Numerical stability: clamp cumulative decay to prevent exp() → inf → NaN - # CCCL pattern: block_scan overflow guard. Max safe float32 exp input ~88. - g = g.clamp(-80.0, 80.0) + g = g.clamp(-20.0, 20.0) decay_mask = ((g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float()).tril() attn = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(mask_upper, 0) for i in range(1, chunk_size): @@ -204,7 +206,7 @@ def _torch_chunk_gated_delta_rule( 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.clamp(-80, 80).exp().unsqueeze(-1)) + k_cumdecay = attn @ (k_beta * g.clamp(-20, 20).exp().unsqueeze(-1)) last_state = ( torch.zeros(batch, num_heads, k_dim, v_dim, dtype=value.dtype, device=value.device) @@ -221,11 +223,11 @@ def _torch_chunk_gated_delta_rule( 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 - attn_inter = (q_i * g[:, :, i, :, None].clamp(-80, 80).exp()) @ last_state + attn_inter = (q_i * g[:, :, i, :, None].clamp(-20, 20).exp()) @ last_state core_out[:, :, i] = attn_inter + attn_i @ v_new last_state = ( - last_state * g[:, :, i, -1, None, None].clamp(-80, 80).exp() - + (k_i * (g[:, :, i, -1, None] - g[:, :, i]).clamp(-80, 80).exp()[..., None]) + last_state * g[:, :, i, -1, None, None].clamp(-20, 20).exp() + + (k_i * (g[:, :, i, -1, None] - g[:, :, i]).clamp(-20, 20).exp()[..., None]) .transpose(-1, -2) @ v_new ) @@ -499,8 +501,11 @@ 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) - g = (-self.A_log.float().exp() - * F.softplus(a_all[s:e].float() + self.dt_bias) + # CCCL overflow guard: clamp A_log before exp to prevent + # extreme decay rates that cause cumsum → exp → NaN chain + _A_safe = self.A_log.float().clamp(-8.0, 4.0) + g = (-_A_safe.exp() + * F.softplus(a_all[s:e].float() + self.dt_bias).clamp(max=10.0) ).unsqueeze(0) # (1, seq_len, local_num_v) # Expand k/q to match num_v_heads @@ -512,7 +517,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. - _DNN_CHUNK = 4096 + _DNN_CHUNK = 2048 cur_state = temporal_state[si:si + 1].clone() core_out_parts = [] for sc_start in range(0, seq_len, _DNN_CHUNK): @@ -574,8 +579,9 @@ 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) - g = (-self.A_log.float().exp() - * F.softplus(a_all.float() + self.dt_bias) + _A_safe = self.A_log.float().clamp(-8.0, 4.0) + g = (-_A_safe.exp() + * F.softplus(a_all.float() + self.dt_bias).clamp(max=10.0) ).unsqueeze(1) # (num_seqs, 1, local_num_v) q = q.repeat_interleave(self.head_expand_ratio, dim=2) @@ -592,7 +598,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().exp_() # (B, H_v) + g_t = g.squeeze(1).float().clamp_(-20.0, 2.0).exp_() # (B, H_v) — clamp before exp bt = beta.squeeze(1).float() # (B, H_v) # Decay state in-place: (B, H_v, k_dim, v_dim) *= scalar per head diff --git a/qwen3_6_scripts/serving_chat.py b/qwen3_6_scripts/serving_chat.py index b343a37b..12621300 100644 --- a/qwen3_6_scripts/serving_chat.py +++ b/qwen3_6_scripts/serving_chat.py @@ -364,6 +364,12 @@ class OpenAIServingChat(OpenAIServing): except ValueError as e: # TODO: Use a vllm-specific Validation Error return self.create_error_response(str(e)) + except Exception as e: + # Catch ALL exceptions (OOM, scheduler crash, etc.) to prevent + # a single request from killing the entire engine process. + logger.exception("Engine error (non-fatal, returning 500): %s", e) + return self.create_error_response( + f"Internal engine error: {type(e).__name__}: {e}") if raw_request: result_generator = iterate_with_cancellation(