From 97bdd604e9a00362dda4a74da248b624334cb878 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Mon, 18 May 2026 20:09:19 +0000 Subject: [PATCH] Fix scale assembly: reshape swizzled output to 2D --- cutedsl/shared_expert_pipeline.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/cutedsl/shared_expert_pipeline.py b/cutedsl/shared_expert_pipeline.py index 6c098178..9c9d4ef0 100644 --- a/cutedsl/shared_expert_pipeline.py +++ b/cutedsl/shared_expert_pipeline.py @@ -158,14 +158,21 @@ class CuTeDSLSharedExpertRunner: 1. Zero the padded buffer 2. Copy x_sf into the top rows 3. Apply pad_and_swizzle_single (pads to 128 rows + Blackwell swizzle) + 4. Reshape back to 2D (kernel expects 2D scale_a) - Returns the swizzled 1D scale tensor. + Same as assemble_raw_scales_2d3d_2d_side but for a single group + (no cat of multiple expert scales). """ padded_x_sf = padded_x_sf_buf padded_x_sf.zero_() num_rows, num_cols = x_sf.shape padded_x_sf[:num_rows, :num_cols] = x_sf - return pad_and_swizzle_single(padded_x_sf) + swizzled_flat = pad_and_swizzle_single(padded_x_sf) + # pad_and_swizzle_single returns 1D flattened; reshape to 2D + # Total rows = round_up(num_rows, 128), cols = round_up(num_cols, 4) + total_rows = cutedsl_ceil_div(num_rows, 128) * 128 + total_cols = cutedsl_ceil_div(num_cols, 4) * 4 + return swizzled_flat.reshape(total_rows, total_cols) def compute_activation_global_scales(self, hidden_states_sample): """Compute activation global scales from a warmup forward pass.