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