InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
linear.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_LINEAR_H_
2#define INFINI_OPS_BASE_LINEAR_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class Linear : public Operator<Linear> {
11 public:
12 Linear(const Tensor input, const Tensor weight, std::optional<Tensor> bias,
13 Tensor out)
14 : input_shape_{input.shape()},
15 input_strides_{input.strides()},
16 input_type_{input.dtype()},
17 weight_shape_{weight.shape()},
18 weight_strides_{weight.strides()},
19 weight_type_{weight.dtype()},
20 out_shape_{out.shape()},
21 out_strides_{out.strides()},
22 out_type_{out.dtype()},
23 rows_{Rows(input)},
24 has_bias_{bias.has_value()} {
25 assert(input.ndim() >= 1 && "operator `Linear` requires non-scalar input");
26 assert(weight.ndim() == 2 && "operator `Linear` requires 2D `weight`");
27 assert(input.size(-1) == weight.size(-1) &&
28 "operator `Linear` input features must match `weight`");
29 assert(out.ndim() == input.ndim() &&
30 "operator `Linear` output rank must match input rank");
31 assert(out.size(-1) == weight.size(0) &&
32 "operator `Linear` output features must match `weight`");
33
34 for (Tensor::Size axis = 0; axis + 1 < input.ndim(); ++axis) {
35 assert(input.size(axis) == out.size(axis) &&
36 "operator `Linear` output leading dimensions must match input");
37 }
38
39 assert(input.dtype() == weight.dtype() &&
40 "operator `Linear` requires input and weight to have the same "
41 "dtype");
42 assert(input.dtype() == out.dtype() &&
43 "operator `Linear` requires output to have the input dtype");
44 if (has_bias_) {
45 assert(bias->ndim() == 1 && bias->size(0) == weight.size(0) &&
46 "operator `Linear` bias must have shape `[out_features]`");
47 assert(bias->dtype() == out.dtype() &&
48 "operator `Linear` requires bias to have the output dtype");
49 }
50 }
51
52 virtual void operator()(const Tensor input, const Tensor weight,
53 std::optional<Tensor> bias, Tensor out) const = 0;
54
55 protected:
56 static Tensor::Size Rows(const Tensor input) {
57 Tensor::Size rows = 1;
58
59 for (Tensor::Size axis = 0; axis + 1 < input.ndim(); ++axis) {
60 rows *= input.size(axis);
61 }
62
63 return input.ndim() == 0 ? 0 : rows;
64 }
65
66 Tensor::Shape input_shape_;
67
68 Tensor::Strides input_strides_;
69
70 DataType input_type_;
71
72 Tensor::Shape weight_shape_;
73
74 Tensor::Strides weight_strides_;
75
76 DataType weight_type_;
77
78 Tensor::Shape out_shape_;
79
80 Tensor::Strides out_strides_;
81
82 DataType out_type_;
83
84 Tensor::Size rows_{0};
85
86 bool has_bias_{false};
87};
88
89} // namespace infini::ops
90
91#endif
Definition linear.h:10
static Tensor::Size Rows(const Tensor input)
Definition linear.h:56
DataType out_type_
Definition linear.h:82
Tensor::Strides out_strides_
Definition linear.h:80
DataType input_type_
Definition linear.h:70
Tensor::Shape out_shape_
Definition linear.h:78
Tensor::Shape input_shape_
Definition linear.h:66
DataType weight_type_
Definition linear.h:76
Tensor::Shape weight_shape_
Definition linear.h:72
Tensor::Strides input_strides_
Definition linear.h:68
bool has_bias_
Definition linear.h:86
Tensor::Strides weight_strides_
Definition linear.h:74
Linear(const Tensor input, const Tensor weight, std::optional< Tensor > bias, Tensor out)
Definition linear.h:12
Tensor::Size rows_
Definition linear.h:84
virtual void operator()(const Tensor input, const Tensor weight, std::optional< Tensor > bias, Tensor out) const =0
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8