Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions scripts/run_host_overhead_control.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,7 @@ def gemm():
ops.gemm(
gemm_a,
gemm_b,
None,
1.0,
0.0,
False,
Expand Down
96 changes: 69 additions & 27 deletions src/base/gemm.h
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
#define INFINI_OPS_BASE_GEMM_H_

#include <algorithm>
#include <cassert>
#include <optional>

#include "operator.h"
Expand All @@ -10,59 +11,92 @@ namespace infini::ops {

class Gemm : public Operator<Gemm> {
public:
Gemm(const Tensor a, const Tensor b, std::optional<float> alpha,
std::optional<float> beta, std::optional<int> trans_a,
std::optional<int> trans_b, Tensor c)
Gemm(const Tensor a, const Tensor b, const std::optional<Tensor> c,
std::optional<float> alpha, std::optional<float> beta,
std::optional<int> trans_a, std::optional<int> trans_b, Tensor y)
: alpha_{alpha.value_or(1.0)},
beta_{beta.value_or(1.0)},
trans_a_{static_cast<bool>(trans_a.value_or(false))},
trans_b_{static_cast<bool>(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()},
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(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} {
// TODO: Check constraints.
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 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<Tensor> c,
std::optional<float> alpha, std::optional<float> beta,
std::optional<int> trans_a,
std::optional<int> trans_b, Tensor c) const = 0;
std::optional<int> 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<float> alpha, std::optional<float> beta,
Tensor c) const {
return operator()(a, b, alpha, beta, std::nullopt, std::nullopt, c);
std::nullopt, std::nullopt, y);
}

template <typename TensorLike>
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 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<Tensor>& c,
std::optional<float> beta) const {
return c ? beta.value_or(beta_) : 0.0F;
}

float alpha_{1.0};

float beta_{1.0};
Expand All @@ -81,27 +115,35 @@ class Gemm : public Operator<Gemm> {

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};

Tensor::Stride batch_stride_a_{0};

Tensor::Stride batch_stride_b_{0};

Tensor::Stride batch_stride_c_{0};
Tensor::Stride batch_stride_y_{0};
};

} // namespace infini::ops
Expand Down
44 changes: 30 additions & 14 deletions src/native/ascend/ops/gemm/kernel.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,21 +15,33 @@ namespace infini::ops {
template <>
class Operator<Gemm, Device::Type::kAscend> : public Gemm {
public:
Operator(const Tensor a, const Tensor b, std::optional<float> alpha,
std::optional<float> beta, std::optional<int> trans_a,
std::optional<int> trans_b, Tensor c)
: Gemm(a, b, alpha, beta, trans_a, trans_b, c),
Operator(const Tensor a, const Tensor b, const std::optional<Tensor> c,
std::optional<float> alpha, std::optional<float> beta,
std::optional<int> trans_a, std::optional<int> 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_{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 y)
: Operator{a,
b,
std::nullopt,
std::nullopt,
std::nullopt,
std::nullopt,
std::nullopt,
y} {}

using Gemm::operator();

~Operator() {
if (!ascend::IsAclRuntimeAlive()) return;

Expand All @@ -43,15 +55,19 @@ class Operator<Gemm, Device::Type::kAscend> : public Gemm {
if (beta_scalar_) aclDestroyScalar(beta_scalar_);
}

void operator()(const Tensor a, const Tensor b, std::optional<float> alpha,
std::optional<float> beta, std::optional<int> trans_a,
std::optional<int> trans_b, Tensor c) const override {
void operator()(const Tensor a, const Tensor b, const std::optional<Tensor> c,
std::optional<float> alpha, std::optional<float> beta,
std::optional<int> trans_a, std::optional<int> trans_b,
Tensor y) const override {
auto stream = static_cast<aclrtStream>(stream_);

auto t_self = self_cache_.get(c.data());
const auto self = c ? *c : y;
auto self_ptr = const_cast<void*>(self.data());
auto y_ptr = y.data();
auto t_self = self_cache_.get(self_ptr);
auto t_a = a_cache_.get(const_cast<void*>(a.data()));
auto t_b = b_cache_.get(const_cast<void*>(b.data()));
auto t_out = out_cache_.get(c.data());
auto t_out = out_cache_.get(y_ptr);

if (!executor_) {
if (batched_) {
Expand All @@ -65,10 +81,10 @@ class Operator<Gemm, Device::Type::kAscend> : 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<void*>(a.data()));
aclSetInputTensorAddr(executor_, 2, t_b, const_cast<void*>(b.data()));
aclSetOutputTensorAddr(executor_, 0, t_out, c.data());
aclSetOutputTensorAddr(executor_, 0, t_out, y_ptr);
}

auto& arena = ascend::GetWorkspacePool().Ensure(stream, ws_size_);
Expand Down
Loading
Loading