diff --git a/mlx/backend/cuda/quantized/qqmm.cpp b/mlx/backend/cuda/quantized/qqmm.cpp index e7260da8c1..21e41cf0ed 100644 --- a/mlx/backend/cuda/quantized/qqmm.cpp +++ b/mlx/backend/cuda/quantized/qqmm.cpp @@ -46,7 +46,7 @@ inline array ensure_row_contiguous_matrix( return x_copy; } -array pad_and_repack_scales( +array pad_and_swizzle_scales( const array& scale, cu::CommandEncoder& encoder, const Stream& s) { @@ -64,14 +64,12 @@ array pad_and_repack_scales( cu::malloc_async(pad_outer * pad_inner, encoder), Shape{pad_outer, pad_inner}, scale.dtype()); - repack_scales(scale, scale_tiled, encoder, s); + swizzle_scales(scale, scale_tiled, encoder, s); encoder.add_temporary(scale_tiled); return scale_tiled; } -} // namespace -namespace { void qqmm_impl( cu::CommandEncoder& encoder, int M, @@ -176,8 +174,8 @@ void QQMatmul::eval_gpu(const std::vector& inputs, array& out) { int K = K_packed * (32 / bits_); // Repack scales from linear to tiled layout for tensor cores - array scale_x = pad_and_repack_scales(scale_x_pre, encoder, s); - array scale_w = pad_and_repack_scales(scale_w_pre, encoder, s); + array scale_x = pad_and_swizzle_scales(scale_x_pre, encoder, s); + array scale_w = pad_and_swizzle_scales(scale_w_pre, encoder, s); bool x_transposed = false; bool w_transposed = true; // always transposed diff --git a/mlx/backend/cuda/quantized/qqmm_utils.cu b/mlx/backend/cuda/quantized/qqmm_utils.cu index ff19057b08..c8764709b9 100644 --- a/mlx/backend/cuda/quantized/qqmm_utils.cu +++ b/mlx/backend/cuda/quantized/qqmm_utils.cu @@ -10,6 +10,11 @@ namespace mlx::core { namespace cg = cooperative_groups; +constexpr int TILE_ROWS = 128; +constexpr int TILE_COLS = 4; +constexpr int TILES_PER_LANE = 1; +constexpr int LANES_PER_BLOCK = 32; + // To pass scales to tensor cores, they need to be repacked into a tiled layout // https://docs.nvidia.com/cuda/cublas/index.html#d-block-scaling-factors-layout // Tiled layout for scale factors is very well described in CUTLASS @@ -43,118 +48,171 @@ namespace cg = cooperative_groups; // [252, 253, 254, 255], // [380, 381, 382, 383], // [508, 509, 510, 511]]]]], -__device__ size_t -scale_tiled_offset(size_t scale_index, size_t num_rows, size_t num_scale_cols) { - // Compute the tiled layout offset for scale factors used in tensor cores - // This function maps from a linear scale index to the tiled layout expected - // by tensor cores (and cublaslt). - // - // Input: linear scale index (e.g., for a matrix M x K with group_size, - // scale_index ranges from 0 to (M * K/group_size - 1)) - // - // The tiled layout organizes scales into tiles of 128 rows x 4 columns, - // where each tile is subdivided into 4 sub-blocks of 32 rows x 4 columns. - size_t row = scale_index / num_scale_cols; - size_t col = scale_index % num_scale_cols; - - constexpr size_t rows_per_tile = 128; - constexpr size_t rows_per_sub_block = 32; - constexpr size_t cols_per_sub_block = 4; - constexpr size_t sub_blocks_per_tile = 4; // Vertically stacked - - // Decompose row position - size_t tile_row = row / rows_per_tile; // Which tile row - size_t row_in_tile = row % rows_per_tile; // Row within tile - size_t sub_block_row = - row_in_tile / rows_per_sub_block; // Sub-block within tile - size_t row_in_sub_block = - row_in_tile % rows_per_sub_block; // Row in sub-block - - // Decompose column position - size_t col_tile = col / cols_per_sub_block; // Which column tile - size_t col_in_sub_block = col % cols_per_sub_block; // Column within sub-block - - // Compute tile index and offset within tile - size_t num_col_tiles = cuda::ceil_div(num_scale_cols, cols_per_sub_block); - size_t tile_idx = tile_row * num_col_tiles + col_tile; - - size_t offset_in_tile = - (row_in_sub_block * sub_blocks_per_tile * cols_per_sub_block) + - (sub_block_row * cols_per_sub_block) + col_in_sub_block; - - constexpr size_t tile_size = rows_per_tile * cols_per_sub_block; - return tile_idx * tile_size + offset_in_tile; + +inline std::tuple get_swizzle_launch_args( + size_t M_swizzled, + size_t K_swizzled) { + constexpr int tiles_per_block = LANES_PER_BLOCK * TILES_PER_LANE; + constexpr int warps_per_block = TILE_ROWS / 4; // 128 / 4 = 32 + + const int num_tiles_k = K_swizzled / TILE_COLS; + const int num_tiles_m = M_swizzled / TILE_ROWS; + + dim3 grid; + grid.x = cuda::ceil_div(num_tiles_k, tiles_per_block); + grid.y = num_tiles_m; + grid.z = 1; + // Block is always (32, 32) = 1024 threads + dim3 block(LANES_PER_BLOCK, warps_per_block, 1); + + return std::make_tuple(grid, block); } namespace cu { -__global__ void repack_scales( +__global__ void swizzle_scales( const uint8_t* scales_linear, - uint8_t* scales_tiled, - size_t input_rows, - size_t input_cols, - size_t output_rows, - size_t output_cols) { + uint8_t* scales_swizzled, + const size_t M, + const size_t K, + const size_t M_swizzled, + const size_t K_swizzled) { + constexpr int tile_size = TILE_ROWS * TILE_COLS; + constexpr int num_tile_rows_per_thread = 4; + constexpr int max_tiles_per_block = LANES_PER_BLOCK * TILES_PER_LANE; + + constexpr int tile_stride = tile_size / 16; // 32 int4s per tile + + // Each thread loads 16 scales from 4 rows (stride 32) and packs them into + // int4. For example: thread (0, 0) loads scales at rows 0,32,64,96 of tile 0, + // thread (1, 0) loads rows 0,32,64,96 of of tile 1, etc. + // The store is strided within a warp (stride 32 int4s), so we first + // write to shared memory, then do a coalesced store from shared to global auto block_size = cg::this_thread_block().dim_threads(); auto block_idx = cg::this_thread_block().group_index(); auto idx_in_block = cg::this_thread_block().thread_index(); - auto tidx = block_idx.x * block_size.x + idx_in_block.x; - auto tidy = block_idx.y * block_size.y + idx_in_block.y; - - auto grid_dim_x = - cg::this_grid().dim_blocks().x * cg::this_grid().block_index().x; - - size_t output_index = tidx + grid_dim_x * size_t(tidy); - size_t output_size = output_rows * output_cols; - - if (output_index >= output_size) { - return; + auto tidx = idx_in_block.x; + auto tidy = idx_in_block.y; + auto linear_tid = tidy * block_size.x + tidx; + + const int bid_x = block_idx.x; + const int bid_y = block_idx.y; + + const int K_int = K_swizzled / 4; + + const size_t output_offset = static_cast(bid_y) * TILE_ROWS * K_int + + static_cast(bid_x) * max_tiles_per_block * tile_size / 4; + int* output_block = reinterpret_cast(scales_swizzled) + output_offset; + + const int grid_dim_x = cg::this_grid().dim_blocks().x; + const int grid_dim_y = cg::this_grid().dim_blocks().y; + + int remaining = K_int - bid_x * max_tiles_per_block; + int tiles_in_block = min(remaining, max_tiles_per_block); + bool valid_tile = tidx * TILES_PER_LANE < tiles_in_block; + + __shared__ int4 strided_scales_thread[max_tiles_per_block * tile_stride]; + + // Initialize to zero for padding + int thread_tile_rows[num_tile_rows_per_thread] = {0}; + + if (valid_tile) { + const size_t col_base = + static_cast(bid_x) * max_tiles_per_block * TILE_COLS + + tidx * TILE_COLS; + + const bool aligned_k = (K % 4 == 0); + + if (aligned_k) { + // fast path: K is aligned, use vectorized loads with stride K/4 + const int K_stride = K / 4; + const size_t block_offset = + static_cast(bid_y) * TILE_ROWS * K_stride + + static_cast(bid_x) * max_tiles_per_block; + const int* input_block = + reinterpret_cast(scales_linear) + block_offset; +// load +#pragma unroll + for (int i = 0; i < num_tile_rows_per_thread; i++) { + const size_t row = + static_cast(bid_y) * TILE_ROWS + i * block_size.x + tidy; + const int thread_offset = + (i * block_size.x + tidy) * K_stride + tidx * TILES_PER_LANE; + if (row < M && col_base + TILE_COLS <= K) { + thread_tile_rows[i] = __ldg(input_block + thread_offset); + } else if (row < M) { +// partial tile at K boundary: load byte-by-byte +#pragma unroll + for (int c = 0; c < TILE_COLS; c++) { + if (col_base + c < K) { + reinterpret_cast(&thread_tile_rows[i])[c] = + scales_linear[row * K + col_base + c]; + } + } + } + } + } else { +#pragma unroll + for (int i = 0; i < num_tile_rows_per_thread; i++) { + const size_t row = + static_cast(bid_y) * TILE_ROWS + i * block_size.x + tidy; + if (row < M) { + const size_t row_start = row * K; +#pragma unroll + for (int c = 0; c < TILE_COLS; c++) { + if (col_base + c < K) { + reinterpret_cast(&thread_tile_rows[i])[c] = + scales_linear[row_start + col_base + c]; + } + } + } + } + } + // store to shared with XOR swizzle to avoid bank conflicts + int base_idx = tidx * tile_stride + tidy; + int xor_bits = (tidy >> 3) & 0x3; + int swizzled_idx = base_idx ^ xor_bits; + strided_scales_thread[swizzled_idx] = + *reinterpret_cast(thread_tile_rows); } - size_t tiled_offset = - scale_tiled_offset(output_index, output_rows, output_cols); - - size_t row = output_index / output_cols; - size_t col = output_index % output_cols; - - // Probably this can be done better with 2 separated paths for valid and - // padding - if (row < input_rows && col < input_cols) { - size_t input_index = row * input_cols + col; - scales_tiled[tiled_offset] = scales_linear[input_index]; - } else { - // Zero-fill padding region - scales_tiled[tiled_offset] = 0; + cg::thread_block block = cg::this_thread_block(); + cg::sync(block); + + const int total_int4s = tiles_in_block * tile_stride; +#pragma unroll + for (int i = linear_tid; i < total_int4s; i += block_size.x * block_size.y) { + int tile_idx = i / tile_stride; + int row_idx = i % tile_stride; + int base_idx = tile_idx * tile_stride + row_idx; + int xor_bits = (row_idx >> 3) & 0x3; + int swizzled_idx = base_idx ^ xor_bits; + reinterpret_cast(output_block)[i] = + strided_scales_thread[swizzled_idx]; } } - } // namespace cu -void repack_scales( +void swizzle_scales( const array& scales, array& scales_tiled, cu::CommandEncoder& enc, const Stream& s) { enc.set_input_array(scales); enc.set_output_array(scales_tiled); - // Note: scales_tiled is padded to full tiles so if num_rows or num_cols - // are not multiples of tile sizes, the extra space is filled with zeros - + // are not multiples of tile sizes size_t input_rows = scales.shape(-2); size_t input_cols = scales.shape(-1); size_t output_rows = scales_tiled.shape(-2); size_t output_cols = scales_tiled.shape(-1); - size_t output_size = output_rows * output_cols; - - bool large = output_size > UINT_MAX; - auto [num_blocks, block_dims] = get_launch_args( - output_size, scales_tiled.shape(), scales_tiled.strides(), large); + auto [num_blocks, block_dims] = + get_swizzle_launch_args(output_rows, output_cols); enc.add_kernel_node( - cu::repack_scales, + cu::swizzle_scales, num_blocks, block_dims, 0, diff --git a/mlx/backend/cuda/quantized/qqmm_utils.h b/mlx/backend/cuda/quantized/qqmm_utils.h index 126cc298b2..0a9a78f70c 100644 --- a/mlx/backend/cuda/quantized/qqmm_utils.h +++ b/mlx/backend/cuda/quantized/qqmm_utils.h @@ -21,7 +21,7 @@ inline std::pair get_padded_scale_dims(int num_rows, int num_cols) { return {padded_rows, padded_cols}; } -void repack_scales( +void swizzle_scales( const array& scales, array& scales_tiled, cu::CommandEncoder& enc,