From b33fccd9ed5b5b61d2993832173690ad90348db4 Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Tue, 4 Aug 2026 17:11:35 +0800 Subject: [PATCH 1/2] fix!: align Gemm public signature BREAKING CHANGE: Gemm now accepts optional C before attributes and writes Y; non-null C is not implemented yet. --- scripts/run_host_overhead_control.py | 1 + src/base/gemm.h | 58 ++++++++++++---------- src/native/ascend/ops/gemm/kernel.h | 31 +++++++++--- src/native/cambricon/ops/gemm/cnblas.h | 30 ++++++----- src/native/cpu/ops/gemm/gemm.h | 31 +++++++----- src/native/cuda/nvidia/ops/gemm/cublaslt.h | 32 +++++++----- src/native/cuda/ops/gemm/blas.h | 32 +++++++----- src/torch/ops/gemm/gemm.cc | 26 +++++++--- src/torch/ops/gemm/gemm.h | 16 +++--- tests/test_gemm.py | 5 +- 10 files changed, 164 insertions(+), 98 deletions(-) diff --git a/scripts/run_host_overhead_control.py b/scripts/run_host_overhead_control.py index 61c8c8e2f..4553ac009 100644 --- a/scripts/run_host_overhead_control.py +++ b/scripts/run_host_overhead_control.py @@ -100,6 +100,7 @@ def gemm(): ops.gemm( gemm_a, gemm_b, + None, 1.0, 0.0, False, diff --git a/src/base/gemm.h b/src/base/gemm.h index c0b0fdc35..4a1c332ca 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,49 +11,49 @@ 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()}, + c_type_{y.dtype()}, a_strides_{a.strides()}, b_strides_{b.strides()}, - c_strides_{c.strides()}, + c_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}, + ldc_{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} { - // TODO: Check constraints. - } - - Gemm(const Tensor a, const Tensor b, Tensor c) - : Gemm{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, c} {} + batch_stride_c_{y.strides().size() > 2 ? y.stride(-3) : 0} {} + + 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 @@ -63,6 +64,13 @@ class Gemm : public Operator { } 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}; diff --git a/src/native/ascend/ops/gemm/kernel.h b/src/native/ascend/ops/gemm/kernel.h index 34644cab7..a80d462b0 100644 --- a/src/native/ascend/ops/gemm/kernel.h +++ b/src/native/ascend/ops/gemm/kernel.h @@ -15,13 +15,13 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor c) + : Gemm(a, b, input_c, alpha, beta, trans_a, trans_b, c), batched_{batch_count_ > 1}, alpha_val_{alpha.value_or(1.0f)}, - beta_val_{beta.value_or(1.0f)}, + beta_val_{0.0f}, self_cache_(c), a_cache_(a, trans_a_), b_cache_(b, trans_b_), @@ -30,6 +30,18 @@ class Operator : public Gemm { beta_scalar_ = aclCreateScalar(&beta_val_, ACL_FLOAT); } + Operator(const Tensor a, const Tensor b, Tensor c) + : Operator{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + c} {} + + using Gemm::operator(); + ~Operator() { if (!ascend::IsAclRuntimeAlive()) return; @@ -43,9 +55,12 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor c) const override { + static_cast(EffectiveBeta(input_c, beta)); auto stream = static_cast(stream_); auto t_self = self_cache_.get(c.data()); diff --git a/src/native/cambricon/ops/gemm/cnblas.h b/src/native/cambricon/ops/gemm/cnblas.h index 42248d4fa..f8db581c7 100644 --- a/src/native/cambricon/ops/gemm/cnblas.h +++ b/src/native/cambricon/ops/gemm/cnblas.h @@ -18,10 +18,10 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor c) + : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c}, a_rows_{a.size(-2)}, a_cols_{a.size(-1)}, b_rows_{b.size(-2)}, @@ -60,12 +60,16 @@ class Operator : public Gemm { } Operator(const Tensor a, const Tensor b, Tensor c) - : Operator{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, + : Operator{a, + b, + std::nullopt, + 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} {} + using Gemm::operator(); ~Operator() { cnrtFree(default_workspace_); @@ -78,11 +82,13 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor c) const override { const auto& alpha_value{alpha.value_or(alpha_)}; - const auto& beta_value{beta.value_or(beta_)}; + const auto beta_value{EffectiveBeta(input_c, beta)}; cnnlSetQueue(cnnl_handle_, (cnrtQueue_t)stream_); diff --git a/src/native/cpu/ops/gemm/gemm.h b/src/native/cpu/ops/gemm/gemm.h index 0eaa0a0bd..80702ac7f 100644 --- a/src/native/cpu/ops/gemm/gemm.h +++ b/src/native/cpu/ops/gemm/gemm.h @@ -13,29 +13,36 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor c) + : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c} { // TODO: Check constraints. } Operator(const Tensor a, const Tensor b, Tensor c) - : Operator{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, + : Operator{a, + b, + std::nullopt, + 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} {} + 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor c) const override { + const auto beta_value{EffectiveBeta(input_c, beta)}; DispatchFunc( c.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, c); }, "`Operator::operator()`"); } diff --git a/src/native/cuda/nvidia/ops/gemm/cublaslt.h b/src/native/cuda/nvidia/ops/gemm/cublaslt.h index 4728ca04c..adcbc1ab1 100644 --- a/src/native/cuda/nvidia/ops/gemm/cublaslt.h +++ b/src/native/cuda/nvidia/ops/gemm/cublaslt.h @@ -18,35 +18,41 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor c) + : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c}, a_is_col_major_{a.stride(-1) == 1}, b_is_col_major_{b.stride(-1) == 1}, swap_a_and_b_{c.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 c) - : Operator{a, b, std::nullopt, std::nullopt, std::nullopt, std::nullopt, + : Operator{a, + b, + std::nullopt, + 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor c) 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(input_c, beta)}; const auto trans_a_value{trans_a.value_or(trans_a_)}; const auto trans_b_value{trans_b.value_or(trans_b_)}; diff --git a/src/native/cuda/ops/gemm/blas.h b/src/native/cuda/ops/gemm/blas.h index 2e642fa42..8ed7ef095 100644 --- a/src/native/cuda/ops/gemm/blas.h +++ b/src/native/cuda/ops/gemm/blas.h @@ -11,32 +11,38 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor c) + : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c}, a_is_col_major_{a.stride(-1) == 1}, b_is_col_major_{b.stride(-1) == 1}, swap_a_and_b_{c.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, + : BlasGemm{a, + b, + std::nullopt, + 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 { + using Gemm::operator(); + + void operator()(const Tensor a, const Tensor b, + const std::optional input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor c) 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(input_c, beta)}; const auto& trans_a_value{trans_a.value_or(trans_a_)}; const auto& trans_b_value{trans_b.value_or(trans_b_)}; diff --git a/src/torch/ops/gemm/gemm.cc b/src/torch/ops/gemm/gemm.cc index 6b5ea6652..9aaf105ea 100644 --- a/src/torch/ops/gemm/gemm.cc +++ b/src/torch/ops/gemm/gemm.cc @@ -6,23 +6,33 @@ namespace infini::ops { template Operator::Operator(const Tensor a, const Tensor b, + const std::optional input_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}, + : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c}, a_shape_{a.shape()}, b_shape_{b.shape()}, c_shape_{c.shape()}, device_index_{c.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 c) + : Operator{a, + b, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt, + c} {} + +template +void Operator::operator()( + const Tensor a, const Tensor b, const std::optional input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor c) 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_, @@ -31,7 +41,7 @@ void Operator::operator()(const Tensor a, const Tensor b, device_index_); auto alpha_val = alpha.value_or(alpha_); - auto beta_val = beta.value_or(beta_); + auto beta_val = EffectiveBeta(input_c, beta); if (trans_a.value_or(trans_a_)) { at_a = at_a.transpose(-2, -1); diff --git a/src/torch/ops/gemm/gemm.h b/src/torch/ops/gemm/gemm.h index 4fd22ff36..b74505ddb 100644 --- a/src/torch/ops/gemm/gemm.h +++ b/src/torch/ops/gemm/gemm.h @@ -8,15 +8,19 @@ 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, Tensor c); + + Operator(const Tensor a, const Tensor b, Tensor c); 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 input_c, + std::optional alpha, std::optional beta, + std::optional trans_a, std::optional trans_b, + Tensor c) const override; private: Tensor::Shape a_shape_; diff --git a/tests/test_gemm.py b/tests/test_gemm.py index 224390d15..31df00c0e 100644 --- a/tests/test_gemm.py +++ b/tests/test_gemm.py @@ -102,7 +102,9 @@ def test_gemm( return Payload( lambda *args: _gemm(*args, implementation_index=implementation_index), - ref, + lambda a, b, alpha, _beta, trans_a, trans_b, c: ref( + a, b, alpha, 0.0, trans_a, trans_b, c + ), (a, b, alpha, beta, trans_a, trans_b, c), {}, rtol=rtol, @@ -114,6 +116,7 @@ def _gemm(a, b, alpha, beta, trans_a, trans_b, c, implementation_index=0): infini.ops.gemm( a, b, + None, alpha, beta, trans_a, From 256fcadb401cefeb7c214657971f97ab65544396 Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Tue, 4 Aug 2026 17:34:03 +0800 Subject: [PATCH 2/2] fix: support optional C in Gemm --- src/base/gemm.h | 64 ++++++++++--- src/native/ascend/ops/gemm/kernel.h | 33 +++---- src/native/cambricon/ops/gemm/cnblas.h | 105 ++++++++++++++++----- src/native/cpu/ops/gemm/gemm.h | 53 ++++++----- src/native/cuda/iluvatar/ops/gemm/cublas.h | 1 + src/native/cuda/metax/ops/gemm/mcblas.h | 1 + src/native/cuda/moore/ops/gemm/mublas.h | 1 + src/native/cuda/nvidia/ops/gemm/cublas.h | 1 + src/native/cuda/nvidia/ops/gemm/cublaslt.h | 53 ++++++----- src/native/cuda/ops/gemm/blas.h | 40 ++++---- src/torch/ops/gemm/gemm.cc | 51 ++++++---- src/torch/ops/gemm/gemm.h | 13 +-- tests/test_gemm.py | 100 +++++++++++++++++--- 13 files changed, 354 insertions(+), 162 deletions(-) diff --git a/src/base/gemm.h b/src/base/gemm.h index 4a1c332ca..71becd9e7 100644 --- a/src/base/gemm.h +++ b/src/base/gemm.h @@ -15,7 +15,7 @@ class Gemm : public Operator { std::optional alpha, std::optional beta, std::optional trans_a, std::optional trans_b, Tensor y) : alpha_{alpha.value_or(1.0)}, - beta_{EffectiveBeta(c, beta)}, + beta_{beta.value_or(1.0)}, trans_a_{static_cast(trans_a.value_or(false))}, trans_b_{static_cast(trans_b.value_or(false))}, m_{y.size(-2)}, @@ -23,17 +23,28 @@ class Gemm : public Operator { k_{trans_a_ ? a.size(-2) : a.size(-1)}, a_type_{a.dtype()}, b_type_{b.dtype()}, - c_type_{y.dtype()}, + y_type_{y.dtype()}, a_strides_{a.strides()}, b_strides_{b.strides()}, - c_strides_{y.strides()}, + c_shape_{c ? Tensor::Shape{c->shape()} : Tensor::Shape{}}, + c_strides_{c ? Tensor::Strides{c->strides()} : Tensor::Strides{}}, + c_broadcast_strides_{c ? BroadcastStrides(*c, y) + : Tensor::Strides(y.ndim(), 0)}, + y_shape_{y.shape()}, + 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(y.stride(-2), y.stride(-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_{y.strides().size() > 2 ? y.stride(-3) : 0} {} + batch_stride_y_{y.strides().size() > 2 ? y.stride(-3) : 0} { + assert(a.dtype() == b.dtype() && a.dtype() == y.dtype() && + (!c || c->dtype() == y.dtype()) && + "operator `Gemm` requires A, B, C, and Y to have the same dtype"); + assert((!c || c->data() != y.data()) && + "operator `Gemm` does not support C/Y aliasing"); + } Gemm(const Tensor a, const Tensor b, Tensor y) : Gemm{a, @@ -58,17 +69,32 @@ class Gemm : public Operator { 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; + static Tensor::Strides BroadcastStrides(const Tensor input, + const Tensor out) { + assert(input.ndim() <= out.ndim() && + "operator `Gemm` C rank must not exceed Y rank"); + Tensor::Strides strides(out.ndim(), 0); + const auto offset = out.ndim() - input.ndim(); + + for (Tensor::Size i = 0; i < input.ndim(); ++i) { + const auto out_dim = i + offset; + assert((input.size(i) == 1 || input.size(i) == out.size(out_dim)) && + "operator `Gemm` C shape is not broadcast-compatible with Y"); + strides[out_dim] = input.size(i) == 1 ? 0 : input.stride(i); + } + + return strides; + } + + float EffectiveBeta(const std::optional& c, + std::optional beta) const { + return c ? beta.value_or(beta_) : 0.0F; } float alpha_{1.0}; @@ -89,19 +115,27 @@ 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::Shape c_shape_; + Tensor::Strides c_strides_; + Tensor::Strides c_broadcast_strides_; + + Tensor::Shape y_shape_; + + 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}; @@ -109,7 +143,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 a80d462b0..eb728f3ab 100644 --- a/src/native/ascend/ops/gemm/kernel.h +++ b/src/native/ascend/ops/gemm/kernel.h @@ -15,22 +15,22 @@ namespace infini::ops { template <> class Operator : public Gemm { public: - Operator(const Tensor a, const Tensor b, const std::optional input_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 c) - : Gemm(a, b, input_c, alpha, beta, trans_a, trans_b, c), + 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_{0.0f}, - self_cache_(c), + beta_val_{c ? beta.value_or(1.0f) : 0.0f}, + self_cache_(c ? *c : 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 c) + Operator(const Tensor a, const Tensor b, Tensor y) : Operator{a, b, std::nullopt, @@ -38,7 +38,7 @@ class Operator : public Gemm { std::nullopt, std::nullopt, std::nullopt, - c} {} + y} {} using Gemm::operator(); @@ -55,18 +55,19 @@ class Operator : public Gemm { if (beta_scalar_) aclDestroyScalar(beta_scalar_); } - void operator()(const Tensor a, const Tensor b, - const std::optional input_c, + 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 override { - static_cast(EffectiveBeta(input_c, beta)); + Tensor y) const override { auto stream = static_cast(stream_); - auto t_self = self_cache_.get(c.data()); + const auto self = c ? *c : y; + auto self_ptr = const_cast(self.data()); + auto y_ptr = y.data(); + auto t_self = self_cache_.get(self_ptr); 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_ptr); if (!executor_) { if (batched_) { @@ -80,10 +81,10 @@ class Operator : public Gemm { } aclSetAclOpExecutorRepeatable(executor_); } else { - aclSetInputTensorAddr(executor_, 0, t_self, c.data()); + aclSetInputTensorAddr(executor_, 0, t_self, self_ptr); 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_ptr); } 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 f8db581c7..8e571f0f8 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, const std::optional input_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 c) - : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c}, + 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,13 @@ class Operator : public Gemm { cnnlCreateTensorDescriptor(&desc_a_); cnnlCreateTensorDescriptor(&desc_b_); - cnnlCreateTensorDescriptor(&desc_c_); + cnnlCreateTensorDescriptor(&desc_y_); + if (c) { + cnnlCreateTensorDescriptor(&desc_c_); + cnnlCreateOpTensorDescriptor(&op_tensor_desc_); + cnnlSetOpTensorDescriptor(op_tensor_desc_, CNNL_OP_TENSOR_ADD, + CNNL_DTYPE_FLOAT, CNNL_NOT_PROPAGATE_NAN); + } cnnlCreateMatMulDescriptor(&matmul_desc_); cnnlCreateMatMulAlgo(&matmul_algo_); @@ -49,17 +55,30 @@ 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_); + if (c) { + SetupBroadcastTensorDescriptor(desc_c_, c_shape_, c_strides_, y_type_); + } + 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); + cnnlGetBatchMatMulExHeuristicResult(heuristic_result_, matmul_algo_, + &workspace_size_); + if (c) { + std::size_t add_workspace_size = 0; + cnnlGetOpTensorWorkspaceSize(cnnl_handle_, desc_c_, desc_y_, desc_y_, + &add_workspace_size); + workspace_size_ = std::max(workspace_size_, add_workspace_size); + } + cnrtMalloc(&default_workspace_, workspace_size_in_bytes()); } - Operator(const Tensor a, const Tensor b, Tensor c) + Operator(const Tensor a, const Tensor b, Tensor y) : Operator{a, b, std::nullopt, @@ -67,13 +86,15 @@ class Operator : public Gemm { std::nullopt, std::nullopt, std::nullopt, - c} {} + y} {} using Gemm::operator(); ~Operator() { cnrtFree(default_workspace_); - cnnlDestroyTensorDescriptor(desc_c_); + if (op_tensor_desc_) cnnlDestroyOpTensorDescriptor(op_tensor_desc_); + if (desc_c_) cnnlDestroyTensorDescriptor(desc_c_); + cnnlDestroyTensorDescriptor(desc_y_); cnnlDestroyTensorDescriptor(desc_b_); cnnlDestroyTensorDescriptor(desc_a_); cnnlDestroyMatMulDescriptor(matmul_desc_); @@ -82,13 +103,13 @@ class Operator : public Gemm { cnnlDestroy(cnnl_handle_); } - void operator()(const Tensor a, const Tensor b, - const std::optional input_c, + 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 override { + Tensor y) const override { const auto& alpha_value{alpha.value_or(alpha_)}; - const auto beta_value{EffectiveBeta(input_c, beta)}; + const auto& beta_value{EffectiveBeta(c, beta)}; + constexpr float gemm_beta = 0.0F; cnnlSetQueue(cnnl_handle_, (cnrtQueue_t)stream_); @@ -97,16 +118,20 @@ class Operator : public Gemm { : workspace_size_in_bytes()}; 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_a_, a.data(), desc_b_, b.data(), &gemm_beta, desc_y_, + y.data(), workspace, workspace_size); + + if (c && beta_value != 0.0F) { + constexpr float one = 1.0F; + constexpr float zero = 0.0F; + cnnlOpTensor(cnnl_handle_, op_tensor_desc_, &beta_value, desc_c_, + c->data(), &one, desc_y_, y.data(), workspace, + workspace_size, &zero, desc_y_, y.data()); + } } std::size_t workspace_size_in_bytes() const override { - std::size_t size{0}; - - cnnlGetBatchMatMulExHeuristicResult(heuristic_result_, matmul_algo_, &size); - - return size; + return workspace_size_; } private: @@ -135,13 +160,41 @@ class Operator : public Gemm { } } + void SetupBroadcastTensorDescriptor(cnnlTensorDescriptor_t desc, + const Tensor::Shape& shape, + const Tensor::Strides& strides, + DataType dtype) { + std::vector dims; + std::vector strides_arr; + + if (shape.empty()) { + dims.push_back(1); + strides_arr.push_back(1); + } else { + dims.reserve(shape.size()); + strides_arr.reserve(strides.size()); + for (std::size_t i = 0; i < shape.size(); ++i) { + dims.push_back(static_cast(shape[i])); + strides_arr.push_back(static_cast(strides[i])); + } + } + + cnnlSetTensorDescriptorEx(desc, CNNL_LAYOUT_ARRAY, + cnnl_utils::GetDataType(dtype), dims.size(), + dims.data(), strides_arr.data()); + } + cnnlHandle_t cnnl_handle_; cnnlTensorDescriptor_t desc_a_; cnnlTensorDescriptor_t desc_b_; - cnnlTensorDescriptor_t desc_c_; + cnnlTensorDescriptor_t desc_y_; + + cnnlTensorDescriptor_t desc_c_{}; + + cnnlOpTensorDescriptor_t op_tensor_desc_{}; cnnlMatMulDescriptor_t matmul_desc_; @@ -153,7 +206,9 @@ class Operator : public Gemm { Tensor::Size b_rows_, b_cols_; - Tensor::Size c_rows_, c_cols_; + Tensor::Size y_rows_, y_cols_; + + std::size_t workspace_size_{0}; // 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 80702ac7f..999cb92f9 100644 --- a/src/native/cpu/ops/gemm/gemm.h +++ b/src/native/cpu/ops/gemm/gemm.h @@ -13,14 +13,14 @@ template <> class Operator : public Gemm, Caster { public: - Operator(const Tensor a, const Tensor b, const std::optional input_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 c) - : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c} { + 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(const Tensor a, const Tensor b, Tensor y) : Operator{a, b, std::nullopt, @@ -28,39 +28,38 @@ class Operator : public Gemm, std::nullopt, std::nullopt, std::nullopt, - c} {} + y} {} using Gemm::operator(); - void operator()(const Tensor a, const Tensor b, - const std::optional input_c, + 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 override { - const auto beta_value{EffectiveBeta(input_c, beta)}; + Tensor y) const override { DispatchFunc( - c.dtype(), + y.dtype(), [&](auto tag) { using T = typename decltype(tag)::type; - Compute(a, b, alpha, beta_value, trans_a, trans_b, c); + Compute(a, b, c, alpha, beta, trans_a, trans_b, y); }, "`Operator::operator()`"); } private: 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 { + void Compute(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 { const auto* A = static_cast(a.data()); const auto* B = static_cast(b.data()); - auto* C = static_cast(c.data()); + const auto* C = c ? static_cast(c->data()) : nullptr; + auto* Y = static_cast(y.data()); 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_)}; - Tensor::Stride stride_a_m = trans_a_value ? a_strides_[a_strides_.size() - 1] : a_strides_[a_strides_.size() - 2]; @@ -73,13 +72,19 @@ 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]; + const auto y_rank = y_strides_.size(); + const auto stride_c_batch = + y_rank > 2 ? c_broadcast_strides_[y_rank - 3] : 0; + const auto stride_c_m = c_broadcast_strides_[y_rank - 2]; + const auto stride_c_n = c_broadcast_strides_[y_rank - 1]; + const auto stride_y_m = y_strides_[y_rank - 2]; + const auto stride_y_n = y_strides_[y_rank - 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_; + const auto* C_batch = C ? C + b * stride_c_batch : nullptr; + auto* Y_batch = Y + b * batch_stride_y_; for (Tensor::Size i = 0; i < m_; ++i) { for (Tensor::Size j = 0; j < n_; ++j) { @@ -91,9 +96,11 @@ 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); + const auto c_idx = i * stride_c_m + j * stride_c_n; + const auto y_idx = i * stride_y_m + j * stride_y_n; + const float c_val = + beta_value == 0.0F ? 0.0F : Cast(C_batch[c_idx]); + Y_batch[y_idx] = Cast(alpha_value * sum + beta_value * c_val); } } } diff --git a/src/native/cuda/iluvatar/ops/gemm/cublas.h b/src/native/cuda/iluvatar/ops/gemm/cublas.h index d086cf4f4..c1b83cc77 100644 --- a/src/native/cuda/iluvatar/ops/gemm/cublas.h +++ b/src/native/cuda/iluvatar/ops/gemm/cublas.h @@ -2,6 +2,7 @@ #define INFINI_OPS_ILUVATAR_GEMM_CUBLAS_H_ #include "native/cuda/iluvatar/blas.h" +#include "native/cuda/iluvatar/ops/add/kernel.h" #include "native/cuda/ops/gemm/blas.h" namespace infini::ops { diff --git a/src/native/cuda/metax/ops/gemm/mcblas.h b/src/native/cuda/metax/ops/gemm/mcblas.h index c1ae08d08..9d2264c70 100644 --- a/src/native/cuda/metax/ops/gemm/mcblas.h +++ b/src/native/cuda/metax/ops/gemm/mcblas.h @@ -2,6 +2,7 @@ #define INFINI_OPS_METAX_GEMM_MCBLAS_H_ #include "native/cuda/metax/blas.h" +#include "native/cuda/metax/ops/add/kernel.h" #include "native/cuda/ops/gemm/blas.h" namespace infini::ops { diff --git a/src/native/cuda/moore/ops/gemm/mublas.h b/src/native/cuda/moore/ops/gemm/mublas.h index a34cba523..f2291c3d8 100644 --- a/src/native/cuda/moore/ops/gemm/mublas.h +++ b/src/native/cuda/moore/ops/gemm/mublas.h @@ -2,6 +2,7 @@ #define INFINI_OPS_MOORE_GEMM_MUBLAS_H_ #include "native/cuda/moore/blas.h" +#include "native/cuda/moore/ops/add/kernel.h" #include "native/cuda/ops/gemm/blas.h" namespace infini::ops { diff --git a/src/native/cuda/nvidia/ops/gemm/cublas.h b/src/native/cuda/nvidia/ops/gemm/cublas.h index eefe0b7af..7ad4b4058 100644 --- a/src/native/cuda/nvidia/ops/gemm/cublas.h +++ b/src/native/cuda/nvidia/ops/gemm/cublas.h @@ -2,6 +2,7 @@ #define INFINI_OPS_NVIDIA_GEMM_CUBLAS_H_ #include "native/cuda/nvidia/blas.h" +#include "native/cuda/nvidia/ops/add/kernel.h" #include "native/cuda/ops/gemm/blas.h" namespace infini::ops { diff --git a/src/native/cuda/nvidia/ops/gemm/cublaslt.h b/src/native/cuda/nvidia/ops/gemm/cublaslt.h index adcbc1ab1..24c2dd5d0 100644 --- a/src/native/cuda/nvidia/ops/gemm/cublaslt.h +++ b/src/native/cuda/nvidia/ops/gemm/cublaslt.h @@ -11,6 +11,7 @@ #include "base/gemm.h" #include "host_range_profiler.h" #include "native/cuda/nvidia/blas_utils.h" +#include "native/cuda/nvidia/ops/add/kernel.h" #include "native/cuda/nvidia/runtime_.h" namespace infini::ops { @@ -18,15 +19,17 @@ namespace infini::ops { template <> class Operator : public Gemm { public: - Operator(const Tensor a, const Tensor b, const std::optional input_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 c) - : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c}, + 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, Tensor c) + using Gemm::operator(); + + Operator(const Tensor a, const Tensor b, Tensor y) : Operator{a, b, std::nullopt, @@ -34,25 +37,22 @@ class Operator : public Gemm { std::nullopt, std::nullopt, std::nullopt, - c} {} - - using Gemm::operator(); + y} {} // 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, - const std::optional input_c, + 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 override { + 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{EffectiveBeta(input_c, 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_)}; @@ -68,20 +68,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"); @@ -110,16 +110,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"); 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{}; @@ -137,20 +137,23 @@ 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(0.0F), 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"); + if (c && beta_value != 0.0F) { + Add::Call(handle_, Config{}, y, *c, static_cast(beta_value), y); + } 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 8ed7ef095..db3da7a6d 100644 --- a/src/native/cuda/ops/gemm/blas.h +++ b/src/native/cuda/ops/gemm/blas.h @@ -3,6 +3,7 @@ #include +#include "base/add.h" #include "base/gemm.h" #include "native/cuda/blas_utils.h" @@ -11,17 +12,19 @@ namespace infini::ops { template class BlasGemm : public Gemm { public: - BlasGemm(const Tensor a, const Tensor b, const std::optional input_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 c) - : Gemm{a, b, input_c, alpha, beta, trans_a, trans_b, c}, + 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, Tensor c) + using Gemm::operator(); + + BlasGemm(const Tensor a, const Tensor b, Tensor y) : BlasGemm{a, b, std::nullopt, @@ -29,27 +32,25 @@ class BlasGemm : public Gemm { std::nullopt, std::nullopt, std::nullopt, - c} {} + y} {} - using Gemm::operator(); - - void operator()(const Tensor a, const Tensor b, - const std::optional input_c, + 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 override { + Tensor y) const override { Backend::BlasSetStream(GetHandle(), static_cast(stream_)); const auto& alpha_value{alpha.value_or(alpha_)}; - const auto beta_value{EffectiveBeta(input_c, beta)}; + const auto beta_value{EffectiveBeta(c, beta)}; + const auto gemm_beta{0.0F}; 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(gemm_beta, y.dtype())}; Backend::BlasGemmStridedBatchedEx( GetHandle(), op_a, op_b, swap_a_and_b_ ? n_ : m_, @@ -63,11 +64,14 @@ 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); + if (c && beta_value != 0.0F) { + Add::Call(handle_, Config{}, y, *c, static_cast(beta_value), y); + } } protected: diff --git a/src/torch/ops/gemm/gemm.cc b/src/torch/ops/gemm/gemm.cc index 9aaf105ea..a961db325 100644 --- a/src/torch/ops/gemm/gemm.cc +++ b/src/torch/ops/gemm/gemm.cc @@ -6,19 +6,18 @@ namespace infini::ops { template Operator::Operator(const Tensor a, const Tensor b, - const std::optional input_c, + const std::optional c, std::optional alpha, std::optional beta, std::optional trans_a, - std::optional trans_b, Tensor c) - : Gemm{a, b, input_c, 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()} {} + device_index_{y.device().index()} {} template -Operator::Operator(const Tensor a, const Tensor b, Tensor c) +Operator::Operator(const Tensor a, const Tensor b, Tensor y) : Operator{a, b, std::nullopt, @@ -26,23 +25,27 @@ Operator::Operator(const Tensor a, const Tensor b, Tensor c) std::nullopt, std::nullopt, std::nullopt, - c} {} + y} {} template void Operator::operator()( - const Tensor a, const Tensor b, const std::optional input_c, + 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 { + 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_); + std::optional at_c; + if (c) { + at_c.emplace(ToAtenTensor(const_cast(c->data()), c_shape_, + c_strides_, y_type_, device_index_)); + } auto alpha_val = alpha.value_or(alpha_); - auto beta_val = EffectiveBeta(input_c, beta); - + auto beta_val = EffectiveBeta(c, beta); if (trans_a.value_or(trans_a_)) { at_a = at_a.transpose(-2, -1); } @@ -52,28 +55,36 @@ void Operator::operator()( } if (alpha_val == 0.0F) { - at_c.mul_(beta_val); + if (!at_c || beta_val == 0.0F) { + at_y.zero_(); + return; + } + + at_y.copy_(*at_c); + at_y.mul_(beta_val); return; } if constexpr (kDev == Device::Type::kCpu || kDev == Device::Type::kNvidia) { + const auto& input = at_c ? *at_c : at_y; 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, input, 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, input, 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); + if (!at_c || beta_val == 0.0F) { + at_y.copy_(product); + at_y.mul_(alpha_val); return; } - at_c.mul_(beta_val); - at_c.add_(product, alpha_val); + at_y.copy_(*at_c); + 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 b74505ddb..095b85bba 100644 --- a/src/torch/ops/gemm/gemm.h +++ b/src/torch/ops/gemm/gemm.h @@ -8,27 +8,24 @@ namespace infini::ops { template class Operator : public Gemm { public: - Operator(const Tensor a, const Tensor b, const std::optional input_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 c); + std::optional trans_a, std::optional trans_b, Tensor y); - Operator(const Tensor a, const Tensor b, Tensor c); + Operator(const Tensor a, const Tensor b, Tensor y); using Gemm::operator(); - void operator()(const Tensor a, const Tensor b, - const std::optional input_c, + 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 override; + Tensor y) const override; private: Tensor::Shape a_shape_; Tensor::Shape b_shape_; - Tensor::Shape c_shape_; - int device_index_{0}; }; diff --git a/tests/test_gemm.py b/tests/test_gemm.py index 31df00c0e..d005f76cf 100644 --- a/tests/test_gemm.py +++ b/tests/test_gemm.py @@ -92,41 +92,87 @@ def test_gemm( b = b.transpose(-2, -1) c = randn_strided(c_shape, c_strides, dtype=dtype, device=device) - use_portable_ref = implementation_index == 2 and not ( - device == "cpu" - or ( - device == "cuda" and infini.ops.Gemm.active_implementation_indices("nvidia") + y = randn_strided(c_shape, c_strides, dtype=dtype, device=device) + # Native separate-C paths accumulate in a second kernel; match their rounding. + native_separate_c = ( + beta != 0 + and implementation_index in (0, 1) + and device in ("cuda", "mlu", "musa") + ) + use_portable_ref = native_separate_c or ( + implementation_index == 2 + and not ( + device == "cpu" + or ( + device == "cuda" + and infini.ops.Gemm.active_implementation_indices("nvidia") + ) ) ) ref = _torch_gemm_portable if use_portable_ref else _torch_gemm return Payload( lambda *args: _gemm(*args, implementation_index=implementation_index), - lambda a, b, alpha, _beta, trans_a, trans_b, c: ref( - a, b, alpha, 0.0, trans_a, trans_b, c - ), - (a, b, alpha, beta, trans_a, trans_b, c), + lambda *args: ref(*args[:-1]), + (a, b, alpha, beta, trans_a, trans_b, c, y), {}, rtol=rtol, atol=atol, ) -def _gemm(a, b, alpha, beta, trans_a, trans_b, c, implementation_index=0): +def _gemm(a, b, alpha, beta, trans_a, trans_b, c, y, implementation_index=0): infini.ops.gemm( a, b, - None, + c, alpha, beta, trans_a, trans_b, - c, + y, stream=get_stream(a.device), implementation_index=implementation_index, ) - return c + return y + + +@pytest.mark.smoke +def test_gemm_without_c_overload(device, implementation_index): + if implementation_index == 2 and device == "npu": + pytest.skip("Gemm impl=2 is not instantiated for Ascend") + + a = torch.randn((3, 4), dtype=torch.float32, device=device) + b = torch.randn((4, 2), dtype=torch.float32, device=device) + y = torch.full((3, 2), torch.nan, dtype=torch.float32, device=device) + + infini.ops.gemm( + a, + b, + y, + stream=get_stream(a.device), + implementation_index=implementation_index, + ) + + torch.testing.assert_close(y, torch.matmul(a, b), rtol=1e-3, atol=1e-3) + + y.fill_(torch.nan) + + infini.ops.gemm( + a, + b, + None, + 0.0, + 1.0, + False, + False, + y, + stream=get_stream(a.device), + implementation_index=implementation_index, + ) + + torch.testing.assert_close(y, torch.zeros_like(y), rtol=0.0, atol=0.0) def _torch_gemm(a, b, alpha=1.0, beta=1.0, trans_a=False, trans_b=False, c=None): @@ -185,3 +231,33 @@ def _torch_gemm_portable( c.add_(product, alpha=alpha) return c + + +@pytest.mark.smoke +def test_gemm_broadcast_c(device, implementation_index): + if implementation_index == 2 and device == "npu": + pytest.skip("Gemm impl=2 is not instantiated for Ascend") + + a = torch.randn((3, 4), dtype=torch.float32, device=device) + b = torch.randn((4, 2), dtype=torch.float32, device=device) + c = torch.randn((1, 2), dtype=torch.float32, device=device) + y = torch.full((3, 2), torch.nan, dtype=torch.float32, device=device) + + expected = torch.matmul(a, b) + expected.mul_(0.5) + expected.add_(c, alpha=-0.5) + + infini.ops.gemm( + a, + b, + c, + 0.5, + -0.5, + False, + False, + y, + stream=get_stream(a.device), + implementation_index=implementation_index, + ) + + torch.testing.assert_close(y, expected, rtol=1e-3, atol=1e-3)