feat(SM70): wire up FlashQLA GDN kernel dispatch in prefill path

GDN forward dispatch chain:
1. CoreX fused kernel (if packaged) → fastest
2. FlashQLA SM70 CUDA kernel (prefill only) → verified on BI-V100
3. Pure PyTorch with NaN clamp → fallback

FlashQLA SM70 verified on real BI-V100:
- Compiled with clang++ --cuda-gpu-arch=ivcore10
- gdn_forward returns correct shapes, zero NaN
- 4 kernels: prefill, varlen prefill, decode global, decode ddtree

Also: apt ninja-build instead of pip ninja (pip version has no binary)
This commit is contained in:
Claude
2026-08-10 01:42:16 +00:00
parent 47958c4ed2
commit c077736968
2 changed files with 108 additions and 2 deletions

View File

@@ -57,8 +57,11 @@ for P in /usr/local/lib/python3.10/site-packages/transformers/models \
done
if [ -n "$TMODELS" ]; then
# Base engine requires transformers 4.55.3 for Qwen3_5Config support
pip install transformers==4.55.3 ninja -i https://pypi.tuna.tsinghua.edu.cn/simple --timeout 30 2>&1 || \
pip install transformers==4.55.3 -i https://pypi.tuna.tsinghua.edu.cn/simple --timeout 30 2>&1 || \
echo "[patch_ops] WARNING: pip install failed (may already be correct versions)"
# ninja-build required for torch.utils.cpp_extension CUDA compilation
apt-get update -qq && apt-get install -y -qq ninja-build 2>&1 || \
echo "[patch_ops] WARNING: ninja-build install failed — CUDA kernel will not compile"
cp -r ./qwen3_5 "$TMODELS/" 2>/dev/null && echo "[patch_ops] qwen3_5 config copied" || true
cp -r ./qwen3_5_moe "$TMODELS/" 2>/dev/null && echo "[patch_ops] qwen3_5_moe config copied" || true
python3 ./patch_transformers_qwen3_5.py 2>&1 || echo "[patch_ops] WARNING: transformers patch failed (non-fatal)"

View File

@@ -466,12 +466,115 @@ class GatedDeltaNet(nn.Module):
except Exception as e:
if self.layer_idx == 0:
logger.warning(
"CoreX GDN forward failed (%s), falling back to PyTorch permanently", e)
"CoreX GDN forward failed (%s), falling back", e)
self._use_corex_gdn = False # permanent fallback
# FlashQLA SM70 dispatch: fused CUDA kernel for prefill
# Decode stays PyTorch (SM70 decode kernel needs different state layout)
if _flash_qla_available and attn_metadata.num_prefill_tokens > 0:
try:
return self._flash_qla_prefill(
hidden_states, attn_metadata, conv_state, temporal_state)
except Exception as e:
if self.layer_idx == 0:
logger.warning(
"FlashQLA SM70 prefill failed (%s), falling back to PyTorch", e)
# Don't disable permanently — may work for different shapes
return self._pytorch_forward(
hidden_states, attn_metadata, conv_state, temporal_state)
def _flash_qla_prefill(
self,
hidden_states: torch.Tensor,
attn_metadata: AttentionMetadata,
conv_state: torch.Tensor,
temporal_state: torch.Tensor,
) -> torch.Tensor:
"""Prefill using FlashQLA SM70 fused CUDA kernel."""
tp_size = get_tensor_model_parallel_world_size()
local_key_dim = self.key_dim // tp_size
local_val_dim = self.value_dim // tp_size
local_num_v = self.num_v_heads // tp_size
local_num_k = self.num_k_heads // tp_size
local_conv_dim = self.conv_dim // tp_size
# Project all tokens
mixed_qkv_all, _ = self.in_proj_qkv(hidden_states)
z_all, _ = self.in_proj_z(hidden_states)
b_all, _ = self.in_proj_b(hidden_states)
a_all, _ = self.in_proj_a(hidden_states)
seq_starts = attn_metadata.query_start_loc.tolist()
outputs = []
for i in range(len(seq_starts) - 1):
s, e = seq_starts[i], seq_starts[i + 1]
L = e - s
if L == 0:
continue
mixed = mixed_qkv_all[s:e] # (L, local_conv_dim)
z_seq = z_all[s:e]
b_seq = torch.sigmoid(b_all[s:e]) # (L, local_num_v)
dt = F.softplus(a_all[s:e] + self.dt_bias) # (L, local_num_v)
gate = -dt * self.A_log.exp() # (L, local_num_v) — decay
# Conv1d
conv_out = F.conv1d(
F.pad(mixed.unsqueeze(0).transpose(1, 2),
(self.conv_kernel_size - 1, 0)),
self.conv1d_weight, groups=local_conv_dim
).transpose(1, 2).squeeze(0)
# Split into q, k, v
qkv = conv_out.view(L, local_num_k + local_num_k + local_num_v,
self.head_k_dim)
q_raw = qkv[:, :local_num_k, :]
k_raw = qkv[:, local_num_k:2*local_num_k, :]
v_raw = qkv[:, 2*local_num_k:, :local_val_dim // local_num_v]
# L2 normalize q, k
q = _l2norm(q_raw)
k = _l2norm(k_raw)
# Reshape to [1, L, H, D] for SM70 kernel
q_4d = q.unsqueeze(0) # (1, L, Hk, K)
k_4d = k.unsqueeze(0) # (1, L, Hk, K)
v_4d = v_raw.unsqueeze(0) # (1, L, Hv, V)
g_3d = gate.unsqueeze(0) # (1, L, Hv)
beta_3d = b_seq.unsqueeze(0) # (1, L, Hv)
# Initial state from temporal_state
init_state = temporal_state[i:i+1] # (1, Hv, K, V)
# Call SM70 fused kernel
output_4d, final_state = chunk_gated_delta_rule_fwd_sm70(
q_4d, k_4d, v_4d, g_3d, beta_3d,
scale=1.0, # q already normalized
initial_state=init_state,
output_final_state=True,
gate_is_exp=False,
)
# Update temporal state
if final_state is not None:
temporal_state[i] = final_state[0]
# output_4d: (1, L, Hv, V) → (L, local_val_dim)
out_seq = output_4d.squeeze(0).reshape(L, local_val_dim)
# Apply gated RMSNorm + z gate
z_seq_heads = z_seq.view(L, local_num_v, self.head_v_dim)
out_heads = out_seq.view(L, local_num_v, self.head_v_dim)
normed = self.norm(out_heads, z_seq_heads)
normed_flat = normed.reshape(L, local_val_dim)
proj_out, _ = self.out_proj(normed_flat)
outputs.append(proj_out)
return torch.cat(outputs, dim=0)
def _pytorch_forward(
self,
hidden_states: torch.Tensor,