Skip to main content

next_plaid/
utils.rs

1//! Utility functions for next-plaid
2
3use std::fs::{self, File};
4use std::io::{self, Write};
5use std::path::{Path, PathBuf};
6use std::time::{SystemTime, UNIX_EPOCH};
7
8use ndarray::{Array1, Array2, ArrayView1, Axis};
9
10use crate::error::Result;
11
12/// Write a file through a same-directory temporary file, then atomically rename it into place.
13///
14/// This avoids leaving critical index files truncated if a process is interrupted after
15/// opening with `File::create` but before the full payload is written.
16pub fn atomic_write_file<F>(path: &Path, write_fn: F) -> Result<()>
17where
18    F: FnOnce(&mut File) -> Result<()>,
19{
20    let parent = path.parent().unwrap_or_else(|| Path::new("."));
21    fs::create_dir_all(parent)?;
22
23    let mut temp_path = atomic_temp_path(path);
24    let mut attempts = 0u32;
25    let mut file = loop {
26        match File::options()
27            .write(true)
28            .create_new(true)
29            .open(&temp_path)
30        {
31            Ok(file) => break file,
32            Err(err) if err.kind() == io::ErrorKind::AlreadyExists && attempts < 8 => {
33                attempts += 1;
34                temp_path = atomic_temp_path_with_attempt(path, attempts);
35            }
36            Err(err) => return Err(err.into()),
37        }
38    };
39
40    let write_result = write_fn(&mut file).and_then(|_| {
41        file.flush()?;
42        file.sync_all()?;
43        Ok(())
44    });
45
46    if let Err(err) = write_result {
47        drop(file);
48        let _ = fs::remove_file(&temp_path);
49        return Err(err);
50    }
51
52    drop(file);
53    fs::rename(&temp_path, path)?;
54
55    if let Ok(parent_dir) = File::open(parent) {
56        let _ = parent_dir.sync_all();
57    }
58
59    Ok(())
60}
61
62fn atomic_temp_path(path: &Path) -> PathBuf {
63    atomic_temp_path_with_attempt(path, 0)
64}
65
66fn atomic_temp_path_with_attempt(path: &Path, attempt: u32) -> PathBuf {
67    let file_name = path
68        .file_name()
69        .and_then(|name| name.to_str())
70        .unwrap_or("atomic-write");
71    let nanos = SystemTime::now()
72        .duration_since(UNIX_EPOCH)
73        .map(|duration| duration.as_nanos())
74        .unwrap_or_default();
75    let tmp_name = format!(
76        ".{}.{}.{}.tmp",
77        file_name,
78        std::process::id(),
79        nanos + attempt as u128
80    );
81    path.with_file_name(tmp_name)
82}
83
84/// Compute the k-th quantile of a 1D array using linear interpolation.
85///
86/// # Arguments
87///
88/// * `arr` - Input array (will be sorted)
89/// * `q` - Quantile to compute (between 0.0 and 1.0)
90///
91/// # Returns
92///
93/// The quantile value
94pub fn quantile(arr: &Array1<f32>, q: f64) -> f32 {
95    if arr.is_empty() {
96        return 0.0;
97    }
98
99    let mut sorted: Vec<f32> = arr.iter().copied().collect();
100    sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
101
102    let n = sorted.len();
103    let idx_float = q * (n - 1) as f64;
104    let lower_idx = idx_float.floor() as usize;
105    let upper_idx = idx_float.ceil() as usize;
106
107    if lower_idx == upper_idx {
108        sorted[lower_idx]
109    } else {
110        let weight = (idx_float - lower_idx as f64) as f32;
111        sorted[lower_idx] * (1.0 - weight) + sorted[upper_idx] * weight
112    }
113}
114
115/// Compute multiple quantiles efficiently.
116///
117/// # Arguments
118///
119/// * `arr` - Input array
120/// * `quantiles` - Array of quantiles to compute
121///
122/// # Returns
123///
124/// Array of quantile values
125pub fn quantiles(arr: &Array1<f32>, qs: &[f64]) -> Vec<f32> {
126    if arr.is_empty() {
127        return vec![0.0; qs.len()];
128    }
129
130    let mut sorted: Vec<f32> = arr.iter().copied().collect();
131    sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
132
133    let n = sorted.len();
134
135    qs.iter()
136        .map(|&q| {
137            let idx_float = q * (n - 1) as f64;
138            let lower_idx = idx_float.floor() as usize;
139            let upper_idx = idx_float.ceil() as usize;
140
141            if lower_idx == upper_idx {
142                sorted[lower_idx]
143            } else {
144                let weight = (idx_float - lower_idx as f64) as f32;
145                sorted[lower_idx] * (1.0 - weight) + sorted[upper_idx] * weight
146            }
147        })
148        .collect()
149}
150
151/// Normalize rows of a 2D array to unit length.
152///
153/// # Arguments
154///
155/// * `arr` - Input array of shape `[N, dim]`
156///
157/// # Returns
158///
159/// Normalized array
160pub fn normalize_rows(arr: &Array2<f32>) -> Array2<f32> {
161    let mut result = arr.clone();
162    for mut row in result.axis_iter_mut(Axis(0)) {
163        let norm = row.dot(&row).sqrt().max(1e-12);
164        row /= norm;
165    }
166    result
167}
168
169/// Compute L2 norm of each row.
170///
171/// # Arguments
172///
173/// * `arr` - Input array of shape `[N, dim]`
174///
175/// # Returns
176///
177/// Array of norms of shape `[N]`
178pub fn row_norms(arr: &Array2<f32>) -> Array1<f32> {
179    arr.axis_iter(Axis(0))
180        .map(|row| row.dot(&row).sqrt())
181        .collect()
182}
183
184/// Pack bits into bytes (big-endian).
185///
186/// # Arguments
187///
188/// * `bits` - Array of bits (0 or 1)
189///
190/// # Returns
191///
192/// Packed bytes
193pub fn packbits(bits: &[u8]) -> Vec<u8> {
194    bits.chunks(8)
195        .map(|chunk| {
196            let mut byte = 0u8;
197            for (i, &bit) in chunk.iter().enumerate() {
198                byte |= bit << (7 - i);
199            }
200            byte
201        })
202        .collect()
203}
204
205/// Unpack bytes into bits (big-endian).
206///
207/// # Arguments
208///
209/// * `bytes` - Packed bytes
210///
211/// # Returns
212///
213/// Unpacked bits
214pub fn unpackbits(bytes: &[u8]) -> Vec<u8> {
215    let mut bits = Vec::with_capacity(bytes.len() * 8);
216    for &byte in bytes {
217        for i in (0..8).rev() {
218            bits.push((byte >> i) & 1);
219        }
220    }
221    bits
222}
223
224/// Create a boolean mask from sequence lengths.
225///
226/// # Arguments
227///
228/// * `lengths` - Array of sequence lengths
229/// * `max_len` - Maximum sequence length
230///
231/// # Returns
232///
233/// Boolean mask of shape `[batch_size, max_len]`
234pub fn create_mask(lengths: &ArrayView1<i64>, max_len: usize) -> Array2<bool> {
235    let batch_size = lengths.len();
236    let mut mask = Array2::from_elem((batch_size, max_len), false);
237
238    for (i, &len) in lengths.iter().enumerate() {
239        for j in 0..(len as usize).min(max_len) {
240            mask[[i, j]] = true;
241        }
242    }
243
244    mask
245}
246
247/// Pad sequences to uniform length.
248///
249/// # Arguments
250///
251/// * `sequences` - List of sequence arrays
252/// * `pad_value` - Value to use for padding
253///
254/// # Returns
255///
256/// Tuple of (padded array, lengths)
257pub fn pad_sequences(sequences: &[Array2<f32>], pad_value: f32) -> (Array2<f32>, Array1<i64>) {
258    if sequences.is_empty() {
259        return (Array2::zeros((0, 0)), Array1::zeros(0));
260    }
261
262    let max_len = sequences.iter().map(|s| s.nrows()).max().unwrap_or(0);
263    let dim = sequences[0].ncols();
264    let batch_size = sequences.len();
265
266    let mut padded = Array2::from_elem((batch_size * max_len, dim), pad_value);
267    let mut lengths = Array1::<i64>::zeros(batch_size);
268
269    for (i, seq) in sequences.iter().enumerate() {
270        let len = seq.nrows();
271        lengths[i] = len as i64;
272        for j in 0..len {
273            for k in 0..dim {
274                padded[[i * max_len + j, k]] = seq[[j, k]];
275            }
276        }
277    }
278
279    (padded, lengths)
280}
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285    use crate::error::Error;
286    use std::io::Write;
287
288    #[test]
289    fn test_quantile() {
290        let arr = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0]);
291        assert!((quantile(&arr, 0.5) - 3.0).abs() < 1e-6);
292        assert!((quantile(&arr, 0.0) - 1.0).abs() < 1e-6);
293        assert!((quantile(&arr, 1.0) - 5.0).abs() < 1e-6);
294    }
295
296    #[test]
297    fn test_packbits_unpackbits() {
298        let bits = vec![1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1, 1, 0, 0, 0, 0];
299        let packed = packbits(&bits);
300        assert_eq!(packed, vec![0b10101010, 0b11110000]);
301
302        let unpacked = unpackbits(&packed);
303        assert_eq!(unpacked, bits);
304    }
305
306    #[test]
307    fn test_normalize_rows() {
308        let arr = Array2::from_shape_vec((2, 3), vec![3.0, 0.0, 4.0, 0.0, 5.0, 0.0]).unwrap();
309        let normalized = normalize_rows(&arr);
310
311        // First row: [3, 0, 4] / 5 = [0.6, 0, 0.8]
312        assert!((normalized[[0, 0]] - 0.6).abs() < 1e-6);
313        assert!((normalized[[0, 2]] - 0.8).abs() < 1e-6);
314
315        // Second row: [0, 5, 0] / 5 = [0, 1, 0]
316        assert!((normalized[[1, 1]] - 1.0).abs() < 1e-6);
317    }
318
319    #[test]
320    fn atomic_write_failure_preserves_original_file() {
321        let dir = tempfile::tempdir().unwrap();
322        let path = dir.path().join("centroids.npy");
323        std::fs::write(&path, b"original").unwrap();
324
325        let result = atomic_write_file(&path, |file| {
326            file.write_all(b"partial")?;
327            Err(Error::Update("forced failure".to_string()))
328        });
329
330        assert!(result.is_err());
331        assert_eq!(std::fs::read(&path).unwrap(), b"original");
332        let temp_entries: Vec<_> = std::fs::read_dir(dir.path())
333            .unwrap()
334            .filter_map(|entry| entry.ok())
335            .filter(|entry| entry.file_name().to_string_lossy().ends_with(".tmp"))
336            .collect();
337        assert!(temp_entries.is_empty());
338    }
339}