audioadapter_compat_ndarray/
lib.rs1use core::marker::PhantomData;
32
33use ndarray::{ArrayBase, Axis, Data, DataMut, Ix2};
34
35use audioadapter::{Adapter, AdapterMut};
36
37mod sealed {
38 pub trait Sealed {}
39}
40
41pub trait AxisOrder: sealed::Sealed {
43 const CHANNEL_AXIS: usize;
45 fn channels(dim: (usize, usize)) -> usize;
47 fn frames(dim: (usize, usize)) -> usize;
49 fn index(channel: usize, frame: usize) -> (usize, usize);
51}
52
53pub struct ChannelsFrames;
55impl sealed::Sealed for ChannelsFrames {}
56impl AxisOrder for ChannelsFrames {
57 const CHANNEL_AXIS: usize = 0;
58 fn channels(dim: (usize, usize)) -> usize {
59 dim.0
60 }
61 fn frames(dim: (usize, usize)) -> usize {
62 dim.1
63 }
64 fn index(channel: usize, frame: usize) -> (usize, usize) {
65 (channel, frame)
66 }
67}
68
69pub struct FramesChannels;
71impl sealed::Sealed for FramesChannels {}
72impl AxisOrder for FramesChannels {
73 const CHANNEL_AXIS: usize = 1;
74 fn channels(dim: (usize, usize)) -> usize {
75 dim.1
76 }
77 fn frames(dim: (usize, usize)) -> usize {
78 dim.0
79 }
80 fn index(channel: usize, frame: usize) -> (usize, usize) {
81 (frame, channel)
82 }
83}
84
85pub struct NdarrayAdapter<U, O> {
91 array: U,
92 _order: PhantomData<O>,
93}
94
95impl<S> NdarrayAdapter<ArrayBase<S, Ix2>, ChannelsFrames>
96where
97 S: Data,
98{
99 pub fn new_channels_frames(array: ArrayBase<S, Ix2>) -> Self {
101 Self {
102 array,
103 _order: PhantomData,
104 }
105 }
106}
107
108impl<S> NdarrayAdapter<ArrayBase<S, Ix2>, FramesChannels>
109where
110 S: Data,
111{
112 pub fn new_frames_channels(array: ArrayBase<S, Ix2>) -> Self {
114 Self {
115 array,
116 _order: PhantomData,
117 }
118 }
119}
120
121impl<U, O> NdarrayAdapter<U, O> {
122 pub fn into_inner(self) -> U {
124 self.array
125 }
126
127 pub fn inner(&self) -> &U {
129 &self.array
130 }
131}
132
133unsafe impl<S, O> Adapter<S::Elem> for NdarrayAdapter<ArrayBase<S, Ix2>, O>
134where
135 S: Data,
136 S::Elem: Clone,
137 O: AxisOrder,
138{
139 fn channels(&self) -> usize {
140 O::channels(self.array.dim())
141 }
142
143 fn frames(&self) -> usize {
144 O::frames(self.array.dim())
145 }
146
147 unsafe fn read_sample_unchecked(&self, channel: usize, frame: usize) -> S::Elem {
148 unsafe { self.array.uget(O::index(channel, frame)) }.clone()
149 }
150
151 fn copy_from_channel_to_slice(
152 &self,
153 channel: usize,
154 skip: usize,
155 slice: &mut [S::Elem],
156 ) -> usize {
157 if channel >= self.channels() || skip >= self.frames() {
158 return 0;
159 }
160 let view = self.array.index_axis(Axis(O::CHANNEL_AXIS), channel);
161 let available = view.len() - skip;
162 let to_copy = available.min(slice.len());
163 if let Some(contiguous) = view.as_slice() {
164 slice[..to_copy].clone_from_slice(&contiguous[skip..skip + to_copy]);
165 } else {
166 for (out, sample) in slice.iter_mut().zip(view.iter().skip(skip)).take(to_copy) {
167 *out = sample.clone();
168 }
169 }
170 to_copy
171 }
172}
173
174unsafe impl<S, O> AdapterMut<S::Elem> for NdarrayAdapter<ArrayBase<S, Ix2>, O>
175where
176 S: DataMut,
177 S::Elem: Clone,
178 O: AxisOrder,
179{
180 unsafe fn write_sample_unchecked(
181 &mut self,
182 channel: usize,
183 frame: usize,
184 value: &S::Elem,
185 ) -> bool {
186 unsafe { *self.array.uget_mut(O::index(channel, frame)) = value.clone() };
187 false
188 }
189
190 fn copy_from_slice_to_channel(
191 &mut self,
192 channel: usize,
193 skip: usize,
194 slice: &[S::Elem],
195 ) -> (usize, usize) {
196 if channel >= Adapter::channels(self) || skip >= Adapter::frames(self) {
197 return (0, 0);
198 }
199 let mut view = self.array.index_axis_mut(Axis(O::CHANNEL_AXIS), channel);
200 let available = view.len() - skip;
201 let to_copy = available.min(slice.len());
202 if let Some(contiguous) = view.as_slice_mut() {
203 contiguous[skip..skip + to_copy].clone_from_slice(&slice[..to_copy]);
204 } else {
205 for (dest, sample) in view.iter_mut().skip(skip).zip(slice.iter()).take(to_copy) {
206 *dest = sample.clone();
207 }
208 }
209 (to_copy, 0)
210 }
211}
212
213#[cfg(test)]
220mod tests {
221 use super::*;
222 use ndarray::array;
223
224 #[test]
225 fn channels_frames_dimensions_and_read() {
226 let data = array![[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]];
227 let adapter = NdarrayAdapter::new_channels_frames(data.view());
228 assert_eq!(adapter.channels(), 2);
229 assert_eq!(adapter.frames(), 3);
230 assert_eq!(adapter.read_sample(0, 0), Some(1.0));
231 assert_eq!(adapter.read_sample(1, 2), Some(6.0));
232 assert_eq!(adapter.read_sample(2, 0), None);
233 assert_eq!(adapter.read_sample(0, 3), None);
234 }
235
236 #[test]
237 fn frames_channels_dimensions_and_read() {
238 let data = array![[1.0, 4.0], [2.0, 5.0], [3.0, 6.0]];
240 let adapter = NdarrayAdapter::new_frames_channels(data.view());
241 assert_eq!(adapter.channels(), 2);
242 assert_eq!(adapter.frames(), 3);
243 assert_eq!(adapter.read_sample(0, 0), Some(1.0));
244 assert_eq!(adapter.read_sample(1, 2), Some(6.0));
245 }
246
247 #[test]
248 fn copy_channel_to_slice_contiguous() {
249 let data = array![[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]];
251 let adapter = NdarrayAdapter::new_channels_frames(data.view());
252 let mut out = [0.0; 2];
253 let copied = adapter.copy_from_channel_to_slice(1, 1, &mut out);
254 assert_eq!(copied, 2);
255 assert_eq!(out, [5.0, 6.0]);
256 }
257
258 #[test]
259 fn copy_channel_to_slice_strided() {
260 let data = array![[1.0, 4.0], [2.0, 5.0], [3.0, 6.0]];
262 let adapter = NdarrayAdapter::new_frames_channels(data.view());
263 let mut out = [0.0; 3];
264 let copied = adapter.copy_from_channel_to_slice(1, 0, &mut out);
265 assert_eq!(copied, 3);
266 assert_eq!(out, [4.0, 5.0, 6.0]);
267 }
268
269 #[test]
270 fn write_and_copy_from_slice() {
271 let mut data = array![[0.0, 0.0, 0.0], [0.0, 0.0, 0.0]];
272 let mut adapter = NdarrayAdapter::new_channels_frames(data.view_mut());
273 assert_eq!(adapter.write_sample(0, 1, &9.0), Some(false));
274 assert_eq!(adapter.read_sample(0, 1), Some(9.0));
275 let (copied, clipped) = adapter.copy_from_slice_to_channel(1, 0, &[7.0, 8.0]);
276 assert_eq!((copied, clipped), (2, 0));
277 assert_eq!(adapter.read_sample(1, 0), Some(7.0));
278 assert_eq!(adapter.read_sample(1, 1), Some(8.0));
279 assert_eq!(adapter.read_sample(1, 2), Some(0.0));
280 }
281}