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 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 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 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 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 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}