diff --git a/scripts/run_host_overhead_control.py b/scripts/run_host_overhead_control.py index 61c8c8e2f..0bb884f5c 100644 --- a/scripts/run_host_overhead_control.py +++ b/scripts/run_host_overhead_control.py @@ -94,17 +94,18 @@ def add(): gemm_a = torch.randn((4, 48, 64), dtype=torch.float32, device=device) gemm_b = torch.randn((4, 64, 6), dtype=torch.float32, device=device) - gemm_c = torch.empty((4, 48, 6), dtype=torch.float32, device=device) + gemm_y = torch.empty((4, 48, 6), dtype=torch.float32, device=device) def gemm(): ops.gemm( gemm_a, gemm_b, + None, 1.0, 0.0, False, False, - gemm_c, + gemm_y, stream=stream, implementation_index=1, ) @@ -125,7 +126,7 @@ def gemm(): { "a_shape": [4, 48, 64], "b_shape": [4, 64, 6], - "c_shape": [4, 48, 6], + "y_shape": [4, 48, 6], "dtype": "float32", "implementation_index": 1, }, diff --git a/src/base/gemm.h b/src/base/gemm.h index c0b0fdc35..fe6df2655 100644 --- a/src/base/gemm.h +++ b/src/base/gemm.h @@ -2,6 +2,7 @@ #define INFINI_OPS_BASE_GEMM_H_ #include +#include #include #include "operator.h" @@ -10,59 +11,68 @@ namespace infini::ops { class Gemm : public Operator { public: - Gemm(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) + Gemm(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y) : alpha_{alpha.value_or(1.0)}, - beta_{beta.value_or(1.0)}, + beta_{EffectiveBeta(c, beta)}, trans_a_{static_cast(trans_a.value_or(false))}, trans_b_{static_cast(trans_b.value_or(false))}, - m_{c.size(-2)}, - n_{c.size(-1)}, + m_{y.size(-2)}, + n_{y.size(-1)}, k_{trans_a_ ? a.size(-2) : a.size(-1)}, a_type_{a.dtype()}, b_type_{b.dtype()}, - c_type_{c.dtype()}, + y_type_{y.dtype()}, a_strides_{a.strides()}, b_strides_{b.strides()}, - c_strides_{c.strides()}, + y_strides_{y.strides()}, lda_{std::max(a.stride(-2), a.stride(-1))}, ldb_{std::max(b.stride(-2), b.stride(-1))}, - ldc_{std::max(c.stride(-2), c.stride(-1))}, - batch_count_{c.strides().size() > 2 ? c.size(-3) : 1}, + ldy_{std::max(y.stride(-2), y.stride(-1))}, + batch_count_{y.strides().size() > 2 ? y.size(-3) : 1}, batch_stride_a_{a.strides().size() > 2 ? a.stride(-3) : 0}, batch_stride_b_{b.strides().size() > 2 ? b.stride(-3) : 0}, - batch_stride_c_{c.strides().size() > 2 ? c.stride(-3) : 0} { + batch_stride_y_{y.strides().size() > 2 ? y.stride(-3) : 0} { // TODO: Check constraints. } - Gemm(const Tensor a, const Tensor b, Tensor c) - : Gemm{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, c} {} + Gemm(const Tensor a, const Tensor b, Tensor y) + : Gemm{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + y} {} virtual void operator()(const Tensor a, const Tensor b, + const std::optional c, std::optional alpha, std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const = 0; + std::optional trans_b, Tensor y) const = 0; - virtual void operator()(const Tensor a, const Tensor b, Tensor c) const { + virtual void operator()(const Tensor a, const Tensor b, Tensor y) const { return operator()(a, b, std::nullopt, std::nullopt, std::nullopt, - std::nullopt, c); - } - - virtual void operator()(const Tensor a, const Tensor b, - std::optional alpha, std::optional beta, - Tensor c) const { - return operator()(a, b, alpha, beta, std::nullopt, std::nullopt, c); + std::nullopt, std::nullopt, y); } template static auto MakeReturnValue(const TensorLike& a, const TensorLike& b) { - Tensor::Shape c_shape{a.shape()[a.shape().size() - 2], + Tensor::Shape y_shape{a.shape()[a.shape().size() - 2], b.shape()[b.shape().size() - 1]}; - return TensorLike::Empty(c_shape, a.dtype(), a.device()); + return TensorLike::Empty(y_shape, a.dtype(), a.device()); } protected: + static float EffectiveBeta(const std::optional& c, + std::optional beta) { + static_cast(beta); + assert(!c && "operator Gemm C input is not supported yet"); + return 0.0F; + } + float alpha_{1.0}; float beta_{1.0}; @@ -81,19 +91,19 @@ class Gemm : public Operator { const DataType b_type_; - const DataType c_type_; + const DataType y_type_; Tensor::Strides a_strides_; Tensor::Strides b_strides_; - Tensor::Strides c_strides_; + Tensor::Strides y_strides_; Tensor::Stride lda_{0}; Tensor::Stride ldb_{0}; - Tensor::Stride ldc_{0}; + Tensor::Stride ldy_{0}; Tensor::Size batch_count_{1}; @@ -101,7 +111,7 @@ class Gemm : public Operator { Tensor::Stride batch_stride_b_{0}; - Tensor::Stride batch_stride_c_{0}; + Tensor::Stride batch_stride_y_{0}; }; } // namespace infini::ops diff --git a/src/native/ascend/ops/gemm/kernel.h b/src/native/ascend/ops/gemm/kernel.h index 34644cab7..983042ca0 100644 --- a/src/native/ascend/ops/gemm/kernel.h +++ b/src/native/ascend/ops/gemm/kernel.h @@ -15,21 +15,33 @@ namespace infini::ops { template <> class Operator : public Gemm { public: - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) - : Gemm(a, b, alpha, beta, trans_a, trans_b, c), + Operator(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y) + : Gemm(a, b, c, alpha, beta, trans_a, trans_b, y), batched_{batch_count_ > 1}, alpha_val_{alpha.value_or(1.0f)}, - beta_val_{beta.value_or(1.0f)}, - self_cache_(c), + beta_val_{0.0f}, + self_cache_(y), a_cache_(a, trans_a_), b_cache_(b, trans_b_), - out_cache_(c) { + out_cache_(y) { alpha_scalar_ = aclCreateScalar(&alpha_val_, ACL_FLOAT); beta_scalar_ = aclCreateScalar(&beta_val_, ACL_FLOAT); } + Operator(const Tensor a, const Tensor b, Tensor y) + : Operator{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + y} {} + + using Gemm::operator(); + ~Operator() { if (!ascend::IsAclRuntimeAlive()) return; @@ -43,15 +55,17 @@ class Operator : public Gemm { if (beta_scalar_) aclDestroyScalar(beta_scalar_); } - void operator()(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const override { + void operator()(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor y) const override { + static_cast(EffectiveBeta(c, beta)); auto stream = static_cast(stream_); - auto t_self = self_cache_.get(c.data()); + auto t_self = self_cache_.get(y.data()); auto t_a = a_cache_.get(const_cast(a.data())); auto t_b = b_cache_.get(const_cast(b.data())); - auto t_out = out_cache_.get(c.data()); + auto t_out = out_cache_.get(y.data()); if (!executor_) { if (batched_) { @@ -65,10 +79,10 @@ class Operator : public Gemm { } aclSetAclOpExecutorRepeatable(executor_); } else { - aclSetInputTensorAddr(executor_, 0, t_self, c.data()); + aclSetInputTensorAddr(executor_, 0, t_self, y.data()); aclSetInputTensorAddr(executor_, 1, t_a, const_cast(a.data())); aclSetInputTensorAddr(executor_, 2, t_b, const_cast(b.data())); - aclSetOutputTensorAddr(executor_, 0, t_out, c.data()); + aclSetOutputTensorAddr(executor_, 0, t_out, y.data()); } auto& arena = ascend::GetWorkspacePool().Ensure(stream, ws_size_); diff --git a/src/native/cambricon/ops/gemm/cnblas.h b/src/native/cambricon/ops/gemm/cnblas.h index 42248d4fa..4ad98039b 100644 --- a/src/native/cambricon/ops/gemm/cnblas.h +++ b/src/native/cambricon/ops/gemm/cnblas.h @@ -18,16 +18,16 @@ namespace infini::ops { template <> class Operator : public Gemm { public: - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) - : Gemm{a, b, alpha, beta, trans_a, trans_b, c}, + Operator(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y) + : Gemm{a, b, c, alpha, beta, trans_a, trans_b, y}, a_rows_{a.size(-2)}, a_cols_{a.size(-1)}, b_rows_{b.size(-2)}, b_cols_{b.size(-1)}, - c_rows_{c.size(-2)}, - c_cols_{c.size(-1)} { + y_rows_{y.size(-2)}, + y_cols_{y.size(-1)} { assert(!trans_a_ && "`trans_a` is not currently supported"); assert(!trans_b_ && "`trans_b` is not currently supported"); @@ -35,7 +35,7 @@ class Operator : public Gemm { cnnlCreateTensorDescriptor(&desc_a_); cnnlCreateTensorDescriptor(&desc_b_); - cnnlCreateTensorDescriptor(&desc_c_); + cnnlCreateTensorDescriptor(&desc_y_); cnnlCreateMatMulDescriptor(&matmul_desc_); cnnlCreateMatMulAlgo(&matmul_algo_); @@ -49,27 +49,31 @@ class Operator : public Gemm { batch_count_, batch_stride_a_); SetupTensorDescriptor(desc_b_, b_strides_, b_type_, b_rows_, b_cols_, batch_count_, batch_stride_b_); - SetupTensorDescriptor(desc_c_, c_strides_, c_type_, c_rows_, c_cols_, - batch_count_, batch_stride_c_); + SetupTensorDescriptor(desc_y_, y_strides_, y_type_, y_rows_, y_cols_, + batch_count_, batch_stride_y_); int count = 0; cnnlGetBatchMatMulExAlgoHeuristic(cnnl_handle_, matmul_desc_, desc_a_, - desc_b_, desc_c_, NULL, 1, + desc_b_, desc_y_, NULL, 1, &heuristic_result_, &count); cnrtMalloc(&default_workspace_, workspace_size_in_bytes()); } - Operator(const Tensor a, const Tensor b, Tensor c) - : Operator{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, - c} {} + Operator(const Tensor a, const Tensor b, Tensor y) + : Operator{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + y} {} - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, Tensor c) - : Operator{a, b, alpha, beta, std::nullopt, std::nullopt, c} {} + using Gemm::operator(); ~Operator() { cnrtFree(default_workspace_); - cnnlDestroyTensorDescriptor(desc_c_); + cnnlDestroyTensorDescriptor(desc_y_); cnnlDestroyTensorDescriptor(desc_b_); cnnlDestroyTensorDescriptor(desc_a_); cnnlDestroyMatMulDescriptor(matmul_desc_); @@ -78,11 +82,12 @@ class Operator : public Gemm { cnnlDestroy(cnnl_handle_); } - void operator()(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const override { + void operator()(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor y) const override { const auto& alpha_value{alpha.value_or(alpha_)}; - const auto& beta_value{beta.value_or(beta_)}; + const auto beta_value{EffectiveBeta(c, beta)}; cnnlSetQueue(cnnl_handle_, (cnrtQueue_t)stream_); @@ -92,7 +97,7 @@ class Operator : public Gemm { cnnlBatchMatMulEx(cnnl_handle_, matmul_desc_, matmul_algo_, &alpha_value, desc_a_, a.data(), desc_b_, b.data(), &beta_value, - desc_c_, c.data(), workspace, workspace_size); + desc_y_, y.data(), workspace, workspace_size); } std::size_t workspace_size_in_bytes() const override { @@ -135,7 +140,7 @@ class Operator : public Gemm { cnnlTensorDescriptor_t desc_b_; - cnnlTensorDescriptor_t desc_c_; + cnnlTensorDescriptor_t desc_y_; cnnlMatMulDescriptor_t matmul_desc_; @@ -147,7 +152,7 @@ class Operator : public Gemm { Tensor::Size b_rows_, b_cols_; - Tensor::Size c_rows_, c_cols_; + Tensor::Size y_rows_, y_cols_; // TODO: Remove the following member after default workspace mechanism has // been introduced globally. diff --git a/src/native/cpu/ops/gemm/gemm.h b/src/native/cpu/ops/gemm/gemm.h index 0eaa0a0bd..5bc1e4b89 100644 --- a/src/native/cpu/ops/gemm/gemm.h +++ b/src/native/cpu/ops/gemm/gemm.h @@ -13,29 +13,35 @@ template <> class Operator : public Gemm, Caster { public: - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) - : Gemm{a, b, alpha, beta, trans_a, trans_b, c} { + Operator(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y) + : Gemm{a, b, c, alpha, beta, trans_a, trans_b, y} { // TODO: Check constraints. } - Operator(const Tensor a, const Tensor b, Tensor c) - : Operator{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, - c} {} - - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, Tensor c) - : Operator{a, b, alpha, beta, std::nullopt, std::nullopt, c} {} - - void operator()(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const override { + Operator(const Tensor a, const Tensor b, Tensor y) + : Operator{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + y} {} + + using Gemm::operator(); + + void operator()(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor y) const override { + const auto beta_value{EffectiveBeta(c, beta)}; DispatchFunc( - c.dtype(), + y.dtype(), [&](auto tag) { using T = typename decltype(tag)::type; - Compute(a, b, alpha, beta, trans_a, trans_b, c); + Compute(a, b, alpha, beta_value, trans_a, trans_b, y); }, "`Operator::operator()`"); } @@ -44,10 +50,10 @@ class Operator : public Gemm, template void Compute(const Tensor a, const Tensor b, std::optional alpha, std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const { + std::optional trans_b, Tensor y) const { const auto* A = static_cast(a.data()); const auto* B = static_cast(b.data()); - auto* C = static_cast(c.data()); + auto* Y = static_cast(y.data()); const auto& alpha_value{alpha.value_or(alpha_)}; const auto& beta_value{beta.value_or(beta_)}; @@ -66,13 +72,13 @@ class Operator : public Gemm, Tensor::Stride stride_b_n = trans_b_value ? b_strides_[b_strides_.size() - 2] : b_strides_[b_strides_.size() - 1]; - Tensor::Stride stride_c_m = c_strides_[c_strides_.size() - 2]; - Tensor::Stride stride_c_n = c_strides_[c_strides_.size() - 1]; + Tensor::Stride stride_y_m = y_strides_[y_strides_.size() - 2]; + Tensor::Stride stride_y_n = y_strides_[y_strides_.size() - 1]; for (Tensor::Size b = 0; b < batch_count_; ++b) { const auto* A_batch = A + b * batch_stride_a_; const auto* B_batch = B + b * batch_stride_b_; - auto* C_batch = C + b * batch_stride_c_; + auto* Y_batch = Y + b * batch_stride_y_; for (Tensor::Size i = 0; i < m_; ++i) { for (Tensor::Size j = 0; j < n_; ++j) { @@ -84,9 +90,9 @@ class Operator : public Gemm, sum += a_val * b_val; } - Tensor::Size idx = i * stride_c_m + j * stride_c_n; - float c_val = beta_value == 0.0f ? 0.0f : Cast(C_batch[idx]); - C_batch[idx] = Cast(alpha_value * sum + beta_value * c_val); + Tensor::Size y_idx = i * stride_y_m + j * stride_y_n; + float y_val = beta_value == 0.0f ? 0.0f : Cast(Y_batch[y_idx]); + Y_batch[y_idx] = Cast(alpha_value * sum + beta_value * y_val); } } } diff --git a/src/native/cuda/nvidia/ops/gemm/cublaslt.h b/src/native/cuda/nvidia/ops/gemm/cublaslt.h index 4728ca04c..774fba1bd 100644 --- a/src/native/cuda/nvidia/ops/gemm/cublaslt.h +++ b/src/native/cuda/nvidia/ops/gemm/cublaslt.h @@ -18,35 +18,40 @@ namespace infini::ops { template <> class Operator : public Gemm { public: - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) - : Gemm{a, b, alpha, beta, trans_a, trans_b, c}, + Operator(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y) + : Gemm{a, b, c, alpha, beta, trans_a, trans_b, y}, a_is_col_major_{a.stride(-1) == 1}, b_is_col_major_{b.stride(-1) == 1}, - swap_a_and_b_{c.stride(-1) == 1} {} + swap_a_and_b_{y.stride(-1) == 1} {} - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, Tensor c) - : Operator{a, b, alpha, beta, std::nullopt, std::nullopt, c} {} + Operator(const Tensor a, const Tensor b, Tensor y) + : Operator{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + y} {} - Operator(const Tensor a, const Tensor b, Tensor c) - : Operator{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, - c} {} + using Gemm::operator(); // TODO: Refactor to move initialization/setup logic to the constructor // and cleanup/teardown logic to the destructor, rather than executing // everything within the computation step. // TODO: Replace the current return value checks with utility functions // (e.g., `CheckCublasLt`). - void operator()(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const override { + void operator()(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor y) const override { [[maybe_unused]] HostRangeScope host_range_backend_submit{ HostRangeLayer::kBackendSubmit}; const auto alpha_value{alpha.value_or(alpha_)}; - const auto beta_value{beta.value_or(beta_)}; + const auto beta_value{EffectiveBeta(c, beta)}; const auto trans_a_value{trans_a.value_or(trans_a_)}; const auto trans_b_value{trans_b.value_or(trans_b_)}; @@ -62,20 +67,20 @@ class Operator : public Gemm { swap_a_and_b_ ? b.dtype() : a.dtype())}; const auto b_dtype{BlasUtils::GetDataType( swap_a_and_b_ ? a.dtype() : b.dtype())}; - const auto c_dtype{ - BlasUtils::GetDataType(c.dtype())}; + const auto y_dtype{ + BlasUtils::GetDataType(y.dtype())}; const auto a_ld{static_cast(swap_a_and_b_ ? ldb_ : lda_)}; const auto b_ld{static_cast(swap_a_and_b_ ? lda_ : ldb_)}; - const auto c_ld{static_cast(ldc_)}; + const auto y_ld{static_cast(ldy_)}; const auto a_batch_stride{static_cast( swap_a_and_b_ ? batch_stride_b_ : batch_stride_a_)}; const auto b_batch_stride{static_cast( swap_a_and_b_ ? batch_stride_a_ : batch_stride_b_)}; - const auto c_batch_stride{static_cast(batch_stride_c_)}; + const auto y_batch_stride{static_cast(batch_stride_y_)}; cublasLtMatmulDesc_t op_desc{}; auto status = cublasLtMatmulDescCreate( - &op_desc, BlasUtils::GetComputeType(c.dtype()), + &op_desc, BlasUtils::GetComputeType(y.dtype()), CUDA_R_32F); assert(status == CUBLAS_STATUS_SUCCESS && "failed to create cuBLASLt matmul descriptor"); @@ -104,16 +109,16 @@ class Operator : public Gemm { assert(status == CUBLAS_STATUS_SUCCESS && "failed to create cuBLASLt B layout"); - cublasLtMatrixLayout_t c_layout{}; - status = cublasLtMatrixLayoutCreate(&c_layout, c_dtype, matmul_m, matmul_n, - c_ld); + cublasLtMatrixLayout_t y_layout{}; + status = cublasLtMatrixLayoutCreate(&y_layout, y_dtype, matmul_m, matmul_n, + y_ld); assert(status == CUBLAS_STATUS_SUCCESS && - "failed to create cuBLASLt C layout"); + "failed to create cuBLASLt Y layout"); if (batch_count_ > 1) { SetStridedBatchAttributes(a_layout, a_batch_stride); SetStridedBatchAttributes(b_layout, b_batch_stride); - SetStridedBatchAttributes(c_layout, c_batch_stride); + SetStridedBatchAttributes(y_layout, y_batch_stride); } cublasLtMatmulPreference_t preference{}; @@ -131,20 +136,20 @@ class Operator : public Gemm { cublasLtMatmulHeuristicResult_t heuristic{}; int returned_results{0}; status = cublasLtMatmulAlgoGetHeuristic( - GetHandle(), op_desc, a_layout, b_layout, c_layout, c_layout, + GetHandle(), op_desc, a_layout, b_layout, y_layout, y_layout, preference, 1, &heuristic, &returned_results); assert(status == CUBLAS_STATUS_SUCCESS && returned_results > 0 && "failed to find a cuBLASLt GEMM algorithm"); status = cublasLtMatmul( GetHandle(), op_desc, GetAlphaPtr(alpha_value), a_ptr, a_layout, b_ptr, - b_layout, GetBetaPtr(beta_value), c.data(), c_layout, c.data(), - c_layout, &heuristic.algo, workspace_, workspace_size_in_bytes_, + b_layout, GetBetaPtr(beta_value), y.data(), y_layout, y.data(), + y_layout, &heuristic.algo, workspace_, workspace_size_in_bytes_, static_cast::Stream>(stream_)); assert(status == CUBLAS_STATUS_SUCCESS && "cuBLASLt GEMM launch failed"); cublasLtMatmulPreferenceDestroy(preference); - cublasLtMatrixLayoutDestroy(c_layout); + cublasLtMatrixLayoutDestroy(y_layout); cublasLtMatrixLayoutDestroy(b_layout); cublasLtMatrixLayoutDestroy(a_layout); cublasLtMatmulDescDestroy(op_desc); diff --git a/src/native/cuda/ops/gemm/blas.h b/src/native/cuda/ops/gemm/blas.h index 2e642fa42..ec28a0e1e 100644 --- a/src/native/cuda/ops/gemm/blas.h +++ b/src/native/cuda/ops/gemm/blas.h @@ -11,39 +11,44 @@ namespace infini::ops { template class BlasGemm : public Gemm { public: - BlasGemm(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) - : Gemm{a, b, alpha, beta, trans_a, trans_b, c}, + BlasGemm(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y) + : Gemm{a, b, c, alpha, beta, trans_a, trans_b, y}, a_is_col_major_{a.stride(-1) == 1}, b_is_col_major_{b.stride(-1) == 1}, - swap_a_and_b_{c.stride(-1) == 1} { + swap_a_and_b_{y.stride(-1) == 1} { // TODO: Check constraints. } - BlasGemm(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, Tensor c) - : BlasGemm{a, b, alpha, beta, std::nullopt, std::nullopt, c} {} - - BlasGemm(const Tensor a, const Tensor b, Tensor c) - : BlasGemm{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, - c} {} - - void operator()(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const override { + BlasGemm(const Tensor a, const Tensor b, Tensor y) + : BlasGemm{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + y} {} + + using Gemm::operator(); + + void operator()(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor y) const override { Backend::BlasSetStream(GetHandle(), static_cast(stream_)); const auto& alpha_value{alpha.value_or(alpha_)}; - const auto& beta_value{beta.value_or(beta_)}; + const auto beta_value{EffectiveBeta(c, beta)}; const auto& trans_a_value{trans_a.value_or(trans_a_)}; const auto& trans_b_value{trans_b.value_or(trans_b_)}; auto op_a{GetOpA(trans_a_value, trans_b_value)}; auto op_b{GetOpB(trans_a_value, trans_b_value)}; - const void* alpha_ptr{GetAlphaPtr(alpha_value, c.dtype())}; - const void* beta_ptr{GetBetaPtr(beta_value, c.dtype())}; + const void* alpha_ptr{GetAlphaPtr(alpha_value, y.dtype())}; + const void* beta_ptr{GetBetaPtr(beta_value, y.dtype())}; Backend::BlasGemmStridedBatchedEx( GetHandle(), op_a, op_b, swap_a_and_b_ ? n_ : m_, @@ -57,10 +62,10 @@ class BlasGemm : public Gemm { BlasUtils::GetDataType(swap_a_and_b_ ? a.dtype() : b.dtype()), swap_a_and_b_ ? lda_ : ldb_, - swap_a_and_b_ ? batch_stride_a_ : batch_stride_b_, beta_ptr, c.data(), - BlasUtils::GetDataType(c.dtype()), ldc_, - batch_stride_c_, batch_count_, - BlasUtils::GetComputeType(c.dtype()), + swap_a_and_b_ ? batch_stride_a_ : batch_stride_b_, beta_ptr, y.data(), + BlasUtils::GetDataType(y.dtype()), ldy_, + batch_stride_y_, batch_count_, + BlasUtils::GetComputeType(y.dtype()), Backend::BLAS_GEMM_DEFAULT); } diff --git a/src/torch/ops/gemm/gemm.cc b/src/torch/ops/gemm/gemm.cc index 6b5ea6652..01a4a2a8f 100644 --- a/src/torch/ops/gemm/gemm.cc +++ b/src/torch/ops/gemm/gemm.cc @@ -6,32 +6,42 @@ namespace infini::ops { template Operator::Operator(const Tensor a, const Tensor b, + const std::optional c, std::optional alpha, std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) - : Gemm{a, b, alpha, beta, trans_a, trans_b, c}, + std::optional trans_b, Tensor y) + : Gemm{a, b, c, alpha, beta, trans_a, trans_b, y}, a_shape_{a.shape()}, b_shape_{b.shape()}, - c_shape_{c.shape()}, - device_index_{c.device().index()} {} + y_shape_{y.shape()}, + device_index_{y.device().index()} {} template -void Operator::operator()(const Tensor a, const Tensor b, - std::optional alpha, - std::optional beta, - std::optional trans_a, - std::optional trans_b, - Tensor c) const { +Operator::Operator(const Tensor a, const Tensor b, Tensor y) + : Operator{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + y} {} + +template +void Operator::operator()( + const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y) const { auto at_a = ToAtenTensor(const_cast(a.data()), a_shape_, a_strides_, a_type_, device_index_); auto at_b = ToAtenTensor(const_cast(b.data()), b_shape_, b_strides_, b_type_, device_index_); - auto at_c = ToAtenTensor(c.data(), c_shape_, c_strides_, c_type_, + auto at_y = ToAtenTensor(y.data(), y_shape_, y_strides_, y_type_, device_index_); auto alpha_val = alpha.value_or(alpha_); - auto beta_val = beta.value_or(beta_); + auto beta_val = EffectiveBeta(c, beta); if (trans_a.value_or(trans_a_)) { at_a = at_a.transpose(-2, -1); @@ -42,28 +52,28 @@ void Operator::operator()(const Tensor a, const Tensor b, } if (alpha_val == 0.0F) { - at_c.mul_(beta_val); + at_y.mul_(beta_val); return; } if constexpr (kDev == Device::Type::kCpu || kDev == Device::Type::kNvidia) { if (at_a.dim() == 2) { - at::addmm_out(at_c, at_c, at_a, at_b, beta_val, alpha_val); + at::addmm_out(at_y, at_y, at_a, at_b, beta_val, alpha_val); } else { - at::baddbmm_out(at_c, at_c, at_a, at_b, beta_val, alpha_val); + at::baddbmm_out(at_y, at_y, at_a, at_b, beta_val, alpha_val); } return; } auto product = at::matmul(at_a, at_b); if (beta_val == 0.0F) { - at_c.copy_(product); - at_c.mul_(alpha_val); + at_y.copy_(product); + at_y.mul_(alpha_val); return; } - at_c.mul_(beta_val); - at_c.add_(product, alpha_val); + at_y.mul_(beta_val); + at_y.add_(product, alpha_val); } template class Operator; diff --git a/src/torch/ops/gemm/gemm.h b/src/torch/ops/gemm/gemm.h index 4fd22ff36..efd34a672 100644 --- a/src/torch/ops/gemm/gemm.h +++ b/src/torch/ops/gemm/gemm.h @@ -8,22 +8,25 @@ namespace infini::ops { template class Operator : public Gemm { public: - Operator(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c); + Operator(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor y); + + Operator(const Tensor a, const Tensor b, Tensor y); using Gemm::operator(); - void operator()(const Tensor a, const Tensor b, std::optional alpha, - std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) const override; + void operator()(const Tensor a, const Tensor b, const std::optional c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor y) const override; private: Tensor::Shape a_shape_; Tensor::Shape b_shape_; - Tensor::Shape c_shape_; + Tensor::Shape y_shape_; int device_index_{0}; }; diff --git a/tests/conftest.py b/tests/conftest.py index 09122d62e..6df2e9e8a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -386,10 +386,10 @@ def _is_smoke_gemm_case(params): params, "a_shape", "b_shape", - "c_shape", + "y_shape", "a_strides", "b_strides", - "c_strides", + "y_strides", "alpha", "beta", ) diff --git a/tests/test_gemm.py b/tests/test_gemm.py index 224390d15..a79fc5678 100644 --- a/tests/test_gemm.py +++ b/tests/test_gemm.py @@ -7,7 +7,7 @@ @pytest.mark.auto_act_and_assert @pytest.mark.parametrize( - "a_shape, b_shape, c_shape, a_strides, b_strides, c_strides", + "a_shape, b_shape, y_shape, a_strides, b_strides, y_strides", ( ((1, 2048), (2048, 2048), (1, 2048), None, None, None), ((2, 4, 2048), (2, 2048, 2048), (2, 4, 2048), None, None, None), @@ -31,10 +31,10 @@ def test_gemm( a_shape, b_shape, - c_shape, + y_shape, a_strides, b_strides, - c_strides, + y_strides, alpha, beta, trans_a, @@ -91,7 +91,7 @@ def test_gemm( if trans_b: b = b.transpose(-2, -1) - c = randn_strided(c_shape, c_strides, dtype=dtype, device=device) + y = randn_strided(y_shape, y_strides, dtype=dtype, device=device) use_portable_ref = implementation_index == 2 and not ( device == "cpu" or ( @@ -103,82 +103,94 @@ def test_gemm( return Payload( lambda *args: _gemm(*args, implementation_index=implementation_index), ref, - (a, b, alpha, beta, trans_a, trans_b, c), + (a, b, None, alpha, beta, trans_a, trans_b, y), {}, rtol=rtol, atol=atol, ) -def _gemm(a, b, alpha, beta, trans_a, trans_b, c, implementation_index=0): +def _gemm(a, b, c, alpha, beta, trans_a, trans_b, y, implementation_index=0): infini.ops.gemm( a, b, + c, alpha, beta, trans_a, trans_b, - c, + y, stream=get_stream(a.device), implementation_index=implementation_index, ) - return c + return y -def _torch_gemm(a, b, alpha=1.0, beta=1.0, trans_a=False, trans_b=False, c=None): +def _torch_gemm(a, b, c, alpha, beta, trans_a, trans_b, y): if trans_a: a = a.transpose(-2, -1) if trans_b: b = b.transpose(-2, -1) + if c is None: + beta = 0.0 + y.zero_() + c = y + # PyTorch `baddbmm`/`addmm` ignores `beta` when `alpha=0.0`. if alpha == 0: - c.mul_(beta) + y.copy_(c) + y.mul_(beta) - return c + return y # Some backends (e.g. `torch_musa`) may reject `addmm`/`baddbmm(out=...)` # for certain strided outputs. Fall back to `matmul` plus fused `alpha`/`beta` # update to keep reference coverage. try: if a.ndim == 2: - return torch.addmm(c, a, b, beta=beta, alpha=alpha, out=c) + return torch.addmm(c, a, b, beta=beta, alpha=alpha, out=y) - return torch.baddbmm(c, a, b, beta=beta, alpha=alpha, out=c) + return torch.baddbmm(c, a, b, beta=beta, alpha=alpha, out=y) except RuntimeError: # Fallback for backends that don't support `addmm`/`baddbmm` (e.g. CPU `float16`/`bfloat16`): # compute in float32 and cast back. c_original = c.float() result = torch.matmul(a.float(), b.float()) - c.copy_((alpha * result + beta * c_original).to(c.dtype)) + y.copy_((alpha * result + beta * c_original).to(y.dtype)) - return c + return y -def _torch_gemm_portable( - a, b, alpha=1.0, beta=1.0, trans_a=False, trans_b=False, c=None -): +def _torch_gemm_portable(a, b, c, alpha, beta, trans_a, trans_b, y): if trans_a: a = a.transpose(-2, -1) if trans_b: b = b.transpose(-2, -1) + if c is None: + beta = 0.0 + y.zero_() + c = y + if alpha == 0: - c.mul_(beta) + y.copy_(c) + y.mul_(beta) - return c + return y product = torch.matmul(a, b) if beta == 0: - c.copy_(product) - c.mul_(alpha) + y.copy_(product) + y.mul_(alpha) - return c + return y - c.mul_(beta) - c.add_(product, alpha=alpha) + y.copy_(c) + y.mul_(beta) + y.add_(product, alpha=alpha) - return c + return y