MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Initialize

Method Initialize

tensorflow/stream_executor/cuda/cuda_fft.cc:76–231  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

74} // namespace
75
76port::Status CUDAFftPlan::Initialize(
77 GpuExecutor *parent, Stream *stream, int rank, uint64 *elem_count,
78 uint64 *input_embed, uint64 input_stride, uint64 input_distance,
79 uint64 *output_embed, uint64 output_stride, uint64 output_distance,
80 fft::Type type, int batch_count, ScratchAllocator *scratch_allocator) {
81 if (IsInitialized()) {
82 LOG(FATAL) << "Try to repeatedly initialize.";
83 }
84 is_initialized_ = true;
85 cuda::ScopedActivateExecutorContext sac(parent);
86 int elem_count_[3], input_embed_[3], output_embed_[3];
87 for (int i = 0; i < rank; ++i) {
88 elem_count_[i] = elem_count[i];
89 if (input_embed) {
90 input_embed_[i] = input_embed[i];
91 }
92 if (output_embed) {
93 output_embed_[i] = output_embed[i];
94 }
95 }
96 parent_ = parent;
97 fft_type_ = type;
98 if (batch_count == 1 && input_embed == nullptr && output_embed == nullptr) {
99 cufftResult_t ret;
100 if (scratch_allocator == nullptr) {
101 switch (rank) {
102 case 1:
103 // cufftPlan1d
104 ret = cufftPlan1d(&plan_, elem_count_[0], CUDAFftType(type),
105 1 /* = batch */);
106 if (ret != CUFFT_SUCCESS) {
107 LOG(ERROR) << "failed to create cuFFT 1d plan:" << ret;
108 return port::Status(port::error::INTERNAL,
109 "Failed to create cuFFT 1d plan.");
110 }
111 return port::Status::OK();
112 case 2:
113 // cufftPlan2d
114 ret = cufftPlan2d(&plan_, elem_count_[0], elem_count_[1],
115 CUDAFftType(type));
116 if (ret != CUFFT_SUCCESS) {
117 LOG(ERROR) << "failed to create cuFFT 2d plan:" << ret;
118 return port::Status(port::error::INTERNAL,
119 "Failed to create cuFFT 2d plan.");
120 }
121 return port::Status::OK();
122 case 3:
123 // cufftPlan3d
124 ret = cufftPlan3d(&plan_, elem_count_[0], elem_count_[1],
125 elem_count_[2], CUDAFftType(type));
126 if (ret != CUFFT_SUCCESS) {
127 LOG(ERROR) << "failed to create cuFFT 3d plan:" << ret;
128 return port::Status(port::error::INTERNAL,
129 "Failed to create cuFFT 3d plan.");
130 }
131 return port::Status::OK();
132 default:
133 LOG(ERROR) << "Invalid rank value for cufftPlan. "

Calls 3

CUDAFftTypeFunction · 0.85
StatusEnum · 0.50
InitializeFunction · 0.50

Tested by

no test coverage detected