use std::collections::hash_map::DefaultHasher;
use std::ffi::c_void;
use std::hash::{Hash, Hasher};
use std::io::Write;
use std::mem::size_of;
use std::sync::Arc;
use tenferro_gpu::cuda::interop::{
alloc_device_bytes, with_raw_cuda_stream, CudaExternalUseReadLease, CudaExternalUseWriteLease,
DeviceByteBuffer,
};
use tenferro_gpu::cuda::CudaRuntime;
use tenferro_runtime::ExtensionCacheKey;
use tenferro_tensor::{Tensor, TensorScalar, TypedTensor};
use super::descriptor::{CufftDirection, CufftPlanDescriptor, CufftPlanKey, CufftTransformKind};
use super::error::CudaFftError;
use super::ffi::{
map_cufft_status, CufftApi, CufftHandle, CufftLibrary, CufftStatus, CUFFT_C2C, CUFFT_C2R,
CUFFT_D2Z, CUFFT_FORWARD, CUFFT_INVERSE, CUFFT_R2C, CUFFT_Z2D, CUFFT_Z2Z,
};
use crate::FFT_EXTENSION_FAMILY_ID;
const OP: &str = "cuda_fft";
pub(crate) const CUFFT_CACHE_NAMESPACE: &str = "cufft-plans";
pub(crate) trait CufftCleanup {
fn set_current(&self) -> Result<(), CudaFftError>;
fn synchronize(&self) -> Result<(), CudaFftError>;
}
impl CufftCleanup for CudaRuntime {
fn set_current(&self) -> Result<(), CudaFftError> {
self.set_current_cuda_context("cufft_plan_cleanup")
.map_err(|source| CudaFftError::interop("cufft_plan_cleanup_context", source))
}
fn synchronize(&self) -> Result<(), CudaFftError> {
CudaRuntime::synchronize(self)
.map_err(|source| CudaFftError::interop("cufft_plan_cleanup_stream", source))
}
}
pub(crate) trait CufftWorkspaceOwner: Sized {
fn empty() -> Self;
fn with_ptr(&self, f: impl FnOnce(*mut c_void));
}
pub(crate) struct CufftWorkspace {
_owner: DeviceByteBuffer,
bytes: usize,
}
impl CufftWorkspace {
fn from_device(owner: DeviceByteBuffer, bytes: usize) -> Self {
Self {
_owner: owner,
bytes,
}
}
fn bytes(&self) -> usize {
self.bytes
}
}
impl CufftWorkspaceOwner for CufftWorkspace {
fn empty() -> Self {
Self::from_device(DeviceByteBuffer::none(), 0)
}
fn with_ptr(&self, f: impl FnOnce(*mut c_void)) {
self._owner.with_ptr(f);
}
}
#[derive(Default)]
pub(crate) struct CleanupFailures {
pub(crate) synchronization: Vec<CudaFftError>,
pub(crate) destroy: Option<CudaFftError>,
pub(crate) resources_deferred: bool,
}
impl CleanupFailures {
#[cfg(test)]
pub(crate) fn is_empty(&self) -> bool {
self.synchronization.is_empty() && self.destroy.is_none() && !self.resources_deferred
}
}
pub(crate) fn retire_handle<R, W>(
runtime: &R,
library: &Arc<CufftLibrary>,
handle: &mut Option<CufftHandle>,
workspace: &mut Option<W>,
) -> CleanupFailures
where
R: CufftCleanup + Clone,
W: CufftWorkspaceOwner,
{
let mut failures = CleanupFailures::default();
if handle.is_some() || workspace.is_some() {
if let Err(error) = runtime.set_current() {
failures.synchronization.push(error);
failures.resources_deferred = true;
defer_resources(runtime, library, handle, workspace);
return failures;
}
if let Err(error) = runtime.synchronize() {
failures.synchronization.push(error);
failures.resources_deferred = true;
defer_resources(runtime, library, handle, workspace);
return failures;
}
}
if let Some(plan) = *handle {
let status = unsafe { (library.api.destroy)(plan) };
if let Err(error) = map_cufft_status("cufftDestroy", status) {
failures.destroy = Some(error);
failures.resources_deferred = true;
defer_resources(runtime, library, handle, workspace);
return failures;
}
handle.take();
}
let workspace = workspace.take();
drop(workspace);
failures
}
struct DeferredCufftResources<R, W> {
_handle: Option<CufftHandle>,
_workspace: Option<W>,
_library: Arc<CufftLibrary>,
_runtime: R,
}
fn defer_resources<R, W>(
runtime: &R,
library: &Arc<CufftLibrary>,
handle: &mut Option<CufftHandle>,
workspace: &mut Option<W>,
) where
R: Clone,
W: CufftWorkspaceOwner,
{
std::mem::forget(DeferredCufftResources {
_handle: handle.take(),
_workspace: workspace.take(),
_library: Arc::clone(library),
_runtime: runtime.clone(),
});
}
pub(crate) fn retire_entry_resources<R, W>(
runtime: &R,
library: &Arc<CufftLibrary>,
plan: CufftHandle,
workspace: W,
) -> CleanupFailures
where
R: CufftCleanup + Clone,
W: CufftWorkspaceOwner,
{
let mut handle = Some(plan);
let mut workspace = Some(workspace);
retire_handle(runtime, library, &mut handle, &mut workspace)
}
#[cold]
pub(crate) fn report_cleanup_failures(failures: CleanupFailures) {
let mut stderr = std::io::stderr();
let resources_deferred = failures.resources_deferred;
for error in failures.synchronization {
let _ = writeln!(
stderr,
"tenferro-fft: cuFFT cleanup synchronization failed: {error}"
);
}
if let Some(error) = failures.destroy {
let _ = writeln!(
stderr,
"tenferro-fft: cuFFT plan destroy failed during cleanup: {error}"
);
}
if resources_deferred {
let _ = writeln!(
stderr,
"tenferro-fft: cuFFT plan and lifetime witnesses were intentionally retained after cleanup failure"
);
}
}
struct CufftConstructionGuard<R, W>
where
R: CufftCleanup + Clone,
W: CufftWorkspaceOwner,
{
library: Arc<CufftLibrary>,
runtime: R,
handle: Option<CufftHandle>,
workspace: Option<W>,
}
impl<R, W> CufftConstructionGuard<R, W>
where
R: CufftCleanup + Clone,
W: CufftWorkspaceOwner,
{
fn new(library: Arc<CufftLibrary>, runtime: R) -> Self {
Self {
library,
runtime,
handle: None,
workspace: None,
}
}
fn disarm(mut self) -> Result<(CufftHandle, W), CudaFftError> {
let Some(handle) = self.handle.take() else {
return Err(CudaFftError::internal(
"cuFFT construction completed without a handle",
));
};
let Some(workspace) = self.workspace.take() else {
self.handle = Some(handle);
return Err(CudaFftError::internal(
"cuFFT construction completed without workspace",
));
};
Ok((handle, workspace))
}
}
impl<R, W> Drop for CufftConstructionGuard<R, W>
where
R: CufftCleanup + Clone,
W: CufftWorkspaceOwner,
{
fn drop(&mut self) {
let failures = retire_handle(
&self.runtime,
&self.library,
&mut self.handle,
&mut self.workspace,
);
report_cleanup_failures(failures);
}
}
pub(crate) fn build_plan<R, W, F>(
library: Arc<CufftLibrary>,
runtime: R,
mut descriptor: CufftPlanDescriptor,
mut allocate_workspace: F,
) -> Result<(CufftHandle, W), CudaFftError>
where
R: CufftCleanup + Clone,
W: CufftWorkspaceOwner,
F: FnMut(usize) -> Result<W, CudaFftError>,
{
let mut guard = CufftConstructionGuard::<R, W>::new(Arc::clone(&library), runtime);
let mut handle = 0;
let status = unsafe { (library.api.create)(&mut handle) };
map_cufft_status("cufftCreate", status)?;
guard.handle = Some(handle);
let status = unsafe { (library.api.set_auto_allocation)(handle, 0) };
map_cufft_status("cufftSetAutoAllocation", status)?;
let mut workspace_size = 0usize;
let status = unsafe {
(library.api.make_plan_many_64)(
handle,
descriptor.rank,
descriptor.n.as_mut_ptr(),
descriptor.inembed.as_mut_ptr(),
descriptor.istride,
descriptor.idist,
descriptor.onembed.as_mut_ptr(),
descriptor.ostride,
descriptor.odist,
cufft_type(descriptor.kind),
descriptor.batch,
&mut workspace_size,
)
};
map_cufft_status("cufftMakePlanMany64", status)?;
let workspace = if workspace_size == 0 {
W::empty()
} else {
allocate_workspace(workspace_size)?
};
guard.workspace = Some(workspace);
let mut work_area_result = Ok(());
if let Some(workspace) = guard.workspace.as_ref() {
workspace.with_ptr(|ptr| {
let status = unsafe { (library.api.set_work_area)(handle, ptr) };
work_area_result = map_cufft_status("cufftSetWorkArea", status);
});
}
work_area_result?;
guard.disarm()
}
fn cufft_type(kind: CufftTransformKind) -> i32 {
match kind {
CufftTransformKind::C2c32 => CUFFT_C2C,
CufftTransformKind::C2c64 => CUFFT_Z2Z,
CufftTransformKind::R2c32 => CUFFT_R2C,
CufftTransformKind::R2c64 => CUFFT_D2Z,
CufftTransformKind::C2r32 => CUFFT_C2R,
CufftTransformKind::C2r64 => CUFFT_Z2D,
}
}
fn direction(direction: CufftDirection) -> i32 {
match direction {
CufftDirection::Forward => CUFFT_FORWARD,
CufftDirection::Inverse => CUFFT_INVERSE,
}
}
pub(crate) fn bind_plan_to_stream(
library: &CufftLibrary,
plan: CufftHandle,
stream: u64,
) -> Result<(), CudaFftError> {
let stream = usize::try_from(stream)
.map_err(|_| CudaFftError::InvalidConfiguration { field: "stream" })?
as *mut c_void;
let status = unsafe { (library.api.set_stream)(plan, stream) };
map_cufft_status("cufftSetStream", status)
}
pub(crate) trait CufftExecutionScopes {
fn with_stream(&self, callback: impl FnOnce(u64)) -> Result<(), CudaFftError>;
fn synchronize(&self) -> Result<(), CudaFftError>;
}
pub(crate) trait CufftExternalUseLease {
fn with_ptr(&self, callback: impl FnOnce(*mut c_void)) -> Result<(), CudaFftError>;
}
impl CufftExternalUseLease for CudaExternalUseReadLease {
fn with_ptr(&self, callback: impl FnOnce(*mut c_void)) -> Result<(), CudaFftError> {
self.with_device_ptr(callback)
.map_err(|source| CudaFftError::interop("cufft_execute_pointer", source))
}
}
impl CufftExternalUseLease for CudaExternalUseWriteLease {
fn with_ptr(&self, callback: impl FnOnce(*mut c_void)) -> Result<(), CudaFftError> {
self.with_device_ptr(callback)
.map_err(|source| CudaFftError::interop("cufft_execute_pointer", source))
}
}
struct TypedExecutionScopes<'a> {
runtime: &'a CudaRuntime,
}
impl CufftExecutionScopes for TypedExecutionScopes<'_> {
fn with_stream(&self, callback: impl FnOnce(u64)) -> Result<(), CudaFftError> {
with_raw_cuda_stream(self.runtime, OP, callback)
.map_err(|source| CudaFftError::interop("cufft_execute_stream", source))
}
fn synchronize(&self) -> Result<(), CudaFftError> {
self.runtime
.synchronize()
.map_err(|source| CudaFftError::interop("cufft_execute_synchronize", source))
}
}
pub(crate) fn enqueue_plan_execution<S, I, O, C>(
scopes: &S,
input_lease: I,
output_lease: O,
library: &CufftLibrary,
plan: CufftHandle,
mut call: C,
function: &'static str,
) -> Result<(), CudaFftError>
where
S: CufftExecutionScopes,
I: CufftExternalUseLease,
O: CufftExternalUseLease,
C: FnMut(&CufftApi, CufftHandle, *mut c_void, *mut c_void) -> CufftStatus,
{
let mut execution_error = None;
let mut synchronization_error = None;
scopes.with_stream(|stream| {
if let Err(error) = bind_plan_to_stream(library, plan, stream) {
execution_error = Some(error);
} else {
let mut pointer_error = None;
if let Err(error) = input_lease.with_ptr(|input_ptr| {
if let Err(error) = output_lease.with_ptr(|output_ptr| {
let status = call(&library.api, plan, input_ptr, output_ptr);
if let Err(error) = map_cufft_status(function, status) {
execution_error = Some(error);
}
if let Err(error) = scopes.synchronize() {
synchronization_error = Some(error);
}
}) {
pointer_error = Some(error);
}
}) {
execution_error = Some(error);
} else if let Some(error) = pointer_error {
execution_error = Some(error);
}
}
})?;
if synchronization_error.is_some() {
std::mem::forget(input_lease);
std::mem::forget(output_lease);
}
match (execution_error, synchronization_error) {
(Some(primary), Some(suppressed)) => {
Err(CudaFftError::with_suppressed(primary, suppressed))
}
(Some(error), None) | (None, Some(error)) => Err(error),
(None, None) => Ok(()),
}
}
pub(crate) struct CufftPlanEntry {
pub(crate) library: Arc<CufftLibrary>,
pub(crate) plan: CufftHandle,
pub(crate) workspace: CufftWorkspace,
pub(crate) runtime: CudaRuntime,
pub(crate) key: CufftPlanKey,
retained_bytes: usize,
}
unsafe impl Send for CufftPlanEntry {}
unsafe impl Sync for CufftPlanEntry {}
pub(crate) fn retained_bytes_for_workspace(workspace_bytes: usize) -> usize {
size_of::<Arc<CufftLibrary>>()
.saturating_add(size_of::<CudaRuntime>())
.saturating_add(size_of::<usize>())
.saturating_add(size_of::<CufftPlanKey>())
.saturating_add(size_of::<CufftHandle>())
.saturating_add(size_of::<CufftWorkspace>())
.saturating_add(workspace_bytes)
}
pub(crate) fn with_cufft_plan_for_batch<T>(
batch: usize,
load_and_create: impl FnOnce() -> Result<T, CudaFftError>,
) -> Result<Option<T>, CudaFftError> {
if batch == 0 {
Ok(None)
} else {
load_and_create().map(Some)
}
}
impl CufftPlanEntry {
pub(crate) fn create(
runtime: &CudaRuntime,
key: CufftPlanKey,
descriptor: CufftPlanDescriptor,
) -> Result<Self, CudaFftError> {
let library = CufftLibrary::load()?;
Self::create_with_library(runtime, library, key, descriptor)
}
fn create_with_library(
runtime: &CudaRuntime,
library: Arc<CufftLibrary>,
key: CufftPlanKey,
descriptor: CufftPlanDescriptor,
) -> Result<Self, CudaFftError> {
runtime
.set_current_cuda_context(OP)
.map_err(|source| CudaFftError::interop("cufft_plan_context", source))?;
let runtime_for_cleanup = runtime.clone();
let (plan, workspace) = build_plan(
Arc::clone(&library),
runtime_for_cleanup,
descriptor,
|bytes| {
alloc_device_bytes(runtime, bytes, OP)
.map(|workspace| CufftWorkspace::from_device(workspace, bytes))
.map_err(|source| CudaFftError::interop("cufft_workspace_allocate", source))
},
)?;
let retained_bytes = retained_bytes_for_workspace(workspace.bytes());
Ok(Self {
library,
plan,
workspace,
runtime: runtime.clone(),
key,
retained_bytes,
})
}
pub(crate) fn execute(
&mut self,
input: &Tensor,
output: &mut Tensor,
) -> Result<(), CudaFftError> {
self.runtime
.set_current_cuda_context(OP)
.map_err(|source| CudaFftError::interop("cufft_execute_context", source))?;
match self.key.kind {
CufftTransformKind::C2c32 => match (input, output) {
(Tensor::C32(input), Tensor::C32(output)) => self.execute_pair(
input,
output,
|api, plan, input, output| {
unsafe {
(api.exec_c2c)(plan, input, output, direction(self.key.direction))
}
},
"cufftExecC2C",
),
_ => Err(CudaFftError::InvalidConfiguration { field: "dtype" }),
},
CufftTransformKind::C2c64 => match (input, output) {
(Tensor::C64(input), Tensor::C64(output)) => self.execute_pair(
input,
output,
|api, plan, input, output| {
unsafe {
(api.exec_z2z)(plan, input, output, direction(self.key.direction))
}
},
"cufftExecZ2Z",
),
_ => Err(CudaFftError::InvalidConfiguration { field: "dtype" }),
},
CufftTransformKind::R2c32 => match (input, output) {
(Tensor::F32(input), Tensor::C32(output)) => self.execute_pair(
input,
output,
|api, plan, input, output| {
unsafe { (api.exec_r2c)(plan, input, output) }
},
"cufftExecR2C",
),
_ => Err(CudaFftError::InvalidConfiguration { field: "dtype" }),
},
CufftTransformKind::R2c64 => match (input, output) {
(Tensor::F64(input), Tensor::C64(output)) => self.execute_pair(
input,
output,
|api, plan, input, output| {
unsafe { (api.exec_d2z)(plan, input, output) }
},
"cufftExecD2Z",
),
_ => Err(CudaFftError::InvalidConfiguration { field: "dtype" }),
},
CufftTransformKind::C2r32 => match (input, output) {
(Tensor::C32(input), Tensor::F32(output)) => self.execute_pair(
input,
output,
|api, plan, input, output| {
unsafe { (api.exec_c2r)(plan, input, output) }
},
"cufftExecC2R",
),
_ => Err(CudaFftError::InvalidConfiguration { field: "dtype" }),
},
CufftTransformKind::C2r64 => match (input, output) {
(Tensor::C64(input), Tensor::F64(output)) => self.execute_pair(
input,
output,
|api, plan, input, output| {
unsafe { (api.exec_z2d)(plan, input, output) }
},
"cufftExecZ2D",
),
_ => Err(CudaFftError::InvalidConfiguration { field: "dtype" }),
},
}
}
fn execute_pair<T, U>(
&self,
input: &TypedTensor<T>,
output: &mut TypedTensor<U>,
call: impl FnMut(&CufftApi, CufftHandle, *mut c_void, *mut c_void) -> CufftStatus,
function: &'static str,
) -> Result<(), CudaFftError>
where
T: TensorScalar + 'static,
U: TensorScalar + 'static,
{
let input_lease = CudaExternalUseReadLease::new(&self.runtime, input, OP)
.map_err(|source| CudaFftError::interop("cufft_execute_input_lease", source))?;
let output_lease = CudaExternalUseWriteLease::new(&self.runtime, output, OP)
.map_err(|source| CudaFftError::interop("cufft_execute_output_lease", source))?;
let scopes = TypedExecutionScopes {
runtime: &self.runtime,
};
enqueue_plan_execution(
&scopes,
input_lease,
output_lease,
&self.library,
self.plan,
call,
function,
)
}
pub(crate) fn retained_bytes(&self) -> usize {
self.retained_bytes
}
pub(crate) fn matches_key(&self, key: &CufftPlanKey) -> bool {
plan_key_discriminator_matches(&self.key, key)
}
}
impl Drop for CufftPlanEntry {
fn drop(&mut self) {
let workspace = std::mem::replace(&mut self.workspace, CufftWorkspace::empty());
let failures = retire_entry_resources(&self.runtime, &self.library, self.plan, workspace);
report_cleanup_failures(failures);
}
}
pub(crate) fn extension_plan_key_for_runtime(key: &CufftPlanKey) -> ExtensionCacheKey {
let mut hasher = DefaultHasher::new();
key.hash(&mut hasher);
key.runtime_identity.hash(&mut hasher);
ExtensionCacheKey::new(
FFT_EXTENSION_FAMILY_ID,
CUFFT_CACHE_NAMESPACE,
hasher.finish(),
)
}
pub(crate) fn plan_key_discriminator_matches<I: PartialEq>(
stored: &CufftPlanKey<I>,
requested: &CufftPlanKey<I>,
) -> bool {
stored == requested
}