Skip to main content

singe_cutensor/mg/
copy.rs

1use std::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    pub fn workspace(&self) -> Result<Workspace> {
74        let mut device_workspace_size = vec![0; self.context.device_count()];
75        let mut host_workspace_size = 0;
76        unsafe {
77            try_ffi!(sys::cutensorMgCopyGetWorkspace(
78                self.context.as_raw(),
79                self.handle,
80                device_workspace_size.as_mut_ptr(),
81                &raw mut host_workspace_size,
82            ))?;
83        }
84        Workspace::from_raw(device_workspace_size, host_workspace_size)
85    }
86
87    pub(crate) fn context(&self) -> &ContextRef {
88        &self.context
89    }
90
91    pub const fn as_raw(&self) -> sys::cutensorMgCopyDescriptor_t {
92        self.handle
93    }
94}
95
96impl CopyPlan {
97    pub fn create(context: &Context, desc: &CopyDescriptor, workspace: Workspace) -> Result<Self> {
98        validate_same_context(&context.as_context_ref(), desc.context(), "desc")?;
99        workspace.validate_for_context(desc.context())?;
100
101        let mut handle = ptr::null_mut();
102        unsafe {
103            try_ffi!(sys::cutensorMgCreateCopyPlan(
104                context.as_raw(),
105                &raw mut handle,
106                desc.as_raw(),
107                workspace.device_sizes_ptr(),
108                workspace.host_size(),
109            ))?;
110        }
111
112        if handle.is_null() {
113            return Err(Error::NullHandle);
114        }
115
116        Ok(Self {
117            handle,
118            context: context.as_context_ref(),
119            dst_pointer_count: desc.dst_pointer_count,
120            src_pointer_count: desc.src_pointer_count,
121            workspace,
122        })
123    }
124
125    pub fn workspace(&self) -> &Workspace {
126        &self.workspace
127    }
128
129    /// Executes this copy plan with raw distributed tensor pointers.
130    ///
131    /// # Safety
132    ///
133    /// The pointer arrays must match the tensor descriptors used to create the plan.
134    /// Device workspace pointers and streams must be ordered like the MG context devices.
135    pub unsafe fn copy_raw(
136        &self,
137        dst: &mut [*mut ()],
138        src: &[*const ()],
139        device_workspace: &mut [*mut ()],
140        host_workspace: *mut (),
141        streams: &[StreamBinding],
142    ) -> Result<()> {
143        self.validate_execution_slices(
144            dst.len(),
145            src.len(),
146            device_workspace.len(),
147            streams.len(),
148        )?;
149        let mut src = src.to_vec();
150        let mut streams: Vec<_> = streams.iter().map(StreamBinding::as_raw).collect();
151        unsafe {
152            try_ffi!(sys::cutensorMgCopy(
153                self.context.as_raw(),
154                self.handle,
155                dst.as_mut_ptr() as _,
156                src.as_mut_ptr() as _,
157                device_workspace.as_mut_ptr() as _,
158                host_workspace as _,
159                streams.as_mut_ptr() as _,
160            ))?;
161        }
162        Ok(())
163    }
164
165    fn validate_execution_slices(
166        &self,
167        dst_count: usize,
168        src_count: usize,
169        device_workspace_count: usize,
170        stream_count: usize,
171    ) -> Result<()> {
172        validate_count("dst", self.dst_pointer_count, dst_count)?;
173        validate_count("src", self.src_pointer_count, src_count)?;
174        validate_count(
175            "device_workspace",
176            self.context.device_count(),
177            device_workspace_count,
178        )?;
179        validate_count("streams", self.context.device_count(), stream_count)
180    }
181
182    pub const fn as_raw(&self) -> sys::cutensorMgCopyPlan_t {
183        self.handle
184    }
185}
186
187impl Drop for CopyDescriptor {
188    fn drop(&mut self) {
189        unsafe {
190            if let Err(err) = try_ffi!(sys::cutensorMgDestroyCopyDescriptor(self.handle)) {
191                #[cfg(debug_assertions)]
192                eprintln!("failed to destroy cutensormg copy descriptor: {err}");
193            }
194        }
195    }
196}
197
198impl Drop for CopyPlan {
199    fn drop(&mut self) {
200        unsafe {
201            if let Err(err) = try_ffi!(sys::cutensorMgDestroyCopyPlan(self.handle)) {
202                #[cfg(debug_assertions)]
203                eprintln!("failed to destroy cutensormg copy plan: {err}");
204            }
205        }
206    }
207}
208
209fn validate_modes(desc: &TensorDescriptor, modes: &[Mode], _name: &str) -> Result<()> {
210    if desc.rank() as usize != modes.len() {
211        return Err(Error::TensorModeMismatch {
212            rank: desc.rank(),
213            mode_length: modes.len(),
214        });
215    }
216    Ok(())
217}
218
219pub(crate) fn validate_count(name: &str, expected: usize, actual: usize) -> Result<()> {
220    if expected != actual {
221        return Err(Error::PointerCountMismatch {
222            name: name.into(),
223            expected,
224            actual,
225        });
226    }
227    Ok(())
228}