InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
smooth_l1_loss.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SMOOTH_L1_LOSS_H_
2#define INFINI_OPS_BASE_SMOOTH_L1_LOSS_H_
3
4#include <optional>
5#include <string>
6
7#include "common/op_utils/reduction.h"
8#include "operator.h"
9
10namespace infini::ops {
11
12class SmoothL1Loss : public Operator<SmoothL1Loss> {
13 public:
14 SmoothL1Loss(const Tensor input, const Tensor target,
15 const std::optional<bool> size_average,
16 const std::optional<bool> reduce, const std::string reduction,
17 const double beta, Tensor out)
18 : input_shape_{input.shape()},
19 input_strides_{input.strides()},
20 input_type_{input.dtype()},
21 target_shape_{target.shape()},
22 target_strides_{target.strides()},
23 target_type_{target.dtype()},
24 out_shape_{out.shape()},
25 out_strides_{out.strides()},
26 out_type_{out.dtype()},
27 reduction_{reduction_detail::FromPythonArguments(size_average, reduce,
28 reduction)},
29 beta_{beta},
30 device_index_{out.device().index()} {}
31
34 [[deprecated("Use the Python-compatible reduction overload instead.")]]
35 SmoothL1Loss(const Tensor input, const Tensor target, const int64_t reduction,
36 const double beta, Tensor out)
37 : input_shape_{input.shape()},
38 input_strides_{input.strides()},
39 input_type_{input.dtype()},
40 target_shape_{target.shape()},
41 target_strides_{target.strides()},
42 target_type_{target.dtype()},
43 out_shape_{out.shape()},
44 out_strides_{out.strides()},
45 out_type_{out.dtype()},
46 reduction_{reduction},
47 beta_{beta},
48 device_index_{out.device().index()} {}
49
50 void operator()(const Tensor input, const Tensor target,
51 const std::optional<bool> size_average,
52 const std::optional<bool> reduce, const std::string reduction,
53 const double beta, Tensor out) const {
54 (*this)(
55 input, target,
56 reduction_detail::FromPythonArguments(size_average, reduce, reduction),
57 beta, out);
58 }
59
62 [[deprecated("Use the Python-compatible reduction overload instead.")]]
63 virtual void operator()(const Tensor input, const Tensor target,
64 const int64_t reduction, const double beta,
65 Tensor out) const = 0;
66
67 protected:
68 Tensor::Shape input_shape_;
69
70 Tensor::Strides input_strides_;
71
72 DataType input_type_;
73
74 Tensor::Shape target_shape_;
75
76 Tensor::Strides target_strides_;
77
78 DataType target_type_;
79
80 Tensor::Shape out_shape_;
81
82 Tensor::Strides out_strides_;
83
84 DataType out_type_;
85
86 int64_t reduction_{};
87
88 double beta_{};
89
91};
92
93} // namespace infini::ops
94
95#endif
Definition generated/include/operator.h:282
Definition smooth_l1_loss.h:12
double beta_
Definition smooth_l1_loss.h:88
SmoothL1Loss(const Tensor input, const Tensor target, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, const double beta, Tensor out)
Definition smooth_l1_loss.h:14
Tensor::Shape target_shape_
Definition smooth_l1_loss.h:74
Tensor::Strides input_strides_
Definition smooth_l1_loss.h:70
Tensor::Shape input_shape_
Definition smooth_l1_loss.h:68
int64_t reduction_
Definition smooth_l1_loss.h:86
DataType input_type_
Definition smooth_l1_loss.h:72
DataType target_type_
Definition smooth_l1_loss.h:78
Tensor::Strides out_strides_
Definition smooth_l1_loss.h:82
int device_index_
Definition smooth_l1_loss.h:90
void operator()(const Tensor input, const Tensor target, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, const double beta, Tensor out) const
Definition smooth_l1_loss.h:50
Tensor::Strides target_strides_
Definition smooth_l1_loss.h:76
virtual void operator()(const Tensor input, const Tensor target, const int64_t reduction, const double beta, Tensor out) const =0
Tensor::Shape out_shape_
Definition smooth_l1_loss.h:80
SmoothL1Loss(const Tensor input, const Tensor target, const int64_t reduction, const double beta, Tensor out)
Definition smooth_l1_loss.h:35
DataType out_type_
Definition smooth_l1_loss.h:84
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8