1#ifndef INFINI_OPS_BASE_HARDSHRINK_BACKWARD_H_
2#define INFINI_OPS_BASE_HARDSHRINK_BACKWARD_H_
11 const double lambd,
Tensor grad_input)
25 const double lambd,
Tensor grad_input)
const = 0;
Definition hardshrink_backward.h:8
Tensor::Shape input_shape_
Definition hardshrink_backward.h:34
DataType input_type_
Definition hardshrink_backward.h:38
Tensor::Shape grad_out_shape_
Definition hardshrink_backward.h:28
double lambd_
Definition hardshrink_backward.h:46
Tensor::Strides grad_out_strides_
Definition hardshrink_backward.h:30
virtual void operator()(const Tensor grad_out, const Tensor input, const double lambd, Tensor grad_input) const =0
Tensor::Strides input_strides_
Definition hardshrink_backward.h:36
HardshrinkBackward(const Tensor grad_out, const Tensor input, const double lambd, Tensor grad_input)
Definition hardshrink_backward.h:10
DataType grad_input_type_
Definition hardshrink_backward.h:44
Tensor::Strides grad_input_strides_
Definition hardshrink_backward.h:42
int device_index_
Definition hardshrink_backward.h:48
Tensor::Shape grad_input_shape_
Definition hardshrink_backward.h:40
DataType grad_out_type_
Definition hardshrink_backward.h:32
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8