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