1#ifndef INFINI_OPS_BASE_MULTILABEL_MARGIN_LOSS_FORWARD_H_
2#define INFINI_OPS_BASE_MULTILABEL_MARGIN_LOSS_FORWARD_H_
9 :
public Operator<MultilabelMarginLossForward> {
12 const int64_t reduction,
Tensor output,
30 const int64_t reduction,
Tensor output,
31 Tensor is_target)
const = 0;
Definition multilabel_margin_loss_forward.h:9
int device_index_
Definition multilabel_margin_loss_forward.h:60
Tensor::Strides is_target_strides_
Definition multilabel_margin_loss_forward.h:54
MultilabelMarginLossForward(const Tensor input, const Tensor target, const int64_t reduction, Tensor output, Tensor is_target)
Definition multilabel_margin_loss_forward.h:11
DataType is_target_type_
Definition multilabel_margin_loss_forward.h:56
DataType output_type_
Definition multilabel_margin_loss_forward.h:50
Tensor::Shape output_shape_
Definition multilabel_margin_loss_forward.h:46
Tensor::Shape target_shape_
Definition multilabel_margin_loss_forward.h:40
DataType target_type_
Definition multilabel_margin_loss_forward.h:44
Tensor::Strides target_strides_
Definition multilabel_margin_loss_forward.h:42
Tensor::Strides output_strides_
Definition multilabel_margin_loss_forward.h:48
DataType input_type_
Definition multilabel_margin_loss_forward.h:38
virtual void operator()(const Tensor input, const Tensor target, const int64_t reduction, Tensor output, Tensor is_target) const =0
Tensor::Shape is_target_shape_
Definition multilabel_margin_loss_forward.h:52
Tensor::Strides input_strides_
Definition multilabel_margin_loss_forward.h:36
int64_t reduction_
Definition multilabel_margin_loss_forward.h:58
Tensor::Shape input_shape_
Definition multilabel_margin_loss_forward.h:34
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8