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..66217cb522 --- /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.encoding" + @check "sparse_tensor.assemble" + @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