Fix scale assembly: reshape swizzled output to 2D

This commit is contained in:
2026-05-18 20:09:19 +00:00
parent c1aa4af123
commit 97bdd604e9

View File

@@ -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.