Fix scale assembly: reshape swizzled output to 2D
This commit is contained in:
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user