InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
fft_rfftn.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_FFT_RFFTN_H_
2#define INFINI_OPS_BASE_FFT_RFFTN_H_
3
4#include <optional>
5#include <string>
6#include <vector>
7
8#include "operator.h"
9
10namespace infini::ops::fft {
11
12class Rfftn : public Operator<Rfftn> {
13 public:
14 Rfftn(const Tensor input, const std::optional<std::vector<int64_t>> s,
15 const std::optional<std::vector<int64_t>> dim,
16 const std::optional<std::string> norm, Tensor out)
17 : input_shape_{input.shape()},
18 input_strides_{input.strides()},
19 input_type_{input.dtype()},
20 out_shape_{out.shape()},
21 out_strides_{out.strides()},
22 out_type_{out.dtype()},
23 s_{s},
24 dim_{dim},
25 norm_{norm},
26 device_index_{out.device().index()} {}
27
28 virtual void operator()(const Tensor input,
29 const std::optional<std::vector<int64_t>> s,
30 const std::optional<std::vector<int64_t>> dim,
31 const std::optional<std::string> norm,
32 Tensor out) const = 0;
33
34 protected:
35 Tensor::Shape input_shape_;
36
37 Tensor::Strides input_strides_;
38
39 DataType input_type_;
40
41 Tensor::Shape out_shape_;
42
43 Tensor::Strides out_strides_;
44
45 DataType out_type_;
46
47 std::optional<std::vector<int64_t>> s_{};
48
49 std::optional<std::vector<int64_t>> dim_{};
50
51 std::optional<std::string> norm_{};
52
54};
55
56} // namespace infini::ops::fft
57
58#endif
Definition generated/include/operator.h:282
Definition fft_rfftn.h:12
int device_index_
Definition fft_rfftn.h:53
DataType input_type_
Definition fft_rfftn.h:39
std::optional< std::string > norm_
Definition fft_rfftn.h:51
Tensor::Strides input_strides_
Definition fft_rfftn.h:37
virtual void operator()(const Tensor input, const std::optional< std::vector< int64_t > > s, const std::optional< std::vector< int64_t > > dim, const std::optional< std::string > norm, Tensor out) const =0
std::optional< std::vector< int64_t > > s_
Definition fft_rfftn.h:47
DataType out_type_
Definition fft_rfftn.h:45
Tensor::Shape out_shape_
Definition fft_rfftn.h:41
Tensor::Shape input_shape_
Definition fft_rfftn.h:35
Rfftn(const Tensor input, const std::optional< std::vector< int64_t > > s, const std::optional< std::vector< int64_t > > dim, const std::optional< std::string > norm, Tensor out)
Definition fft_rfftn.h:14
Tensor::Strides out_strides_
Definition fft_rfftn.h:43
std::optional< std::vector< int64_t > > dim_
Definition fft_rfftn.h:49
Definition fft_fft.h:9
infini::rt::TensorView Tensor
Definition tensor.h:8