InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
gemm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_GEMM_H_
2#define INFINI_OPS_BASE_GEMM_H_
3
4#include <algorithm>
5#include <cassert>
6#include <optional>
7
8#include "operator.h"
9
10namespace infini::ops {
11
12class Gemm : public Operator<Gemm> {
13 public:
14 Gemm(const Tensor a, const Tensor b, const std::optional<Tensor> c,
15 std::optional<float> alpha, std::optional<float> beta,
16 std::optional<int> trans_a, std::optional<int> trans_b, Tensor y)
17 : alpha_{alpha.value_or(1.0)},
18 beta_{EffectiveBeta(c, beta)},
19 trans_a_{static_cast<bool>(trans_a.value_or(false))},
20 trans_b_{static_cast<bool>(trans_b.value_or(false))},
21 m_{y.size(-2)},
22 n_{y.size(-1)},
23 k_{trans_a_ ? a.size(-2) : a.size(-1)},
24 a_type_{a.dtype()},
25 b_type_{b.dtype()},
26 y_type_{y.dtype()},
27 a_strides_{a.strides()},
28 b_strides_{b.strides()},
29 y_strides_{y.strides()},
30 lda_{std::max(a.stride(-2), a.stride(-1))},
31 ldb_{std::max(b.stride(-2), b.stride(-1))},
32 ldy_{std::max(y.stride(-2), y.stride(-1))},
33 batch_count_{y.strides().size() > 2 ? y.size(-3) : 1},
34 batch_stride_a_{a.strides().size() > 2 ? a.stride(-3) : 0},
35 batch_stride_b_{b.strides().size() > 2 ? b.stride(-3) : 0},
36 batch_stride_y_{y.strides().size() > 2 ? y.stride(-3) : 0} {
37 // TODO: Check constraints.
38 }
39
40 Gemm(const Tensor a, const Tensor b, Tensor y)
41 : Gemm{a,
42 b,
43 std::nullopt,
44 std::nullopt,
45 std::nullopt,
46 std::nullopt,
47 std::nullopt,
48 y} {}
49
50 virtual void operator()(const Tensor a, const Tensor b,
51 const std::optional<Tensor> c,
52 std::optional<float> alpha, std::optional<float> beta,
53 std::optional<int> trans_a,
54 std::optional<int> trans_b, Tensor y) const = 0;
55
56 virtual void operator()(const Tensor a, const Tensor b, Tensor y) const {
57 return operator()(a, b, std::nullopt, std::nullopt, std::nullopt,
58 std::nullopt, std::nullopt, y);
59 }
60
61 template <typename TensorLike>
62 static auto MakeReturnValue(const TensorLike& a, const TensorLike& b) {
63 Tensor::Shape y_shape{a.shape()[a.shape().size() - 2],
64 b.shape()[b.shape().size() - 1]};
65 return TensorLike::Empty(y_shape, a.dtype(), a.device());
66 }
67
68 protected:
69 static float EffectiveBeta(const std::optional<Tensor>& c,
70 std::optional<float> beta) {
71 static_cast<void>(beta);
72 assert(!c && "operator Gemm C input is not supported yet");
73 return 0.0F;
74 }
75
76 float alpha_{1.0};
77
78 float beta_{1.0};
79
80 bool trans_a_{false};
81
82 bool trans_b_{false};
83
84 Tensor::Size m_{0};
85
86 Tensor::Size n_{0};
87
88 Tensor::Size k_{0};
89
90 const DataType a_type_;
91
92 const DataType b_type_;
93
94 const DataType y_type_;
95
96 Tensor::Strides a_strides_;
97
98 Tensor::Strides b_strides_;
99
100 Tensor::Strides y_strides_;
101
102 Tensor::Stride lda_{0};
103
104 Tensor::Stride ldb_{0};
105
106 Tensor::Stride ldy_{0};
107
108 Tensor::Size batch_count_{1};
109
110 Tensor::Stride batch_stride_a_{0};
111
112 Tensor::Stride batch_stride_b_{0};
113
114 Tensor::Stride batch_stride_y_{0};
115};
116
117} // namespace infini::ops
118
119#endif
Definition gemm.h:12
const DataType b_type_
Definition gemm.h:92
const DataType y_type_
Definition gemm.h:94
Tensor::Stride ldy_
Definition gemm.h:106
Tensor::Size n_
Definition gemm.h:86
Tensor::Strides b_strides_
Definition gemm.h:98
Tensor::Size batch_count_
Definition gemm.h:108
Tensor::Strides a_strides_
Definition gemm.h:96
Tensor::Size k_
Definition gemm.h:88
Tensor::Stride batch_stride_b_
Definition gemm.h:112
float alpha_
Definition gemm.h:76
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 y) const =0
static float EffectiveBeta(const std::optional< Tensor > &c, std::optional< float > beta)
Definition gemm.h:69
virtual void operator()(const Tensor a, const Tensor b, Tensor y) const
Definition gemm.h:56
Gemm(const Tensor a, const Tensor b, Tensor y)
Definition gemm.h:40
Tensor::Stride lda_
Definition gemm.h:102
const DataType a_type_
Definition gemm.h:90
Tensor::Stride batch_stride_a_
Definition gemm.h:110
bool trans_a_
Definition gemm.h:80
Tensor::Strides y_strides_
Definition gemm.h:100
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)
Definition gemm.h:14
Tensor::Size m_
Definition gemm.h:84
bool trans_b_
Definition gemm.h:82
Tensor::Stride batch_stride_y_
Definition gemm.h:114
Tensor::Stride ldb_
Definition gemm.h:104
float beta_
Definition gemm.h:78
static auto MakeReturnValue(const TensorLike &a, const TensorLike &b)
Definition gemm.h:62
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
template void bool
Definition operator_call_instantiations.h:111
infini::rt::TensorView Tensor
Definition tensor.h:8