diff --git a/src/enzyme_ad/jax/Dialect/EnzymeXLAOps.td b/src/enzyme_ad/jax/Dialect/EnzymeXLAOps.td index 0e23fe9b52..932673c4a9 100644 --- a/src/enzyme_ad/jax/Dialect/EnzymeXLAOps.td +++ b/src/enzyme_ad/jax/Dialect/EnzymeXLAOps.td @@ -628,6 +628,38 @@ def SymmOp : EnzymeXLA_Op<"blas.symm", [Pure, SameOperandsAndResultElementType]> }]; } +def HemmOp : EnzymeXLA_Op<"blas.hemm", [Pure, SameOperandsAndResultElementType]> { + let summary = "Multiplication involving an Hermitian matrix"; + + let description = [{ + C := alpha*A*B + beta*C, or C := alpha*B*A + beta*C, where alpha and beta are scalars, A is an Hermitian matrix" + }]; + + let arguments = (ins + TensorFloat:$alpha, + HLO_Tensor:$A, + HLO_Tensor:$B, + TensorFloat:$beta, + HLO_Tensor:$C, + EnzymeXLA_LapackSideAttr:$side, + EnzymeXLA_LapackUploAttr:$uplo + ); + + let results = (outs + HLO_Tensor: $output + ); + + let assemblyFormat = [{ + $alpha `,` $A `,` $B `,` $beta `,` $C attr-dict `:` functional-type(operands, results) + }]; + + let hasVerifier = 1; + + let builders = [ + OpBuilder<(ins "Value":$A, "Value":$B, "enzymexla::LapackSide":$side, "enzymexla::LapackUplo":$uplo)> + ]; +} + def SyrkOp: EnzymeXLA_Op<"blas.syrk", [Pure, SameOperandsAndResultElementType]> { let summary = "Multiplication involving a symmetric matrix"; diff --git a/src/enzyme_ad/jax/Dialect/Ops.cpp b/src/enzyme_ad/jax/Dialect/Ops.cpp index 7a4e52f847..2f119a2a62 100644 --- a/src/enzyme_ad/jax/Dialect/Ops.cpp +++ b/src/enzyme_ad/jax/Dialect/Ops.cpp @@ -1262,6 +1262,116 @@ void GemmOp::build(OpBuilder &builder, OperationState &result, Value A, Value B, builder.getContext(), transb)); } +LogicalResult enzymexla::HemmOp::verify() { + auto alpha = op.getAlpha(); + auto a = op.getA(); + auto b = op.getB(); + auto beta = op.getBeta(); + auto c = op.getC(); + auto side = op.getSide(); + + auto type_alpha = cast(alpha.getType()); + auto type_A = cast(A.getType()); + auto type_B = cast(B.getType()); + auto type_beta = cast(beta.getType()); + auto type_C = cast(C.getType()); + + auto shape_A = type_A.getShape(); + auto shape_B = type_B.getShape(); + auto shape_C = type_C.getShape(); + + auto type_element = type_alpha.getElementType(); + auto rank = type_A.getRank(); + + auto inner_dim_A = + side == enzymexla::LapackSide::left ? rank - 1 : rank - 2; + auto inner_dim_B = + side == enzymexla::LapackSide::left ? rank - 2 : rank - 1; + + auto outer_dim_A = + side == enzymexla::LapackSide::left ? rank - 2 : rank - 1; + auto outer_dim_B = + side == enzymexla::LapackSide::left ? rank - 1 : rank - 2; + + if (type_A.getElementType() != type_element || + type_B.getElementType() != type_element || + type_C.getElementType() != type_element || + type_beta.getElementType() != type_element) { + return emitOpError("Element types of alpha, A, B and C must match"); + } + + if (!isa(type_element)) { + return emitOpError("Hemm works only with complex element type") + } + + if (type_A.getRank() != type_B.getRank() || + type_A.getRank() != type_C.getRank()) { + return emitOpError("Ranks of A, B and C must match"); + } + + if (shape_A.drop_back(2) != shape_B.drop_back(2) || + shape_A.drop_back(2) != shape_C.drop_back(2)) { + return emitOpError("Batch dimensions of A, B and C must match"); + } + + if (shape_A[inner_dim_A] != shape_B[inner_dim_B]) { + return emitOpError("Inner dimensions of A and B must match"); + } + + if (shape_A[outer_dim_A] != shape_C[rank - 2] || + shape_B[outer_dim_B] != shape_C[rank - 1]) { + return emitOpError( + "Outer dimensions of A and B must match corresponding dimensions of C"); + } + + if (getResult().getType() != type_C) { + return emitOpError("Result type must match C's type"); + } + + return success(); +} + +void HemmOp::build(OpBuilder &builder, OperationState &result, Value A, Value B, + enzymexla::LapackSide side, enzymexla::LapackUplo uplo) { + auto type_A = cast(A.getType()); + auto type_B = cast(B.getType()); + + auto element_type = type_A.getElementType(); + auto rank = type_A.getRank(); + + auto outer_dim_A = + side == enzymexla::LapackTranspose::left ? rank - 2 : rank - 1; + auto outer_dim_B = + uplo == enzymexla::LapackTranspose::right ? rank - 1 : rank - 2; + + auto shape_a = type_A.getShape(); + auto shape_b = type_B.getShape(); + SmallVector shape_c; + for (int i = 0; i < rank - 2; i++) { + shape_c.push_back(shape_a[i]); + } + shape_c.push_back(shape_a[outer_dim_A]); + shape_c.push_back(shape_b[outer_dim_B]); + + auto type_scalar = RankedTensorType::get({}, element_type); + auto alpha = stablehlo::ConstantOp::create( + builder, result.location, type_scalar, + cast(makeAttr(type_scalar, 1))); + auto beta = stablehlo::ConstantOp::create( + builder, result.location, type_scalar, + cast(makeAttr(type_scalar, 0))); + + auto type_C = RankedTensorType::get(shape_c, element_type); + auto C = + stablehlo::ConstantOp::create(builder, result.location, type_C, + cast(makeAttr(type_C, 0))); + + result.addTypes(type_C); + result.addOperands({alpha, A, B, beta, C}); + result.addAttribute("side", enzymexla::LapackSideAttr::get(builder.getContext(), side)); + result.addAttribute("uplo", enzymexla::LapackUploAttr::get(builder.getContext(), uplo)); +} + LogicalResult enzymexla::SyrkOp::verify() { auto CType = cast(getC().getType()); bool isComplex = false; diff --git a/src/enzyme_ad/jax/Passes/LowerEnzymeXLABLAS.cpp b/src/enzyme_ad/jax/Passes/LowerEnzymeXLABLAS.cpp index 1156e27004..9343eed441 100644 --- a/src/enzyme_ad/jax/Passes/LowerEnzymeXLABLAS.cpp +++ b/src/enzyme_ad/jax/Passes/LowerEnzymeXLABLAS.cpp @@ -527,6 +527,171 @@ struct SymmOpLowering : public OpRewritePattern { } }; +struct HemmOpLowering : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + std::string backend; + int64_t blasIntWidth; + + HemmOpLowering(std::string backend, int64_t blasIntWidth, + MLIRContext *context, PatternBenefit benefit = 1) + : OpRewritePattern(context, benefit), backend(backend), + blasIntWidth(blasIntWidth) {} + + LogicalResult matchAndRewrite(enzymexla::HemmOp op, + PatternRewriter &rewriter) const override { + auto AType = cast(op.getA().getType()); + auto nBatchDims = AType.getRank() - 2; + + if (nBatchDims == 0) { + if (backend == "cpu") { + return matchAndRewriteCPU(op, rewriter); + } else if (backend == "cuda") { + return matchAndRewriteCUDA(op, rewriter); + } + } + + return matchAndRewriteFallback(op, rewriter); + } + + LogicalResult matchAndRewriteCPU(enzymexla::HemmOp op, + PatternRewriter &rewriter) const { + auto ctx = op->getContext(); + LLVMTypeConverter typeConverter(ctx); + + auto alpha = op.getAlpha(); + auto a = op.getA(); + auto b = op.getB(); + auto beta = op.getBeta(); + auto c = op.getC(); + auto side = op.getSide().getBlasChar(); + auto uplo = op.getUplo().getBlasChar(); + + auto type_a = cast(a.getType()); + auto type_b = cast(b.getType()); + auto type_c = cast(c.getType()); + auto type_elem = type_a.getElementType(); + + if (!type_a || !type_b || !type_c) + return rewriter.notifyMatchFailure( + op, "expected ranked tensor types"); + + if (type_a.getRank() != 2 || type_b.getRank() > 2 || type_c.getRank() > 2) + return rewriter.notifyMatchFailure(op,"only 2D matrices supported for HemmOp"); + + std::string fn; + if (auto prefix = lapackPrecisionPrefix(elementType)) { + fn = *prefix + "hemm_"; + } else { + op->emitOpError() << "Unsupported element type: " << elementType; + return rewriter.notifyMatchFailure(op, "unsupported element type"); + } + std::string bind_fn = "enzymexla_blas_bind_" + fn; + std::string wrapper_fn = "enzymexla_blas_wrap_" + fn; + + auto type_char = rewriter.getI8Type(); + auto type_blas_int = rewriter.getIntegerType(blasIntWidth); + auto type_llvm_char = typeConverter.convertType(type_llvm_char); + auto type_llvm_blas_int = typeConverter.convertType(type_blas_int); + auto type_llvm_ptr = LLVM::LLVMPointerType::get(ctx); + auto type_llvm_void = LLVM::LLVMVoidType::get(ctx); + auto type_llvm_elem = typeConverter.convertType(type_elem); + + auto moduleOp = op->getParentOfType(); + + if (!moduleOp.lookupSymbol(blasFn)) { + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPointToStart(moduleOp.getBody()); + auto funcType = LLVM::LLVMFunctionType::get(type_llvm_void, + { + type_llvm_char, // side + type_llvm_char, // uplo + type_llvm_blas_int, // m + type_llvm_blas_int, // n + type_llvm_ptr, // alpha + type_llvm_ptr, // A + type_llvm_blas_int, // lda + type_llvm_ptr, // B + type_llvm_blas_int, // ldb + type_llvm_ptr, // beta + type_llvm_ptr, // C + type_llvm_blas_int, // ldc + }, + false + ); + LLVM::LLVMFuncOp::create(rewriter, op.getLoc(), bind_fn, funcType, + LLVM::Linkage::External); + } + + if (!moduleOp.lookupSymbol(wrapper_fn)) { + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPointToStart(moduleOp.getBody()); + + auto funcType = LLVM::LLVMFunctionType::get(type_llvm_void, {/*side*/type_llvm_ptr, /*uplo*/type_llvm_ptr, /*m*/type_llvm_ptr, /*n*/type_llvm_ptr, /*alpha*/type_llvm_ptr, /*A*/type_llvm_ptr, /*lda*/type_llvm_ptr, /*B*/type_llvm_ptr, /*ldb*/type_llvm_ptr, /*beta*/type_llvm_ptr, /*C*/type_llvm_ptr, /*ldc*/type_llvm_ptr}, false); + auto funcOp = LLVM::LLVMFuncOp::create(rewriter, op.getLoc(), wrapper_fn, funcType, LLVM::Linkage::Private); + rewriter.setInsertionPointToStart(funcOp.addEntryBlock(rewriter)); + + auto opn_side = LLVM::LoadOp(rewriter, op.getLoc(), type_llvm_char, funcOp.getArgument(0)); + auto opn_uplo = LLVM::LoadOp(rewriter, op.getLoc(), type_llvm_char, funcOp.getArgument(1)); + auto opn_m = funcOp.getArgument(2); + auto opn_n = funcOp.getArgument(3); + auto opn_alpha = funcOp.getArgument(4); + auto opn_a = funcOp.getArgument(5); + auto opn_lda = LLVM::LoadOp(rewriter, op.getLoc(), type_llvm_blas_int, funcOp.getArgument(6)); + auto opn_b = funcOp.getArgument(7); + auto opn_ldb = LLVM::LoadOp(rewriter, op.getLoc(), type_llvm_blas_int, funcOp.getArgument(8)); + auto opn_beta = funcOp.getArgument(9); + auto opn_c = LLVM::LoadOp(rewriter, op.getLoc(), type_llvm_blas_int, funcOp.getArgument(10)); + auto opn_ldc = LLVM::LoadOp(rewriter, op.getLoc(), type_llvm_blas_int, funcOp.getArgument(11)); + + LLVM::CallOp::create(rewriter, op.getLoc(), TypeRange{}, SymbolRefAttr::get(ctx, bind_fn), {opn_side, opn_uplo, opn_m, opn_n, opn_alpha, opn_a, opn_lda, opn_b, opn_ldb, opn_beta, opn_c, opn_ldc}); + LLVM::ReturnOp::create(rewriter, op.getLoc(), ValueRange{}); + } + + SmallVector isColMajorArr(12, true); + SmallVector operandRanks = {0, 0, 0, 0, 0, type_a.getRank(), 0, type_b.getRank(), 0, 0, type_c.getRank(), 0}; + SmallVector outputRanks = {type_c.getRank()}; + auto operandLayouts = getSHLOLayout(rewriter, operandRanks, isColMajorArr, 2); + auto resultLayouts = getSHLOLayout(rewriter, outputRanks, {true}, 2); + + SmallVector aliases = {stablehlo::OutputOperandAliasAttr::get(ctx, {}, 10, {})}; + + auto side_value = stablehlo::ConstantOp::create(rewriter, op.getLoc(), type_char, cast(makeAttr(type_char, side))); + auto side_value = stablehlo::ConstantOp::create(rewriter, op.getLoc(), type_char, cast(makeAttr(type_char, uplo))); + auto m = stablehlo::ConvertOp::create(rewriter, op.getLoc(), type_blas_int, stablehlo::GetDimensionSizeOp::create(rewriter, op.getLoc(), C, type_c.getRank() - 2)); + auto n = stablehlo::ConvertOp::create(rewriter, op.getLoc(), type_blas_int, stablehlo::GetDimensionSizeOp::create(rewriter, op.getLoc(), C, type_c.getRank() - 1)); + auto lda = stablehlo::ConvertOp::create(rewriter, op.getLoc(), type_blas_int, stablehlo::GetDimensionSizeOp::create(rewriter, op.getLoc(), A, type_a.getRank() - 2)); + auto ldb = stablehlo::ConvertOp::create(rewriter, op.getLoc(), type_blas_int, stablehlo::GetDimensionSizeOp::create(rewriter, op.getLoc(), B, type_b.getRank() - 2)); + auto ldc = stablehlo::ConvertOp::create(rewriter, op.getLoc(), type_blas_int, stablehlo::GetDimensionSizeOp::create(rewriter, op.getLoc(), C, type_c.getRank() - 2)); + + auto jitCall = enzymexla::JITCallOp::create( + rewriter, op.getLoc(), TypeRange{type_c}, + mlir::FlatSymbolRefAttr::get(ctx, wrapper_fn), + ValueRange{side_value, uplo_value, m, n, alpha, A, lda, B, ldb, beta, C, ldc}, + rewriter.getStringAttr(""), + /*operand_layouts=*/operandLayouts, + /*result_layouts=*/resultLayouts, + /*arg_attrs=*/nullptr, + /*res_attrs=*/nullptr, + /*output_operand_aliases=*/rewriter.getArrayAttr(aliases), + /*xla_side_effect_free=*/rewriter.getUnitAttr()); + + rewriter.replaceOp(op, jitCall); + + return success(); + }t + + LogicalResult matchAndRewriteCUDA(enzymexla::HemmOp op, + PatternRewriter &rewriter) const { + return success(); + } + + LogicalResult matchAndRewriteFallback(enzymexla::HemmOp op, + PatternRewriter &rewriter) const { + return success(); + } +}; + struct SyrkOpLowering : public OpRewritePattern { using OpRewritePattern::OpRewritePattern; @@ -1052,7 +1217,7 @@ struct LowerEnzymeXLABLASPass // We need to run SymmOp lowering first, since it conditionally lowers // to a custom call that needs us to detect potential Syrk Ops RewritePatternSet patternsSet1(context); - patternsSet1.add(backend, blasIntWidth, context); + patternsSet1.add(backend, blasIntWidth, context); if (failed(applyPatternsGreedily(getOperation(), std::move(patternsSet1), config))) { signalPassFailure(); diff --git a/src/enzyme_ad/jax/xla_ffi/cuda/blas.cc b/src/enzyme_ad/jax/xla_ffi/cuda/blas.cc index cda12f3c78..d985c3acfc 100644 --- a/src/enzyme_ad/jax/xla_ffi/cuda/blas.cc +++ b/src/enzyme_ad/jax/xla_ffi/cuda/blas.cc @@ -130,6 +130,13 @@ ffi::Error Symm(cublasHandle_t handle, cublasSideMode_t side, return ffi::Error::InvalidArgument("Unsupported type for symm"); } +template +ffi::Error Hemm(cublasHandle_t handle, cublasSideMode_t side, + cublasFillMode_t uplo, int m, int n, const T *alpha, const T *A, + int lda, const T *B, int ldb, const T *beta, T *C, int ldc) { + return ffi::Error::InvalidArgument("Unsupported type for hemm"); +} + #define SYRK_SPECIALIZATION(T, cublas_func) \ template <> \ ffi::Error Syrk(cublasHandle_t handle, cublasFillMode_t uplo, \ @@ -165,6 +172,22 @@ SYMM_SPECIALIZATION(cuDoubleComplex, cublasZsymm) #undef SYMM_SPECIALIZATION +#define HEMM_SPECIALIZATION(T, cublas_func) \ + template <> \ + ffi::Error Hemm(cublasHandle_t handle, cublasSideMode_t side, \ + cublasFillMode_t uplo, int m, int n, const T *alpha, \ + const T *A, int lda, const T *B, int ldb, const T *beta, \ + T *C, int ldc) { \ + cublasStatus_t status = cublas_func(handle, side, uplo, m, n, alpha, A, \ + lda, B, ldb, beta, C, ldc); \ + return CublasStatusToError(status, #cublas_func); \ + } + +HEMM_SPECIALIZATION(cuComplex, cublasChemm) +HEMM_SPECIALIZATION(cuDoubleComplex, cublasZhemm) + +#undef HEMM_SPECIALIZATION + } // namespace blas template @@ -476,6 +499,186 @@ XLA_FFI_DEFINE_HANDLER( .Ret() // c_out ); +template +ffi::Error HemmImpl(CUstream stream, bool side_, bool uplo_, ffi::AnyBuffer a, + ffi::AnyBuffer b, const T *alpha, const T *beta, + ffi::Result c_out) { + FFI_ASSIGN_OR_RETURN((auto [batch, rows, cols]), + SplitBatch2D(b.dimensions())); + int a_size = side_ ? cols : rows; // cols if right, rows if left + + FFI_RETURN_IF_ERROR( + CheckShape(a.dimensions(), {batch, a_size, a_size}, "a", "hemm")); + // C should have same shape as B + FFI_RETURN_IF_ERROR( + CheckShape(c_out->dimensions(), {batch, rows, cols}, "c_out", "hemm")); + + FFI_ASSIGN_OR_RETURN(auto n, MaybeCastNoOverflow(rows)); + FFI_ASSIGN_OR_RETURN(auto m, MaybeCastNoOverflow(cols)); + + // We flip uplo here because A is passed in row-major format. + // Row-major A is equivalent to A^T in column-major, and since A is + // hemmetric, this means we need to swap upper/lower triangular. + cublasFillMode_t uplo = + uplo_ ? CUBLAS_FILL_MODE_LOWER : CUBLAS_FILL_MODE_UPPER; + // We can swap side since (A*B)^T = B^T*A, where B^T is also the column-major + // interpretation of B + cublasSideMode_t side = side_ ? CUBLAS_SIDE_LEFT : CUBLAS_SIDE_RIGHT; + + const T *a_data = static_cast(a.untyped_data()); + const T *b_data = static_cast(b.untyped_data()); + T *c_out_data = static_cast(c_out->untyped_data()); + + FFI_ASSIGN_OR_RETURN(auto handle, BlasHandlePool::Borrow(stream)); + // lda is the leading dimension of a, etc. + int lda = side == CUBLAS_SIDE_LEFT ? m : n; + int ldb = m; + int ldc = m; + for (int i = 0; i < batch; ++i) { + FFI_RETURN_IF_ERROR(blas::Hemm(handle.get(), side, uplo, m, n, alpha, + a_data, lda, b_data, ldb, beta, + c_out_data, ldc)); + a_data += lda * lda; + b_data += m * n; + c_out_data += m * n; + } + return ffi::Error::Success(); +} + +template +ffi::Error HemmImpl(CUstream stream, bool side_, bool uplo_, ffi::AnyBuffer a, + ffi::AnyBuffer b, ffi::AnyBuffer c_in, const T *alpha, + const T *beta, ffi::Result c_out) { + FFI_ASSIGN_OR_RETURN((auto [batch, rows, cols]), + SplitBatch2D(b.dimensions())); + int a_size = side_ ? cols : rows; + + FFI_RETURN_IF_ERROR( + CheckShape(a.dimensions(), {batch, a_size, a_size}, "a", "hemm")); + FFI_RETURN_IF_ERROR( + CheckShape(c_out->dimensions(), {batch, rows, cols}, "c_out", "hemm")); + + T *c_data = static_cast(c_in.untyped_data()); + T *c_out_data = static_cast(c_out->untyped_data()); + + if (c_data != c_out_data) { + cudaError_t err = cudaMemcpyAsync(c_out_data, c_data, c_in.size_bytes(), + cudaMemcpyDeviceToDevice, stream); + if (err != cudaSuccess) { + return ffi::Error::InvalidArgument(absl::StrFormat( + "cudaMemcpyAsync failed: %s", cudaGetErrorString(err))); + } + } + return HemmImpl(stream, side_, uplo_, a, b, alpha, beta, c_out); +} + +template +ffi::Error +HemmImpl(CUstream stream, bool side, bool uplo, bool use_alpha_attribute, + double alpha_real, double alpha_imag, bool use_beta_attribute, + double beta_real, double beta_imag, ffi::AnyBuffer a, ffi::AnyBuffer b, + ffi::AnyBuffer c_in, ffi::AnyBuffer alpha_, ffi::AnyBuffer beta_, + ffi::Result c_out) { + T host_alpha, host_beta; + FFI_RETURN_IF_ERROR(GetHostScalar(stream, use_alpha_attribute, alpha_real, + alpha_imag, alpha_, &host_alpha)); + FFI_RETURN_IF_ERROR(GetHostScalar(stream, use_beta_attribute, beta_real, + beta_imag, beta_, &host_beta)); + return HemmImpl(stream, side, uplo, a, b, c_in, &host_alpha, &host_beta, + c_out); +} + +template +ffi::Error HemmImpl(CUstream stream, bool side, bool uplo, + bool use_alpha_attribute, double alpha_real, + double alpha_imag, ffi::AnyBuffer a, ffi::AnyBuffer b, + ffi::AnyBuffer alpha_, ffi::Result c_out) { + T host_alpha, host_beta; + FFI_RETURN_IF_ERROR(GetHostScalar(stream, use_alpha_attribute, alpha_real, + alpha_imag, alpha_, &host_alpha)); + FFI_RETURN_IF_ERROR(GetHostScalar(0.0, 0.0, &host_beta)); + return HemmImpl(stream, side, uplo, a, b, &host_alpha, &host_beta, c_out); +} + +ffi::Error HemmDispatch(CUstream stream, bool side, bool uplo, + bool use_alpha_attribute, double alpha_real, + double alpha_imag, bool use_beta_attribute, + double beta_real, double beta_imag, ffi::AnyBuffer a, + ffi::AnyBuffer b, ffi::AnyBuffer c_in, + ffi::AnyBuffer alpha_, ffi::AnyBuffer beta_, + ffi::Result c_out) { + auto dataType = c_in.element_type(); + switch (dataType) { + case ffi::C64: + return HemmImpl(stream, side, uplo, use_alpha_attribute, + alpha_real, alpha_imag, use_beta_attribute, + beta_real, beta_imag, a, b, c_in, alpha_, beta_, + c_out); + case ffi::C128: + return HemmImpl(stream, side, uplo, use_alpha_attribute, + alpha_real, alpha_imag, use_beta_attribute, + beta_real, beta_imag, a, b, c_in, alpha_, beta_, + c_out); + default: + return ffi::Error::InvalidArgument(absl::StrFormat( + "Unsupported dtype %s in hemm", absl::FormatStreamed(dataType))); + } +} + +ffi::Error HemmNoCDispatch(CUstream stream, bool side, bool uplo, + bool use_alpha_attribute, double alpha_real, + double alpha_imag, ffi::AnyBuffer a, + ffi::AnyBuffer b, ffi::AnyBuffer alpha_, + ffi::Result c_out) { + auto dataType = a.element_type(); + switch (dataType) { + case ffi::C64: + return HemmImpl(stream, side, uplo, use_alpha_attribute, + alpha_real, alpha_imag, a, b, alpha_, c_out); + case ffi::C128: + return HemmImpl(stream, side, uplo, use_alpha_attribute, + alpha_real, alpha_imag, a, b, alpha_, c_out); + default: + return ffi::Error::InvalidArgument(absl::StrFormat( + "Unsupported dtype %s in hemm", absl::FormatStreamed(dataType))); + } +} + +XLA_FFI_DEFINE_HANDLER( + HemmFfi, HemmDispatch, + xla::ffi::Ffi::Bind() + .Ctx>() + .Attr("side") // side + .Attr("uplo") // uplo + .Attr("use_alpha_attribute") // use_alpha_attribute + .Attr("alpha_real") // alpha_real + .Attr("alpha_imag") // alpha_imag + .Attr("use_beta_attribute") // use_beta_attribute + .Attr("beta_real") // beta_real + .Attr("beta_imag") // beta_imag + .Arg() // a + .Arg() // b + .Arg() // c_in + .Arg() // alpha + .Arg() // beta + .Ret() // c_out +); + +XLA_FFI_DEFINE_HANDLER( + HemmNoCFfi, HemmNoCDispatch, + xla::ffi::Ffi::Bind() + .Ctx>() + .Attr("side") // side + .Attr("uplo") // uplo + .Attr("use_alpha_attribute") // use_alpha_attribute + .Attr("alpha_real") // alpha_real + .Attr("alpha_imag") // alpha_imag + .Arg() // a + .Arg() // b + .Arg() // alpha + .Ret() // c_out +); + #undef SOLVER_BLAS_DISPATCH_IMPL void registerEnzymeJaXXLACudaBlasFFI() { @@ -489,6 +692,11 @@ void registerEnzymeJaXXLACudaBlasFFI() { XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), "enzymejax_cublas_symm_no_c_ffi", "CUDA", SymmNoCFfi); + XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), + "enzymejax_cublas_hemm_ffi", "CUDA", HemmFfi); + XLA_FFI_REGISTER_HANDLER(xla::ffi::GetXlaFfiApi(), + "enzymejax_cublas_hemm_no_c_ffi", "CUDA", + HemmNoCFfi); } } // namespace ffi_internal