Skip to main content

singe_cutensor/mg/
copy.rs

1use std::{mem::ManuallyDrop, ptr};
2
3use singe_cuda::stream::StreamBinding;
4
5use crate::{
6    error::{Error, Result},
7    mg::{
8        context::{Context, ContextRef, validate_same_context},
9        tensor::TensorDescriptor,
10        workspace::Workspace,
11    },
12    sys, try_ffi,
13    types::Mode,
14    utility::modes_to_i32_vec,
15};
16
17#[derive(Debug)]
18pub struct CopyDescriptor {
19    handle: sys::cutensorMgCopyDescriptor_t,
20    context: ContextRef,
21    dst_pointer_count: usize,
22    src_pointer_count: usize,
23}
24
25#[derive(Debug)]
26pub struct CopyPlan {
27    handle: sys::cutensorMgCopyPlan_t,
28    context: ContextRef,
29    dst_pointer_count: usize,
30    src_pointer_count: usize,
31    workspace: Workspace,
32}
33
34impl CopyDescriptor {
35    pub fn create(
36        context: &Context,
37        dst: &TensorDescriptor,
38        modes_dst: &[Mode],
39        src: &TensorDescriptor,
40        modes_src: &[Mode],
41    ) -> Result<Self> {
42        validate_same_context(&context.as_context_ref(), dst.context(), "dst")?;
43        validate_same_context(&context.as_context_ref(), src.context(), "src")?;
44        validate_modes(dst, modes_dst, "modes_dst")?;
45        validate_modes(src, modes_src, "modes_src")?;
46
47        let modes_dst = modes_to_i32_vec(modes_dst);
48        let modes_src = modes_to_i32_vec(modes_src);
49        let mut handle = ptr::null_mut();
50        unsafe {
51            try_ffi!(sys::cutensorMgCreateCopyDescriptor(
52                context.as_raw(),
53                &raw mut handle,
54                dst.as_raw(),
55                modes_dst.as_ptr(),
56                src.as_raw(),
57                modes_src.as_ptr(),
58            ))?;
59        }
60
61        if handle.is_null() {
62            return Err(Error::NullHandle);
63        }
64
65        Ok(Self {
66            handle,
67            context: context.as_context_ref(),
68            dst_pointer_count: dst.devices().len(),
69            src_pointer_count: src.devices().len(),
70        })
71    }
72
73    /// Wraps an existing cuTENSORMg copy descriptor.
74    ///
75    /// # Safety
76    ///
77    /// `handle` must be a valid cuTENSORMg copy descriptor associated with
78    /// `context` and created for `dst_pointer_count` destination pointers and
79    /// `src_pointer_count` source pointers. The returned value takes ownership
80    /// of `handle` and destroys it on drop.
81    pub unsafe fn from_raw(
82        handle: sys::cutensorMgCopyDescriptor_t,
83        context: &Context,
84        dst_pointer_count: usize,
85        src_pointer_count: usize,
86    ) -> Result<Self> {
87        if handle.is_null() {
88            return Err(Error::NullHandle);
89        }
90
91        Ok(Self {
92            handle,
93            context: context.as_context_ref(),
94            dst_pointer_count,
95            src_pointer_count,
96        })
97    }
98
99    pub fn workspace(&self) -> Result<Workspace> {
100        let mut device_workspace_size = vec![0; self.context.device_count()];
101        let mut host_workspace_size = 0;
102        unsafe {
103            try_ffi!(sys::cutensorMgCopyGetWorkspace(
104                self.context.as_raw(),
105                self.handle,
106                device_workspace_size.as_mut_ptr(),
107                &raw mut host_workspace_size,
108            ))?;
109        }
110        Workspace::from_raw(device_workspace_size, host_workspace_size)
111    }
112
113    pub(crate) fn context(&self) -> &ContextRef {
114        &self.context
115    }
116
117    pub const fn as_raw(&self) -> sys::cutensorMgCopyDescriptor_t {
118        self.handle
119    }
120
121    /// Consumes this descriptor and returns the owned raw cuTENSORMg copy descriptor.
122    ///
123    /// The caller becomes responsible for destroying the descriptor.
124    pub fn into_raw(self) -> sys::cutensorMgCopyDescriptor_t {
125        let this = ManuallyDrop::new(self);
126        this.handle
127    }
128}
129
130impl CopyPlan {
131    pub fn create(context: &Context, desc: &CopyDescriptor, workspace: Workspace) -> Result<Self> {
132        validate_same_context(&context.as_context_ref(), desc.context(), "desc")?;
133        workspace.validate_for_context(desc.context())?;
134
135        let mut handle = ptr::null_mut();
136        unsafe {
137            try_ffi!(sys::cutensorMgCreateCopyPlan(
138                context.as_raw(),
139                &raw mut handle,
140                desc.as_raw(),
141                workspace.device_sizes_ptr(),
142                workspace.host_size(),
143            ))?;
144        }
145
146        if handle.is_null() {
147            return Err(Error::NullHandle);
148        }
149
150        Ok(Self {
151            handle,
152            context: context.as_context_ref(),
153            dst_pointer_count: desc.dst_pointer_count,
154            src_pointer_count: desc.src_pointer_count,
155            workspace,
156        })
157    }
158
159    /// Wraps an existing cuTENSORMg copy plan.
160    ///
161    /// # Safety
162    ///
163    /// `handle` must be a valid cuTENSORMg copy plan associated with `context`
164    /// and `workspace`. Pointer counts must match the descriptors used to
165    /// create the plan. The returned value takes ownership of `handle` and
166    /// destroys it on drop.
167    pub unsafe fn from_raw(
168        handle: sys::cutensorMgCopyPlan_t,
169        context: &Context,
170        dst_pointer_count: usize,
171        src_pointer_count: usize,
172        workspace: Workspace,
173    ) -> Result<Self> {
174        if handle.is_null() {
175            return Err(Error::NullHandle);
176        }
177
178        Ok(Self {
179            handle,
180            context: context.as_context_ref(),
181            dst_pointer_count,
182            src_pointer_count,
183            workspace,
184        })
185    }
186
187    pub fn workspace(&self) -> &Workspace {
188        &self.workspace
189    }
190
191    /// Executes this copy plan with raw distributed tensor pointers.
192    ///
193    /// # Safety
194    ///
195    /// The pointer arrays must match the tensor descriptors used to create the plan.
196    /// Device workspace pointers and streams must be ordered like the MG context devices.
197    pub unsafe fn copy_raw(
198        &self,
199        dst: &mut [*mut ()],
200        src: &[*const ()],
201        device_workspace: &mut [*mut ()],
202        host_workspace: *mut (),
203        streams: &[StreamBinding],
204    ) -> Result<()> {
205        self.validate_execution_slices(
206            dst.len(),
207            src.len(),
208            device_workspace.len(),
209            streams.len(),
210        )?;
211        let mut src = src.to_vec();
212        let mut streams: Vec<_> = streams.iter().map(StreamBinding::as_raw).collect();
213        unsafe {
214            try_ffi!(sys::cutensorMgCopy(
215                self.context.as_raw(),
216                self.handle,
217                dst.as_mut_ptr() as _,
218                src.as_mut_ptr() as _,
219                device_workspace.as_mut_ptr() as _,
220                host_workspace as _,
221                streams.as_mut_ptr() as _,
222            ))?;
223        }
224        Ok(())
225    }
226
227    fn validate_execution_slices(
228        &self,
229        dst_count: usize,
230        src_count: usize,
231        device_workspace_count: usize,
232        stream_count: usize,
233    ) -> Result<()> {
234        validate_count("dst", self.dst_pointer_count, dst_count)?;
235        validate_count("src", self.src_pointer_count, src_count)?;
236        validate_count(
237            "device_workspace",
238            self.context.device_count(),
239            device_workspace_count,
240        )?;
241        validate_count("streams", self.context.device_count(), stream_count)
242    }
243
244    pub const fn as_raw(&self) -> sys::cutensorMgCopyPlan_t {
245        self.handle
246    }
247
248    /// Consumes this plan and returns the owned raw cuTENSORMg copy plan.
249    ///
250    /// The caller becomes responsible for destroying the plan.
251    pub fn into_raw(self) -> sys::cutensorMgCopyPlan_t {
252        let this = ManuallyDrop::new(self);
253        this.handle
254    }
255}
256
257impl Drop for CopyDescriptor {
258    fn drop(&mut self) {
259        unsafe {
260            if let Err(err) = try_ffi!(sys::cutensorMgDestroyCopyDescriptor(self.handle)) {
261                #[cfg(debug_assertions)]
262                eprintln!("failed to destroy cutensormg copy descriptor: {err}");
263            }
264        }
265    }
266}
267
268impl Drop for CopyPlan {
269    fn drop(&mut self) {
270        unsafe {
271            if let Err(err) = try_ffi!(sys::cutensorMgDestroyCopyPlan(self.handle)) {
272                #[cfg(debug_assertions)]
273                eprintln!("failed to destroy cutensormg copy plan: {err}");
274            }
275        }
276    }
277}
278
279fn validate_modes(desc: &TensorDescriptor, modes: &[Mode], _name: &str) -> Result<()> {
280    if desc.rank() as usize != modes.len() {
281        return Err(Error::TensorModeMismatch {
282            rank: desc.rank(),
283            mode_length: modes.len(),
284        });
285    }
286    Ok(())
287}
288
289pub(crate) fn validate_count(name: &str, expected: usize, actual: usize) -> Result<()> {
290    if expected != actual {
291        return Err(Error::PointerCountMismatch {
292            name: name.into(),
293            expected,
294            actual,
295        });
296    }
297    Ok(())
298}