InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
relu.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_RELU_H_
2#define INFINI_OPS_BASE_RELU_H_
3
4#include <cassert>
5#include <cstdint>
6#include <utility>
7
8#include "operator.h"
9
10namespace infini::ops {
11
13 ConcatType<ConcatType<AllFloatTypes, IntTypes>, List<DataType::kUInt8>>;
14
15class Relu : public Operator<Relu> {
16 public:
17 Relu(const Tensor input, Tensor out)
18 : ndim_{out.ndim()},
19 output_size_{out.numel()},
20 input_type_{input.dtype()},
21 out_type_{out.dtype()},
22 input_shape_{input.shape()},
23 out_shape_{out.shape()},
24 input_strides_{input.strides()},
25 out_strides_{out.strides()},
26 is_input_contiguous_{input.IsContiguous()},
27 is_out_contiguous_{out.IsContiguous()} {
28 assert(input_shape_ == out_shape_ &&
29 "operator `Relu` requires matching input and output shapes");
30 assert(input_type_ == out_type_ &&
31 "operator `Relu` requires matching input and output dtypes");
32 assert(input.device() == out.device() &&
33 "operator `Relu` requires input and output on the same device");
35 "operator `Relu` received an unsupported dtype");
36 assert(!out.HasBroadcastDim() &&
37 "operator `Relu` output must not have broadcasted dimensions");
38 }
39
40 virtual void operator()(const Tensor input, Tensor out) const = 0;
41
42 protected:
43 bool NeedsInputCopy(const Tensor input, const Tensor out) const {
44 if (output_size_ == 0 ||
45 (input.data() == out.data() && input.strides() == out.strides())) {
46 return false;
47 }
48
49 const auto [input_begin, input_end] = StorageByteRange(input);
50 const auto [out_begin, out_end] = StorageByteRange(out);
51
52 return input_begin < out_end && out_begin < input_end;
53 }
54
55 static std::pair<std::intptr_t, std::intptr_t> StorageByteRange(
56 const Tensor tensor) {
57 Tensor::Stride min_offset = 0;
58 Tensor::Stride max_offset = 0;
59
60 for (Tensor::Size i = 0; i < tensor.ndim(); ++i) {
61 const auto extent =
62 static_cast<Tensor::Stride>(tensor.size(i) - 1) * tensor.stride(i);
63
64 if (extent < 0) {
65 min_offset += extent;
66 } else {
67 max_offset += extent;
68 }
69 }
70
71 const auto address = reinterpret_cast<std::intptr_t>(tensor.data());
72 const auto element_size = static_cast<std::intptr_t>(tensor.element_size());
73
74 return {address + min_offset * element_size,
75 address + (max_offset + 1) * element_size};
76 }
77
78 Tensor::Size ndim_{0};
79
80 Tensor::Size output_size_{0};
81
82 const DataType input_type_;
83
84 const DataType out_type_;
85
86 Tensor::Shape input_shape_;
87
88 Tensor::Shape out_shape_;
89
90 Tensor::Strides input_strides_;
91
92 Tensor::Strides out_strides_;
93
95
96 bool is_out_contiguous_{false};
97};
98
99} // namespace infini::ops
100
101#endif
Definition generated/include/operator.h:282
Definition relu.h:15
const DataType input_type_
Definition relu.h:82
bool NeedsInputCopy(const Tensor input, const Tensor out) const
Definition relu.h:43
Tensor::Shape out_shape_
Definition relu.h:88
Tensor::Size ndim_
Definition relu.h:78
Tensor::Shape input_shape_
Definition relu.h:86
Tensor::Strides out_strides_
Definition relu.h:92
bool is_input_contiguous_
Definition relu.h:94
static std::pair< std::intptr_t, std::intptr_t > StorageByteRange(const Tensor tensor)
Definition relu.h:55
Tensor::Strides input_strides_
Definition relu.h:90
bool is_out_contiguous_
Definition relu.h:96
virtual void operator()(const Tensor input, Tensor out) const =0
const DataType out_type_
Definition relu.h:84
Tensor::Size output_size_
Definition relu.h:80
Relu(const Tensor input, Tensor out)
Definition relu.h:17
bool ListContains(ValueType value, List< values... >)
Definition generated/include/operator.h:83
Definition generated/include/operator.h:28
ConcatType< ConcatType< AllFloatTypes, IntTypes >, List< DataType::kUInt8 > > ReluDataTypes
Definition relu.h:13
infini::rt::TensorView Tensor
Definition tensor.h:8