InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
clamp.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CLAMP_H_
2#define INFINI_OPS_BASE_CLAMP_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class Clamp : public Operator<Clamp> {
11 public:
12 Clamp(const Tensor input, const std::optional<double> min,
13 const std::optional<double> max, Tensor out)
14 : input_shape_{input.shape()},
15 input_strides_{input.strides()},
16 input_type_{input.dtype()},
17 out_shape_{out.shape()},
18 out_strides_{out.strides()},
19 out_type_{out.dtype()},
20 min_{min},
21 max_{max},
22 device_index_{out.device().index()} {}
23
24 Clamp(const Tensor input, const std::optional<Tensor> min,
25 const std::optional<Tensor> max, Tensor out)
26 : input_shape_{input.shape()},
27 input_strides_{input.strides()},
28 input_type_{input.dtype()},
29 out_shape_{out.shape()},
30 out_strides_{out.strides()},
31 out_type_{out.dtype()},
32 has_min_{min.has_value()},
33 min_shape_{min ? Tensor::Shape{min->shape()} : Tensor::Shape{}},
34 min_strides_{min ? Tensor::Strides{min->strides()} : Tensor::Strides{}},
35 min_type_{min ? min->dtype() : DataType::kFloat32},
36 has_max_{max.has_value()},
37 max_shape_{max ? Tensor::Shape{max->shape()} : Tensor::Shape{}},
38 max_strides_{max ? Tensor::Strides{max->strides()} : Tensor::Strides{}},
39 max_type_{max ? max->dtype() : DataType::kFloat32},
40 device_index_{out.device().index()} {}
41
42 virtual void operator()(const Tensor input, const std::optional<double> min,
43 const std::optional<double> max,
44 Tensor out) const = 0;
45
46 virtual void operator()(const Tensor input, const std::optional<Tensor> min,
47 const std::optional<Tensor> max,
48 Tensor out) const = 0;
49
50 protected:
51 Tensor::Shape input_shape_;
52
53 Tensor::Strides input_strides_;
54
55 DataType input_type_;
56
57 Tensor::Shape out_shape_;
58
59 Tensor::Strides out_strides_;
60
61 DataType out_type_;
62
63 std::optional<double> min_{};
64
65 std::optional<double> max_{};
66
67 bool has_min_{false};
68
69 Tensor::Shape min_shape_;
70
71 Tensor::Strides min_strides_;
72
73 DataType min_type_{DataType::kFloat32};
74
75 bool has_max_{false};
76
77 Tensor::Shape max_shape_;
78
79 Tensor::Strides max_strides_;
80
81 DataType max_type_{DataType::kFloat32};
82
84};
85
86} // namespace infini::ops
87
88#endif
Definition clamp.h:10
DataType min_type_
Definition clamp.h:73
Tensor::Strides out_strides_
Definition clamp.h:59
virtual void operator()(const Tensor input, const std::optional< double > min, const std::optional< double > max, Tensor out) const =0
Clamp(const Tensor input, const std::optional< double > min, const std::optional< double > max, Tensor out)
Definition clamp.h:12
int device_index_
Definition clamp.h:83
std::optional< double > max_
Definition clamp.h:65
Clamp(const Tensor input, const std::optional< Tensor > min, const std::optional< Tensor > max, Tensor out)
Definition clamp.h:24
bool has_max_
Definition clamp.h:75
bool has_min_
Definition clamp.h:67
virtual void operator()(const Tensor input, const std::optional< Tensor > min, const std::optional< Tensor > max, Tensor out) const =0
Tensor::Shape out_shape_
Definition clamp.h:57
DataType max_type_
Definition clamp.h:81
Tensor::Shape input_shape_
Definition clamp.h:51
std::optional< double > min_
Definition clamp.h:63
Tensor::Strides max_strides_
Definition clamp.h:79
Tensor::Strides input_strides_
Definition clamp.h:53
DataType out_type_
Definition clamp.h:61
DataType input_type_
Definition clamp.h:55
Tensor::Shape min_shape_
Definition clamp.h:69
Tensor::Shape max_shape_
Definition clamp.h:77
Tensor::Strides min_strides_
Definition clamp.h:71
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8