sync: align with enginex-vllm-bi100-qwen36 baseline (1902c81f)

Synced files from EngineX baseline zip (2026-06-30):
- ADD paged_attn.py (root): production paged attention with PyTorch fallback
- ADD launch_service: BI-V100 server startup script with env configuration
- SYNC computility-run.yaml: gpu_memory=0.9, batched_tokens=8192, seq_capture=32768
- SYNC qwen3_6_scripts/paged_attn.py: +311 lines, Triton bypass docs, _forward_decode_pytorch shape docs
- SYNC qwen3_6_scripts/qwen3_5.py: -72 lines, revert optimized MoE prefill to baseline (untested on BI-V100)
- KEEP Dockerfile: repo version has V2/Triton/head256 optimization patches not in baseline

Baseline commit: 1902c81fdd373943f17f5983eb8750758c7f4a69
Source: enginex-vllm-bi100-qwen36-main.zip (dev.modelhub.org.cn)
This commit is contained in:
Claude
2026-07-31 09:43:58 +00:00
parent de7ee4383e
commit c5a0d61851
5 changed files with 535 additions and 169 deletions

View File

@@ -795,66 +795,18 @@ class Qwen3_5MoeSparseBlock(nn.Module):
# 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)
# Optimized prefill MoE: sort tokens by expert for contiguous access.
# Pattern from CCCL segmented sort: partition input by key (expert),
# process each segment on contiguous memory, scatter results back.
#
# Before: for eid in unique_eids: hidden_states[tok_ids] (scattered gather)
# After: sort by expert → each F.linear gets contiguous input
# activation computed in ONE fused op across all token×expert pairs
# index_add_ for scatter-back is one kernel
T_local = hidden_states.shape[0]
two_I = w13.shape[1]
# Flatten token-expert assignments: each token appears top_k times
flat_ids = topk_ids.view(-1) # (T*K,)
flat_weights_f = topk_weights.view(-1) # (T*K,)
flat_tok_idx = torch.arange(
T_local, device=hidden_states.device
).unsqueeze(1).expand(-1, self.top_k).reshape(-1) # (T*K,)
# Sort by expert ID → contiguous memory per expert
sorted_order = flat_ids.argsort(stable=True)
sorted_eids = flat_ids[sorted_order]
sorted_tok_idx = flat_tok_idx[sorted_order]
sorted_weights_f = flat_weights_f[sorted_order]
sorted_tokens = hidden_states[sorted_tok_idx] # (T*K, H) contiguous groups
# Expert boundaries via unique_consecutive
unique_eids_t, counts = torch.unique_consecutive(
sorted_eids, return_counts=True)
cum_counts = counts.cumsum(0)
starts = torch.cat([torch.zeros(1, dtype=torch.long, device=hidden_states.device),
cum_counts[:-1]])
# Grouped up-projection: each expert's tokens are contiguous
TK = T_local * self.top_k
gate_up_all = torch.empty(TK, two_I, dtype=hidden_states.dtype,
device=hidden_states.device)
for idx_e in range(len(unique_eids_t)):
eid_val = unique_eids_t[idx_e].item()
s = starts[idx_e].item()
e = cum_counts[idx_e].item()
gate_up_all[s:e] = F.linear(sorted_tokens[s:e], w13[eid_val])
# Fused activation across ALL token×expert pairs (one kernel)
gate, up = gate_up_all.chunk(2, dim=-1)
act = F.silu(gate) * up
# Grouped down-projection
down_all = torch.empty(TK, hidden_states.shape[-1],
dtype=hidden_states.dtype,
device=hidden_states.device)
for idx_e in range(len(unique_eids_t)):
eid_val = unique_eids_t[idx_e].item()
s = starts[idx_e].item()
e = cum_counts[idx_e].item()
down_all[s:e] = F.linear(act[s:e], w2[eid_val])
# Weighted scatter-add (one kernel)
weighted = down_all * sorted_weights_f.unsqueeze(-1)
out.index_add_(0, sorted_tok_idx, weighted.to(out.dtype))
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()