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