InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
huber_loss.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_HUBER_LOSS_H_
2#define INFINI_OPS_BASE_HUBER_LOSS_H_
3
4#include <cassert>
5#include <optional>
6#include <string>
7
8#include "common/op_utils/reduction.h"
9#include "operator.h"
10
11namespace infini::ops {
12
13class HuberLoss : public Operator<HuberLoss> {
14 public:
15 HuberLoss(const Tensor input, const Tensor target,
16 const std::optional<Tensor> weight, const std::string reduction,
17 const double delta, 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::FromString(reduction)},
28 delta_{delta},
29 device_index_{out.device().index()} {
30 assert(!weight.has_value() &&
31 "The current ATen Huber loss ABI does not support `weight`; pass "
32 "`std::nullopt`.");
33 }
34
37 [[deprecated("Use the Python-compatible constructor instead.")]]
38 HuberLoss(const Tensor input, const Tensor target, const int64_t reduction,
39 const double delta, Tensor out)
40 : input_shape_{input.shape()},
41 input_strides_{input.strides()},
42 input_type_{input.dtype()},
43 target_shape_{target.shape()},
44 target_strides_{target.strides()},
45 target_type_{target.dtype()},
46 out_shape_{out.shape()},
47 out_strides_{out.strides()},
48 out_type_{out.dtype()},
49 reduction_{reduction},
50 delta_{delta},
51 device_index_{out.device().index()} {}
52
53 void operator()(const Tensor input, const Tensor target,
54 const std::optional<Tensor> weight,
55 const std::string reduction, const double delta,
56 Tensor out) const {
57 assert(!weight.has_value() &&
58 "The current ATen Huber loss ABI does not support `weight`; pass "
59 "`std::nullopt`.");
60 (*this)(input, target, reduction_detail::FromString(reduction), delta, out);
61 }
62
65 [[deprecated("Use the Python-compatible call overload instead.")]]
66 virtual void operator()(const Tensor input, const Tensor target,
67 const int64_t reduction, const double delta,
68 Tensor out) const = 0;
69
70 protected:
71 Tensor::Shape input_shape_;
72
73 Tensor::Strides input_strides_;
74
75 DataType input_type_;
76
77 Tensor::Shape target_shape_;
78
79 Tensor::Strides target_strides_;
80
81 DataType target_type_;
82
83 Tensor::Shape out_shape_;
84
85 Tensor::Strides out_strides_;
86
87 DataType out_type_;
88
89 int64_t reduction_{};
90
91 double delta_{};
92
94};
95
96} // namespace infini::ops
97
98#endif
Definition huber_loss.h:13
Tensor::Shape out_shape_
Definition huber_loss.h:83
void operator()(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::string reduction, const double delta, Tensor out) const
Definition huber_loss.h:53
Tensor::Shape input_shape_
Definition huber_loss.h:71
DataType input_type_
Definition huber_loss.h:75
int64_t reduction_
Definition huber_loss.h:89
DataType out_type_
Definition huber_loss.h:87
int device_index_
Definition huber_loss.h:93
HuberLoss(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::string reduction, const double delta, Tensor out)
Definition huber_loss.h:15
Tensor::Strides target_strides_
Definition huber_loss.h:79
double delta_
Definition huber_loss.h:91
virtual void operator()(const Tensor input, const Tensor target, const int64_t reduction, const double delta, Tensor out) const =0
HuberLoss(const Tensor input, const Tensor target, const int64_t reduction, const double delta, Tensor out)
Definition huber_loss.h:38
Tensor::Strides out_strides_
Definition huber_loss.h:85
DataType target_type_
Definition huber_loss.h:81
Tensor::Shape target_shape_
Definition huber_loss.h:77
Tensor::Strides input_strides_
Definition huber_loss.h:73
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8