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); 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/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 = [] diff --git a/deps/ReactantExtra/xla_ffi.cpp b/deps/ReactantExtra/xla_ffi.cpp index 85f4189a0c..596d608fef 100644 --- a/deps/ReactantExtra/xla_ffi.cpp +++ b/deps/ReactantExtra/xla_ffi.cpp @@ -91,12 +91,502 @@ XLA_FFI_DEFINE_HANDLER( "callback_ptr")); #endif +// ============================================================================ +// CSR sparse matrix products (spmv / spmm) via cuSPARSE / hipSPARSE. +// +// 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) +#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) + +// 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, + "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"); + } + + 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; + 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() || + (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) { + 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; +} + +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>() + .Ctx() + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Ret() // out + .Attr("m") + .Attr("n") + .Attr("transpose") + .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) +#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) + +// 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, + "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"); + } + + 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; + 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() || + (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) { + 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; +} + +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>() + .Ctx() + .Arg() // rowptr + .Arg() // colind + .Arg() // nzval + .Arg() // dense + .Ret() // out + .Attr("m") + .Attr("n") + .Attr("transpose") + .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() { 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); + 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/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 diff --git a/ext/ReactantSparseArraysExt/CSR.jl b/ext/ReactantSparseArraysExt/CSR.jl new file mode 100644 index 0000000000..47553f50e8 --- /dev/null +++ b/ext/ReactantSparseArraysExt/CSR.jl @@ -0,0 +1,27 @@ +# 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 + # `sparse_tensor` positions/coordinates are 0-based + return Reactant.CSRMatrix{T,Ti,Vector{T},Vector{Ti}}( + size(A, 1), size(A, 2), At.colptr .- one(Ti), At.rowval .- one(Ti), 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 (0-based) CSR buffers of A are the CSC representation of Aᵀ + At = SparseMatrixCSC{T,Ti}( + 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/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..28cdc440fe --- /dev/null +++ b/src/SparseTensors.jl @@ -0,0 +1,219 @@ +# 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 an `enzymexla.sparse.spmm` (alpha * A * B + beta * C). The +# 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 +# `mul!`) are fused into a single library call. + +""" + 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 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 +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_spmm(alpha, A::TracedCSRMatrix, B, beta, C) + +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-enzymexla-sparse` pass; when `alpha`/`beta` trace to constants the +scaling and accumulation are fused into that call. +""" +function sparse_csr_spmm( + alpha::TracedRNumber{T}, + A::TracedCSRMatrix{T,Ti}, + 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( + MLIR.IR.Value[A.rowptr.mlir_data, A.colind.mlir_data], + A.nzval.mlir_data; + result=sparse_type, + location, + ) + + res = MLIR.IR.result( + MLIR.Dialects.enzymexla.sparse_spmm( + alpha.mlir_data, + MLIR.IR.result(asm, 1), + B.mlir_data, + beta.mlir_data, + C.mlir_data; + output=MLIR.IR.TensorType(ressize, MLIR.IR.Type(T)), + 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))")) + + # 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}, β) + + 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 + + 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)) + return C +end + +function _sparse_mul(A::TracedCSRMatrix{T}, B::AbstractVecOrMat) where {T} + T2 = Base.promote_op(*, T, unwrapped_eltype(eltype(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) +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..843855a71a 100644 --- a/src/compiler/Compiler.jl +++ b/src/compiler/Compiler.jl @@ -27,6 +27,23 @@ include("OptimizationPasses.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}}}() @@ -430,6 +447,13 @@ function compile_mlir!( legal_to_run_shardy_passes = compile_options.optimization_passes === :all + # 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-enzymexla-sparse", "lower_enzymexla_sparse") + 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/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/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..1d6e4dd54b --- /dev/null +++ b/test/core/sparse.jl @@ -0,0 +1,128 @@ +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 + 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 + 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 "enzymexla.sparse.spmm" + hlo + end + + hlo = @code_hlo optimize = :none spmm(A_ra, B_ra) + @test @filecheck begin + @check "sparse_tensor.assemble" + @check "enzymexla.sparse.spmm" + 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))) + @test @filecheck begin + @check "stablehlo.custom_call" + @check "reactant_csr_matmul" + hlo + 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 + 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 + + # 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