1#ifndef INFINI_OPS_BASE_LOGIT_BACKWARD_H_
2#define INFINI_OPS_BASE_LOGIT_BACKWARD_H_
13 const std::optional<double> eps,
Tensor grad_input)
27 const std::optional<double> eps,
28 Tensor grad_input)
const = 0;
49 std::optional<double>
eps_{};
Definition logit_backward.h:10
virtual void operator()(const Tensor grad_output, const Tensor input, const std::optional< double > eps, Tensor grad_input) const =0
Tensor::Shape grad_output_shape_
Definition logit_backward.h:31
DataType input_type_
Definition logit_backward.h:41
Tensor::Shape input_shape_
Definition logit_backward.h:37
Tensor::Shape grad_input_shape_
Definition logit_backward.h:43
int device_index_
Definition logit_backward.h:51
LogitBackward(const Tensor grad_output, const Tensor input, const std::optional< double > eps, Tensor grad_input)
Definition logit_backward.h:12
Tensor::Strides grad_input_strides_
Definition logit_backward.h:45
std::optional< double > eps_
Definition logit_backward.h:49
Tensor::Strides input_strides_
Definition logit_backward.h:39
DataType grad_input_type_
Definition logit_backward.h:47
Tensor::Strides grad_output_strides_
Definition logit_backward.h:33
DataType grad_output_type_
Definition logit_backward.h:35
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8