1use core::{
13 marker::PhantomData,
14 mem::MaybeUninit,
15 ops::{Add, Mul},
16};
17
18use crate::map_view::map_raw_into;
19use crate::ops_view::{copy_conj, copy_into, copy_scale};
20use crate::raw_ops::{apply_fused_pair, fuse_pair_layout, FusedPairLayout};
21use crate::{
22 ElementOpApply, Identity, MaybeSendSync, RawStridedMut, RawStridedRef, Result, StridedError,
23};
24
25#[cfg(feature = "parallel")]
29type AxisVec<T> = smallvec::SmallVec<[T; crate::RAW_FUSED_RANK_LIMIT]>;
30#[cfg(not(feature = "parallel"))]
31type AxisVec<T> = Vec<T>;
32
33pub(crate) trait OverwriteWriter<T> {
34 fn dims(&self) -> &[usize];
35 fn strides(&self) -> &[isize];
36 fn offset(&self) -> isize;
37 unsafe fn data_ptr(&mut self) -> *mut T;
41 unsafe fn write_at(&mut self, offset: isize, value: T);
44}
45
46pub(crate) trait ReadModifyWrite<T>: OverwriteWriter<T> {
47 unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
51 where
52 F: FnOnce(T, T) -> T;
53}
54
55impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, T> {
56 fn dims(&self) -> &[usize] {
57 self.dims()
58 }
59 fn strides(&self) -> &[isize] {
60 self.strides()
61 }
62 fn offset(&self) -> isize {
63 self.offset()
64 }
65 unsafe fn data_ptr(&mut self) -> *mut T {
66 self.data_mut().as_mut_ptr()
67 }
68 unsafe fn write_at(&mut self, offset: isize, value: T) {
69 unsafe { self.data_mut().as_mut_ptr().offset(offset).write(value) }
71 }
72}
73
74impl<'a, T> ReadModifyWrite<T> for RawStridedMut<'a, T>
75where
76 T: Add<Output = T>,
77{
78 #[inline(always)]
79 unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
80 where
81 F: FnOnce(T, T) -> T,
82 {
83 unsafe {
85 let ptr = self.data_mut().as_mut_ptr().offset(offset);
86 ptr.write(combine(ptr.read(), value));
87 }
88 }
89}
90
91impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, MaybeUninit<T>> {
92 fn dims(&self) -> &[usize] {
93 self.dims()
94 }
95 fn strides(&self) -> &[isize] {
96 self.strides()
97 }
98 fn offset(&self) -> isize {
99 self.offset()
100 }
101 unsafe fn data_ptr(&mut self) -> *mut T {
102 self.data_mut().as_mut_ptr().cast()
103 }
104 unsafe fn write_at(&mut self, offset: isize, value: T) {
105 unsafe {
107 self.data_mut()
108 .as_mut_ptr()
109 .offset(offset)
110 .write(MaybeUninit::new(value))
111 }
112 }
113}
114
115pub(crate) struct InitializedRawDest<'a, T> {
116 ptr: *mut T,
117 extent: usize,
118 dims: &'a [usize],
119 strides: &'a [isize],
120 offset: isize,
121 _marker: PhantomData<&'a mut [MaybeUninit<T>]>,
122}
123
124impl<'a, T> OverwriteWriter<T> for InitializedRawDest<'a, T> {
125 fn dims(&self) -> &[usize] {
126 self.dims
127 }
128 fn strides(&self) -> &[isize] {
129 self.strides
130 }
131 fn offset(&self) -> isize {
132 self.offset
133 }
134 unsafe fn data_ptr(&mut self) -> *mut T {
135 self.ptr
136 }
137 unsafe fn write_at(&mut self, offset: isize, value: T) {
138 debug_assert!(offset >= 0 && (offset as usize) < self.extent);
139 unsafe { self.ptr.offset(offset).write(value) }
141 }
142}
143
144impl<'a, T> ReadModifyWrite<T> for InitializedRawDest<'a, T>
145where
146 T: Add<Output = T>,
147{
148 #[inline(always)]
149 unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
150 where
151 F: FnOnce(T, T) -> T,
152 {
153 debug_assert!(offset >= 0 && (offset as usize) < self.extent);
154 unsafe {
156 let ptr = self.ptr.offset(offset);
157 ptr.write(combine(ptr.read(), value));
158 }
159 }
160}
161
162#[derive(Clone, Debug)]
194pub struct CopyPlan {
195 dims: AxisVec<usize>,
196 dst_strides: AxisVec<isize>,
197 src_strides: AxisVec<isize>,
198 fused: Option<FusedPairLayout>,
201}
202
203impl CopyPlan {
204 pub(crate) fn execute_uninit_then<'a, T, R>(
205 &self,
206 dest: &'a mut RawStridedMut<'a, MaybeUninit<T>>,
207 src: &RawStridedRef<'_, T>,
208 f: impl for<'b> FnOnce(InitializedRawDest<'b, T>) -> R,
209 ) -> Result<R>
210 where
211 T: Copy + MaybeSendSync,
212 {
213 self.execute_uninit(dest, src)?;
214 let data = dest.data_mut();
215 let receipt = InitializedRawDest {
216 ptr: data.as_mut_ptr().cast(),
217 extent: data.len(),
218 dims: dest.dims(),
219 strides: dest.strides(),
220 offset: dest.offset(),
221 _marker: PhantomData,
222 };
223 Ok(f(receipt))
224 }
225
226 pub fn compile(dims: &[usize], dst_strides: &[isize], src_strides: &[isize]) -> Result<Self> {
239 if dims.len() != dst_strides.len() || dims.len() != src_strides.len() {
240 return Err(StridedError::StrideLengthMismatch);
241 }
242 crate::kernel::total_len(dims)?;
245 if !crate::layout_check::is_injective_layout(dims, dst_strides) {
246 return Err(StridedError::NonInjectiveOutputLayout);
247 }
248 Ok(Self {
249 dims: dims.into(),
250 dst_strides: dst_strides.into(),
251 src_strides: src_strides.into(),
252 fused: fuse_pair_layout(dims, dst_strides, src_strides),
253 })
254 }
255
256 fn check_call<D, S>(
262 &self,
263 dest: &RawStridedMut<'_, D>,
264 src: &RawStridedRef<'_, S>,
265 ) -> Result<()> {
266 if dest.dims() != &self.dims[..]
267 || src.dims() != &self.dims[..]
268 || dest.strides() != &self.dst_strides[..]
269 || src.strides() != &self.src_strides[..]
270 {
271 return Err(StridedError::PlanLayoutMismatch);
272 }
273 Ok(())
274 }
275
276 pub fn execute_uninit<T>(
281 &self,
282 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
283 src: &RawStridedRef<'_, T>,
284 ) -> Result<()>
285 where
286 T: Copy + MaybeSendSync,
287 {
288 self.check_call(dest, src)?;
289 match &self.fused {
290 Some(layout) => {
291 apply_fused_pair(
292 dest,
293 src,
294 layout,
295 |dst, value| {
296 dst.write(value);
297 },
298 |value| value,
299 );
300 Ok(())
301 }
302 None => map_raw_into::<MaybeUninit<T>, T, Identity>(dest, src, MaybeUninit::new),
303 }
304 }
305
306 pub fn execute<T>(
308 &self,
309 dest: &mut RawStridedMut<'_, T>,
310 src: &RawStridedRef<'_, T>,
311 ) -> Result<()>
312 where
313 T: Copy + MaybeSendSync,
314 {
315 self.check_call(dest, src)?;
316 match &self.fused {
317 Some(layout) => {
318 apply_fused_pair(
319 dest,
320 src,
321 layout,
322 |dst, value| *dst = value,
323 |value: T| value,
324 );
325 Ok(())
326 }
327 None => copy_into(&mut dest.as_view_mut(), &src.as_view()),
328 }
329 }
330
331 pub fn execute_scale<T>(
334 &self,
335 dest: &mut RawStridedMut<'_, T>,
336 src: &RawStridedRef<'_, T>,
337 scale: T,
338 ) -> Result<()>
339 where
340 T: Copy + Mul<T, Output = T> + MaybeSendSync,
341 {
342 self.check_call(dest, src)?;
343 match &self.fused {
344 Some(layout) => {
345 apply_fused_pair(
346 dest,
347 src,
348 layout,
349 |dst, value| *dst = value,
350 |value: T| scale * value,
351 );
352 Ok(())
353 }
354 None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
355 }
356 }
357
358 pub fn execute_conj<T>(
361 &self,
362 dest: &mut RawStridedMut<'_, T>,
363 src: &RawStridedRef<'_, T>,
364 ) -> Result<()>
365 where
366 T: Copy + ElementOpApply + MaybeSendSync,
367 {
368 self.check_call(dest, src)?;
369 match &self.fused {
370 Some(layout) => {
371 apply_fused_pair(
372 dest,
373 src,
374 layout,
375 |dst, value| *dst = value,
376 |value: T| value.conj(),
377 );
378 Ok(())
379 }
380 None => copy_conj(&mut dest.as_view_mut(), &src.as_view()),
381 }
382 }
383}
384
385#[cfg(test)]
386#[path = "copy_plan/tests/tests.rs"]
387mod tests;