From 0bdbfdb8387a97db0ec401f2c7c580356b96ae03 Mon Sep 17 00:00:00 2001 From: Long Yixing Date: Sun, 15 Mar 2026 09:33:55 +0800 Subject: [PATCH] [CUDA] Implement MaskedScatter (#3151) --- benchmarks/python/masked_scatter.py | 34 ++++++-- mlx/backend/cuda/device/scatter.cuh | 87 +++++++++++++++++++ mlx/backend/cuda/indexing.cpp | 127 ++++++++++++++++++++++++++++ mlx/backend/cuda/primitives.cpp | 1 - mlx/backend/cuda/scan.cu | 75 +++++++++------- mlx/backend/{metal => gpu}/scan.h | 0 mlx/backend/metal/indexing.cpp | 2 +- mlx/backend/metal/scan.cpp | 2 +- python/tests/cuda_skip.py | 4 - tests/autograd_tests.cpp | 5 -- tests/ops_tests.cpp | 5 -- 11 files changed, 289 insertions(+), 53 deletions(-) rename mlx/backend/{metal => gpu}/scan.h (100%) diff --git a/benchmarks/python/masked_scatter.py b/benchmarks/python/masked_scatter.py index 71857c54..e1c84ee6 100644 --- a/benchmarks/python/masked_scatter.py +++ b/benchmarks/python/masked_scatter.py @@ -1,5 +1,6 @@ import math import os +import platform import subprocess import time from copy import copy @@ -17,9 +18,6 @@ RESULTS_DIR = "./results" if not os.path.isdir(RESULTS_DIR): os.mkdir(RESULTS_DIR) -DEVICE_NAME = subprocess.check_output(["sysctl", "-n", "machdep.cpu.brand_string"]) -DEVICE_NAME = DEVICE_NAME.decode("utf-8").strip("\n") - TORCH_DEVICE = torch.device( "mps" if torch.backends.mps.is_available() @@ -27,11 +25,36 @@ TORCH_DEVICE = torch.device( ) +def get_device_name(): + if TORCH_DEVICE.type == "cuda": + try: + out = subprocess.check_output( + ["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"], + stderr=subprocess.DEVNULL, + ) + return out.decode("utf-8").splitlines()[0].strip() + except Exception: + return "CUDA_GPU" + if TORCH_DEVICE.type == "mps": + try: + out = subprocess.check_output( + ["sysctl", "-n", "machdep.cpu.brand_string"], + stderr=subprocess.DEVNULL, + ) + return out.decode("utf-8").strip() + except Exception: + return "Apple_Silicon" + return platform.processor() or platform.machine() or "CPU" + + +DEVICE_NAME = get_device_name() + + N_WARMUP = 5 N_ITER_BENCH = 50 N_ITER_FUNC = 20 -VECTOR_LENGTHS = [4096 * (2**i) for i in range(10)] +VECTOR_LENGTHS = [4096 * (2**i) for i in range(12)] MASK_DENSITIES = [0.01, 0.1, 0.25, 0.5] D_TYPES = ("float32", "float16") @@ -202,9 +225,10 @@ def main(): ) output_path = os.path.join( RESULTS_DIR, - f"{DEVICE_NAME.replace(' ', '_')}_masked_scatter_{dtype}.pdf", + f"{DEVICE_NAME.replace(' ', '_')}_masked_scatter_{dtype}.png", ) fig.savefig(output_path) + print(f"Saved benchmark image: {output_path}") plt.close(fig) diff --git a/mlx/backend/cuda/device/scatter.cuh b/mlx/backend/cuda/device/scatter.cuh index b2f64035..c5378bfc 100644 --- a/mlx/backend/cuda/device/scatter.cuh +++ b/mlx/backend/cuda/device/scatter.cuh @@ -65,4 +65,91 @@ __global__ void scatter( Op{}(out + out_idx, upd[upd_loc]); } +template +__global__ void masked_scatter( + const T* dst, + const bool* mask, + const int32_t* scatter_offsets, + const T* src, + T* out, + IdxT size, + IdxT src_batch_size, + IdxT mask_batch_size, + const __grid_constant__ Shape dst_shape, + const __grid_constant__ Strides dst_strides, + int32_t dst_ndim, + const __grid_constant__ Shape src_shape, + const __grid_constant__ Strides src_strides, + int32_t src_ndim) { + IdxT index = cg::this_grid().thread_rank(); + if (index >= size) { + return; + } + + T dst_val; + if constexpr (DstContiguous) { + dst_val = dst[index]; + } else { + IdxT dst_loc = + elem_to_loc(index, dst_shape.data(), dst_strides.data(), dst_ndim); + dst_val = dst[dst_loc]; + } + + if (mask[index]) { + IdxT src_index = static_cast(scatter_offsets[index]); + if (src_index < src_batch_size) { + IdxT batch_idx = index / mask_batch_size; + if constexpr (SrcContiguous) { + out[index] = src[batch_idx * src_batch_size + src_index]; + } else { + IdxT src_elem = batch_idx * src_batch_size + src_index; + IdxT src_loc = elem_to_loc( + src_elem, src_shape.data(), src_strides.data(), src_ndim); + out[index] = src[src_loc]; + } + return; + } + } + + out[index] = dst_val; +} + +template +__global__ void masked_scatter_vec_contiguous( + const T* dst, + const bool* mask, + const int32_t* scatter_offsets, + const T* src, + T* out, + IdxT size, + IdxT src_batch_size, + IdxT mask_batch_size) { + IdxT vec_index = cg::this_grid().thread_rank(); + IdxT base = vec_index * N_READS; + if (base >= size) { + return; + } + + auto out_vec = load_vector(dst, vec_index, size, static_cast(0)); + auto mask_vec = load_vector(mask, vec_index, size, false); + auto offset_vec = load_vector(scatter_offsets, vec_index, size, 0); + +#pragma unroll + for (int i = 0; i < N_READS; ++i) { + IdxT index = base + i; + if (index >= size) { + break; + } + if (mask_vec[i]) { + IdxT src_index = static_cast(offset_vec[i]); + if (src_index < src_batch_size) { + IdxT batch_idx = index / mask_batch_size; + out_vec[i] = src[batch_idx * src_batch_size + src_index]; + } + } + } + + store_vector(out, vec_index, out_vec, size); +} + } // namespace mlx::core::cu diff --git a/mlx/backend/cuda/indexing.cpp b/mlx/backend/cuda/indexing.cpp index 424566d2..a84de113 100644 --- a/mlx/backend/cuda/indexing.cpp +++ b/mlx/backend/cuda/indexing.cpp @@ -5,6 +5,7 @@ #include "mlx/backend/cuda/jit_module.h" #include "mlx/backend/cuda/kernel_utils.cuh" #include "mlx/backend/gpu/copy.h" +#include "mlx/backend/gpu/scan.h" #include "mlx/dtype_utils.h" #include "mlx/primitives.h" @@ -435,4 +436,130 @@ void ScatterAxis::eval_gpu(const std::vector& inputs, array& out) { kernel, num_blocks, block_dims, {}, 0, args.args()); } +void MaskedScatter::eval_gpu(const std::vector& inputs, array& out) { + nvtx3::scoped_range r("MaskedScatter::eval_gpu"); + assert(inputs.size() == 3); + + const array& dst = inputs[0]; + const array& mask = inputs[1]; + const array& src = inputs[2]; + + auto& s = stream(); + auto& encoder = cu::get_command_encoder(s); + + const size_t total = mask.size(); + out.set_data(cu::malloc_async(out.nbytes(), encoder)); + if (total == 0) { + return; + } + + array mask_flat = flatten_in_eval(mask, 1, -1, s); + if (mask_flat.data() != mask.data()) { + encoder.add_temporary(mask_flat); + } + if (!mask_flat.flags().row_contiguous) { + mask_flat = contiguous_copy_gpu(mask_flat, s); + encoder.add_temporary(mask_flat); + } + + array scatter_offsets(mask_flat.shape(), int32, nullptr, {}); + scatter_offsets.set_data(cu::malloc_async(scatter_offsets.nbytes(), encoder)); + encoder.add_temporary(scatter_offsets); + + scan_gpu_inplace( + mask_flat, + scatter_offsets, + Scan::Sum, + /* axis= */ 1, + /* reverse= */ false, + /* inclusive= */ false, + s); + + const size_t batch_count = mask.shape(0); + const size_t mask_batch_size = mask_flat.size() / batch_count; + const size_t src_batch_size = src.size() / src.shape(0); + bool large = total > INT32_MAX || src.size() > INT32_MAX; + bool vectorized = src.flags().row_contiguous && dst.flags().row_contiguous; + constexpr int kMaskedScatterVecSize = 16; + constexpr int kMaskedScatterVecBlockDim = 256; + + std::string module_name = + fmt::format("masked_scatter_{}", dtype_to_string(out.dtype())); + cu::JitModule& mod = cu::get_jit_module(s.device, module_name, [&]() { + std::vector kernel_names; + for (int src_contiguous = 0; src_contiguous <= 1; ++src_contiguous) { + for (int dst_contiguous = 0; dst_contiguous <= 1; ++dst_contiguous) { + for (int use_large = 0; use_large <= 1; ++use_large) { + kernel_names.push_back( + fmt::format( + "mlx::core::cu::masked_scatter<{}, {}, {}, {}>", + dtype_to_cuda_type(out.dtype()), + src_contiguous ? "true" : "false", + dst_contiguous ? "true" : "false", + use_large ? "int64_t" : "int32_t")); + } + } + } + for (int use_large = 0; use_large <= 1; ++use_large) { + kernel_names.push_back( + fmt::format( + "mlx::core::cu::masked_scatter_vec_contiguous<{}, {}, {}>", + dtype_to_cuda_type(out.dtype()), + use_large ? "int64_t" : "int32_t", + kMaskedScatterVecSize)); + } + return std::make_tuple(false, jit_source_scatter, std::move(kernel_names)); + }); + + cu::KernelArgs args; + args.append(dst); + args.append(mask_flat); + args.append(scatter_offsets); + args.append(src); + args.append(out); + if (large) { + args.append(mask_flat.size()); + args.append(src_batch_size); + args.append(mask_batch_size); + } else { + args.append(mask_flat.size()); + args.append(src_batch_size); + args.append(mask_batch_size); + } + if (!vectorized) { + args.append_ndim(dst.shape()); + args.append_ndim(dst.strides()); + args.append(dst.ndim()); + args.append_ndim(src.shape()); + args.append_ndim(src.strides()); + args.append(src.ndim()); + } + + encoder.set_input_array(dst); + encoder.set_input_array(mask_flat); + encoder.set_input_array(scatter_offsets); + encoder.set_input_array(src); + encoder.set_output_array(out); + + std::string kernel_name = vectorized + ? fmt::format( + "mlx::core::cu::masked_scatter_vec_contiguous<{}, {}, {}>", + dtype_to_cuda_type(out.dtype()), + large ? "int64_t" : "int32_t", + kMaskedScatterVecSize) + : fmt::format( + "mlx::core::cu::masked_scatter<{}, {}, {}, {}>", + dtype_to_cuda_type(out.dtype()), + src.flags().row_contiguous ? "true" : "false", + dst.flags().row_contiguous ? "true" : "false", + large ? "int64_t" : "int32_t"); + auto kernel = mod.get_kernel(kernel_name); + auto [num_blocks, block_dims] = vectorized + ? get_launch_args( + mask_flat, large, kMaskedScatterVecSize, kMaskedScatterVecBlockDim) + : get_launch_args(mask_flat, large); + encoder.add_kernel_node_raw( + kernel, num_blocks, block_dims, {}, 0, args.args()); +} + } // namespace mlx::core diff --git a/mlx/backend/cuda/primitives.cpp b/mlx/backend/cuda/primitives.cpp index 0bf45214..98dca570 100644 --- a/mlx/backend/cuda/primitives.cpp +++ b/mlx/backend/cuda/primitives.cpp @@ -33,7 +33,6 @@ NO_GPU(Inverse) NO_GPU(Cholesky) NO_GPU_MULTI(Eig) NO_GPU_MULTI(Eigh) -NO_GPU(MaskedScatter) namespace distributed { NO_GPU_MULTI(Send) diff --git a/mlx/backend/cuda/scan.cu b/mlx/backend/cuda/scan.cu index bd25084c..206419e4 100644 --- a/mlx/backend/cuda/scan.cu +++ b/mlx/backend/cuda/scan.cu @@ -5,6 +5,7 @@ #include "mlx/backend/cuda/kernel_utils.cuh" #include "mlx/backend/cuda/reduce/reduce_ops.cuh" #include "mlx/backend/gpu/copy.h" +#include "mlx/backend/gpu/scan.h" #include "mlx/dtype_utils.h" #include "mlx/primitives.h" @@ -362,51 +363,38 @@ constexpr bool supports_scan_op() { } } -void Scan::eval_gpu(const std::vector& inputs, array& out) { - nvtx3::scoped_range r("Scan::eval_gpu"); - assert(inputs.size() == 1); - auto in = inputs[0]; - auto& s = stream(); +void scan_gpu_inplace( + array in, + array& out, + Scan::ReduceType reduce_type, + int axis, + bool reverse, + bool inclusive, + const Stream& s) { auto& encoder = cu::get_command_encoder(s); - - if (in.flags().contiguous && in.strides()[axis_] != 0) { - if (in.is_donatable() && in.itemsize() == out.itemsize()) { - out.copy_shared_buffer(in); - } else { - out.set_data( - cu::malloc_async(in.data_size() * out.itemsize(), encoder), - in.data_size(), - in.strides(), - in.flags()); - } - } else { - in = contiguous_copy_gpu(in, s); - out.copy_shared_buffer(in); - } - constexpr int N_READS = 4; - int32_t axis_size = in.shape(axis_); - bool contiguous = in.strides()[axis_] == 1; + int32_t axis_size = in.shape(axis); + bool contiguous = in.strides()[axis] == 1; encoder.set_input_array(in); encoder.set_output_array(out); dispatch_all_types(in.dtype(), [&](auto type_tag) { using T = cuda_type_t; - dispatch_scan_ops(reduce_type_, [&](auto scan_op_tag) { + dispatch_scan_ops(reduce_type, [&](auto scan_op_tag) { using Op = MLX_GET_TYPE(scan_op_tag); if constexpr (supports_scan_op()) { using U = typename cu::ScanResult::type; - dispatch_bool(inclusive_, [&](auto inclusive) { - dispatch_bool(reverse_, [&](auto reverse) { + dispatch_bool(inclusive, [&](auto inclusive_tag) { + dispatch_bool(reverse, [&](auto reverse_tag) { if (contiguous) { auto kernel = cu::contiguous_scan< T, U, Op, N_READS, - inclusive.value, - reverse.value>; + inclusive_tag.value, + reverse_tag.value>; int block_dim = cuda::ceil_div(axis_size, N_READS); block_dim = cuda::ceil_div(block_dim, WARP_SIZE) * WARP_SIZE; block_dim = std::min(block_dim, WARP_SIZE * WARP_SIZE); @@ -427,9 +415,9 @@ void Scan::eval_gpu(const std::vector& inputs, array& out) { N_READS, BM, BN, - inclusive.value, - reverse.value>; - int64_t stride = in.strides()[axis_]; + inclusive_tag.value, + reverse_tag.value>; + int64_t stride = in.strides()[axis]; int64_t stride_blocks = cuda::ceil_div(stride, BN); dim3 num_blocks = get_2d_grid_dims( in.shape(), in.strides(), axis_size * stride); @@ -463,4 +451,29 @@ void Scan::eval_gpu(const std::vector& inputs, array& out) { }); } +void Scan::eval_gpu(const std::vector& inputs, array& out) { + nvtx3::scoped_range r("Scan::eval_gpu"); + assert(inputs.size() == 1); + auto in = inputs[0]; + auto& s = stream(); + auto& encoder = cu::get_command_encoder(s); + + if (in.flags().contiguous && in.strides()[axis_] != 0) { + if (in.is_donatable() && in.itemsize() == out.itemsize()) { + out.copy_shared_buffer(in); + } else { + out.set_data( + cu::malloc_async(in.data_size() * out.itemsize(), encoder), + in.data_size(), + in.strides(), + in.flags()); + } + } else { + in = contiguous_copy_gpu(in, s); + out.copy_shared_buffer(in); + } + + scan_gpu_inplace(in, out, reduce_type_, axis_, reverse_, inclusive_, s); +} + } // namespace mlx::core diff --git a/mlx/backend/metal/scan.h b/mlx/backend/gpu/scan.h similarity index 100% rename from mlx/backend/metal/scan.h rename to mlx/backend/gpu/scan.h diff --git a/mlx/backend/metal/indexing.cpp b/mlx/backend/metal/indexing.cpp index 48311dc6..e0ebb790 100644 --- a/mlx/backend/metal/indexing.cpp +++ b/mlx/backend/metal/indexing.cpp @@ -5,11 +5,11 @@ #include "mlx/backend/common/compiled.h" #include "mlx/backend/common/utils.h" #include "mlx/backend/gpu/copy.h" +#include "mlx/backend/gpu/scan.h" #include "mlx/backend/metal/device.h" #include "mlx/backend/metal/jit/includes.h" #include "mlx/backend/metal/jit/indexing.h" #include "mlx/backend/metal/kernels.h" -#include "mlx/backend/metal/scan.h" #include "mlx/backend/metal/utils.h" #include "mlx/dtype.h" #include "mlx/primitives.h" diff --git a/mlx/backend/metal/scan.cpp b/mlx/backend/metal/scan.cpp index b48ec41c..5d269813 100644 --- a/mlx/backend/metal/scan.cpp +++ b/mlx/backend/metal/scan.cpp @@ -4,9 +4,9 @@ #include #include "mlx/backend/gpu/copy.h" +#include "mlx/backend/gpu/scan.h" #include "mlx/backend/metal/device.h" #include "mlx/backend/metal/kernels.h" -#include "mlx/backend/metal/scan.h" #include "mlx/backend/metal/utils.h" #include "mlx/primitives.h" diff --git a/python/tests/cuda_skip.py b/python/tests/cuda_skip.py index 6de59455..fe042da8 100644 --- a/python/tests/cuda_skip.py +++ b/python/tests/cuda_skip.py @@ -39,8 +39,4 @@ cuda_skip = { "TestQuantized.test_throw", "TestQuantized.test_vjp_scales_biases", "TestExportImport.test_export_quantized_model", - # Masked scatter - "TestOps.test_masked_scatter", - "TestVmap.test_vmap_masked_scatter", - "TestArray.test_setitem_with_boolean_mask", } diff --git a/tests/autograd_tests.cpp b/tests/autograd_tests.cpp index ff8d986b..25c871cd 100644 --- a/tests/autograd_tests.cpp +++ b/tests/autograd_tests.cpp @@ -1357,11 +1357,6 @@ TEST_CASE("test grad dynamic slices") { } TEST_CASE("test masked_scatter autograd") { - if (cu::is_available()) { - INFO("Skipping masked_scatter cuda autograd tests"); - return; - } - // Test jvp { auto self = array({10.f, 20.f, 30.f, 40.f}, {4}); diff --git a/tests/ops_tests.cpp b/tests/ops_tests.cpp index 6a924cfb..51df7dcd 100644 --- a/tests/ops_tests.cpp +++ b/tests/ops_tests.cpp @@ -2441,11 +2441,6 @@ TEST_CASE("test scatter") { } TEST_CASE("test masked_scatter") { - if (cu::is_available()) { - INFO("Skipping masked_scatter cuda ops tests"); - return; - } - // Wrong mask dtype CHECK_THROWS(masked_scatter(array({1, 2}), array({1, 2}), array({1, 2})));