1#ifndef INFINI_OPS_BASE_MULTI_MARGIN_LOSS_H_
2#define INFINI_OPS_BASE_MULTI_MARGIN_LOSS_H_
7#include "common/op_utils/reduction.h"
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,
33 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
34 p_{static_cast<double>(p)},
36 reduction_{reduction_detail::FromPythonArguments(size_average, reduce,
44 "Use the overload with `weight` before integer `p` and string "
47 const double margin,
const std::optional<Tensor> weight,
48 const int64_t reduction,
Tensor out)
63 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
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,
75 input, target,
static_cast<double>(p), margin, weight,
76 reduction_detail::FromPythonArguments(size_average, reduce, reduction),
83 "Use the overload with `weight` before integer `p` and string "
86 const double p,
const double margin,
87 const std::optional<Tensor> weight,
88 const int64_t reduction,
Tensor out)
const = 0;
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