| 386 | } |
| 387 | |
| 388 | std::unique_ptr<fft::Plan> CUDAFft::CreateBatchedPlan( |
| 389 | Stream *stream, int rank, uint64 *elem_count, uint64 *input_embed, |
| 390 | uint64 input_stride, uint64 input_distance, uint64 *output_embed, |
| 391 | uint64 output_stride, uint64 output_distance, fft::Type type, |
| 392 | bool in_place_fft, int batch_count) { |
| 393 | std::unique_ptr<CUDAFftPlan> fft_plan_ptr{new CUDAFftPlan()}; |
| 394 | port::Status status = fft_plan_ptr->Initialize( |
| 395 | parent_, stream, rank, elem_count, input_embed, input_stride, |
| 396 | input_distance, output_embed, output_stride, output_distance, type, |
| 397 | batch_count, /*scratch_allocator=*/nullptr); |
| 398 | if (!status.ok()) { |
| 399 | LOG(ERROR) << "Initialize Params: rank: " << rank |
| 400 | << " elem_count: " << *elem_count |
| 401 | << " input_embed: " << *input_embed |
| 402 | << " input_stride: " << input_stride |
| 403 | << " input_distance: " << input_distance |
| 404 | << " output_embed: " << *output_embed |
| 405 | << " output_stride: " << output_stride |
| 406 | << " output_distance: " << output_distance |
| 407 | << " batch_count: " << batch_count; |
| 408 | LOG(FATAL) << "failed to initialize batched cufft plan: " |
| 409 | << status.error_message(); |
| 410 | } |
| 411 | |
| 412 | return std::move(fft_plan_ptr); |
| 413 | } |
| 414 | |
| 415 | std::unique_ptr<fft::Plan> CUDAFft::CreateBatchedPlanWithScratchAllocator( |
| 416 | Stream *stream, int rank, uint64 *elem_count, uint64 *input_embed, |
nothing calls this directly
no test coverage detected