InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
mul.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MUL_H_
2#define INFINI_OPS_BASE_MUL_H_
3
4#include "operator.h"
5
6namespace infini::ops {
7
8class Mul : public Operator<Mul> {
9 public:
10 Mul(const Tensor input, const Tensor other, Tensor out)
11 : ndim_{out.ndim()},
12 output_size_{out.numel()},
13 input_type_{input.dtype()},
14 other_type_{other.dtype()},
15 out_type_{out.dtype()},
16 input_shape_{out.shape()},
17 other_shape_{out.shape()},
18 out_shape_{out.shape()},
21 out_strides_{out.strides()},
22 is_input_contiguous_{input.shape() == out.shape() &&
23 input.IsContiguous()},
24 is_other_contiguous_{other.shape() == out.shape() &&
25 other.IsContiguous()},
26 is_out_contiguous_{out.IsContiguous()} {
27 assert(!out.HasBroadcastDim() &&
28 "the output of `Mul` should NOT have broadcasted dim!");
30 "operator `Mul` requires all input and output tensors to have the "
31 "same dtype");
32 ValidateBroadcast(input, other, out);
33 }
34
35 virtual void operator()(const Tensor input, const Tensor other,
36 Tensor out) const = 0;
37
38 protected:
39 static Tensor::Strides BroadcastStrides(const Tensor input,
40 const Tensor out) {
41 assert(input.ndim() <= out.ndim() &&
42 "operator `Mul` input rank must not exceed output rank");
43 Tensor::Strides strides(out.ndim(), 0);
44 auto offset = out.ndim() - input.ndim();
45
46 for (Tensor::Size i = 0; i < input.ndim(); ++i) {
47 auto out_dim = i + offset;
48 assert((input.size(i) == 1 || input.size(i) == out.size(out_dim)) &&
49 "operator `Mul` input shape is not broadcast-compatible with "
50 "output shape");
51 strides[out_dim] = input.size(i) == 1 ? 0 : input.stride(i);
52 }
53
54 return strides;
55 }
56
57 static void ValidateBroadcast(const Tensor input, const Tensor other,
58 const Tensor out) {
59 for (Tensor::Size i = 0; i < out.ndim(); ++i) {
60 auto input_dim = i < out.ndim() - input.ndim()
61 ? 1
62 : input.size(i + input.ndim() - out.ndim());
63 auto other_dim = i < out.ndim() - other.ndim()
64 ? 1
65 : other.size(i + other.ndim() - out.ndim());
66 [[maybe_unused]] auto broadcast_dim =
67 input_dim == 1 ? other_dim : input_dim;
68 assert(out.size(i) == broadcast_dim &&
69 "operator `Mul` output shape must equal the broadcasted input "
70 "shape");
71 }
72 }
73
74 Tensor::Size ndim_{0};
75
76 Tensor::Size output_size_{0};
77
78 const DataType input_type_;
79
80 const DataType other_type_;
81
82 const DataType out_type_;
83
84 Tensor::Shape input_shape_;
85
86 Tensor::Shape other_shape_;
87
88 Tensor::Shape out_shape_;
89
90 Tensor::Strides input_strides_;
91
92 Tensor::Strides other_strides_;
93
94 Tensor::Strides out_strides_;
95
97
99
101};
102
103} // namespace infini::ops
104
105#endif
Definition mul.h:8
Tensor::Size ndim_
Definition mul.h:74
Tensor::Strides out_strides_
Definition mul.h:94
Tensor::Strides input_strides_
Definition mul.h:90
bool is_input_contiguous_
Definition mul.h:96
Tensor::Shape input_shape_
Definition mul.h:84
Tensor::Shape other_shape_
Definition mul.h:86
Tensor::Strides other_strides_
Definition mul.h:92
bool is_out_contiguous_
Definition mul.h:100
virtual void operator()(const Tensor input, const Tensor other, Tensor out) const =0
Tensor::Size output_size_
Definition mul.h:76
bool is_other_contiguous_
Definition mul.h:98
const DataType other_type_
Definition mul.h:80
static void ValidateBroadcast(const Tensor input, const Tensor other, const Tensor out)
Definition mul.h:57
const DataType out_type_
Definition mul.h:82
static Tensor::Strides BroadcastStrides(const Tensor input, const Tensor out)
Definition mul.h:39
const DataType input_type_
Definition mul.h:78
Mul(const Tensor input, const Tensor other, Tensor out)
Definition mul.h:10
Tensor::Shape out_shape_
Definition mul.h:88
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8