Skip to main content

audioadapter_compat_ndarray/
lib.rs

1//! # [ndarray](https://crates.io/crates/ndarray) crate compatibility
2//!
3//! This module implements the `audioadapter` traits for two-dimensional
4//! [ndarray](https://crates.io/crates/ndarray) arrays (`ArrayBase<_, Ix2>`,
5//! which includes `Array2<T>` and array views).
6//!
7//! Because both channel-major and frame-major storage are common, the axis
8//! order is selected explicitly when the adapter is created:
9//!
10//! * [`NdarrayAdapter::new_channels_frames`] for arrays shaped `(channels, frames)`.
11//! * [`NdarrayAdapter::new_frames_channels`] for arrays shaped `(frames, channels)`.
12//!
13//! Sample access uses ndarray indexing, so it is correct for any memory layout.
14//! The bulk copy helpers take a fast path that copies directly from a contiguous
15//! slice when the relevant axis is contiguous (the common standard-layout case),
16//! and fall back to element-wise access otherwise.
17//!
18//! ```
19//! use ndarray::array;
20//! use audioadapter::Adapter;
21//! use audioadapter_compat_ndarray::NdarrayAdapter;
22//!
23//! // Two channels, three frames, stored channel-major.
24//! let data = array![[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]];
25//! let adapter = NdarrayAdapter::new_channels_frames(data.view());
26//! assert_eq!(adapter.channels(), 2);
27//! assert_eq!(adapter.frames(), 3);
28//! assert_eq!(adapter.read_sample(1, 2), Some(6.0));
29//! ```
30
31use 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
41/// Axis-order marker describing how the two array axes map to channels and frames.
42pub trait AxisOrder: sealed::Sealed {
43    /// The array axis (0 or 1) along which the channel index runs.
44    const CHANNEL_AXIS: usize;
45    /// Number of channels for an array of the given shape.
46    fn channels(dim: (usize, usize)) -> usize;
47    /// Number of frames for an array of the given shape.
48    fn frames(dim: (usize, usize)) -> usize;
49    /// Map a `(channel, frame)` pair to an ndarray `(row, column)` index.
50    fn index(channel: usize, frame: usize) -> (usize, usize);
51}
52
53/// Axis order for arrays shaped `(channels, frames)` (channel-major).
54pub 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
69/// Axis order for arrays shaped `(frames, channels)` (frame-major).
70pub 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
85/// A wrapper implementing the `audioadapter` traits for a two-dimensional
86/// ndarray array.
87///
88/// The axis order `O` records whether the array is channel-major
89/// ([`ChannelsFrames`]) or frame-major ([`FramesChannels`]).
90pub 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    /// Wrap an array shaped `(channels, frames)`.
100    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    /// Wrap an array shaped `(frames, channels)`.
113    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    /// Consume the adapter and return the wrapped array.
123    pub fn into_inner(self) -> U {
124        self.array
125    }
126
127    /// Get a reference to the wrapped array.
128    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//   _____         _
214//  |_   _|__  ___| |_ ___
215//    | |/ _ \/ __| __/ __|
216//    | |  __/\__ \ |_\__ \
217//    |_|\___||___/\__|___/
218
219#[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        // Same logical audio, stored frame-major (each row is a frame).
239        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        // Channel-major standard layout: each channel is a contiguous row.
250        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        // Frame-major: a channel runs down a column, which is strided.
261        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}