fix(NO-FALLBACK): eliminate all silent fallbacks — crash or succeed
Policy: fallback = 0 score = same as crash. Better to crash with clear
error log so we can diagnose.
Changes:
1. corex_gdn.py: COMPLETE REWRITE (374 lines)
- CoreXGDN.forward() now implements full GDN layer forward
- Accepts all 13 args from qwen3_5.py (hidden_states, attn_metadata,
conv_state, temporal_state, in_proj_qkv/z/b/a, conv1d_weight,
A_log, dt_bias, norm, out_proj)
- Prefill: causal conv1d → split q/k/v → chunk_gated_delta_rule
(fp32 accumulation, xllm-aligned cumsum+difference form)
- Decode: causal_conv1d_update → single-step recurrent with
bmm/baddbmm_ (ixformer accelerated)
- NO FALLBACK — if something fails, it crashes
2. qwen3_5.py: Remove all try/except fallbacks
- GatedDeltaNet.__init__: CoreXGDN init MUST succeed (no try/except)
- GatedDeltaNet.forward: CoreXGDN.forward() called directly, no catch
- MoE init: raise RuntimeError if moe_forward missing
3. patch_ops.sh: MUST deploy all three corex modules
- Reverted previous 'don't overwrite' — base image produces NaN
- corex_gdn.py + corex_moe.py + corex_fa2.py all deployed unconditionally
This commit is contained in:
@@ -459,38 +459,19 @@ class GatedDeltaNet(nn.Module):
|
||||
self.norm = Qwen3_5RMSNormGated(self.head_v_dim,
|
||||
eps=text_cfg.rms_norm_eps)
|
||||
|
||||
# CoreX dispatch: try to create fused GDN operator from base image
|
||||
# CoreX dispatch — our corex_gdn.py is deployed, init MUST succeed
|
||||
self._use_corex_gdn = False
|
||||
if _corex_gdn_available and _corex_gdn_module is not None:
|
||||
try:
|
||||
# Try base image's CoreXGDN signature first (may differ from ours)
|
||||
self._corex_gdn_obj = _corex_gdn_module.CoreXGDN(
|
||||
num_v_heads=self.num_v_heads // tp_size,
|
||||
num_k_heads=self.num_k_heads // tp_size,
|
||||
head_k_dim=self.head_k_dim,
|
||||
head_v_dim=self.head_v_dim,
|
||||
conv_kernel_size=self.conv_kernel_size,
|
||||
layer_idx=layer_idx,
|
||||
)
|
||||
self._use_corex_gdn = True
|
||||
except TypeError:
|
||||
# Fallback: simpler signature
|
||||
try:
|
||||
self._corex_gdn_obj = _corex_gdn_module.CoreXGDN(
|
||||
self.num_v_heads // tp_size,
|
||||
self.head_k_dim,
|
||||
layer_idx=layer_idx,
|
||||
)
|
||||
self._use_corex_gdn = True
|
||||
except Exception as e2:
|
||||
logger.warning(
|
||||
"GatedDeltaNet layer %d: CoreX GDN init failed (%s), PyTorch",
|
||||
layer_idx, e2)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"GatedDeltaNet layer %d: CoreX GDN init failed (%s), PyTorch",
|
||||
layer_idx, e)
|
||||
if self._use_corex_gdn and layer_idx == 0:
|
||||
self._corex_gdn_obj = _corex_gdn_module.CoreXGDN(
|
||||
num_v_heads=self.num_v_heads // tp_size,
|
||||
num_k_heads=self.num_k_heads // tp_size,
|
||||
head_k_dim=self.head_k_dim,
|
||||
head_v_dim=self.head_v_dim,
|
||||
conv_kernel_size=self.conv_kernel_size,
|
||||
layer_idx=layer_idx,
|
||||
)
|
||||
self._use_corex_gdn = True
|
||||
if layer_idx == 0:
|
||||
logger.info("GatedDeltaNet: CoreX fused GDN enabled")
|
||||
|
||||
def _conv1d_weight_loader(self, param: torch.Tensor,
|
||||
@@ -517,22 +498,16 @@ class GatedDeltaNet(nn.Module):
|
||||
conv_state: torch.Tensor, # (batch, local_conv_dim, kernel-1) in-place
|
||||
temporal_state: torch.Tensor, # (batch, local_v_heads, k_dim, v_dim) in-place
|
||||
) -> torch.Tensor:
|
||||
# CoreX dispatch: try fused GDN kernel first (CCCL env_dispatch pattern)
|
||||
# CoreX dispatch — NO FALLBACK. 0 score with fallback = same as crash.
|
||||
if self._use_corex_gdn:
|
||||
try:
|
||||
return self._corex_gdn_obj.forward(
|
||||
hidden_states, attn_metadata,
|
||||
conv_state, temporal_state,
|
||||
self.in_proj_qkv, self.in_proj_z,
|
||||
self.in_proj_b, self.in_proj_a,
|
||||
self.conv1d_weight, self.A_log, self.dt_bias,
|
||||
self.norm, self.out_proj,
|
||||
)
|
||||
except Exception as e:
|
||||
if self.layer_idx == 0:
|
||||
logger.warning(
|
||||
"CoreX GDN forward failed (%s), falling back", e)
|
||||
self._use_corex_gdn = False # permanent fallback
|
||||
return self._corex_gdn_obj.forward(
|
||||
hidden_states, attn_metadata,
|
||||
conv_state, temporal_state,
|
||||
self.in_proj_qkv, self.in_proj_z,
|
||||
self.in_proj_b, self.in_proj_a,
|
||||
self.conv1d_weight, self.A_log, self.dt_bias,
|
||||
self.norm, self.out_proj,
|
||||
)
|
||||
|
||||
# flash_qla SM70 DISABLED: produces inf on BI-V100 (abs mean=inf from real test)
|
||||
# xllm uses equivalent PyTorch chunked path (qwen3_gated_delta_net_base.cpp)
|
||||
@@ -1094,20 +1069,17 @@ class Qwen3_5MoeSparseBlock(nn.Module):
|
||||
self.shared_expert_gate = ReplicatedLinear(
|
||||
hidden_size, 1, bias=False, quant_config=quant_config)
|
||||
|
||||
# CoreX dispatch: try to use fused MoE kernels from base image
|
||||
# CoreX dispatch — corex_moe.py is deployed, moe_forward MUST exist
|
||||
self._use_corex_moe = False
|
||||
if _corex_moe_available and _corex_moe_module is not None:
|
||||
try:
|
||||
# corex_moe module provides direct forward functions
|
||||
self._corex_moe_forward = getattr(
|
||||
_corex_moe_module, 'moe_forward', None)
|
||||
if self._corex_moe_forward is not None:
|
||||
self._use_corex_moe = True
|
||||
self._corex_moe_forward = getattr(
|
||||
_corex_moe_module, 'moe_forward', None)
|
||||
if self._corex_moe_forward is not None:
|
||||
self._use_corex_moe = True
|
||||
if layer_idx == 0:
|
||||
logger.info("MoE: CoreX fused MoE forward available")
|
||||
else:
|
||||
logger.warning("MoE: corex_moe has no moe_forward, using PyTorch")
|
||||
except Exception as e:
|
||||
logger.warning("MoE: CoreX MoE init failed (%s), using PyTorch", e)
|
||||
else:
|
||||
raise RuntimeError("corex_moe module loaded but moe_forward missing")
|
||||
|
||||
def _pure_pytorch_experts(
|
||||
self,
|
||||
|
||||
Reference in New Issue
Block a user