InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
binary_cross_entropy_backward.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_BINARY_CROSS_ENTROPY_BACKWARD_H_
2#define INFINI_OPS_BASE_BINARY_CROSS_ENTROPY_BACKWARD_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class BinaryCrossEntropyBackward : public Operator<BinaryCrossEntropyBackward> {
11 public:
12 BinaryCrossEntropyBackward(const Tensor grad_output, const Tensor input,
13 const Tensor target,
14 const std::optional<Tensor> weight,
15 const int64_t reduction, Tensor grad_input)
16 : grad_output_shape_{grad_output.shape()},
17 grad_output_strides_{grad_output.strides()},
18 grad_output_type_{grad_output.dtype()},
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 grad_input_shape_{grad_input.shape()},
26 grad_input_strides_{grad_input.strides()},
27 grad_input_type_{grad_input.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 device_index_{grad_input.device().index()} {}
36
37 virtual void operator()(const Tensor grad_output, const Tensor input,
38 const Tensor target,
39 const std::optional<Tensor> weight,
40 const int64_t reduction, Tensor grad_input) const = 0;
41
42 protected:
43 Tensor::Shape grad_output_shape_;
44
45 Tensor::Strides grad_output_strides_;
46
48
49 Tensor::Shape input_shape_;
50
51 Tensor::Strides input_strides_;
52
53 DataType input_type_;
54
55 Tensor::Shape target_shape_;
56
57 Tensor::Strides target_strides_;
58
59 DataType target_type_;
60
61 Tensor::Shape grad_input_shape_;
62
63 Tensor::Strides grad_input_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
78};
79
80} // namespace infini::ops
81
82#endif
Definition binary_cross_entropy_backward.h:10
DataType grad_input_type_
Definition binary_cross_entropy_backward.h:65
int64_t reduction_
Definition binary_cross_entropy_backward.h:75
Tensor::Strides target_strides_
Definition binary_cross_entropy_backward.h:57
virtual void operator()(const Tensor grad_output, const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t reduction, Tensor grad_input) const =0
DataType weight_type_
Definition binary_cross_entropy_backward.h:73
Tensor::Strides weight_strides_
Definition binary_cross_entropy_backward.h:71
Tensor::Shape weight_shape_
Definition binary_cross_entropy_backward.h:69
Tensor::Shape target_shape_
Definition binary_cross_entropy_backward.h:55
DataType grad_output_type_
Definition binary_cross_entropy_backward.h:47
BinaryCrossEntropyBackward(const Tensor grad_output, const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t reduction, Tensor grad_input)
Definition binary_cross_entropy_backward.h:12
DataType target_type_
Definition binary_cross_entropy_backward.h:59
Tensor::Strides input_strides_
Definition binary_cross_entropy_backward.h:51
Tensor::Strides grad_output_strides_
Definition binary_cross_entropy_backward.h:45
DataType input_type_
Definition binary_cross_entropy_backward.h:53
Tensor::Shape grad_output_shape_
Definition binary_cross_entropy_backward.h:43
Tensor::Strides grad_input_strides_
Definition binary_cross_entropy_backward.h:63
Tensor::Shape grad_input_shape_
Definition binary_cross_entropy_backward.h:61
bool has_weight_
Definition binary_cross_entropy_backward.h:67
Tensor::Shape input_shape_
Definition binary_cross_entropy_backward.h:49
int device_index_
Definition binary_cross_entropy_backward.h:77
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8