InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
nll_loss.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_NLL_LOSS_H_
2#define INFINI_OPS_BASE_NLL_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 NllLoss : public Operator<NllLoss> {
13 public:
14 NllLoss(const Tensor input, const Tensor target,
15 const std::optional<Tensor> weight,
16 const std::optional<bool> size_average, const int64_t ignore_index,
17 const std::optional<bool> reduce, const std::string reduction,
18 Tensor out)
19 : input_shape_{input.shape()},
20 input_strides_{input.strides()},
21 input_type_{input.dtype()},
22 target_shape_{target.shape()},
23 target_strides_{target.strides()},
24 target_type_{target.dtype()},
25 out_shape_{out.shape()},
26 out_strides_{out.strides()},
27 out_type_{out.dtype()},
28 has_weight_{weight.has_value()},
29 weight_shape_{weight ? Tensor::Shape{weight->shape()}
30 : Tensor::Shape{}},
31 weight_strides_{weight ? Tensor::Strides{weight->strides()}
32 : Tensor::Strides{}},
33 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
34 reduction_{reduction_detail::FromPythonArguments(size_average, reduce,
35 reduction)},
36 ignore_index_{ignore_index},
37 device_index_{out.device().index()} {}
38
39 [[deprecated(
40 "Use the overload with `ignore_index` before string `reduction`.")]]
41 NllLoss(const Tensor input, const Tensor target,
42 const std::optional<Tensor> weight, const int64_t reduction,
43 const int64_t ignore_index, Tensor out)
44 : input_shape_{input.shape()},
45 input_strides_{input.strides()},
46 input_type_{input.dtype()},
47 target_shape_{target.shape()},
48 target_strides_{target.strides()},
49 target_type_{target.dtype()},
50 out_shape_{out.shape()},
51 out_strides_{out.strides()},
52 out_type_{out.dtype()},
53 has_weight_{weight.has_value()},
54 weight_shape_{weight ? Tensor::Shape{weight->shape()}
55 : Tensor::Shape{}},
56 weight_strides_{weight ? Tensor::Strides{weight->strides()}
57 : Tensor::Strides{}},
58 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
59 reduction_{reduction},
60 ignore_index_{ignore_index},
61 device_index_{out.device().index()} {}
62
63 void operator()(const Tensor input, const Tensor target,
64 const std::optional<Tensor> weight,
65 const std::optional<bool> size_average,
66 const int64_t ignore_index, const std::optional<bool> reduce,
67 const std::string reduction, Tensor out) const {
68 return operator()(
69 input, target, weight,
70 reduction_detail::FromPythonArguments(size_average, reduce, reduction),
71 ignore_index, out);
72 }
73
74 [[deprecated(
75 "Use the overload with `ignore_index` before string `reduction`.")]]
76 virtual void operator()(const Tensor input, const Tensor target,
77 const std::optional<Tensor> weight,
78 const int64_t reduction, const int64_t ignore_index,
79 Tensor out) const = 0;
80
81 protected:
82 Tensor::Shape input_shape_;
83
84 Tensor::Strides input_strides_;
85
86 DataType input_type_;
87
88 Tensor::Shape target_shape_;
89
90 Tensor::Strides target_strides_;
91
92 DataType target_type_;
93
94 Tensor::Shape out_shape_;
95
96 Tensor::Strides out_strides_;
97
98 DataType out_type_;
99
100 bool has_weight_{false};
101
102 Tensor::Shape weight_shape_;
103
104 Tensor::Strides weight_strides_;
105
106 DataType weight_type_{DataType::kFloat32};
107
108 int64_t reduction_{};
109
110 int64_t ignore_index_{};
111
113};
114
115} // namespace infini::ops
116
117#endif
Definition nll_loss.h:12
NllLoss(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t reduction, const int64_t ignore_index, Tensor out)
Definition nll_loss.h:41
int64_t reduction_
Definition nll_loss.h:108
int device_index_
Definition nll_loss.h:112
DataType out_type_
Definition nll_loss.h:98
Tensor::Strides weight_strides_
Definition nll_loss.h:104
DataType input_type_
Definition nll_loss.h:86
Tensor::Shape out_shape_
Definition nll_loss.h:94
void operator()(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::optional< bool > size_average, const int64_t ignore_index, const std::optional< bool > reduce, const std::string reduction, Tensor out) const
Definition nll_loss.h:63
bool has_weight_
Definition nll_loss.h:100
DataType target_type_
Definition nll_loss.h:92
Tensor::Strides target_strides_
Definition nll_loss.h:90
Tensor::Strides out_strides_
Definition nll_loss.h:96
Tensor::Shape target_shape_
Definition nll_loss.h:88
Tensor::Shape weight_shape_
Definition nll_loss.h:102
Tensor::Shape input_shape_
Definition nll_loss.h:82
NllLoss(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::optional< bool > size_average, const int64_t ignore_index, const std::optional< bool > reduce, const std::string reduction, Tensor out)
Definition nll_loss.h:14
int64_t ignore_index_
Definition nll_loss.h:110
virtual void operator()(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t reduction, const int64_t ignore_index, Tensor out) const =0
DataType weight_type_
Definition nll_loss.h:106
Tensor::Strides input_strides_
Definition nll_loss.h:84
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8