InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
binary_cross_entropy.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_BINARY_CROSS_ENTROPY_H_
2#define INFINI_OPS_BASE_BINARY_CROSS_ENTROPY_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 BinaryCrossEntropy : public Operator<BinaryCrossEntropy> {
13 public:
14 BinaryCrossEntropy(const Tensor input, const Tensor target,
15 const std::optional<Tensor> weight,
16 const std::optional<bool> size_average,
17 const std::optional<bool> reduce,
18 const std::string reduction, 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 device_index_{out.device().index()} {}
37
40 [[deprecated("Use the Python-compatible reduction overload instead.")]]
41 BinaryCrossEntropy(const Tensor input, const Tensor target,
42 const std::optional<Tensor> weight,
43 const int64_t reduction, 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 device_index_{out.device().index()} {}
61
62 void operator()(const Tensor input, const Tensor target,
63 const std::optional<Tensor> weight,
64 const std::optional<bool> size_average,
65 const std::optional<bool> reduce, const std::string reduction,
66 Tensor out) const {
67 (*this)(
68 input, target, weight,
69 reduction_detail::FromPythonArguments(size_average, reduce, reduction),
70 out);
71 }
72
75 [[deprecated("Use the Python-compatible reduction overload instead.")]]
76 virtual void operator()(const Tensor input, const Tensor target,
77 const std::optional<Tensor> weight,
78 const int64_t reduction, Tensor out) const = 0;
79
80 protected:
81 Tensor::Shape input_shape_;
82
83 Tensor::Strides input_strides_;
84
85 DataType input_type_;
86
87 Tensor::Shape target_shape_;
88
89 Tensor::Strides target_strides_;
90
91 DataType target_type_;
92
93 Tensor::Shape out_shape_;
94
95 Tensor::Strides out_strides_;
96
97 DataType out_type_;
98
99 bool has_weight_{false};
100
101 Tensor::Shape weight_shape_;
102
103 Tensor::Strides weight_strides_;
104
105 DataType weight_type_{DataType::kFloat32};
106
107 int64_t reduction_{};
108
110};
111
112} // namespace infini::ops
113
114#endif
Definition binary_cross_entropy.h:12
DataType weight_type_
Definition binary_cross_entropy.h:105
Tensor::Strides input_strides_
Definition binary_cross_entropy.h:83
Tensor::Strides weight_strides_
Definition binary_cross_entropy.h:103
DataType target_type_
Definition binary_cross_entropy.h:91
Tensor::Shape target_shape_
Definition binary_cross_entropy.h:87
Tensor::Strides out_strides_
Definition binary_cross_entropy.h:95
Tensor::Strides target_strides_
Definition binary_cross_entropy.h:89
virtual void operator()(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t reduction, Tensor out) const =0
int device_index_
Definition binary_cross_entropy.h:109
void operator()(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out) const
Definition binary_cross_entropy.h:62
BinaryCrossEntropy(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t reduction, Tensor out)
Definition binary_cross_entropy.h:41
DataType out_type_
Definition binary_cross_entropy.h:97
int64_t reduction_
Definition binary_cross_entropy.h:107
DataType input_type_
Definition binary_cross_entropy.h:85
Tensor::Shape weight_shape_
Definition binary_cross_entropy.h:101
Tensor::Shape out_shape_
Definition binary_cross_entropy.h:93
BinaryCrossEntropy(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out)
Definition binary_cross_entropy.h:14
bool has_weight_
Definition binary_cross_entropy.h:99
Tensor::Shape input_shape_
Definition binary_cross_entropy.h:81
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8