From 862dac89e02e461b3b69ccb5be83c8aa8076c8b5 Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Tue, 18 Aug 2026 11:59:52 +0000 Subject: [PATCH 01/11] Add minimal sparse_tensor-dialect CSR spmv/spmm lowered to cuSPARSE/hipSPARSE Introduces an opaque `Reactant.CSRMatrix` (convertible from SparseMatrixCSC via the SparseArrays ext). Inside traced code `A * x`, `A * B`, and `mul!` emit `sparse_tensor.assemble` producing a CSR-encoded tensor consumed by `stablehlo.dot_general`; `Compiler.lower_sparse_ops!` rewrites the pair to `stablehlo.custom_call @reactant_csr_matmul` on the raw buffers before any pass pipeline runs, so XLA only ever sees dense types. The custom call is served by new cuSPARSE ("CUDA") and hipSPARSE ("ROCM") typed-FFI handlers in ReactantExtra (requires a local jll rebuild for execution). `@code_hlo optimize=:none` keeps the sparse_tensor IR for inspection. Co-Authored-By: Claude Fable 5 --- deps/ReactantExtra/BUILD | 7 + deps/ReactantExtra/xla_ffi.cpp | 369 ++++++++++++++++++ ext/ReactantSparseArraysExt/CSR.jl | 22 ++ .../ReactantSparseArraysExt.jl | 8 +- src/Reactant.jl | 4 +- src/SparseTensors.jl | 230 +++++++++++ src/compiler/Compiler.jl | 9 + src/compiler/SparseLowering.jl | 108 +++++ test/Project.toml | 2 + test/core/sparse.jl | 96 +++++ 10 files changed, 853 insertions(+), 2 deletions(-) create mode 100644 ext/ReactantSparseArraysExt/CSR.jl create mode 100644 src/SparseTensors.jl create mode 100644 src/compiler/SparseLowering.jl create mode 100644 test/core/sparse.jl diff --git a/deps/ReactantExtra/BUILD b/deps/ReactantExtra/BUILD index e424274192..d3c6683064 100644 --- a/deps/ReactantExtra/BUILD +++ b/deps/ReactantExtra/BUILD @@ -1390,6 +1390,7 @@ cc_library( "@com_google_absl//absl/log:globals", "@llvm-project//mlir:CAPIIRObjects", "@llvm-project//mlir:CAPILLVMObjects", + "@llvm-project//mlir:CAPISparseTensorObjects", # Broken upstream x/ref https://github.com/jax-ml/jax/issues/33344 # "@jax//jaxlib/mosaic:tpu_dialect_capi_objects", @@ -1427,6 +1428,9 @@ cc_library( "@xla//xla/stream_executor:cuda_platform", "@xla//xla/stream_executor:kernel", "@xla//xla/stream_executor/cuda:all_runtime", + # cuSPARSE (dlopen stub) + headers for the reactant_csr_matmul FFI handler + "@xla//xla/tsl/cuda:cusparse", + "@local_config_cuda//cuda:cuda_headers", ]) + if_rocm([ "@xla//xla/service:gpu_plugin", "@xla//xla/pjrt/c:pjrt_c_api_gpu", @@ -1435,6 +1439,9 @@ cc_library( "@xla//xla/stream_executor:rocm_platform", "@xla//xla/service/gpu:amdgpu_compiler", "@xla//xla/backends/profiler/gpu:device_tracer", + # hipSPARSE + headers for the reactant_csr_matmul FFI handler + "@local_config_rocm//rocm:hipsparse", + "@local_config_rocm//rocm:rocm_headers", ]) + select({ # gloo tcp transport only builds on linux "@xla//xla/tsl:macos": [ diff --git a/deps/ReactantExtra/xla_ffi.cpp b/deps/ReactantExtra/xla_ffi.cpp index 85f4189a0c..9dc7a0b9bd 100644 --- a/deps/ReactantExtra/xla_ffi.cpp +++ b/deps/ReactantExtra/xla_ffi.cpp @@ -91,12 +91,381 @@ XLA_FFI_DEFINE_HANDLER( "callback_ptr")); #endif +// ============================================================================ +// CSR sparse matrix products (spmv / spmm) via cuSPARSE / hipSPARSE. +// +// The Julia side (src/compiler/SparseLowering.jl) emits a +// stablehlo.custom_call targeting "reactant_csr_matmul" with api_version = 4 +// (TYPED_FFI). Operands are (rowptr, colind, nzval, dense) with column-major +// layouts pinned; the result is dense. The backend_config dict carries i64 +// attributes "m", "n", "transpose" (must be 0 for now), and "index_base" +// (0 or 1; the Julia side emits 1-based CSR buffers). +// ============================================================================ + +#if defined(REACTANT_CUDA) +#include +#include + +#define REACTANT_CUSPARSE_RET(expr) \ + do { \ + cusparseStatus_t status__ = (expr); \ + if (status__ != CUSPARSE_STATUS_SUCCESS) { \ + return ffi::Error( \ + ffi::ErrorCode::kInternal, \ + absl::StrFormat("reactant_csr_matmul: %s failed: %s", #expr, \ + cusparseGetErrorString(status__))); \ + } \ + } while (0) + +static ffi::Error csrMatmulCuda(cudaStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, + ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base) { + if (transpose != 0) { + return ffi::Error( + ffi::ErrorCode::kUnimplemented, + "reactant_csr_matmul: transposed products are not supported"); + } + if (colind.element_type() != rowptr.element_type()) { + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: rowptr and colind dtypes must match"); + } + + cusparseIndexType_t index_type; + switch (rowptr.element_type()) { + case ffi::DataType::S32: + index_type = CUSPARSE_INDEX_32I; + break; + case ffi::DataType::S64: + index_type = CUSPARSE_INDEX_64I; + break; + default: + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: index buffers must be i32 or i64"); + } + + static const float alpha_f = 1.0f, beta_f = 0.0f; + static const double alpha_d = 1.0, beta_d = 0.0; + cudaDataType value_type; + const void *alpha, *beta; + switch (nzval.element_type()) { + case ffi::DataType::F32: + value_type = CUDA_R_32F; + alpha = &alpha_f; + beta = &beta_f; + break; + case ffi::DataType::F64: + value_type = CUDA_R_64F; + alpha = &alpha_d; + beta = &beta_d; + break; + default: + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: only f32 and f64 values are supported"); + } + if (dense.element_type() != nzval.element_type() || + out->element_type() != nzval.element_type()) { + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: value, operand and result dtypes must match"); + } + + static thread_local cusparseHandle_t handle = nullptr; + if (handle == nullptr) { + REACTANT_CUSPARSE_RET(cusparseCreate(&handle)); + } + REACTANT_CUSPARSE_RET(cusparseSetStream(handle, stream)); + + int64_t nnz = colind.element_count(); + cusparseSpMatDescr_t mat_a; + REACTANT_CUSPARSE_RET(cusparseCreateCsr( + &mat_a, m, n, nnz, rowptr.untyped_data(), colind.untyped_data(), + nzval.untyped_data(), index_type, index_type, + index_base == 1 ? CUSPARSE_INDEX_BASE_ONE : CUSPARSE_INDEX_BASE_ZERO, + value_type)); + + auto with_workspace = [&](size_t buffer_size, + auto &&compute) -> ffi::Error { + void *workspace = nullptr; + if (buffer_size > 0) { + auto maybe_workspace = scratch.Allocate(buffer_size); + if (!maybe_workspace.has_value()) { + return ffi::Error( + ffi::ErrorCode::kResourceExhausted, + "reactant_csr_matmul: failed to allocate workspace"); + } + workspace = *maybe_workspace; + } + return compute(workspace); + }; + + ffi::Error err = ffi::Error::Success(); + int64_t rank = dense.dimensions().size(); + if (rank == 1) { + cusparseDnVecDescr_t vec_x, vec_y; + REACTANT_CUSPARSE_RET( + cusparseCreateDnVec(&vec_x, n, dense.untyped_data(), value_type)); + REACTANT_CUSPARSE_RET( + cusparseCreateDnVec(&vec_y, m, out->untyped_data(), value_type)); + err = [&]() -> ffi::Error { + size_t buffer_size = 0; + REACTANT_CUSPARSE_RET(cusparseSpMV_bufferSize( + handle, CUSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, vec_x, beta, + vec_y, value_type, CUSPARSE_SPMV_ALG_DEFAULT, &buffer_size)); + return with_workspace(buffer_size, [&](void *workspace) -> ffi::Error { + REACTANT_CUSPARSE_RET(cusparseSpMV( + handle, CUSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, vec_x, beta, + vec_y, value_type, CUSPARSE_SPMV_ALG_DEFAULT, workspace)); + return ffi::Error::Success(); + }); + }(); + cusparseDestroyDnVec(vec_x); + cusparseDestroyDnVec(vec_y); + } else if (rank == 2) { + // Layouts are pinned column-major by the Julia rewrite. + int64_t k = dense.dimensions()[0]; + int64_t c = dense.dimensions()[1]; + cusparseDnMatDescr_t mat_b, mat_c; + REACTANT_CUSPARSE_RET(cusparseCreateDnMat(&mat_b, k, c, /*ld=*/k, + dense.untyped_data(), value_type, + CUSPARSE_ORDER_COL)); + REACTANT_CUSPARSE_RET(cusparseCreateDnMat(&mat_c, m, c, /*ld=*/m, + out->untyped_data(), value_type, + CUSPARSE_ORDER_COL)); + err = [&]() -> ffi::Error { + size_t buffer_size = 0; + REACTANT_CUSPARSE_RET(cusparseSpMM_bufferSize( + handle, CUSPARSE_OPERATION_NON_TRANSPOSE, + CUSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, mat_b, beta, mat_c, + value_type, CUSPARSE_SPMM_ALG_DEFAULT, &buffer_size)); + return with_workspace(buffer_size, [&](void *workspace) -> ffi::Error { + REACTANT_CUSPARSE_RET(cusparseSpMM( + handle, CUSPARSE_OPERATION_NON_TRANSPOSE, + CUSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, mat_b, beta, mat_c, + value_type, CUSPARSE_SPMM_ALG_DEFAULT, workspace)); + return ffi::Error::Success(); + }); + }(); + cusparseDestroyDnMat(mat_b); + cusparseDestroyDnMat(mat_c); + } else { + err = ffi::Error(ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: dense operand must have rank 1 or 2"); + } + cusparseDestroySpMat(mat_a); + return err; +} + +XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerCUDA, csrMatmulCuda, + xla::ffi::Ffi::Bind() + .Ctx>() + .Ctx() + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Ret() // out + .Attr("m") + .Attr("n") + .Attr("transpose") + .Attr("index_base")); +#endif // REACTANT_CUDA + +#if defined(REACTANT_ROCM) +#include +#include + +#define REACTANT_HIPSPARSE_RET(expr) \ + do { \ + hipsparseStatus_t status__ = (expr); \ + if (status__ != HIPSPARSE_STATUS_SUCCESS) { \ + return ffi::Error( \ + ffi::ErrorCode::kInternal, \ + absl::StrFormat("reactant_csr_matmul: %s failed: hipSPARSE status " \ + "%d", \ + #expr, static_cast(status__))); \ + } \ + } while (0) + +static ffi::Error csrMatmulRocm(hipStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, + ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base) { + if (transpose != 0) { + return ffi::Error( + ffi::ErrorCode::kUnimplemented, + "reactant_csr_matmul: transposed products are not supported"); + } + if (colind.element_type() != rowptr.element_type()) { + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: rowptr and colind dtypes must match"); + } + + hipsparseIndexType_t index_type; + switch (rowptr.element_type()) { + case ffi::DataType::S32: + index_type = HIPSPARSE_INDEX_32I; + break; + case ffi::DataType::S64: + index_type = HIPSPARSE_INDEX_64I; + break; + default: + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: index buffers must be i32 or i64"); + } + + static const float alpha_f = 1.0f, beta_f = 0.0f; + static const double alpha_d = 1.0, beta_d = 0.0; + hipDataType value_type; + const void *alpha, *beta; + switch (nzval.element_type()) { + case ffi::DataType::F32: + value_type = HIP_R_32F; + alpha = &alpha_f; + beta = &beta_f; + break; + case ffi::DataType::F64: + value_type = HIP_R_64F; + alpha = &alpha_d; + beta = &beta_d; + break; + default: + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: only f32 and f64 values are supported"); + } + if (dense.element_type() != nzval.element_type() || + out->element_type() != nzval.element_type()) { + return ffi::Error( + ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: value, operand and result dtypes must match"); + } + + static thread_local hipsparseHandle_t handle = nullptr; + if (handle == nullptr) { + REACTANT_HIPSPARSE_RET(hipsparseCreate(&handle)); + } + REACTANT_HIPSPARSE_RET(hipsparseSetStream(handle, stream)); + + int64_t nnz = colind.element_count(); + hipsparseSpMatDescr_t mat_a; + REACTANT_HIPSPARSE_RET(hipsparseCreateCsr( + &mat_a, m, n, nnz, rowptr.untyped_data(), colind.untyped_data(), + nzval.untyped_data(), index_type, index_type, + index_base == 1 ? HIPSPARSE_INDEX_BASE_ONE : HIPSPARSE_INDEX_BASE_ZERO, + value_type)); + + auto with_workspace = [&](size_t buffer_size, + auto &&compute) -> ffi::Error { + void *workspace = nullptr; + if (buffer_size > 0) { + auto maybe_workspace = scratch.Allocate(buffer_size); + if (!maybe_workspace.has_value()) { + return ffi::Error( + ffi::ErrorCode::kResourceExhausted, + "reactant_csr_matmul: failed to allocate workspace"); + } + workspace = *maybe_workspace; + } + return compute(workspace); + }; + + ffi::Error err = ffi::Error::Success(); + int64_t rank = dense.dimensions().size(); + if (rank == 1) { + hipsparseDnVecDescr_t vec_x, vec_y; + REACTANT_HIPSPARSE_RET( + hipsparseCreateDnVec(&vec_x, n, dense.untyped_data(), value_type)); + REACTANT_HIPSPARSE_RET( + hipsparseCreateDnVec(&vec_y, m, out->untyped_data(), value_type)); + err = [&]() -> ffi::Error { + size_t buffer_size = 0; + REACTANT_HIPSPARSE_RET(hipsparseSpMV_bufferSize( + handle, HIPSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, vec_x, beta, + vec_y, value_type, HIPSPARSE_SPMV_ALG_DEFAULT, &buffer_size)); + return with_workspace(buffer_size, [&](void *workspace) -> ffi::Error { + REACTANT_HIPSPARSE_RET(hipsparseSpMV( + handle, HIPSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, vec_x, + beta, vec_y, value_type, HIPSPARSE_SPMV_ALG_DEFAULT, workspace)); + return ffi::Error::Success(); + }); + }(); + hipsparseDestroyDnVec(vec_x); + hipsparseDestroyDnVec(vec_y); + } else if (rank == 2) { + // Layouts are pinned column-major by the Julia rewrite. + int64_t k = dense.dimensions()[0]; + int64_t c = dense.dimensions()[1]; + hipsparseDnMatDescr_t mat_b, mat_c; + REACTANT_HIPSPARSE_RET(hipsparseCreateDnMat(&mat_b, k, c, /*ld=*/k, + dense.untyped_data(), + value_type, + HIPSPARSE_ORDER_COL)); + REACTANT_HIPSPARSE_RET(hipsparseCreateDnMat(&mat_c, m, c, /*ld=*/m, + out->untyped_data(), value_type, + HIPSPARSE_ORDER_COL)); + err = [&]() -> ffi::Error { + size_t buffer_size = 0; + REACTANT_HIPSPARSE_RET(hipsparseSpMM_bufferSize( + handle, HIPSPARSE_OPERATION_NON_TRANSPOSE, + HIPSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, mat_b, beta, mat_c, + value_type, HIPSPARSE_SPMM_ALG_DEFAULT, &buffer_size)); + return with_workspace(buffer_size, [&](void *workspace) -> ffi::Error { + REACTANT_HIPSPARSE_RET(hipsparseSpMM( + handle, HIPSPARSE_OPERATION_NON_TRANSPOSE, + HIPSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, mat_b, beta, + mat_c, value_type, HIPSPARSE_SPMM_ALG_DEFAULT, workspace)); + return ffi::Error::Success(); + }); + }(); + hipsparseDestroyDnMat(mat_b); + hipsparseDestroyDnMat(mat_c); + } else { + err = ffi::Error(ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: dense operand must have rank 1 or 2"); + } + hipsparseDestroySpMat(mat_a); + return err; +} + +XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerROCM, csrMatmulRocm, + xla::ffi::Ffi::Bind() + .Ctx>() + .Ctx() + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Ret() // out + .Attr("m") + .Attr("n") + .Attr("transpose") + .Attr("index_base")); +#endif // REACTANT_ROCM + void registerReactantXLAInternalFFI() { XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_julia_callback", "Host", juliaCallbackHandlerHost); #if defined(REACTANT_CUDA) XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_julia_callback", "CUDA", juliaCallbackHandlerCUDA); + XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_csr_matmul", + "CUDA", csrMatmulHandlerCUDA); +#endif +#if defined(REACTANT_ROCM) + XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_csr_matmul", + "ROCM", csrMatmulHandlerROCM); #endif } diff --git a/ext/ReactantSparseArraysExt/CSR.jl b/ext/ReactantSparseArraysExt/CSR.jl new file mode 100644 index 0000000000..9a87e3dcaf --- /dev/null +++ b/ext/ReactantSparseArraysExt/CSR.jl @@ -0,0 +1,22 @@ +# Conversion between SparseArrays types and the opaque `Reactant.CSRMatrix`. + +function Reactant.CSRMatrix(A::SparseMatrixCSC{T,Ti}) where {T,Ti} + At = copy(transpose(A)) # CSC of Aᵀ is the CSR representation of A + return Reactant.CSRMatrix{T,Ti,Vector{T},Vector{Ti}}( + size(A, 1), size(A, 2), At.colptr, At.rowval, At.nzval + ) +end + +function Reactant.to_rarray(A::SparseMatrixCSC; kwargs...) + return Reactant.to_rarray(Reactant.CSRMatrix(A); kwargs...) +end + +SparseArrays.nnz(A::Reactant.CSRMatrix) = length(A.colind) + +function SparseArrays.SparseMatrixCSC(A::Reactant.CSRMatrix{T,Ti}) where {T,Ti} + # The CSR buffers of A are the CSC representation of Aᵀ + At = SparseMatrixCSC{T,Ti}( + A.n, A.m, Vector{Ti}(A.rowptr), Vector{Ti}(A.colind), Vector{T}(A.nzval) + ) + return copy(transpose(At)) +end diff --git a/ext/ReactantSparseArraysExt/ReactantSparseArraysExt.jl b/ext/ReactantSparseArraysExt/ReactantSparseArraysExt.jl index b782ad62ad..6d224c8179 100644 --- a/ext/ReactantSparseArraysExt/ReactantSparseArraysExt.jl +++ b/ext/ReactantSparseArraysExt/ReactantSparseArraysExt.jl @@ -2,9 +2,15 @@ module ReactantSparseArraysExt using Reactant: Reactant, TracedRNumber using SparseArrays: - SparseArrays, ReadOnly, AbstractSparseArray, CHOLMOD, AbstractSparseMatrixCSC + SparseArrays, + ReadOnly, + AbstractSparseArray, + CHOLMOD, + AbstractSparseMatrixCSC, + SparseMatrixCSC include("Errors.jl") include("ReadOnly.jl") +include("CSR.jl") end diff --git a/src/Reactant.jl b/src/Reactant.jl index 789f352b27..2559716d67 100644 --- a/src/Reactant.jl +++ b/src/Reactant.jl @@ -271,6 +271,7 @@ const TracedType = Union{TracedRArray,TracedRNumber,MissingTracedValue} include("ControlFlow.jl") include("Tracing.jl") +include("SparseTensors.jl") include("compiler/Compiler.jl") @@ -305,7 +306,8 @@ export ConcreteRArray, @code_xla, @jit, @trace, - within_compile + within_compile, + CSRMatrix @static if VERSION ≥ v"1.11" @eval $(Expr(:public, :Periodic, :Binomial)) diff --git a/src/SparseTensors.jl b/src/SparseTensors.jl new file mode 100644 index 0000000000..e61a244a6d --- /dev/null +++ b/src/SparseTensors.jl @@ -0,0 +1,230 @@ +# Minimal sparse-matrix support lowered through the MLIR `sparse_tensor` dialect. +# +# `CSRMatrix` is an opaque wrapper (deliberately not an `AbstractArray`, so it is +# never swept into the dense `AnyTracedRArray` overloads) holding the three CSR +# buffers. At trace time `A * x` / `A * B` / `mul!` emit a +# `sparse_tensor.assemble` producing a `tensor` +# value consumed by a `stablehlo.dot_general`. `Compiler.lower_sparse_ops!` +# rewrites that pair into a `stablehlo.custom_call @reactant_csr_matmul` handled +# by cuSPARSE/hipSPARSE (see deps/ReactantExtra/xla_ffi.cpp) before any pass +# pipeline (and hence the verifier or XLA) sees the sparse-encoded types. + +""" + CSRMatrix{T,Ti}(m, n, rowptr, colind, nzval) + +An opaque `m × n` sparse matrix in CSR format with element type `T` and index +type `Ti`. `rowptr` has length `m + 1` and `colind`/`nzval` have length `nnz`; +all indices are 1-based. + +Construct one from a `SparseArrays.SparseMatrixCSC` via `CSRMatrix(A)` (requires +loading SparseArrays), and pass it through [`Reactant.to_rarray`](@ref) like any +other array. Inside traced functions only `A * x`, `A * B`, and +`LinearAlgebra.mul!` are supported, and execution requires a CUDA or ROCm +backend. +""" +struct CSRMatrix{T,Ti,V<:AbstractVector,Vi<:AbstractVector} + m::Int + n::Int + rowptr::Vi + colind::Vi + nzval::V +end + +function CSRMatrix( + m::Integer, + n::Integer, + rowptr::AbstractVector, + colind::AbstractVector, + nzval::AbstractVector, +) + length(rowptr) == m + 1 || + throw(ArgumentError("rowptr must have length m + 1 = $(m + 1)")) + length(colind) == length(nzval) || + throw(ArgumentError("colind and nzval must have the same length")) + return CSRMatrix{ + unwrapped_eltype(eltype(nzval)), + unwrapped_eltype(eltype(colind)), + typeof(nzval), + typeof(colind), + }( + m, n, rowptr, colind, nzval + ) +end + +const TracedCSRMatrix{T,Ti} = CSRMatrix{T,Ti,TracedRArray{T,1},TracedRArray{Ti,1}} + +Base.size(A::CSRMatrix) = (A.m, A.n) +Base.size(A::CSRMatrix, i::Integer) = i <= 2 ? size(A)[i] : 1 +Base.eltype(::Core.Type{<:CSRMatrix{T}}) where {T} = T +Base.eltype(::CSRMatrix{T}) where {T} = T + +function Base.show(io::IO, A::CSRMatrix{T,Ti}) where {T,Ti} + return print( + io, "$(A.m)×$(A.n) CSRMatrix{$T,$Ti} with $(length(A.colind)) stored entries" + ) +end + +# Tracing +Base.@nospecializeinfer function traced_type_inner( + @nospecialize(_::Core.Type{CSRMatrix{T,Ti,V,Vi}}), + seen, + mode::TraceMode, + @nospecialize(track_numbers::Core.Type), + @nospecialize(ndevices), + @nospecialize(runtime) +) where {T,Ti,V,Vi} + V2 = traced_type_inner(V, seen, mode, track_numbers, ndevices, runtime) + Vi2 = traced_type_inner(Vi, seen, mode, track_numbers, ndevices, runtime) + return CSRMatrix{T,Ti,V2,Vi2} +end + +Base.@nospecializeinfer function make_tracer( + seen, @nospecialize(prev::CSRMatrix), @nospecialize(path), mode; kwargs... +) + return make_tracer_via_immutable_constructor(seen, prev, path, mode; kwargs...) +end + +function use_overlayed_version(A::CSRMatrix) + return use_overlayed_version((A.rowptr, A.colind, A.nzval)) +end + +# IR emission +function _csr_encoding(::Core.Type{Ti}) where {Ti} + width = 8 * sizeof(Ti) + return Base.parse( + MLIR.IR.Attribute, + "#sparse_tensor.encoding<{ map = (d0, d1) -> (d0 : dense, d1 : compressed), posWidth = $width, crdWidth = $width }>", + ) +end + +function _with_nzval_eltype(::Core.Type{T}, A::TracedCSRMatrix{T}) where {T} + return A +end +function _with_nzval_eltype(::Core.Type{T}, A::TracedCSRMatrix{T2,Ti}) where {T,T2,Ti} + nzval = promote_to(TracedRArray{T,1}, A.nzval) + return CSRMatrix{T,Ti,TracedRArray{T,1},TracedRArray{Ti,1}}( + A.m, A.n, A.rowptr, A.colind, nzval + ) +end + +""" + sparse_csr_dot(A::TracedCSRMatrix, B::TracedRArray) + +Emits `sparse_tensor.assemble` + `stablehlo.dot_general` computing `A * B` (spmv +for vector `B`, spmm for matrix `B`) and returns the dense result. The emitted +pair is lowered to a library call by `Compiler.lower_sparse_ops!`. +""" +function sparse_csr_dot( + A::TracedCSRMatrix{T,Ti}, + B::TracedRArray{T}; + location=Ops.mlir_stacktrace("sparse_csr_dot", @__FILE__, @__LINE__), +) where {T,Ti} + ndims(B) in (1, 2) || + throw(ArgumentError("Only vectors and matrices can be multiplied by a CSRMatrix")) + size(B, 1) == A.n || + throw(DimensionMismatch("A has size $(size(A)), B has size $(size(B))")) + ressize = ndims(B) == 1 ? Int[A.m] : Int[A.m, size(B, 2)] + + sparse_type = MLIR.IR.TensorType(Int[A.m, A.n], MLIR.IR.Type(T), _csr_encoding(Ti)) + asm = MLIR.Dialects.sparse_tensor.assemble( + MLIR.IR.Value[A.rowptr.mlir_data, A.colind.mlir_data], + A.nzval.mlir_data; + result=sparse_type, + location, + ) + + ctx = MLIR.IR.current_context() + batching_dimensions = Int64[] + lhs_contracting_dimensions = Int64[1] + rhs_contracting_dimensions = Int64[0] + dot_dimension_numbers = GC.@preserve ctx batching_dimensions lhs_contracting_dimensions rhs_contracting_dimensions begin + MLIR.IR.Attribute( + MLIR.API.stablehloDotDimensionNumbersGet( + ctx, + 0, + batching_dimensions, + 0, + batching_dimensions, + 1, + lhs_contracting_dimensions, + 1, + rhs_contracting_dimensions, + ), + ) + end + + res = MLIR.IR.result( + MLIR.Dialects.stablehlo.dot_general( + MLIR.IR.result(asm, 1), + B.mlir_data; + result_0=MLIR.IR.TensorType(ressize, MLIR.IR.Type(T)), + dot_dimension_numbers, + location, + ), + ) + return TracedRArray{T,length(ressize)}((), res, Tuple(ressize)) +end + +# LinearAlgebra surface +function LinearAlgebra.mul!( + C::TracedRArray{T}, + A::TracedCSRMatrix, + B::AbstractVecOrMat, + α::Number=true, + β::Number=false, +) where {T} + ndims(C) in (1, 2) || throw(ArgumentError("C must be a vector or a matrix")) + B = promote_to(TracedRArray{T}, B) + + size(A, 2) == size(B, 1) || + throw(DimensionMismatch("A has size $(size(A)), B has size $(size(B))")) + size(C, 1) == size(A, 1) || + throw(DimensionMismatch("C has size $(size(C)), A has size $(size(A))")) + size(C, 2) == size(B, 2) || + throw(DimensionMismatch("C has size $(size(C)), B has size $(size(B))")) + + tmp = sparse_csr_dot(_with_nzval_eltype(T, A), B) + + β_is_zero = !(β isa TracedRNumber) && iszero(β) + α_is_one = !(α isa TracedRNumber) && isone(α) + + if α_is_one && β_is_zero + res = tmp + else + α_res = if α_is_one + tmp + else + Ops.multiply(tmp, Ops.fill(promote_to(TracedRNumber{T}, α), size(tmp))) + end + if β_is_zero + res = α_res + else + C_mat = ReactantCore.materialize_traced_array(C) + β_C = Ops.multiply(C_mat, Ops.fill(promote_to(TracedRNumber{T}, β), size(C_mat))) + res = Ops.add(α_res, β_C) + end + end + + if ndims(C) == 2 && size(C, 2) == 1 && ndims(res) == 1 + res = reshape(res, size(C)) + end + + TracedUtils.set_mlir_data!(C, TracedUtils.get_mlir_data(res)) + return C +end + +function _sparse_mul(A::TracedCSRMatrix{T}, B::AbstractVecOrMat) where {T} + T2 = Base.promote_op(*, T, unwrapped_eltype(eltype(B))) + return sparse_csr_dot(_with_nzval_eltype(T2, A), promote_to(TracedRArray{T2}, B)) +end + +Base.:*(A::TracedCSRMatrix, x::AbstractVector) = _sparse_mul(A, x) +Base.:*(A::TracedCSRMatrix, B::AbstractMatrix) = _sparse_mul(A, B) + +for f in (:adjoint, :transpose) + @eval function Base.$(f)(::CSRMatrix) + return error( + "`$($(QuoteNode(f)))` of a `Reactant.CSRMatrix` is not supported yet; only `A * x`, `A * B` and `mul!` are implemented.", + ) + end +end diff --git a/src/compiler/Compiler.jl b/src/compiler/Compiler.jl index c0a28f4a52..40fe75134c 100644 --- a/src/compiler/Compiler.jl +++ b/src/compiler/Compiler.jl @@ -24,6 +24,7 @@ using Reactant_jll: Reactant_jll include("Macros.jl") include("CompilationError.jl") include("OptimizationPasses.jl") +include("SparseLowering.jl") include("Codegen.jl") include("Thunk.jl") @@ -430,6 +431,14 @@ function compile_mlir!( legal_to_run_shardy_passes = compile_options.optimization_passes === :all + # Lower sparse_tensor CSR products to library custom calls before any pass + # pipeline runs: XLA cannot consume sparse-encoded tensors and the + # sparse-encoded ops must never reach the verifier. With `:none` the sparse + # IR is kept as-is for inspection. + if compile_options.optimization_passes !== :none + lower_sparse_ops!(mod) + end + # Raise any triton kernel that might exist as a custom call # We will lower them back later on, but having the full triton IR enables # optimizations / differentiation, so we unconditionally do it diff --git a/src/compiler/SparseLowering.jl b/src/compiler/SparseLowering.jl new file mode 100644 index 0000000000..239f571ea3 --- /dev/null +++ b/src/compiler/SparseLowering.jl @@ -0,0 +1,108 @@ +# Lowers the `sparse_tensor.assemble` + `stablehlo.dot_general` pairs emitted by +# `Reactant.sparse_csr_dot` into `stablehlo.custom_call @reactant_csr_matmul` +# operating on the raw CSR buffers. Must run before any pass pipeline: XLA +# cannot consume sparse-encoded tensor types, and the sparse-encoded +# `dot_general` must never reach the verifier. + +function _walk_operations!(f, op::MLIR.IR.Operation) + f(op) + for region in op, block in region, inner in block + _walk_operations!(f, inner) + end + return nothing +end + +function _is_csr_dot(op::MLIR.IR.Operation) + MLIR.IR.name(op) == "stablehlo.dot_general" || return false + MLIR.IR.noperands(op) == 2 || return false + lhs = MLIR.IR.operand(op, 1) + MLIR.IR.is_op_res(lhs) || return false + return MLIR.IR.name(MLIR.IR.op_owner(lhs)) == "sparse_tensor.assemble" +end + +function _lower_csr_dot!(dot::MLIR.IR.Operation) + asm = MLIR.IR.op_owner(MLIR.IR.operand(dot, 1)) + rowptr = MLIR.IR.operand(asm, 1) + colind = MLIR.IR.operand(asm, 2) + nzval = MLIR.IR.operand(asm, 3) + rhs = MLIR.IR.operand(dot, 2) + + sparse_type = MLIR.IR.type(MLIR.IR.operand(dot, 1)) + m, n = size(sparse_type, 1), size(sparse_type, 2) + result_type = MLIR.IR.type(MLIR.IR.result(dot, 1)) + + cc = MLIR.Dialects.stablehlo.custom_call( + MLIR.IR.Value[rowptr, colind, nzval, rhs]; + result_0=MLIR.IR.Type[result_type], + call_target_name="reactant_csr_matmul", + api_version=Int32(4), + has_side_effect=MLIR.IR.Attribute(false), + backend_config=Dict( + "m" => MLIR.IR.Attribute(Int64(m)), + "n" => MLIR.IR.Attribute(Int64(n)), + "transpose" => MLIR.IR.Attribute(Int64(0)), + "index_base" => MLIR.IR.Attribute(Int64(1)), + ), + operand_layouts=MLIR.IR.Attribute([ + Reactant.Ops._col_major_layout(1), + Reactant.Ops._col_major_layout(1), + Reactant.Ops._col_major_layout(1), + Reactant.Ops._col_major_layout(ndims(MLIR.IR.type(rhs))), + ]), + result_layouts=MLIR.IR.Attribute([Reactant.Ops._col_major_layout(ndims(result_type))]), + location=MLIR.IR.location(dot), + ) + + MLIR.IR.insert_before!(MLIR.IR.block(dot), dot, cc) + MLIR.API.mlirValueReplaceAllUsesOfWith(MLIR.IR.result(dot, 1), MLIR.IR.result(cc, 1)) + MLIR.IR.rmfromparent!(dot) + MLIR.IR.dispose(dot) + return asm +end + +function lower_sparse_ops!(mod::MLIR.IR.Module) + # `stablehlo.custom_call` ops created below must stay detached until we + # insert them next to the `dot_general` they replace. + @assert !MLIR.IR.has_block() + + has_sparse = false + dots = MLIR.IR.Operation[] + for top in collect(MLIR.IR.body(mod)) + _walk_operations!(top) do op + startswith(MLIR.IR.name(op), "sparse_tensor.") && (has_sparse = true) + _is_csr_dot(op) && push!(dots, op) + end + end + has_sparse || return mod + + assembles = MLIR.IR.Operation[] + for dot in dots + push!(assembles, _lower_csr_dot!(dot)) + end + + seen = Set{Ptr{Cvoid}}() + for asm in assembles + ptr = Base.unsafe_convert(MLIR.API.MlirOperation, asm).ptr + ptr in seen && continue + push!(seen, ptr) + if MLIR.IR.first_use(MLIR.IR.result(asm, 1)) === nothing + MLIR.IR.rmfromparent!(asm) + MLIR.IR.dispose(asm) + end + end + + for top in collect(MLIR.IR.body(mod)) + _walk_operations!(top) do op + opname = MLIR.IR.name(op) + if startswith(opname, "sparse_tensor.") + error( + "Unsupported use of the sparse_tensor dialect: `$opname` survived " * + "sparse lowering. Only CSR `A * x` / `A * B` products (as emitted " * + "by `Reactant.sparse_csr_dot`) can be lowered.", + ) + end + end + end + + return mod +end diff --git a/test/Project.toml b/test/Project.toml index 9dc88cf865..b130956b10 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -40,6 +40,7 @@ Random123 = "74087812-796a-5b5d-8853-05524746bad3" Reactant = "3c362404-f566-11ee-1572-e11a4b42c853" Serialization = "9e88b42a-f829-5b0c-bbe9-9e923198166b" Setfield = "efcf1570-3423-57d1-acb7-fd33fddbac46" +SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf" SpecialFunctions = "276daf66-3868-5448-9aa4-cd146d93841b" StableRNGs = "860ef19b-820b-49d6-a774-d7a799459cd3" Static = "aedffcd0-7271-4cad-89d0-dc628f76c6d3" @@ -89,6 +90,7 @@ Random = "1.10" Random123 = "1.7" Serialization = "1.10" Setfield = "1" +SparseArrays = "1.10" SpecialFunctions = "2.4" StableRNGs = "1" Static = "0.8, 1" diff --git a/test/core/sparse.jl b/test/core/sparse.jl new file mode 100644 index 0000000000..34c00a2887 --- /dev/null +++ b/test/core/sparse.jl @@ -0,0 +1,96 @@ +using Reactant, Test, LinearAlgebra, SparseArrays, FileCheck +using Random: Random + +const RunningOnGPU = + contains(string(Reactant.devices()[1]), "CUDA") || + contains(string(Reactant.devices()[1]), "ROCM") + +spmv(A, x) = A * x +spmm(A, B) = A * B +spmv_mul!(C, A, B, α, β) = LinearAlgebra.mul!(C, A, B, α, β) + +@testset "CSRMatrix construction and tracing" begin + rng = Random.MersenneTwister(0) + A = sprand(rng, Float64, 10, 8, 0.3) + + Acsr = Reactant.CSRMatrix(A) + @test size(Acsr) == (10, 8) + @test eltype(Acsr) == Float64 + @test nnz(Acsr) == nnz(A) + @test SparseMatrixCSC(Acsr) == A + + A_ra = Reactant.to_rarray(A) + @test A_ra isa Reactant.CSRMatrix + @test A_ra.rowptr isa ConcreteRArray + @test A_ra.colind isa ConcreteRArray + @test A_ra.nzval isa ConcreteRArray + @test SparseMatrixCSC(A_ra) == A +end + +@testset "sparse_tensor IR" begin + rng = Random.MersenneTwister(0) + A_ra = Reactant.to_rarray(sprand(rng, Float64, 10, 8, 0.3)) + x_ra = Reactant.to_rarray(rand(rng, 8)) + B_ra = Reactant.to_rarray(rand(rng, 8, 3)) + + hlo = @code_hlo optimize = :none spmv(A_ra, x_ra) + @test @filecheck begin + @check "sparse_tensor.assemble" + @check "#sparse_tensor.encoding" + @check "stablehlo.dot_general" + hlo + end + + hlo = @code_hlo optimize = :none spmm(A_ra, B_ra) + @test @filecheck begin + @check "sparse_tensor.assemble" + @check "stablehlo.dot_general" + hlo + end +end + +@testset "lowering to custom_call" begin + rng = Random.MersenneTwister(0) + A_ra = Reactant.to_rarray(sprand(rng, Float64, 10, 8, 0.3)) + x_ra = Reactant.to_rarray(rand(rng, 8)) + B_ra = Reactant.to_rarray(rand(rng, 8, 3)) + C_ra = Reactant.to_rarray(rand(rng, 10)) + + for hlo in ( + @code_hlo(spmv(A_ra, x_ra)), + @code_hlo(spmm(A_ra, B_ra)), + @code_hlo(spmv_mul!(C_ra, A_ra, x_ra, 2.0, 3.0)), + ) + @test @filecheck begin + @check "stablehlo.custom_call" + @check "reactant_csr_matmul" + hlo + end + @test !contains(repr(hlo), "sparse_tensor") + end +end + +@testset "numerical correctness" begin + if !RunningOnGPU + @test_skip "CSR matmul execution requires a CUDA or ROCm backend" + else + rng = Random.MersenneTwister(0) + @testset for T in (Float32, Float64), Ti in (Int32, Int64) + A = SparseMatrixCSC{T,Ti}(sprand(rng, T, 10, 8, 0.3)) + x = rand(rng, T, 8) + B = rand(rng, T, 8, 3) + C = rand(rng, T, 10) + + A_ra = Reactant.to_rarray(A) + x_ra = Reactant.to_rarray(x) + B_ra = Reactant.to_rarray(B) + + @test Array(@jit(spmv(A_ra, x_ra))) ≈ A * x atol = 1e-5 rtol = 1e-5 + @test Array(@jit(spmm(A_ra, B_ra))) ≈ A * B atol = 1e-5 rtol = 1e-5 + + C_ra = Reactant.to_rarray(C) + @jit spmv_mul!(C_ra, A_ra, x_ra, T(2), T(3)) + @test Array(C_ra) ≈ 2 .* (A * x) .+ 3 .* C atol = 1e-5 rtol = 1e-5 + end + end +end From 6e6d784fddfb41c4ee521d73d202262cb57ff72f Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Wed, 19 Aug 2026 09:46:06 +0200 Subject: [PATCH 02/11] fix FileCheck test --- test/core/sparse.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/core/sparse.jl b/test/core/sparse.jl index 34c00a2887..66217cb522 100644 --- a/test/core/sparse.jl +++ b/test/core/sparse.jl @@ -35,8 +35,8 @@ end hlo = @code_hlo optimize = :none spmv(A_ra, x_ra) @test @filecheck begin - @check "sparse_tensor.assemble" @check "#sparse_tensor.encoding" + @check "sparse_tensor.assemble" @check "stablehlo.dot_general" hlo end From 1890046b9ee4c0147df717257ffb717a2aa35cdc Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Wed, 19 Aug 2026 09:01:54 +0000 Subject: [PATCH 03/11] Move sparse CSR lowering into the Enzyme-JAX lower-sparse-csr pass Replaces the Julia-side IR rewrite (src/compiler/SparseLowering.jl) with the new `lower-sparse-csr` MLIR pass in Enzyme-JAX, invoked via run_pass_pipeline! only when the traced module actually contains sparse_tensor ops (so older jlls without the pass keep working for dense code). Also switches the CSR buffers to 0-based indices as required by the sparse_tensor dialect; the pass emits index_base = 0 and the FFI handlers already support both bases. Co-Authored-By: Claude Fable 5 --- deps/ReactantExtra/xla_ffi.cpp | 4 +- ext/ReactantSparseArraysExt/CSR.jl | 11 ++- src/SparseTensors.jl | 12 ++-- src/compiler/Compiler.jl | 29 ++++++-- src/compiler/SparseLowering.jl | 108 ----------------------------- 5 files changed, 38 insertions(+), 126 deletions(-) delete mode 100644 src/compiler/SparseLowering.jl diff --git a/deps/ReactantExtra/xla_ffi.cpp b/deps/ReactantExtra/xla_ffi.cpp index 9dc7a0b9bd..a91db6d84d 100644 --- a/deps/ReactantExtra/xla_ffi.cpp +++ b/deps/ReactantExtra/xla_ffi.cpp @@ -94,12 +94,12 @@ XLA_FFI_DEFINE_HANDLER( // ============================================================================ // CSR sparse matrix products (spmv / spmm) via cuSPARSE / hipSPARSE. // -// The Julia side (src/compiler/SparseLowering.jl) emits a +// The Enzyme-JAX pass `lower-sparse-csr` emits a // stablehlo.custom_call targeting "reactant_csr_matmul" with api_version = 4 // (TYPED_FFI). Operands are (rowptr, colind, nzval, dense) with column-major // layouts pinned; the result is dense. The backend_config dict carries i64 // attributes "m", "n", "transpose" (must be 0 for now), and "index_base" -// (0 or 1; the Julia side emits 1-based CSR buffers). +// (0 or 1; sparse_tensor-dialect buffers are 0-based). // ============================================================================ #if defined(REACTANT_CUDA) diff --git a/ext/ReactantSparseArraysExt/CSR.jl b/ext/ReactantSparseArraysExt/CSR.jl index 9a87e3dcaf..47553f50e8 100644 --- a/ext/ReactantSparseArraysExt/CSR.jl +++ b/ext/ReactantSparseArraysExt/CSR.jl @@ -2,8 +2,9 @@ function Reactant.CSRMatrix(A::SparseMatrixCSC{T,Ti}) where {T,Ti} At = copy(transpose(A)) # CSC of Aᵀ is the CSR representation of A + # `sparse_tensor` positions/coordinates are 0-based return Reactant.CSRMatrix{T,Ti,Vector{T},Vector{Ti}}( - size(A, 1), size(A, 2), At.colptr, At.rowval, At.nzval + size(A, 1), size(A, 2), At.colptr .- one(Ti), At.rowval .- one(Ti), At.nzval ) end @@ -14,9 +15,13 @@ end SparseArrays.nnz(A::Reactant.CSRMatrix) = length(A.colind) function SparseArrays.SparseMatrixCSC(A::Reactant.CSRMatrix{T,Ti}) where {T,Ti} - # The CSR buffers of A are the CSC representation of Aᵀ + # The (0-based) CSR buffers of A are the CSC representation of Aᵀ At = SparseMatrixCSC{T,Ti}( - A.n, A.m, Vector{Ti}(A.rowptr), Vector{Ti}(A.colind), Vector{T}(A.nzval) + A.n, + A.m, + Vector{Ti}(A.rowptr) .+ one(Ti), + Vector{Ti}(A.colind) .+ one(Ti), + Vector{T}(A.nzval), ) return copy(transpose(At)) end diff --git a/src/SparseTensors.jl b/src/SparseTensors.jl index e61a244a6d..aea4d7b24d 100644 --- a/src/SparseTensors.jl +++ b/src/SparseTensors.jl @@ -4,17 +4,17 @@ # never swept into the dense `AnyTracedRArray` overloads) holding the three CSR # buffers. At trace time `A * x` / `A * B` / `mul!` emit a # `sparse_tensor.assemble` producing a `tensor` -# value consumed by a `stablehlo.dot_general`. `Compiler.lower_sparse_ops!` -# rewrites that pair into a `stablehlo.custom_call @reactant_csr_matmul` handled -# by cuSPARSE/hipSPARSE (see deps/ReactantExtra/xla_ffi.cpp) before any pass -# pipeline (and hence the verifier or XLA) sees the sparse-encoded types. +# value consumed by a `stablehlo.dot_general`. The Enzyme-JAX `lower-sparse-csr` +# pass rewrites that pair into a `stablehlo.custom_call @reactant_csr_matmul` +# handled by cuSPARSE/hipSPARSE (see deps/ReactantExtra/xla_ffi.cpp) before +# XLA sees the sparse-encoded types. """ CSRMatrix{T,Ti}(m, n, rowptr, colind, nzval) An opaque `m × n` sparse matrix in CSR format with element type `T` and index type `Ti`. `rowptr` has length `m + 1` and `colind`/`nzval` have length `nnz`; -all indices are 1-based. +all indices are 0-based, as required by the MLIR `sparse_tensor` dialect. Construct one from a `SparseArrays.SparseMatrixCSC` via `CSRMatrix(A)` (requires loading SparseArrays), and pass it through [`Reactant.to_rarray`](@ref) like any @@ -112,7 +112,7 @@ end Emits `sparse_tensor.assemble` + `stablehlo.dot_general` computing `A * B` (spmv for vector `B`, spmm for matrix `B`) and returns the dense result. The emitted -pair is lowered to a library call by `Compiler.lower_sparse_ops!`. +pair is lowered to a library call by the Enzyme-JAX `lower-sparse-csr` pass. """ function sparse_csr_dot( A::TracedCSRMatrix{T,Ti}, diff --git a/src/compiler/Compiler.jl b/src/compiler/Compiler.jl index 40fe75134c..a5f72fcf22 100644 --- a/src/compiler/Compiler.jl +++ b/src/compiler/Compiler.jl @@ -24,10 +24,26 @@ using Reactant_jll: Reactant_jll include("Macros.jl") include("CompilationError.jl") include("OptimizationPasses.jl") -include("SparseLowering.jl") include("Codegen.jl") include("Thunk.jl") +function has_sparse_tensor_ops(mod::MLIR.IR.Module) + found = false + function visit(op) + found && return nothing + startswith(MLIR.IR.name(op), "sparse_tensor.") && (found = true; return nothing) + for region in op, block in region, inner in block + visit(inner) + end + return nothing + end + for op in MLIR.IR.body(mod) + visit(op) + found && break + end + return found +end + const DEBUG_PRINT_CODEGEN = Ref(false) const __module_gc_vector = Dict{MLIR.IR.Module,Vector{Union{TracedRArray,TracedRNumber}}}() @@ -431,12 +447,11 @@ function compile_mlir!( legal_to_run_shardy_passes = compile_options.optimization_passes === :all - # Lower sparse_tensor CSR products to library custom calls before any pass - # pipeline runs: XLA cannot consume sparse-encoded tensors and the - # sparse-encoded ops must never reach the verifier. With `:none` the sparse - # IR is kept as-is for inspection. - if compile_options.optimization_passes !== :none - lower_sparse_ops!(mod) + # Lower sparse_tensor CSR products to library custom calls before anything + # else runs: XLA cannot consume sparse-encoded tensor types. With `:none` + # the sparse IR is kept as-is for inspection. + if compile_options.optimization_passes !== :none && has_sparse_tensor_ops(mod) + run_pass_pipeline!(mod, "lower-sparse-csr", "lower_sparse_csr") end # Raise any triton kernel that might exist as a custom call diff --git a/src/compiler/SparseLowering.jl b/src/compiler/SparseLowering.jl deleted file mode 100644 index 239f571ea3..0000000000 --- a/src/compiler/SparseLowering.jl +++ /dev/null @@ -1,108 +0,0 @@ -# Lowers the `sparse_tensor.assemble` + `stablehlo.dot_general` pairs emitted by -# `Reactant.sparse_csr_dot` into `stablehlo.custom_call @reactant_csr_matmul` -# operating on the raw CSR buffers. Must run before any pass pipeline: XLA -# cannot consume sparse-encoded tensor types, and the sparse-encoded -# `dot_general` must never reach the verifier. - -function _walk_operations!(f, op::MLIR.IR.Operation) - f(op) - for region in op, block in region, inner in block - _walk_operations!(f, inner) - end - return nothing -end - -function _is_csr_dot(op::MLIR.IR.Operation) - MLIR.IR.name(op) == "stablehlo.dot_general" || return false - MLIR.IR.noperands(op) == 2 || return false - lhs = MLIR.IR.operand(op, 1) - MLIR.IR.is_op_res(lhs) || return false - return MLIR.IR.name(MLIR.IR.op_owner(lhs)) == "sparse_tensor.assemble" -end - -function _lower_csr_dot!(dot::MLIR.IR.Operation) - asm = MLIR.IR.op_owner(MLIR.IR.operand(dot, 1)) - rowptr = MLIR.IR.operand(asm, 1) - colind = MLIR.IR.operand(asm, 2) - nzval = MLIR.IR.operand(asm, 3) - rhs = MLIR.IR.operand(dot, 2) - - sparse_type = MLIR.IR.type(MLIR.IR.operand(dot, 1)) - m, n = size(sparse_type, 1), size(sparse_type, 2) - result_type = MLIR.IR.type(MLIR.IR.result(dot, 1)) - - cc = MLIR.Dialects.stablehlo.custom_call( - MLIR.IR.Value[rowptr, colind, nzval, rhs]; - result_0=MLIR.IR.Type[result_type], - call_target_name="reactant_csr_matmul", - api_version=Int32(4), - has_side_effect=MLIR.IR.Attribute(false), - backend_config=Dict( - "m" => MLIR.IR.Attribute(Int64(m)), - "n" => MLIR.IR.Attribute(Int64(n)), - "transpose" => MLIR.IR.Attribute(Int64(0)), - "index_base" => MLIR.IR.Attribute(Int64(1)), - ), - operand_layouts=MLIR.IR.Attribute([ - Reactant.Ops._col_major_layout(1), - Reactant.Ops._col_major_layout(1), - Reactant.Ops._col_major_layout(1), - Reactant.Ops._col_major_layout(ndims(MLIR.IR.type(rhs))), - ]), - result_layouts=MLIR.IR.Attribute([Reactant.Ops._col_major_layout(ndims(result_type))]), - location=MLIR.IR.location(dot), - ) - - MLIR.IR.insert_before!(MLIR.IR.block(dot), dot, cc) - MLIR.API.mlirValueReplaceAllUsesOfWith(MLIR.IR.result(dot, 1), MLIR.IR.result(cc, 1)) - MLIR.IR.rmfromparent!(dot) - MLIR.IR.dispose(dot) - return asm -end - -function lower_sparse_ops!(mod::MLIR.IR.Module) - # `stablehlo.custom_call` ops created below must stay detached until we - # insert them next to the `dot_general` they replace. - @assert !MLIR.IR.has_block() - - has_sparse = false - dots = MLIR.IR.Operation[] - for top in collect(MLIR.IR.body(mod)) - _walk_operations!(top) do op - startswith(MLIR.IR.name(op), "sparse_tensor.") && (has_sparse = true) - _is_csr_dot(op) && push!(dots, op) - end - end - has_sparse || return mod - - assembles = MLIR.IR.Operation[] - for dot in dots - push!(assembles, _lower_csr_dot!(dot)) - end - - seen = Set{Ptr{Cvoid}}() - for asm in assembles - ptr = Base.unsafe_convert(MLIR.API.MlirOperation, asm).ptr - ptr in seen && continue - push!(seen, ptr) - if MLIR.IR.first_use(MLIR.IR.result(asm, 1)) === nothing - MLIR.IR.rmfromparent!(asm) - MLIR.IR.dispose(asm) - end - end - - for top in collect(MLIR.IR.body(mod)) - _walk_operations!(top) do op - opname = MLIR.IR.name(op) - if startswith(opname, "sparse_tensor.") - error( - "Unsupported use of the sparse_tensor dialect: `$opname` survived " * - "sparse lowering. Only CSR `A * x` / `A * B` products (as emitted " * - "by `Reactant.sparse_csr_dot`) can be lowered.", - ) - end - end - end - - return mod -end From 66397abca21b612a75a4feeca4ae79a9a65c49bf Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Wed, 19 Aug 2026 09:35:29 +0000 Subject: [PATCH 04/11] Emit enzymexla.sparse.spmm from mul! for fused sparse alpha/beta `mul!(C, A, B, alpha, beta)` (and `A * B`) on a CSRMatrix now emit the semantic `enzymexla.sparse.spmm alpha, A, B, beta, C` op instead of a bare dot_general, so lower-sparse-csr can fuse constant alpha/beta into a single accumulating cuSPARSE/hipSPARSE call (reactant_csr_matmul_acc, with C aliased to the output and copied on-device when XLA does not donate the buffer). Runtime (traced) alpha/beta keep the unfused custom_call + multiply/add form, now materialized by the pass instead of the Julia frontend. The plain handler gains an f64 "alpha" attribute. Co-Authored-By: Claude Fable 5 --- deps/ReactantExtra/xla_ffi.cpp | 198 ++++++++++++++++++++++++++++----- src/SparseTensors.jl | 106 ++++++++---------- src/mlir/Dialects/EnzymeXLA.jl | 37 ++++++ test/core/sparse.jl | 43 +++++-- 4 files changed, 291 insertions(+), 93 deletions(-) diff --git a/deps/ReactantExtra/xla_ffi.cpp b/deps/ReactantExtra/xla_ffi.cpp index a91db6d84d..c1991df277 100644 --- a/deps/ReactantExtra/xla_ffi.cpp +++ b/deps/ReactantExtra/xla_ffi.cpp @@ -94,12 +94,16 @@ XLA_FFI_DEFINE_HANDLER( // ============================================================================ // CSR sparse matrix products (spmv / spmm) via cuSPARSE / hipSPARSE. // -// The Enzyme-JAX pass `lower-sparse-csr` emits a -// stablehlo.custom_call targeting "reactant_csr_matmul" with api_version = 4 -// (TYPED_FFI). Operands are (rowptr, colind, nzval, dense) with column-major -// layouts pinned; the result is dense. The backend_config dict carries i64 -// attributes "m", "n", "transpose" (must be 0 for now), and "index_base" -// (0 or 1; sparse_tensor-dialect buffers are 0-based). +// The Enzyme-JAX pass `lower-sparse-csr` emits stablehlo.custom_calls with +// api_version = 4 (TYPED_FFI) targeting either +// - "reactant_csr_matmul": out = alpha * A * dense, with operands +// (rowptr, colind, nzval, dense), or +// - "reactant_csr_matmul_acc": out = alpha * A * dense + beta * acc, with +// operands (rowptr, colind, nzval, dense, acc) and acc aliased to out. +// Column-major layouts are pinned; the result is dense. The backend_config +// dict carries i64 attributes "m", "n", "transpose" (must be 0 for now) and +// "index_base" (0 or 1; sparse_tensor-dialect buffers are 0-based), plus f64 +// "alpha" (and "beta" for the accumulating variant). // ============================================================================ #if defined(REACTANT_CUDA) @@ -117,13 +121,18 @@ XLA_FFI_DEFINE_HANDLER( } \ } while (0) -static ffi::Error csrMatmulCuda(cudaStream_t stream, - ffi::ScratchAllocator scratch, - ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, - ffi::AnyBuffer nzval, ffi::AnyBuffer dense, - ffi::Result out, int64_t m, - int64_t n, int64_t transpose, - int64_t index_base) { +// Computes out = alpha_v * A * dense (+ beta_v * *acc when acc != nullptr; +// the accumulated-into operand is aliased to out by the lowering, so it is +// copied into out first if XLA did not reuse the buffer). +static ffi::Error csrMatmulCudaImpl(cudaStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, + ffi::AnyBuffer colind, ffi::AnyBuffer nzval, + ffi::AnyBuffer dense, ffi::AnyBuffer *acc, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base, double alpha_v, + double beta_v) { if (transpose != 0) { return ffi::Error( ffi::ErrorCode::kUnimplemented, @@ -149,18 +158,22 @@ static ffi::Error csrMatmulCuda(cudaStream_t stream, "reactant_csr_matmul: index buffers must be i32 or i64"); } - static const float alpha_f = 1.0f, beta_f = 0.0f; - static const double alpha_d = 1.0, beta_d = 0.0; + const float alpha_f = static_cast(alpha_v), + beta_f = static_cast(beta_v); + const double alpha_d = alpha_v, beta_d = beta_v; cudaDataType value_type; + size_t value_bytes; const void *alpha, *beta; switch (nzval.element_type()) { case ffi::DataType::F32: value_type = CUDA_R_32F; + value_bytes = sizeof(float); alpha = &alpha_f; beta = &beta_f; break; case ffi::DataType::F64: value_type = CUDA_R_64F; + value_bytes = sizeof(double); alpha = &alpha_d; beta = &beta_d; break; @@ -170,11 +183,26 @@ static ffi::Error csrMatmulCuda(cudaStream_t stream, "reactant_csr_matmul: only f32 and f64 values are supported"); } if (dense.element_type() != nzval.element_type() || - out->element_type() != nzval.element_type()) { + out->element_type() != nzval.element_type() || + (acc != nullptr && acc->element_type() != nzval.element_type())) { return ffi::Error( ffi::ErrorCode::kInvalidArgument, "reactant_csr_matmul: value, operand and result dtypes must match"); } + if (acc != nullptr && acc->untyped_data() != out->untyped_data()) { + cudaError_t copy_status = cudaMemcpyAsync( + out->untyped_data(), acc->untyped_data(), + static_cast(out->element_count()) * value_bytes, + cudaMemcpyDeviceToDevice, stream); + if (copy_status != cudaSuccess) { + return ffi::Error( + ffi::ErrorCode::kInternal, + absl::StrFormat( + "reactant_csr_matmul: copying the accumulated-into operand " + "failed: %s", + cudaGetErrorString(copy_status))); + } + } static thread_local cusparseHandle_t handle = nullptr; if (handle == nullptr) { @@ -262,6 +290,31 @@ static ffi::Error csrMatmulCuda(cudaStream_t stream, return err; } +static ffi::Error csrMatmulCuda(cudaStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, + ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base, double alpha) { + return csrMatmulCudaImpl(stream, scratch, rowptr, colind, nzval, dense, + /*acc=*/nullptr, out, m, n, transpose, index_base, + alpha, /*beta_v=*/0.0); +} + +static ffi::Error csrMatmulAccCuda(cudaStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, + ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::AnyBuffer acc, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base, double alpha, + double beta) { + return csrMatmulCudaImpl(stream, scratch, rowptr, colind, nzval, dense, &acc, + out, m, n, transpose, index_base, alpha, beta); +} + XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerCUDA, csrMatmulCuda, xla::ffi::Ffi::Bind() .Ctx>() @@ -274,7 +327,25 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerCUDA, csrMatmulCuda, .Attr("m") .Attr("n") .Attr("transpose") - .Attr("index_base")); + .Attr("index_base") + .Attr("alpha")); + +XLA_FFI_DEFINE_HANDLER(csrMatmulAccHandlerCUDA, csrMatmulAccCuda, + xla::ffi::Ffi::Bind() + .Ctx>() + .Ctx() + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Arg() // acc (aliased to out) + .Ret() // out + .Attr("m") + .Attr("n") + .Attr("transpose") + .Attr("index_base") + .Attr("alpha") + .Attr("beta")); #endif // REACTANT_CUDA #if defined(REACTANT_ROCM) @@ -293,13 +364,18 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerCUDA, csrMatmulCuda, } \ } while (0) -static ffi::Error csrMatmulRocm(hipStream_t stream, - ffi::ScratchAllocator scratch, - ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, - ffi::AnyBuffer nzval, ffi::AnyBuffer dense, - ffi::Result out, int64_t m, - int64_t n, int64_t transpose, - int64_t index_base) { +// Computes out = alpha_v * A * dense (+ beta_v * *acc when acc != nullptr; +// the accumulated-into operand is aliased to out by the lowering, so it is +// copied into out first if XLA did not reuse the buffer). +static ffi::Error csrMatmulRocmImpl(hipStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, + ffi::AnyBuffer colind, ffi::AnyBuffer nzval, + ffi::AnyBuffer dense, ffi::AnyBuffer *acc, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base, double alpha_v, + double beta_v) { if (transpose != 0) { return ffi::Error( ffi::ErrorCode::kUnimplemented, @@ -325,18 +401,22 @@ static ffi::Error csrMatmulRocm(hipStream_t stream, "reactant_csr_matmul: index buffers must be i32 or i64"); } - static const float alpha_f = 1.0f, beta_f = 0.0f; - static const double alpha_d = 1.0, beta_d = 0.0; + const float alpha_f = static_cast(alpha_v), + beta_f = static_cast(beta_v); + const double alpha_d = alpha_v, beta_d = beta_v; hipDataType value_type; + size_t value_bytes; const void *alpha, *beta; switch (nzval.element_type()) { case ffi::DataType::F32: value_type = HIP_R_32F; + value_bytes = sizeof(float); alpha = &alpha_f; beta = &beta_f; break; case ffi::DataType::F64: value_type = HIP_R_64F; + value_bytes = sizeof(double); alpha = &alpha_d; beta = &beta_d; break; @@ -346,11 +426,26 @@ static ffi::Error csrMatmulRocm(hipStream_t stream, "reactant_csr_matmul: only f32 and f64 values are supported"); } if (dense.element_type() != nzval.element_type() || - out->element_type() != nzval.element_type()) { + out->element_type() != nzval.element_type() || + (acc != nullptr && acc->element_type() != nzval.element_type())) { return ffi::Error( ffi::ErrorCode::kInvalidArgument, "reactant_csr_matmul: value, operand and result dtypes must match"); } + if (acc != nullptr && acc->untyped_data() != out->untyped_data()) { + hipError_t copy_status = hipMemcpyAsync( + out->untyped_data(), acc->untyped_data(), + static_cast(out->element_count()) * value_bytes, + hipMemcpyDeviceToDevice, stream); + if (copy_status != hipSuccess) { + return ffi::Error( + ffi::ErrorCode::kInternal, + absl::StrFormat( + "reactant_csr_matmul: copying the accumulated-into operand " + "failed: %s", + hipGetErrorString(copy_status))); + } + } static thread_local hipsparseHandle_t handle = nullptr; if (handle == nullptr) { @@ -439,6 +534,31 @@ static ffi::Error csrMatmulRocm(hipStream_t stream, return err; } +static ffi::Error csrMatmulRocm(hipStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, + ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base, double alpha) { + return csrMatmulRocmImpl(stream, scratch, rowptr, colind, nzval, dense, + /*acc=*/nullptr, out, m, n, transpose, index_base, + alpha, /*beta_v=*/0.0); +} + +static ffi::Error csrMatmulAccRocm(hipStream_t stream, + ffi::ScratchAllocator scratch, + ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, + ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::AnyBuffer acc, + ffi::Result out, int64_t m, + int64_t n, int64_t transpose, + int64_t index_base, double alpha, + double beta) { + return csrMatmulRocmImpl(stream, scratch, rowptr, colind, nzval, dense, &acc, + out, m, n, transpose, index_base, alpha, beta); +} + XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerROCM, csrMatmulRocm, xla::ffi::Ffi::Bind() .Ctx>() @@ -451,7 +571,25 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerROCM, csrMatmulRocm, .Attr("m") .Attr("n") .Attr("transpose") - .Attr("index_base")); + .Attr("index_base") + .Attr("alpha")); + +XLA_FFI_DEFINE_HANDLER(csrMatmulAccHandlerROCM, csrMatmulAccRocm, + xla::ffi::Ffi::Bind() + .Ctx>() + .Ctx() + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Arg() // acc (aliased to out) + .Ret() // out + .Attr("m") + .Attr("n") + .Attr("transpose") + .Attr("index_base") + .Attr("alpha") + .Attr("beta")); #endif // REACTANT_ROCM void registerReactantXLAInternalFFI() { @@ -462,10 +600,14 @@ void registerReactantXLAInternalFFI() { "CUDA", juliaCallbackHandlerCUDA); XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_csr_matmul", "CUDA", csrMatmulHandlerCUDA); + XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_csr_matmul_acc", + "CUDA", csrMatmulAccHandlerCUDA); #endif #if defined(REACTANT_ROCM) XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_csr_matmul", "ROCM", csrMatmulHandlerROCM); + XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "reactant_csr_matmul_acc", + "ROCM", csrMatmulAccHandlerROCM); #endif } diff --git a/src/SparseTensors.jl b/src/SparseTensors.jl index aea4d7b24d..aae9a14800 100644 --- a/src/SparseTensors.jl +++ b/src/SparseTensors.jl @@ -4,10 +4,12 @@ # never swept into the dense `AnyTracedRArray` overloads) holding the three CSR # buffers. At trace time `A * x` / `A * B` / `mul!` emit a # `sparse_tensor.assemble` producing a `tensor` -# value consumed by a `stablehlo.dot_general`. The Enzyme-JAX `lower-sparse-csr` -# pass rewrites that pair into a `stablehlo.custom_call @reactant_csr_matmul` -# handled by cuSPARSE/hipSPARSE (see deps/ReactantExtra/xla_ffi.cpp) before -# XLA sees the sparse-encoded types. +# value consumed by an `enzymexla.sparse.spmm` (alpha * A * B + beta * C). The +# Enzyme-JAX `lower-sparse-csr` pass rewrites that pair into a +# `stablehlo.custom_call @reactant_csr_matmul[_acc]` handled by +# cuSPARSE/hipSPARSE (see deps/ReactantExtra/xla_ffi.cpp) before XLA sees the +# sparse-encoded types; constant alpha/beta (e.g. the scalars of a 5-arg +# `mul!`) are fused into a single library call. """ CSRMatrix{T,Ti}(m, n, rowptr, colind, nzval) @@ -108,22 +110,30 @@ function _with_nzval_eltype(::Core.Type{T}, A::TracedCSRMatrix{T2,Ti}) where {T, end """ - sparse_csr_dot(A::TracedCSRMatrix, B::TracedRArray) + sparse_csr_spmm(alpha, A::TracedCSRMatrix, B, beta, C) -Emits `sparse_tensor.assemble` + `stablehlo.dot_general` computing `A * B` (spmv -for vector `B`, spmm for matrix `B`) and returns the dense result. The emitted -pair is lowered to a library call by the Enzyme-JAX `lower-sparse-csr` pass. +Emits `sparse_tensor.assemble` + `enzymexla.sparse.spmm` computing +`alpha * A * B + beta * C` (spmv for vector `B`, spmm for matrix `B`) and +returns the dense result. The emitted pair is lowered to a library call by the +Enzyme-JAX `lower-sparse-csr` pass; when `alpha`/`beta` trace to constants the +scaling and accumulation are fused into that call. """ -function sparse_csr_dot( +function sparse_csr_spmm( + alpha::TracedRNumber{T}, A::TracedCSRMatrix{T,Ti}, - B::TracedRArray{T}; - location=Ops.mlir_stacktrace("sparse_csr_dot", @__FILE__, @__LINE__), + B::TracedRArray{T}, + beta::TracedRNumber{T}, + C::TracedRArray{T}; + location=Ops.mlir_stacktrace("sparse_csr_spmm", @__FILE__, @__LINE__), ) where {T,Ti} ndims(B) in (1, 2) || throw(ArgumentError("Only vectors and matrices can be multiplied by a CSRMatrix")) size(B, 1) == A.n || throw(DimensionMismatch("A has size $(size(A)), B has size $(size(B))")) ressize = ndims(B) == 1 ? Int[A.m] : Int[A.m, size(B, 2)] + ndims(C) == ndims(B) && size(C) == Tuple(ressize) || throw( + DimensionMismatch("C has size $(size(C)), expected $(Tuple(ressize))") + ) sparse_type = MLIR.IR.TensorType(Int[A.m, A.n], MLIR.IR.Type(T), _csr_encoding(Ti)) asm = MLIR.Dialects.sparse_tensor.assemble( @@ -133,32 +143,14 @@ function sparse_csr_dot( location, ) - ctx = MLIR.IR.current_context() - batching_dimensions = Int64[] - lhs_contracting_dimensions = Int64[1] - rhs_contracting_dimensions = Int64[0] - dot_dimension_numbers = GC.@preserve ctx batching_dimensions lhs_contracting_dimensions rhs_contracting_dimensions begin - MLIR.IR.Attribute( - MLIR.API.stablehloDotDimensionNumbersGet( - ctx, - 0, - batching_dimensions, - 0, - batching_dimensions, - 1, - lhs_contracting_dimensions, - 1, - rhs_contracting_dimensions, - ), - ) - end - res = MLIR.IR.result( - MLIR.Dialects.stablehlo.dot_general( + MLIR.Dialects.enzymexla.sparse_spmm( + alpha.mlir_data, MLIR.IR.result(asm, 1), - B.mlir_data; - result_0=MLIR.IR.TensorType(ressize, MLIR.IR.Type(T)), - dot_dimension_numbers, + B.mlir_data, + beta.mlir_data, + C.mlir_data; + output=MLIR.IR.TensorType(ressize, MLIR.IR.Type(T)), location, ), ) @@ -183,30 +175,20 @@ function LinearAlgebra.mul!( size(C, 2) == size(B, 2) || throw(DimensionMismatch("C has size $(size(C)), B has size $(size(B))")) - tmp = sparse_csr_dot(_with_nzval_eltype(T, A), B) - - β_is_zero = !(β isa TracedRNumber) && iszero(β) - α_is_one = !(α isa TracedRNumber) && isone(α) - - if α_is_one && β_is_zero - res = tmp - else - α_res = if α_is_one - tmp - else - Ops.multiply(tmp, Ops.fill(promote_to(TracedRNumber{T}, α), size(tmp))) - end - if β_is_zero - res = α_res - else - C_mat = ReactantCore.materialize_traced_array(C) - β_C = Ops.multiply(C_mat, Ops.fill(promote_to(TracedRNumber{T}, β), size(C_mat))) - res = Ops.add(α_res, β_C) - end + # Non-traced α/β become constants that `lower-sparse-csr` fuses into a + # single library call; in particular a constant β == 0 never reads C. + alpha = promote_to(TracedRNumber{T}, α) + beta = promote_to(TracedRNumber{T}, β) + + ressize = ndims(B) == 1 ? (size(A, 1),) : (size(A, 1), size(B, 2)) + C_arr = ReactantCore.materialize_traced_array(C) + if ndims(C_arr) != ndims(B) + C_arr = ReactantCore.materialize_traced_array(reshape(C_arr, ressize...)) end - if ndims(C) == 2 && size(C, 2) == 1 && ndims(res) == 1 - res = reshape(res, size(C)) + res = sparse_csr_spmm(alpha, _with_nzval_eltype(T, A), B, beta, C_arr) + if size(res) != size(C) + res = ReactantCore.materialize_traced_array(reshape(res, size(C))) end TracedUtils.set_mlir_data!(C, TracedUtils.get_mlir_data(res)) @@ -215,7 +197,15 @@ end function _sparse_mul(A::TracedCSRMatrix{T}, B::AbstractVecOrMat) where {T} T2 = Base.promote_op(*, T, unwrapped_eltype(eltype(B))) - return sparse_csr_dot(_with_nzval_eltype(T2, A), promote_to(TracedRArray{T2}, B)) + B2 = promote_to(TracedRArray{T2}, B) + ressize = ndims(B2) == 1 ? (A.m,) : (A.m, size(B2, 2)) + return sparse_csr_spmm( + promote_to(TracedRNumber{T2}, 1), + _with_nzval_eltype(T2, A), + B2, + promote_to(TracedRNumber{T2}, 0), + Ops.fill(zero(T2), ressize), + ) end Base.:*(A::TracedCSRMatrix, x::AbstractVector) = _sparse_mul(A, x) diff --git a/src/mlir/Dialects/EnzymeXLA.jl b/src/mlir/Dialects/EnzymeXLA.jl index c0c5dfebef..7e76f6b3a7 100755 --- a/src/mlir/Dialects/EnzymeXLA.jl +++ b/src/mlir/Dialects/EnzymeXLA.jl @@ -1951,6 +1951,43 @@ function math_softplus( ) end +""" +`sparse_spmm` + +output := alpha * A * B + beta * C + +where `A` is a 2-d sparse tensor (e.g. the result of a +`sparse_tensor.assemble`) and `B`, `C` and `output` are dense vectors +(spmv) or matrices (spmm) of matching shapes. `alpha` and `beta` are 0-d +tensors. +""" +function sparse_spmm( + alpha::Value, + A::Value, + B::Value, + beta::Value, + C::Value; + output::IR.Type, + location=Location(), +) + op_ty_results = IR.Type[output,] + operands = Value[alpha, A, B, beta, C] + owned_regions = Region[] + successors = Block[] + attributes = NamedAttribute[] + + return create_operation( + "enzymexla.sparse.spmm", + location; + operands, + owned_regions, + successors, + attributes, + results=op_ty_results, + result_inference=false, + ) +end + function special_sphericalbesselj( nu::Value, z::Value; res=nothing::Union{Nothing,IR.Type}, location=Location() ) diff --git a/test/core/sparse.jl b/test/core/sparse.jl index 66217cb522..91f55cd27f 100644 --- a/test/core/sparse.jl +++ b/test/core/sparse.jl @@ -37,14 +37,14 @@ end @test @filecheck begin @check "#sparse_tensor.encoding" @check "sparse_tensor.assemble" - @check "stablehlo.dot_general" + @check "enzymexla.sparse.spmm" hlo end hlo = @code_hlo optimize = :none spmm(A_ra, B_ra) @test @filecheck begin @check "sparse_tensor.assemble" - @check "stablehlo.dot_general" + @check "enzymexla.sparse.spmm" hlo end end @@ -56,11 +56,7 @@ end B_ra = Reactant.to_rarray(rand(rng, 8, 3)) C_ra = Reactant.to_rarray(rand(rng, 10)) - for hlo in ( - @code_hlo(spmv(A_ra, x_ra)), - @code_hlo(spmm(A_ra, B_ra)), - @code_hlo(spmv_mul!(C_ra, A_ra, x_ra, 2.0, 3.0)), - ) + for hlo in (@code_hlo(spmv(A_ra, x_ra)), @code_hlo(spmm(A_ra, B_ra))) @test @filecheck begin @check "stablehlo.custom_call" @check "reactant_csr_matmul" @@ -68,6 +64,32 @@ end end @test !contains(repr(hlo), "sparse_tensor") end + + # Constant α/β are fused into a single accumulating library call with C + # aliased to the output. + hlo = @code_hlo spmv_mul!(C_ra, A_ra, x_ra, 2.0, 3.0) + @test @filecheck begin + @check "stablehlo.custom_call" + @check "reactant_csr_matmul_acc" + @check "output_operand_alias" + hlo + end + @test !contains(repr(hlo), "sparse_tensor") + + # Traced (runtime) α/β fall back to explicit scaling around the plain + # product. + α_rn = Reactant.ConcreteRNumber(2.0) + β_rn = Reactant.ConcreteRNumber(3.0) + hlo = @code_hlo spmv_mul!(C_ra, A_ra, x_ra, α_rn, β_rn) + @test @filecheck begin + @check "stablehlo.custom_call" + @check "reactant_csr_matmul" + @check "stablehlo.multiply" + @check "stablehlo.add" + hlo + end + @test !contains(repr(hlo), "reactant_csr_matmul_acc") + @test !contains(repr(hlo), "sparse_tensor") end @testset "numerical correctness" begin @@ -91,6 +113,13 @@ end C_ra = Reactant.to_rarray(C) @jit spmv_mul!(C_ra, A_ra, x_ra, T(2), T(3)) @test Array(C_ra) ≈ 2 .* (A * x) .+ 3 .* C atol = 1e-5 rtol = 1e-5 + + # runtime α/β + C_ra = Reactant.to_rarray(C) + α_rn = Reactant.ConcreteRNumber(T(2)) + β_rn = Reactant.ConcreteRNumber(T(3)) + @jit spmv_mul!(C_ra, A_ra, x_ra, α_rn, β_rn) + @test Array(C_ra) ≈ 2 .* (A * x) .+ 3 .* C atol = 1e-5 rtol = 1e-5 end end end From cd135d3a4a572873c92c5a235ca2e179bbaf0635 Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Wed, 19 Aug 2026 10:37:21 +0000 Subject: [PATCH 05/11] Pass ScratchAllocator by reference in CSR matmul impls ffi::ScratchAllocator is move-only; the entry points were copying it into the shared impl. Co-Authored-By: Claude Fable 5 --- deps/ReactantExtra/xla_ffi.cpp | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/deps/ReactantExtra/xla_ffi.cpp b/deps/ReactantExtra/xla_ffi.cpp index c1991df277..f535f2f535 100644 --- a/deps/ReactantExtra/xla_ffi.cpp +++ b/deps/ReactantExtra/xla_ffi.cpp @@ -125,7 +125,7 @@ XLA_FFI_DEFINE_HANDLER( // the accumulated-into operand is aliased to out by the lowering, so it is // copied into out first if XLA did not reuse the buffer). static ffi::Error csrMatmulCudaImpl(cudaStream_t stream, - ffi::ScratchAllocator scratch, + ffi::ScratchAllocator &scratch, ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, ffi::AnyBuffer nzval, ffi::AnyBuffer dense, ffi::AnyBuffer *acc, @@ -368,7 +368,7 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulAccHandlerCUDA, csrMatmulAccCuda, // the accumulated-into operand is aliased to out by the lowering, so it is // copied into out first if XLA did not reuse the buffer). static ffi::Error csrMatmulRocmImpl(hipStream_t stream, - ffi::ScratchAllocator scratch, + ffi::ScratchAllocator &scratch, ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, ffi::AnyBuffer nzval, ffi::AnyBuffer dense, ffi::AnyBuffer *acc, From b150aa99a9a6057c4a6f425906a5d46d352c47e2 Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Thu, 20 Aug 2026 12:45:17 +0000 Subject: [PATCH 06/11] Apply clang-format and JuliaFormatter to sparse CSR sources Co-Authored-By: Claude Fable 5 --- deps/ReactantExtra/xla_ffi.cpp | 171 +++++++++++++++------------------ src/SparseTensors.jl | 5 +- 2 files changed, 77 insertions(+), 99 deletions(-) diff --git a/deps/ReactantExtra/xla_ffi.cpp b/deps/ReactantExtra/xla_ffi.cpp index f535f2f535..596d608fef 100644 --- a/deps/ReactantExtra/xla_ffi.cpp +++ b/deps/ReactantExtra/xla_ffi.cpp @@ -114,25 +114,21 @@ XLA_FFI_DEFINE_HANDLER( do { \ cusparseStatus_t status__ = (expr); \ if (status__ != CUSPARSE_STATUS_SUCCESS) { \ - return ffi::Error( \ - ffi::ErrorCode::kInternal, \ - absl::StrFormat("reactant_csr_matmul: %s failed: %s", #expr, \ - cusparseGetErrorString(status__))); \ + return ffi::Error(ffi::ErrorCode::kInternal, \ + absl::StrFormat("reactant_csr_matmul: %s failed: %s", \ + #expr, \ + cusparseGetErrorString(status__))); \ } \ } while (0) // Computes out = alpha_v * A * dense (+ beta_v * *acc when acc != nullptr; // the accumulated-into operand is aliased to out by the lowering, so it is // copied into out first if XLA did not reuse the buffer). -static ffi::Error csrMatmulCudaImpl(cudaStream_t stream, - ffi::ScratchAllocator &scratch, - ffi::AnyBuffer rowptr, - ffi::AnyBuffer colind, ffi::AnyBuffer nzval, - ffi::AnyBuffer dense, ffi::AnyBuffer *acc, - ffi::Result out, int64_t m, - int64_t n, int64_t transpose, - int64_t index_base, double alpha_v, - double beta_v) { +static ffi::Error csrMatmulCudaImpl( + cudaStream_t stream, ffi::ScratchAllocator &scratch, ffi::AnyBuffer rowptr, + ffi::AnyBuffer colind, ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::AnyBuffer *acc, ffi::Result out, int64_t m, int64_t n, + int64_t transpose, int64_t index_base, double alpha_v, double beta_v) { if (transpose != 0) { return ffi::Error( ffi::ErrorCode::kUnimplemented, @@ -153,9 +149,8 @@ static ffi::Error csrMatmulCudaImpl(cudaStream_t stream, index_type = CUSPARSE_INDEX_64I; break; default: - return ffi::Error( - ffi::ErrorCode::kInvalidArgument, - "reactant_csr_matmul: index buffers must be i32 or i64"); + return ffi::Error(ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: index buffers must be i32 or i64"); } const float alpha_f = static_cast(alpha_v), @@ -190,10 +185,10 @@ static ffi::Error csrMatmulCudaImpl(cudaStream_t stream, "reactant_csr_matmul: value, operand and result dtypes must match"); } if (acc != nullptr && acc->untyped_data() != out->untyped_data()) { - cudaError_t copy_status = cudaMemcpyAsync( - out->untyped_data(), acc->untyped_data(), - static_cast(out->element_count()) * value_bytes, - cudaMemcpyDeviceToDevice, stream); + cudaError_t copy_status = + cudaMemcpyAsync(out->untyped_data(), acc->untyped_data(), + static_cast(out->element_count()) * value_bytes, + cudaMemcpyDeviceToDevice, stream); if (copy_status != cudaSuccess) { return ffi::Error( ffi::ErrorCode::kInternal, @@ -218,15 +213,13 @@ static ffi::Error csrMatmulCudaImpl(cudaStream_t stream, index_base == 1 ? CUSPARSE_INDEX_BASE_ONE : CUSPARSE_INDEX_BASE_ZERO, value_type)); - auto with_workspace = [&](size_t buffer_size, - auto &&compute) -> ffi::Error { + auto with_workspace = [&](size_t buffer_size, auto &&compute) -> ffi::Error { void *workspace = nullptr; if (buffer_size > 0) { auto maybe_workspace = scratch.Allocate(buffer_size); if (!maybe_workspace.has_value()) { - return ffi::Error( - ffi::ErrorCode::kResourceExhausted, - "reactant_csr_matmul: failed to allocate workspace"); + return ffi::Error(ffi::ErrorCode::kResourceExhausted, + "reactant_csr_matmul: failed to allocate workspace"); } workspace = *maybe_workspace; } @@ -283,8 +276,9 @@ static ffi::Error csrMatmulCudaImpl(cudaStream_t stream, cusparseDestroyDnMat(mat_b); cusparseDestroyDnMat(mat_c); } else { - err = ffi::Error(ffi::ErrorCode::kInvalidArgument, - "reactant_csr_matmul: dense operand must have rank 1 or 2"); + err = + ffi::Error(ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: dense operand must have rank 1 or 2"); } cusparseDestroySpMat(mat_a); return err; @@ -302,15 +296,11 @@ static ffi::Error csrMatmulCuda(cudaStream_t stream, alpha, /*beta_v=*/0.0); } -static ffi::Error csrMatmulAccCuda(cudaStream_t stream, - ffi::ScratchAllocator scratch, - ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, - ffi::AnyBuffer nzval, ffi::AnyBuffer dense, - ffi::AnyBuffer acc, - ffi::Result out, int64_t m, - int64_t n, int64_t transpose, - int64_t index_base, double alpha, - double beta) { +static ffi::Error csrMatmulAccCuda( + cudaStream_t stream, ffi::ScratchAllocator scratch, ffi::AnyBuffer rowptr, + ffi::AnyBuffer colind, ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::AnyBuffer acc, ffi::Result out, int64_t m, int64_t n, + int64_t transpose, int64_t index_base, double alpha, double beta) { return csrMatmulCudaImpl(stream, scratch, rowptr, colind, nzval, dense, &acc, out, m, n, transpose, index_base, alpha, beta); } @@ -319,11 +309,11 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerCUDA, csrMatmulCuda, xla::ffi::Ffi::Bind() .Ctx>() .Ctx() - .Arg() // rowptr - .Arg() // colind - .Arg() // nzval - .Arg() // dense - .Ret() // out + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Ret() // out .Attr("m") .Attr("n") .Attr("transpose") @@ -334,12 +324,12 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulAccHandlerCUDA, csrMatmulAccCuda, xla::ffi::Ffi::Bind() .Ctx>() .Ctx() - .Arg() // rowptr - .Arg() // colind - .Arg() // nzval - .Arg() // dense - .Arg() // acc (aliased to out) - .Ret() // out + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Arg() // acc (aliased to out) + .Ret() // out .Attr("m") .Attr("n") .Attr("transpose") @@ -367,15 +357,11 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulAccHandlerCUDA, csrMatmulAccCuda, // Computes out = alpha_v * A * dense (+ beta_v * *acc when acc != nullptr; // the accumulated-into operand is aliased to out by the lowering, so it is // copied into out first if XLA did not reuse the buffer). -static ffi::Error csrMatmulRocmImpl(hipStream_t stream, - ffi::ScratchAllocator &scratch, - ffi::AnyBuffer rowptr, - ffi::AnyBuffer colind, ffi::AnyBuffer nzval, - ffi::AnyBuffer dense, ffi::AnyBuffer *acc, - ffi::Result out, int64_t m, - int64_t n, int64_t transpose, - int64_t index_base, double alpha_v, - double beta_v) { +static ffi::Error csrMatmulRocmImpl( + hipStream_t stream, ffi::ScratchAllocator &scratch, ffi::AnyBuffer rowptr, + ffi::AnyBuffer colind, ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::AnyBuffer *acc, ffi::Result out, int64_t m, int64_t n, + int64_t transpose, int64_t index_base, double alpha_v, double beta_v) { if (transpose != 0) { return ffi::Error( ffi::ErrorCode::kUnimplemented, @@ -396,9 +382,8 @@ static ffi::Error csrMatmulRocmImpl(hipStream_t stream, index_type = HIPSPARSE_INDEX_64I; break; default: - return ffi::Error( - ffi::ErrorCode::kInvalidArgument, - "reactant_csr_matmul: index buffers must be i32 or i64"); + return ffi::Error(ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: index buffers must be i32 or i64"); } const float alpha_f = static_cast(alpha_v), @@ -433,10 +418,10 @@ static ffi::Error csrMatmulRocmImpl(hipStream_t stream, "reactant_csr_matmul: value, operand and result dtypes must match"); } if (acc != nullptr && acc->untyped_data() != out->untyped_data()) { - hipError_t copy_status = hipMemcpyAsync( - out->untyped_data(), acc->untyped_data(), - static_cast(out->element_count()) * value_bytes, - hipMemcpyDeviceToDevice, stream); + hipError_t copy_status = + hipMemcpyAsync(out->untyped_data(), acc->untyped_data(), + static_cast(out->element_count()) * value_bytes, + hipMemcpyDeviceToDevice, stream); if (copy_status != hipSuccess) { return ffi::Error( ffi::ErrorCode::kInternal, @@ -461,15 +446,13 @@ static ffi::Error csrMatmulRocmImpl(hipStream_t stream, index_base == 1 ? HIPSPARSE_INDEX_BASE_ONE : HIPSPARSE_INDEX_BASE_ZERO, value_type)); - auto with_workspace = [&](size_t buffer_size, - auto &&compute) -> ffi::Error { + auto with_workspace = [&](size_t buffer_size, auto &&compute) -> ffi::Error { void *workspace = nullptr; if (buffer_size > 0) { auto maybe_workspace = scratch.Allocate(buffer_size); if (!maybe_workspace.has_value()) { - return ffi::Error( - ffi::ErrorCode::kResourceExhausted, - "reactant_csr_matmul: failed to allocate workspace"); + return ffi::Error(ffi::ErrorCode::kResourceExhausted, + "reactant_csr_matmul: failed to allocate workspace"); } workspace = *maybe_workspace; } @@ -503,10 +486,9 @@ static ffi::Error csrMatmulRocmImpl(hipStream_t stream, int64_t k = dense.dimensions()[0]; int64_t c = dense.dimensions()[1]; hipsparseDnMatDescr_t mat_b, mat_c; - REACTANT_HIPSPARSE_RET(hipsparseCreateDnMat(&mat_b, k, c, /*ld=*/k, - dense.untyped_data(), - value_type, - HIPSPARSE_ORDER_COL)); + REACTANT_HIPSPARSE_RET( + hipsparseCreateDnMat(&mat_b, k, c, /*ld=*/k, dense.untyped_data(), + value_type, HIPSPARSE_ORDER_COL)); REACTANT_HIPSPARSE_RET(hipsparseCreateDnMat(&mat_c, m, c, /*ld=*/m, out->untyped_data(), value_type, HIPSPARSE_ORDER_COL)); @@ -519,16 +501,17 @@ static ffi::Error csrMatmulRocmImpl(hipStream_t stream, return with_workspace(buffer_size, [&](void *workspace) -> ffi::Error { REACTANT_HIPSPARSE_RET(hipsparseSpMM( handle, HIPSPARSE_OPERATION_NON_TRANSPOSE, - HIPSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, mat_b, beta, - mat_c, value_type, HIPSPARSE_SPMM_ALG_DEFAULT, workspace)); + HIPSPARSE_OPERATION_NON_TRANSPOSE, alpha, mat_a, mat_b, beta, mat_c, + value_type, HIPSPARSE_SPMM_ALG_DEFAULT, workspace)); return ffi::Error::Success(); }); }(); hipsparseDestroyDnMat(mat_b); hipsparseDestroyDnMat(mat_c); } else { - err = ffi::Error(ffi::ErrorCode::kInvalidArgument, - "reactant_csr_matmul: dense operand must have rank 1 or 2"); + err = + ffi::Error(ffi::ErrorCode::kInvalidArgument, + "reactant_csr_matmul: dense operand must have rank 1 or 2"); } hipsparseDestroySpMat(mat_a); return err; @@ -546,15 +529,11 @@ static ffi::Error csrMatmulRocm(hipStream_t stream, alpha, /*beta_v=*/0.0); } -static ffi::Error csrMatmulAccRocm(hipStream_t stream, - ffi::ScratchAllocator scratch, - ffi::AnyBuffer rowptr, ffi::AnyBuffer colind, - ffi::AnyBuffer nzval, ffi::AnyBuffer dense, - ffi::AnyBuffer acc, - ffi::Result out, int64_t m, - int64_t n, int64_t transpose, - int64_t index_base, double alpha, - double beta) { +static ffi::Error csrMatmulAccRocm( + hipStream_t stream, ffi::ScratchAllocator scratch, ffi::AnyBuffer rowptr, + ffi::AnyBuffer colind, ffi::AnyBuffer nzval, ffi::AnyBuffer dense, + ffi::AnyBuffer acc, ffi::Result out, int64_t m, int64_t n, + int64_t transpose, int64_t index_base, double alpha, double beta) { return csrMatmulRocmImpl(stream, scratch, rowptr, colind, nzval, dense, &acc, out, m, n, transpose, index_base, alpha, beta); } @@ -563,11 +542,11 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulHandlerROCM, csrMatmulRocm, xla::ffi::Ffi::Bind() .Ctx>() .Ctx() - .Arg() // rowptr - .Arg() // colind - .Arg() // nzval - .Arg() // dense - .Ret() // out + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Ret() // out .Attr("m") .Attr("n") .Attr("transpose") @@ -578,12 +557,12 @@ XLA_FFI_DEFINE_HANDLER(csrMatmulAccHandlerROCM, csrMatmulAccRocm, xla::ffi::Ffi::Bind() .Ctx>() .Ctx() - .Arg() // rowptr - .Arg() // colind - .Arg() // nzval - .Arg() // dense - .Arg() // acc (aliased to out) - .Ret() // out + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Arg() // acc (aliased to out) + .Ret() // out .Attr("m") .Attr("n") .Attr("transpose") diff --git a/src/SparseTensors.jl b/src/SparseTensors.jl index aae9a14800..819fd0562c 100644 --- a/src/SparseTensors.jl +++ b/src/SparseTensors.jl @@ -131,9 +131,8 @@ function sparse_csr_spmm( size(B, 1) == A.n || throw(DimensionMismatch("A has size $(size(A)), B has size $(size(B))")) ressize = ndims(B) == 1 ? Int[A.m] : Int[A.m, size(B, 2)] - ndims(C) == ndims(B) && size(C) == Tuple(ressize) || throw( - DimensionMismatch("C has size $(size(C)), expected $(Tuple(ressize))") - ) + ndims(C) == ndims(B) && size(C) == Tuple(ressize) || + throw(DimensionMismatch("C has size $(size(C)), expected $(Tuple(ressize))")) sparse_type = MLIR.IR.TensorType(Int[A.m, A.n], MLIR.IR.Type(T), _csr_encoding(Ti)) asm = MLIR.Dialects.sparse_tensor.assemble( From e87ab564f91e5084ad47d9e9c15ff417f162f7fb Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Thu, 20 Aug 2026 12:45:17 +0000 Subject: [PATCH 07/11] clang-format API.cpp Pre-existing violations on main that make the format-check-cpp job fail. Co-Authored-By: Claude Fable 5 --- deps/ReactantExtra/API.cpp | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/deps/ReactantExtra/API.cpp b/deps/ReactantExtra/API.cpp index 369a874450..8bcb7a6120 100644 --- a/deps/ReactantExtra/API.cpp +++ b/deps/ReactantExtra/API.cpp @@ -3566,7 +3566,7 @@ REACTANT_ABI void reactantXLAFree(LinkableRuntime **__restrict__ lrtP, return; auto lrt = *lrtP; void *buffer = *(void **)buffer0; - auto erased = lrt->allocations.erase((void*)buffer0); + auto erased = lrt->allocations.erase((void *)buffer0); assert(erased == 1); (void)erased; free(buffer0); @@ -4035,7 +4035,8 @@ REACTANT_ABI void ReactantCreateLLVMMod( llvm::LLVMContext **out_context, size_t *out_off, size_t *out_tmp_buf) { std::string fn(fn_str ? std::string(fn_str, fn_len) : std::string()); - llvm::StringRef source(source_str ? llvm::StringRef(source_str, source_len) : llvm::StringRef()); + llvm::StringRef source(source_str ? llvm::StringRef(source_str, source_len) + : llvm::StringRef()); std::vector> out_shapes; out_shapes.reserve(num_out_shapes); From 5dbced897692a1eb22a5ea81db21714547376a8f Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Thu, 20 Aug 2026 13:47:29 +0000 Subject: [PATCH 08/11] Document CSRMatrix and sparse_csr_spmm in the API reference Documenter aborted with :missing_docs because the two new sparse docstrings were not included in any @docs block. Co-Authored-By: Claude Fable 5 --- docs/src/api/api.md | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/docs/src/api/api.md b/docs/src/api/api.md index 2e25512a57..6831f21079 100644 --- a/docs/src/api/api.md +++ b/docs/src/api/api.md @@ -36,6 +36,13 @@ ConcreteRArray ConcreteRNumber ``` +## Sparse arrays + +```@docs +Reactant.CSRMatrix +Reactant.sparse_csr_spmm +``` + ## Inspect Generated HLO ```@docs From a65eaa5fb3c3a98aa6e6b81613097af8c9a4b36f Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Thu, 20 Aug 2026 19:42:40 +0000 Subject: [PATCH 09/11] Compare CSR round trip approximately in the sparse tests TPUs emulate Float64, so `SparseMatrixCSC(to_rarray(A)) == A` fails there by a few ULPs. Keep nnz exact and compare values with isapprox. Co-Authored-By: Claude Fable 5 --- test/core/sparse.jl | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/test/core/sparse.jl b/test/core/sparse.jl index 91f55cd27f..1d6e4dd54b 100644 --- a/test/core/sparse.jl +++ b/test/core/sparse.jl @@ -24,7 +24,10 @@ spmv_mul!(C, A, B, α, β) = LinearAlgebra.mul!(C, A, B, α, β) @test A_ra.rowptr isa ConcreteRArray @test A_ra.colind isa ConcreteRArray @test A_ra.nzval isa ConcreteRArray - @test SparseMatrixCSC(A_ra) == A + A_rt = SparseMatrixCSC(A_ra) + @test nnz(A_rt) == nnz(A) + # TPUs emulate Float64, so the value round trip is only approximate there. + @test A_rt ≈ A end @testset "sparse_tensor IR" begin From 71f7371f0af94c6ed1c98dbd4f5bdd2efff5c0e6 Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Fri, 21 Aug 2026 10:21:55 +0000 Subject: [PATCH 10/11] Follow the Enzyme-JAX rename to lower-enzymexla-sparse Co-Authored-By: Claude Fable 5 --- src/SparseTensors.jl | 6 +++--- src/compiler/Compiler.jl | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/src/SparseTensors.jl b/src/SparseTensors.jl index 819fd0562c..28cdc440fe 100644 --- a/src/SparseTensors.jl +++ b/src/SparseTensors.jl @@ -5,7 +5,7 @@ # buffers. At trace time `A * x` / `A * B` / `mul!` emit a # `sparse_tensor.assemble` producing a `tensor` # value consumed by an `enzymexla.sparse.spmm` (alpha * A * B + beta * C). The -# Enzyme-JAX `lower-sparse-csr` pass rewrites that pair into a +# Enzyme-JAX `lower-enzymexla-sparse` pass rewrites that pair into a # `stablehlo.custom_call @reactant_csr_matmul[_acc]` handled by # cuSPARSE/hipSPARSE (see deps/ReactantExtra/xla_ffi.cpp) before XLA sees the # sparse-encoded types; constant alpha/beta (e.g. the scalars of a 5-arg @@ -115,7 +115,7 @@ end Emits `sparse_tensor.assemble` + `enzymexla.sparse.spmm` computing `alpha * A * B + beta * C` (spmv for vector `B`, spmm for matrix `B`) and returns the dense result. The emitted pair is lowered to a library call by the -Enzyme-JAX `lower-sparse-csr` pass; when `alpha`/`beta` trace to constants the +Enzyme-JAX `lower-enzymexla-sparse` pass; when `alpha`/`beta` trace to constants the scaling and accumulation are fused into that call. """ function sparse_csr_spmm( @@ -174,7 +174,7 @@ function LinearAlgebra.mul!( size(C, 2) == size(B, 2) || throw(DimensionMismatch("C has size $(size(C)), B has size $(size(B))")) - # Non-traced α/β become constants that `lower-sparse-csr` fuses into a + # Non-traced α/β become constants that `lower-enzymexla-sparse` fuses into a # single library call; in particular a constant β == 0 never reads C. alpha = promote_to(TracedRNumber{T}, α) beta = promote_to(TracedRNumber{T}, β) diff --git a/src/compiler/Compiler.jl b/src/compiler/Compiler.jl index a5f72fcf22..843855a71a 100644 --- a/src/compiler/Compiler.jl +++ b/src/compiler/Compiler.jl @@ -451,7 +451,7 @@ function compile_mlir!( # else runs: XLA cannot consume sparse-encoded tensor types. With `:none` # the sparse IR is kept as-is for inspection. if compile_options.optimization_passes !== :none && has_sparse_tensor_ops(mod) - run_pass_pipeline!(mod, "lower-sparse-csr", "lower_sparse_csr") + run_pass_pipeline!(mod, "lower-enzymexla-sparse", "lower_enzymexla_sparse") end # Raise any triton kernel that might exist as a custom call From 0fc45b70de9ecbef53c9da9976a959502d39e0ad Mon Sep 17 00:00:00 2001 From: Simeon David Schaub Date: Thu, 20 Aug 2026 12:23:28 +0000 Subject: [PATCH 11/11] [do not merge] pin Enzyme-JAX to the sds/sparse_csr fork branch Local builds (deps/build_local.jl) were still fetching the upstream Enzyme-JAX pin, which lacks enzymexla.sparse.spmm and the lower-enzymexla-sparse pass, so the emitted op was unregistered at compile time. Introduces ENZYMEXLA_REPO so the download URL follows the pin. Pinned to simeonschaub/Enzyme-JAX@8508ac2a (sds/sparse_csr), which is rebased on current EnzymeAD/Enzyme-JAX main (includes 51a4cd6, the pin on Reactant main) plus the lower-enzymexla-sparse pass, enzymexla.sparse.spmm, the getSHLOLayout overload disambiguation fix, and clang-format fixes. Co-Authored-By: Claude Fable 5 --- deps/ReactantExtra/WORKSPACE | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/deps/ReactantExtra/WORKSPACE b/deps/ReactantExtra/WORKSPACE index af98b4046e..bbdd91def3 100644 --- a/deps/ReactantExtra/WORKSPACE +++ b/deps/ReactantExtra/WORKSPACE @@ -4,10 +4,15 @@ NSYNC_COMMIT = "82b118aa7ace3132e517e2c467f8732978cf4023" NSYNC_SHA256 = "" -ENZYMEXLA_COMMIT = "51a4cd6c58b1950f28b650167d1e946eec8ec0d5" +# [do not merge] sds/sparse_csr branch of simeonschaub/Enzyme-JAX (see +# ENZYMEXLA_REPO below); switch back to an EnzymeAD/Enzyme-JAX commit once +# https://github.com/EnzymeAD/Enzyme-JAX merges the sparse support. +ENZYMEXLA_COMMIT = "8508ac2af05574b59653c079a6c907e0a985cc9e" ENZYMEXLA_SHA256 = "" +ENZYMEXLA_REPO = "https://github.com/simeonschaub/Enzyme-JAX" + http_archive( name = "nsync", sha256 = NSYNC_SHA256, @@ -48,7 +53,7 @@ sed -i.bak0 "s,//:patches,@enzyme_ad//:patches,g" third_party/*/workspace.bzl ], sha256 = ENZYMEXLA_SHA256, strip_prefix = "Enzyme-JAX-" + ENZYMEXLA_COMMIT, - urls = ["https://github.com/EnzymeAD/Enzyme-JAX/archive/{commit}.tar.gz".format(commit = ENZYMEXLA_COMMIT)], + urls = ["{repo}/archive/{commit}.tar.gz".format(commit = ENZYMEXLA_COMMIT, repo = ENZYMEXLA_REPO)], ) NEW_XLA_PATCHES = []