InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
mse_loss.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MSE_LOSS_H_
2#define INFINI_OPS_BASE_MSE_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 MseLoss : public Operator<MseLoss> {
14 public:
15 MseLoss(const Tensor input, const Tensor target,
16 const std::optional<Tensor> weight,
17 const std::optional<bool> size_average,
18 const std::optional<bool> reduce, const std::string reduction,
19 Tensor out)
20 : input_shape_{input.shape()},
21 input_strides_{input.strides()},
22 input_type_{input.dtype()},
23 target_shape_{target.shape()},
24 target_strides_{target.strides()},
25 target_type_{target.dtype()},
26 out_shape_{out.shape()},
27 out_strides_{out.strides()},
28 out_type_{out.dtype()},
29 reduction_{reduction_detail::FromPythonArguments(size_average, reduce,
30 reduction)},
31 device_index_{out.device().index()} {
32 assert(!weight.has_value() &&
33 "`weight` is unsupported because the current ATen `MseLoss` ABI "
34 "cannot implement the Python wrapper's weighted composition");
35 }
36
39 [[deprecated("Use the PyTorch-compatible overload instead.")]]
40 MseLoss(const Tensor input, const Tensor target, const int64_t reduction,
41 Tensor out)
42 : input_shape_{input.shape()},
43 input_strides_{input.strides()},
44 input_type_{input.dtype()},
45 target_shape_{target.shape()},
46 target_strides_{target.strides()},
47 target_type_{target.dtype()},
48 out_shape_{out.shape()},
49 out_strides_{out.strides()},
50 out_type_{out.dtype()},
51 reduction_{reduction},
52 device_index_{out.device().index()} {}
53
54 void operator()(const Tensor input, const Tensor target,
55 const std::optional<Tensor> weight,
56 const std::optional<bool> size_average,
57 const std::optional<bool> reduce, const std::string reduction,
58 Tensor out) const {
59 assert(!weight.has_value() &&
60 "`weight` is unsupported because the current ATen `MseLoss` ABI "
61 "cannot implement the Python wrapper's weighted composition");
62
63 return operator()(
64 input, target,
65 reduction_detail::FromPythonArguments(size_average, reduce, reduction),
66 out);
67 }
68
71 [[deprecated("Use the PyTorch-compatible overload instead.")]]
72 virtual void operator()(const Tensor input, const Tensor target,
73 const int64_t reduction, Tensor out) const = 0;
74
75 protected:
76 Tensor::Shape input_shape_;
77
78 Tensor::Strides input_strides_;
79
80 DataType input_type_;
81
82 Tensor::Shape target_shape_;
83
84 Tensor::Strides target_strides_;
85
86 DataType target_type_;
87
88 Tensor::Shape out_shape_;
89
90 Tensor::Strides out_strides_;
91
92 DataType out_type_;
93
94 int64_t reduction_{};
95
97};
98
99} // namespace infini::ops
100
101#endif
Definition mse_loss.h:13
Tensor::Shape target_shape_
Definition mse_loss.h:82
Tensor::Strides out_strides_
Definition mse_loss.h:90
MseLoss(const Tensor input, const Tensor target, const int64_t reduction, Tensor out)
Definition mse_loss.h:40
Tensor::Strides input_strides_
Definition mse_loss.h:78
DataType target_type_
Definition mse_loss.h:86
int64_t reduction_
Definition mse_loss.h:94
Tensor::Strides target_strides_
Definition mse_loss.h:84
MseLoss(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out)
Definition mse_loss.h:15
DataType out_type_
Definition mse_loss.h:92
Tensor::Shape input_shape_
Definition mse_loss.h:76
Tensor::Shape out_shape_
Definition mse_loss.h:88
void operator()(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out) const
Definition mse_loss.h:54
DataType input_type_
Definition mse_loss.h:80
int device_index_
Definition mse_loss.h:96
virtual void operator()(const Tensor input, const Tensor target, const int64_t reduction, Tensor out) const =0
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8