InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
clip.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CLIP_H_
2#define INFINI_OPS_BASE_CLIP_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class Clip : public Operator<Clip> {
11 public:
12 Clip(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 Clip(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 clip.h:10
Tensor::Shape min_shape_
Definition clip.h:69
std::optional< double > min_
Definition clip.h:63
int device_index_
Definition clip.h:83
DataType min_type_
Definition clip.h:73
Tensor::Strides min_strides_
Definition clip.h:71
DataType out_type_
Definition clip.h:61
virtual void operator()(const Tensor input, const std::optional< Tensor > min, const std::optional< Tensor > max, Tensor out) const =0
Clip(const Tensor input, const std::optional< double > min, const std::optional< double > max, Tensor out)
Definition clip.h:12
Clip(const Tensor input, const std::optional< Tensor > min, const std::optional< Tensor > max, Tensor out)
Definition clip.h:24
DataType max_type_
Definition clip.h:81
std::optional< double > max_
Definition clip.h:65
Tensor::Shape input_shape_
Definition clip.h:51
Tensor::Strides max_strides_
Definition clip.h:79
bool has_max_
Definition clip.h:75
Tensor::Strides input_strides_
Definition clip.h:53
DataType input_type_
Definition clip.h:55
Tensor::Strides out_strides_
Definition clip.h:59
Tensor::Shape out_shape_
Definition clip.h:57
Tensor::Shape max_shape_
Definition clip.h:77
bool has_min_
Definition clip.h:67
virtual void operator()(const Tensor input, const std::optional< double > min, const std::optional< double > max, Tensor out) const =0
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8