perf: P2 landed (gsa fill elimination). P0/P1 fused SwiGLU disabled — CuTeDSL kernel arg-binding bug.
P0/P1: The fused SwiGLU kernel's warmup_fused_swiglu_compilation() triggers 'TypeError: too many positional arguments' during cute.compile(). The kernel signature doesn't match the positional args being passed. This is a kernel-side fix, not a single_shot fix. Disabled until the fused kernel is debugged. P2: Landed — Nvfp4Linear skips redundant _gsa_buf.fill_() after warmup. SE fused SwiGLU infrastructure (set_fused_swiglu, _run_l1_fused, interleaved weight path) is wired but disabled. Will activate once kernel fix lands.
This commit is contained in:
@@ -254,14 +254,13 @@ class Nvfp4SharedExpert:
|
||||
# Run fused grouped GEMM with 1 group
|
||||
l1_out = run_fused_swiglu_grouped_gemm(
|
||||
mat_a=x_fp4,
|
||||
scale_a=x_sf,
|
||||
global_scale_a=gsa,
|
||||
mat_b=self._l1_mat_b,
|
||||
scale_a=x_sf,
|
||||
scale_b=self._l1_sf_view,
|
||||
global_scale_b=self._l1_gs_view,
|
||||
expert_offsets=torch.tensor([num_tokens], dtype=torch.int64, device=x_fp4.device),
|
||||
global_scale_a=gsa,
|
||||
global_scale_b=self._l1_gs_view,
|
||||
swiglu_limit=self.swiglu_limit if self.swiglu_limit is not None else 0.0,
|
||||
num_tokens=num_tokens,
|
||||
)
|
||||
return l1_out # (num_tokens, intermediate_size) BF16, SwiGLU already applied
|
||||
|
||||
|
||||
@@ -1021,7 +1021,7 @@ def main():
|
||||
intermediate_size=cfg.get("moe_intermediate_size", 3072),
|
||||
top_k=cfg.get("num_experts_per_tok", 6), device=dev)
|
||||
moe.set_swiglu_limit(cfg.get("swiglu_limit", 10.0))
|
||||
moe._fused_swiglu = True # P0: Enable fused SwiGLU kernel — eliminates 8 BF16 launches per MoE per token
|
||||
moe._fused_swiglu = False # P0: Fused SwiGLU kernel has CuTeDSL arg-binding issue — disabled until kernel fix
|
||||
_load_moe_weights_stacked(all_w, li, pfx, dev, moe, cfg)
|
||||
# EAGERLY process stacked weights → K-major + swizzle, free raw tensors
|
||||
moe._ensure_stacked()
|
||||
@@ -1037,7 +1037,8 @@ def main():
|
||||
device=dev, swiglu_limit=cfg.get("swiglu_limit", 10.0))
|
||||
_load_shared_expert_weights(all_w, li, pfx, dev, se, cfg)
|
||||
# P1: Enable fused SwiGLU for shared expert (1-group variant of MoE fused kernel)
|
||||
se.set_fused_swiglu(True)
|
||||
# DISABLED: Same CuTeDSL arg-binding issue as MoE fused kernel
|
||||
se.set_fused_swiglu(False)
|
||||
# EAGERLY process shared expert weights
|
||||
se._ensure_initialized()
|
||||
# Fix activation global scales — _ensure_initialized sets gsa from l1_gs (which is 1.0)
|
||||
|
||||
Reference in New Issue
Block a user