InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
multi_margin_loss.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MULTI_MARGIN_LOSS_H_
2#define INFINI_OPS_BASE_MULTI_MARGIN_LOSS_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 MultiMarginLoss : public Operator<MultiMarginLoss> {
13 public:
14 MultiMarginLoss(const Tensor input, const Tensor target,
15 const std::optional<Tensor> weight, const int64_t p,
16 const double margin, const std::optional<bool> size_average,
17 const std::optional<bool> reduce, const std::string reduction,
18 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 p_{static_cast<double>(p)},
35 margin_{margin},
36 reduction_{reduction_detail::FromPythonArguments(size_average, reduce,
37 reduction)},
38 device_index_{out.device().index()} {}
39
43 [[deprecated(
44 "Use the overload with `weight` before integer `p` and string "
45 "`reduction`.")]]
46 MultiMarginLoss(const Tensor input, const Tensor target, const double p,
47 const double margin, const std::optional<Tensor> weight,
48 const int64_t reduction, Tensor out)
49 : input_shape_{input.shape()},
50 input_strides_{input.strides()},
51 input_type_{input.dtype()},
52 target_shape_{target.shape()},
53 target_strides_{target.strides()},
54 target_type_{target.dtype()},
55 out_shape_{out.shape()},
56 out_strides_{out.strides()},
57 out_type_{out.dtype()},
58 has_weight_{weight.has_value()},
59 weight_shape_{weight ? Tensor::Shape{weight->shape()}
60 : Tensor::Shape{}},
61 weight_strides_{weight ? Tensor::Strides{weight->strides()}
62 : Tensor::Strides{}},
63 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
64 p_{p},
65 margin_{margin},
66 reduction_{reduction},
67 device_index_{out.device().index()} {}
68
69 void operator()(const Tensor input, const Tensor target,
70 const std::optional<Tensor> weight, const int64_t p,
71 const double margin, const std::optional<bool> size_average,
72 const std::optional<bool> reduce, const std::string reduction,
73 Tensor out) const {
74 return operator()(
75 input, target, static_cast<double>(p), margin, weight,
76 reduction_detail::FromPythonArguments(size_average, reduce, reduction),
77 out);
78 }
79
82 [[deprecated(
83 "Use the overload with `weight` before integer `p` and string "
84 "`reduction`.")]]
85 virtual void operator()(const Tensor input, const Tensor target,
86 const double p, const double margin,
87 const std::optional<Tensor> weight,
88 const int64_t reduction, Tensor out) const = 0;
89
90 protected:
91 Tensor::Shape input_shape_;
92
93 Tensor::Strides input_strides_;
94
95 DataType input_type_;
96
97 Tensor::Shape target_shape_;
98
99 Tensor::Strides target_strides_;
100
101 DataType target_type_;
102
103 Tensor::Shape out_shape_;
104
105 Tensor::Strides out_strides_;
106
107 DataType out_type_;
108
109 bool has_weight_{false};
110
111 Tensor::Shape weight_shape_;
112
113 Tensor::Strides weight_strides_;
114
115 DataType weight_type_{DataType::kFloat32};
116
117 double p_{};
118
119 double margin_{};
120
121 int64_t reduction_{};
122
124};
125
126} // namespace infini::ops
127
128#endif
Definition multi_margin_loss.h:12
Tensor::Shape target_shape_
Definition multi_margin_loss.h:97
double p_
Definition multi_margin_loss.h:117
double margin_
Definition multi_margin_loss.h:119
Tensor::Strides out_strides_
Definition multi_margin_loss.h:105
MultiMarginLoss(const Tensor input, const Tensor target, const double p, const double margin, const std::optional< Tensor > weight, const int64_t reduction, Tensor out)
Definition multi_margin_loss.h:46
Tensor::Shape weight_shape_
Definition multi_margin_loss.h:111
Tensor::Strides weight_strides_
Definition multi_margin_loss.h:113
virtual void operator()(const Tensor input, const Tensor target, const double p, const double margin, const std::optional< Tensor > weight, const int64_t reduction, Tensor out) const =0
Tensor::Shape out_shape_
Definition multi_margin_loss.h:103
bool has_weight_
Definition multi_margin_loss.h:109
Tensor::Strides target_strides_
Definition multi_margin_loss.h:99
int device_index_
Definition multi_margin_loss.h:123
void operator()(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t p, const double margin, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out) const
Definition multi_margin_loss.h:69
DataType weight_type_
Definition multi_margin_loss.h:115
DataType input_type_
Definition multi_margin_loss.h:95
Tensor::Strides input_strides_
Definition multi_margin_loss.h:93
Tensor::Shape input_shape_
Definition multi_margin_loss.h:91
DataType target_type_
Definition multi_margin_loss.h:101
int64_t reduction_
Definition multi_margin_loss.h:121
MultiMarginLoss(const Tensor input, const Tensor target, const std::optional< Tensor > weight, const int64_t p, const double margin, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out)
Definition multi_margin_loss.h:14
DataType out_type_
Definition multi_margin_loss.h:107
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8