InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
special_zeta.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SPECIAL_ZETA_H_
2#define INFINI_OPS_BASE_SPECIAL_ZETA_H_
3
4#include "operator.h"
5
6namespace infini::ops::special {
7
8class Zeta : public Operator<Zeta> {
9 public:
10 Zeta(const Tensor input, const Tensor other, Tensor out)
11 : input_shape_{input.shape()},
12 input_strides_{input.strides()},
13 input_type_{input.dtype()},
14 other_shape_{other.shape()},
15 other_strides_{other.strides()},
16 other_type_{other.dtype()},
17 out_shape_{out.shape()},
18 out_strides_{out.strides()},
19 out_type_{out.dtype()},
20 device_index_{out.device().index()} {}
21
22 Zeta(const Tensor input, const double other, Tensor out)
23 : input_shape_{input.shape()},
24 input_strides_{input.strides()},
25 input_type_{input.dtype()},
26 out_shape_{out.shape()},
27 out_strides_{out.strides()},
28 out_type_{out.dtype()},
29 other_{other},
30 device_index_{out.device().index()} {}
31
32 virtual void operator()(const Tensor input, const Tensor other,
33 Tensor out) const = 0;
34
35 virtual void operator()(const Tensor input, const double other,
36 Tensor out) const = 0;
37
38 protected:
39 Tensor::Shape input_shape_;
40
41 Tensor::Strides input_strides_;
42
43 DataType input_type_;
44
45 Tensor::Shape other_shape_;
46
47 Tensor::Strides other_strides_;
48
49 DataType other_type_;
50
51 Tensor::Shape out_shape_;
52
53 Tensor::Strides out_strides_;
54
55 DataType out_type_;
56
57 double other_{};
58
60};
61
62} // namespace infini::ops::special
63
64#endif
Definition generated/include/operator.h:282
Definition special_zeta.h:8
Tensor::Strides out_strides_
Definition special_zeta.h:53
virtual void operator()(const Tensor input, const Tensor other, Tensor out) const =0
Zeta(const Tensor input, const double other, Tensor out)
Definition special_zeta.h:22
DataType out_type_
Definition special_zeta.h:55
DataType input_type_
Definition special_zeta.h:43
Tensor::Strides other_strides_
Definition special_zeta.h:47
int device_index_
Definition special_zeta.h:59
Tensor::Shape other_shape_
Definition special_zeta.h:45
Tensor::Strides input_strides_
Definition special_zeta.h:41
Tensor::Shape out_shape_
Definition special_zeta.h:51
DataType other_type_
Definition special_zeta.h:49
virtual void operator()(const Tensor input, const double other, Tensor out) const =0
Zeta(const Tensor input, const Tensor other, Tensor out)
Definition special_zeta.h:10
Tensor::Shape input_shape_
Definition special_zeta.h:39
double other_
Definition special_zeta.h:57
Definition special_airy_ai.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8