InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
soft_margin_loss.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SOFT_MARGIN_LOSS_H_
2#define INFINI_OPS_BASE_SOFT_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 SoftMarginLoss : public Operator<SoftMarginLoss> {
13 public:
14 SoftMarginLoss(const Tensor input, const Tensor target,
15 const std::optional<bool> size_average,
16 const std::optional<bool> reduce, const std::string reduction,
17 Tensor out)
18 : input_shape_{input.shape()},
19 input_strides_{input.strides()},
20 input_type_{input.dtype()},
21 target_shape_{target.shape()},
22 target_strides_{target.strides()},
23 target_type_{target.dtype()},
24 out_shape_{out.shape()},
25 out_strides_{out.strides()},
26 out_type_{out.dtype()},
27 reduction_{reduction_detail::FromPythonArguments(size_average, reduce,
28 reduction)},
29 device_index_{out.device().index()} {}
30
33 [[deprecated("Use the Python-compatible reduction overload instead.")]]
34 SoftMarginLoss(const Tensor input, const Tensor target,
35 const int64_t reduction, Tensor out)
36 : input_shape_{input.shape()},
37 input_strides_{input.strides()},
38 input_type_{input.dtype()},
39 target_shape_{target.shape()},
40 target_strides_{target.strides()},
41 target_type_{target.dtype()},
42 out_shape_{out.shape()},
43 out_strides_{out.strides()},
44 out_type_{out.dtype()},
45 reduction_{reduction},
46 device_index_{out.device().index()} {}
47
48 void operator()(const Tensor input, const Tensor target,
49 const std::optional<bool> size_average,
50 const std::optional<bool> reduce, const std::string reduction,
51 Tensor out) const {
52 (*this)(
53 input, target,
54 reduction_detail::FromPythonArguments(size_average, reduce, reduction),
55 out);
56 }
57
60 [[deprecated("Use the Python-compatible reduction overload instead.")]]
61 virtual void operator()(const Tensor input, const Tensor target,
62 const int64_t reduction, Tensor out) const = 0;
63
64 protected:
65 Tensor::Shape input_shape_;
66
67 Tensor::Strides input_strides_;
68
69 DataType input_type_;
70
71 Tensor::Shape target_shape_;
72
73 Tensor::Strides target_strides_;
74
75 DataType target_type_;
76
77 Tensor::Shape out_shape_;
78
79 Tensor::Strides out_strides_;
80
81 DataType out_type_;
82
83 int64_t reduction_{};
84
86};
87
88} // namespace infini::ops
89
90#endif
Definition generated/include/operator.h:282
Definition soft_margin_loss.h:12
Tensor::Shape input_shape_
Definition soft_margin_loss.h:65
Tensor::Strides target_strides_
Definition soft_margin_loss.h:73
DataType input_type_
Definition soft_margin_loss.h:69
SoftMarginLoss(const Tensor input, const Tensor target, const int64_t reduction, Tensor out)
Definition soft_margin_loss.h:34
DataType out_type_
Definition soft_margin_loss.h:81
Tensor::Strides input_strides_
Definition soft_margin_loss.h:67
Tensor::Shape out_shape_
Definition soft_margin_loss.h:77
void operator()(const Tensor input, const Tensor target, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out) const
Definition soft_margin_loss.h:48
SoftMarginLoss(const Tensor input, const Tensor target, const std::optional< bool > size_average, const std::optional< bool > reduce, const std::string reduction, Tensor out)
Definition soft_margin_loss.h:14
DataType target_type_
Definition soft_margin_loss.h:75
virtual void operator()(const Tensor input, const Tensor target, const int64_t reduction, Tensor out) const =0
Tensor::Strides out_strides_
Definition soft_margin_loss.h:79
int device_index_
Definition soft_margin_loss.h:85
Tensor::Shape target_shape_
Definition soft_margin_loss.h:71
int64_t reduction_
Definition soft_margin_loss.h:83
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8