From 4d1cea591904e1cb68568a9fc13addf38fe6a413 Mon Sep 17 00:00:00 2001 From: Francesco Bertolotti Date: Mon, 12 Jan 2026 16:19:43 +0100 Subject: [PATCH 1/2] wrong number of dimensions --- unsloth/kernels/swiglu.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index b321f5179e..9e2680e862 100644 --- a/unsloth/kernels/swiglu.py +++ b/unsloth/kernels/swiglu.py @@ -128,7 +128,7 @@ def _DWf_DW_dfg_kernel( def swiglu_DWf_DW_dfg_kernel(DW, e, g): - batch_seq_len, hd = e.shape + batch, seq_len, hd = e.shape n_elements = e.numel() grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) with torch_gpu_device(e.device): From 56900ab2efbc99745c375669445c3cfcd9d0e44e Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 12 Jan 2026 21:32:20 -0800 Subject: [PATCH 2/2] Apply suggestion from @danielhanchen --- unsloth/kernels/swiglu.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index 9e2680e862..b3ae9d40e6 100644 --- a/unsloth/kernels/swiglu.py +++ b/unsloth/kernels/swiglu.py @@ -128,7 +128,7 @@ def _DWf_DW_dfg_kernel( def swiglu_DWf_DW_dfg_kernel(DW, e, g): - batch, seq_len, hd = e.shape + batch_seq_len, hd = e.shape # Flattened to 2D, so 1st dim is bsz * seq_len n_elements = e.numel() grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) with torch_gpu_device(e.device):