diff --git a/mlx/backend/metal/kernels/fp_quantized.h b/mlx/backend/metal/kernels/fp_quantized.h index cc9b68ad..f4bf438d 100644 --- a/mlx/backend/metal/kernels/fp_quantized.h +++ b/mlx/backend/metal/kernels/fp_quantized.h @@ -533,6 +533,7 @@ METAL_FUNC void fp_qvm_impl( device T* y, const int in_vec_size, const int out_vec_size, + const int in_vec_stride, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { @@ -563,7 +564,7 @@ METAL_FUNC void fp_qvm_impl( int out_col = pack_factor * tn * (tid.y * num_simdgroups + simd_gid); ws += out_col * bytes_per_pack / pack_factor + simd_lid * out_vec_size_w; scales += out_col / group_size + simd_lid * out_vec_size_g; - x += tid.x * in_vec_size + simd_lid; + x += tid.x * in_vec_stride + simd_lid; y += tid.x * out_vec_size + out_col; if (out_col >= out_vec_size) { @@ -1122,7 +1123,16 @@ template tid); } fp_qvm_impl( - w, scales, x, y, in_vec_size, out_vec_size, tid, simd_gid, simd_lid); + w, + scales, + x, + y, + in_vec_size, + out_vec_size, + in_vec_size, + tid, + simd_gid, + simd_lid); } template @@ -1164,8 +1174,20 @@ template int in_vec_size_adj = tid.z % split_k == split_k - 1 ? final_block_size : in_vec_size; + // The in_vec_stride is the full K dimension, not the partition size + int in_vec_stride = (split_k - 1) * in_vec_size + final_block_size; + fp_qvm_impl( - w, scales, x, y, in_vec_size_adj, out_vec_size, tid, simd_gid, simd_lid); + w, + scales, + x, + y, + in_vec_size_adj, + out_vec_size, + in_vec_stride, + tid, + simd_gid, + simd_lid); } template < @@ -1423,7 +1445,16 @@ template s_strides, tid); fp_qvm_impl( - w, scales, x, y, in_vec_size, out_vec_size, tid, simd_gid, simd_lid); + w, + scales, + x, + y, + in_vec_size, + out_vec_size, + in_vec_size, + tid, + simd_gid, + simd_lid); } template < diff --git a/mlx/backend/metal/kernels/quantized.h b/mlx/backend/metal/kernels/quantized.h index d4a5d28b..12b5c85c 100644 --- a/mlx/backend/metal/kernels/quantized.h +++ b/mlx/backend/metal/kernels/quantized.h @@ -983,6 +983,7 @@ METAL_FUNC void qvm_impl( device T* y, const int in_vec_size, const int out_vec_size, + const int in_vec_stride, uint3 tid [[threadgroup_position_in_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { @@ -1016,7 +1017,7 @@ METAL_FUNC void qvm_impl( ws += out_col * bytes_per_pack / pack_factor + simd_lid * out_vec_size_w; scales += out_col / group_size + simd_lid * out_vec_size_g; biases += out_col / group_size + simd_lid * out_vec_size_g; - x += tid.x * in_vec_size + simd_lid; + x += tid.x * in_vec_stride + simd_lid; y += tid.x * out_vec_size + out_col; if (out_col >= out_vec_size) { @@ -1643,6 +1644,7 @@ template y, in_vec_size, out_vec_size, + in_vec_size, tid, simd_gid, simd_lid); @@ -1691,6 +1693,9 @@ template int in_vec_size_adj = tid.z % split_k == split_k - 1 ? final_block_size : in_vec_size; + // The in_vec_stride is the full K dimension, not the partition size + int in_vec_stride = (split_k - 1) * in_vec_size + final_block_size; + qvm_impl( w, scales, @@ -1699,6 +1704,7 @@ template y, in_vec_size_adj, out_vec_size, + in_vec_stride, tid, simd_gid, simd_lid); @@ -2077,6 +2083,7 @@ template y, in_vec_size, out_vec_size, + in_vec_size, tid, simd_gid, simd_lid);