InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
fft_irfft2.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_FFT_IRFFT2_H_
2#define INFINI_OPS_BASE_FFT_IRFFT2_H_
3
4#include <optional>
5#include <string>
6#include <vector>
7
8#include "operator.h"
9
10namespace infini::ops::fft {
11
12class Irfft2 : public Operator<Irfft2> {
13 public:
14 Irfft2(const Tensor input, const std::optional<std::vector<int64_t>> s,
15 const std::vector<int64_t> dim, const std::optional<std::string> norm,
16 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::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::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_irfft2.h:12
std::optional< std::vector< int64_t > > s_
Definition fft_irfft2.h:47
std::vector< int64_t > dim_
Definition fft_irfft2.h:49
Tensor::Strides input_strides_
Definition fft_irfft2.h:37
virtual void operator()(const Tensor input, const std::optional< std::vector< int64_t > > s, const std::vector< int64_t > dim, const std::optional< std::string > norm, Tensor out) const =0
Irfft2(const Tensor input, const std::optional< std::vector< int64_t > > s, const std::vector< int64_t > dim, const std::optional< std::string > norm, Tensor out)
Definition fft_irfft2.h:14
int device_index_
Definition fft_irfft2.h:53
Tensor::Strides out_strides_
Definition fft_irfft2.h:43
DataType out_type_
Definition fft_irfft2.h:45
std::optional< std::string > norm_
Definition fft_irfft2.h:51
DataType input_type_
Definition fft_irfft2.h:39
Tensor::Shape input_shape_
Definition fft_irfft2.h:35
Tensor::Shape out_shape_
Definition fft_irfft2.h:41
Definition fft_fft.h:9
infini::rt::TensorView Tensor
Definition tensor.h:8