Skip to main content

singe_cutensor/mg/
context.rs

1use std::{ptr, sync::Arc};
2
3use crate::{
4    error::{Error, Result},
5    sys, try_ffi,
6    utility::to_u32,
7};
8
9#[derive(Debug, Clone)]
10pub struct Context {
11    handle: Arc<Handle>,
12}
13
14#[derive(Debug)]
15struct Handle {
16    raw: sys::cutensorMgHandle_t,
17    devices: Vec<i32>,
18}
19
20#[derive(Debug, Clone)]
21pub struct ContextRef {
22    handle: Arc<Handle>,
23}
24
25// cuTENSORMg contexts are immutable after creation for a fixed device set. The
26// Arc-backed handle can be shared, while wrappers remain freely movable.
27unsafe impl Send for Handle {}
28unsafe impl Sync for Handle {}
29unsafe impl Send for Context {}
30unsafe impl Send for ContextRef {}
31
32impl Context {
33    /// Creates a cuTENSORMg handle for the devices that may participate in later operations.
34    ///
35    /// # Errors
36    ///
37    /// Returns an error if the device count cannot be represented by cuTENSORMg,
38    /// if cuTENSORMg cannot create a handle, or if it returns a null handle.
39    pub fn create(devices: &[i32]) -> Result<Self> {
40        let num_devices = to_u32(devices.len(), "devices")?;
41        let mut handle = ptr::null_mut();
42        unsafe {
43            try_ffi!(sys::cutensorMgCreate(
44                &raw mut handle,
45                num_devices,
46                devices.as_ptr(),
47            ))?;
48        }
49
50        if handle.is_null() {
51            return Err(Error::NullHandle);
52        }
53
54        Ok(Self {
55            handle: Arc::new(Handle {
56                raw: handle,
57                devices: devices.to_vec(),
58            }),
59        })
60    }
61
62    pub fn devices(&self) -> &[i32] {
63        &self.handle.devices
64    }
65
66    pub fn device_count(&self) -> usize {
67        self.handle.devices.len()
68    }
69
70    pub(crate) fn as_context_ref(&self) -> ContextRef {
71        ContextRef {
72            handle: Arc::clone(&self.handle),
73        }
74    }
75
76    pub fn as_raw(&self) -> sys::cutensorMgHandle_t {
77        self.handle.raw
78    }
79}
80
81impl ContextRef {
82    pub fn devices(&self) -> &[i32] {
83        &self.handle.devices
84    }
85
86    pub fn device_count(&self) -> usize {
87        self.handle.devices.len()
88    }
89
90    pub(crate) fn same_handle(&self, other: &Self) -> bool {
91        Arc::ptr_eq(&self.handle, &other.handle)
92    }
93
94    pub fn as_raw(&self) -> sys::cutensorMgHandle_t {
95        self.handle.raw
96    }
97}
98
99impl Drop for Handle {
100    fn drop(&mut self) {
101        unsafe {
102            if let Err(err) = try_ffi!(sys::cutensorMgDestroy(self.raw)) {
103                #[cfg(debug_assertions)]
104                eprintln!("failed to destroy cutensormg context: {err}");
105            }
106        }
107    }
108}
109
110pub(crate) fn validate_same_context(
111    expected: &ContextRef,
112    actual: &ContextRef,
113    name: &str,
114) -> Result<()> {
115    if !expected.same_handle(actual) {
116        return Err(Error::MgContextMismatch { name: name.into() });
117    }
118    Ok(())
119}