InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
nll_loss2d.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_NLL_LOSS2D_H_
2#define INFINI_OPS_BASE_NLL_LOSS2D_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class NllLoss2d : public Operator<NllLoss2d> {
11 public:
12 NllLoss2d(const Tensor input, const Tensor target,
13 const std::optional<Tensor> weight, const int64_t reduction,
14 const int64_t ignore_index, Tensor out)
15 : input_shape_{input.shape()},
16 input_strides_{input.strides()},
17 input_type_{input.dtype()},
18 target_shape_{target.shape()},
19 target_strides_{target.strides()},
20 target_type_{target.dtype()},
21 out_shape_{out.shape()},
22 out_strides_{out.strides()},
23 out_type_{out.dtype()},
24 has_weight_{weight.has_value()},
25 weight_shape_{weight ? Tensor::Shape{weight->shape()}
26 : Tensor::Shape{}},
27 weight_strides_{weight ? Tensor::Strides{weight->strides()}
28 : Tensor::Strides{}},
29 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
30 reduction_{reduction},
31 ignore_index_{ignore_index},
32 device_index_{out.device().index()} {}
33
34 virtual void operator()(const Tensor input, const Tensor target,
35 const std::optional<Tensor> weight,
36 const int64_t reduction, const int64_t ignore_index,
37 Tensor out) const = 0;
38
39 protected:
40 Tensor::Shape input_shape_;
41
42 Tensor::Strides input_strides_;
43
44 DataType input_type_;
45
46 Tensor::Shape target_shape_;
47
48 Tensor::Strides target_strides_;
49
50 DataType target_type_;
51
52 Tensor::Shape out_shape_;
53
54 Tensor::Strides out_strides_;
55
56 DataType out_type_;
57
58 bool has_weight_{false};
59
60 Tensor::Shape weight_shape_;
61
62 Tensor::Strides weight_strides_;
63
64 DataType weight_type_{DataType::kFloat32};
65
66 int64_t reduction_{};
67
68 int64_t ignore_index_{};
69
71};
72
73} // namespace infini::ops
74
75#endif
Definition nll_loss2d.h:10
DataType target_type_
Definition nll_loss2d.h:50
Tensor::Strides target_strides_
Definition nll_loss2d.h:48
int64_t ignore_index_
Definition nll_loss2d.h:68
Tensor::Strides weight_strides_
Definition nll_loss2d.h:62
Tensor::Strides out_strides_
Definition nll_loss2d.h:54
DataType weight_type_
Definition nll_loss2d.h:64
Tensor::Shape out_shape_
Definition nll_loss2d.h:52
int64_t reduction_
Definition nll_loss2d.h:66
NllLoss2d(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t reduction, const int64_t ignore_index, Tensor out)
Definition nll_loss2d.h:12
DataType input_type_
Definition nll_loss2d.h:44
Tensor::Shape weight_shape_
Definition nll_loss2d.h:60
int device_index_
Definition nll_loss2d.h:70
Tensor::Shape input_shape_
Definition nll_loss2d.h:40
bool has_weight_
Definition nll_loss2d.h:58
DataType out_type_
Definition nll_loss2d.h:56
Tensor::Strides input_strides_
Definition nll_loss2d.h:42
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
Tensor::Shape target_shape_
Definition nll_loss2d.h:46
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8