Skip to main content

dino_quant/
lib.rs

1use std::f32::consts::PI;
2use std::{io::Read, path::Path};
3
4const BASE_A: u8 = 0;
5const BASE_C: u8 = 1;
6const BASE_G: u8 = 2;
7const BASE_T: u8 = 3;
8const CONTIG_SEPARATOR: u8 = b'N';
9
10#[derive(Clone, Copy, Debug)]
11pub struct SplitMix64 {
12    state: u64,
13}
14
15impl SplitMix64 {
16    pub const fn new(seed: u64) -> Self {
17        Self { state: seed }
18    }
19
20    pub fn next_u64(&mut self) -> u64 {
21        self.state = self.state.wrapping_add(0x9e37_79b9_7f4a_7c15);
22        mix_u64(self.state)
23    }
24}
25
26#[derive(Clone, Copy, Debug)]
27pub struct QuantizerConfig {
28    pub bits: u8,
29    pub clip_sigma: f32,
30    pub rotation_seed: u64,
31    pub qjl_seed: u64,
32    pub use_qjl_residual: bool,
33}
34
35impl Default for QuantizerConfig {
36    fn default() -> Self {
37        Self {
38            bits: 4,
39            clip_sigma: 4.0,
40            rotation_seed: 0x51d0_51d0_51d0_51d0,
41            qjl_seed: 0x71b0_7171_b071_7171,
42            use_qjl_residual: true,
43        }
44    }
45}
46
47#[derive(Clone, Debug)]
48pub struct QjlResidual {
49    signs: Vec<u64>,
50    dim: usize,
51    norm: f32,
52    seed: u64,
53}
54
55#[derive(Clone, Debug)]
56struct PackedCodes {
57    data: Vec<u8>,
58    len: usize,
59    bits: u8,
60}
61
62#[derive(Clone, Debug)]
63pub struct QuantizedVector {
64    dim: usize,
65    bits: u8,
66    clip: f32,
67    rotation_seed: u64,
68    codes: PackedCodes,
69    qjl_residual: Option<QjlResidual>,
70}
71
72#[derive(Clone, Debug)]
73pub struct PreparedQuantizedQuery {
74    dim: usize,
75    rotation_seed: u64,
76    qjl_seed: Option<u64>,
77    rotated: Vec<f32>,
78    rotated_sum: f32,
79    scalar_4bit_lookup: Vec<[f32; 256]>,
80    rotated_qjl: Option<Vec<f32>>,
81}
82
83#[derive(Clone, Debug)]
84pub struct QuantizedVectorSnapshot {
85    pub dim: usize,
86    pub bits: u8,
87    pub clip: f32,
88    pub rotation_seed: u64,
89    pub codes_data: Vec<u8>,
90    pub codes_len: usize,
91    pub codes_bits: u8,
92    pub qjl_residual: Option<QjlResidualSnapshot>,
93}
94
95#[derive(Clone, Debug)]
96pub struct QjlResidualSnapshot {
97    pub signs: Vec<u64>,
98    pub dim: usize,
99    pub norm: f32,
100    pub seed: u64,
101}
102
103#[derive(Clone, Copy, Debug)]
104pub struct ReconstructionMetrics {
105    pub mse: f32,
106    pub cosine: f32,
107    pub dot_error: f32,
108}
109
110#[derive(Clone, Debug)]
111pub struct SequenceRecord {
112    pub name: String,
113    pub bases: Vec<u8>,
114}
115
116#[derive(Clone, Copy, Debug)]
117pub struct SequenceRecordRef<'a> {
118    pub name: &'a [u8],
119    pub bases: &'a [u8],
120}
121
122#[derive(Clone, Copy, Debug)]
123pub struct ReferenceIndexConfig {
124    pub k: usize,
125    pub dim: usize,
126    pub window_len: usize,
127    pub stride: usize,
128    pub quantizer: QuantizerConfig,
129}
130
131#[derive(Clone, Debug)]
132pub struct SearchHit {
133    pub target_name: String,
134    pub target_start: usize,
135    pub target_end: usize,
136    pub start: usize,
137    pub end: usize,
138    pub score: f32,
139}
140
141#[derive(Clone, Debug)]
142pub struct ReferenceWindowIndex {
143    config: ReferenceIndexConfig,
144    windows: Vec<ReferenceWindow>,
145}
146
147#[derive(Clone, Debug)]
148struct ReferenceWindow {
149    target_name: String,
150    target_start: usize,
151    target_end: usize,
152    start: usize,
153    end: usize,
154    sketch: QuantizedVector,
155}
156
157impl QuantizerConfig {
158    pub fn encode(&self, vector: &[f32]) -> Result<QuantizedVector, String> {
159        self.validate(vector.len())?;
160
161        let dim = vector.len();
162        let clip = self.clip_sigma / (dim as f32).sqrt();
163        let rotated = rotate(vector, self.rotation_seed)?;
164        let codes = encode_scalar_codes(&rotated, self.bits, clip)?;
165
166        let mut quantized = QuantizedVector {
167            dim,
168            bits: self.bits,
169            clip,
170            rotation_seed: self.rotation_seed,
171            codes,
172            qjl_residual: None,
173        };
174
175        if self.use_qjl_residual {
176            let base = quantized.decode()?;
177            let residual = subtract(vector, &base)?;
178            quantized.qjl_residual = QjlResidual::encode(&residual, self.qjl_seed)?;
179        }
180
181        Ok(quantized)
182    }
183
184    fn validate(&self, dim: usize) -> Result<(), String> {
185        if dim == 0 {
186            return Err("vector dimension must be non-zero".to_owned());
187        }
188        if !dim.is_power_of_two() {
189            return Err(
190                "vector dimension must be a power of two for the Hadamard rotation".to_owned(),
191            );
192        }
193        if !(1..=8).contains(&self.bits) {
194            return Err("quantizer bits must be in 1..=8".to_owned());
195        }
196        if !self.clip_sigma.is_finite() || self.clip_sigma <= 0.0 {
197            return Err("clip_sigma must be finite and positive".to_owned());
198        }
199        Ok(())
200    }
201}
202
203impl QuantizedVector {
204    pub fn prepare_approximate_query(
205        query: &[f32],
206        rotation_seed: u64,
207        qjl_seed: Option<u64>,
208    ) -> Result<PreparedQuantizedQuery, String> {
209        let rotated = rotate(query, rotation_seed)?;
210        let rotated_sum = rotated.iter().sum();
211        let scalar_4bit_lookup = build_4bit_lookup(&rotated);
212        let rotated_qjl = qjl_seed.map(|seed| rotate(query, seed)).transpose()?;
213        Ok(PreparedQuantizedQuery {
214            dim: query.len(),
215            rotation_seed,
216            qjl_seed,
217            rotated,
218            rotated_sum,
219            scalar_4bit_lookup,
220            rotated_qjl,
221        })
222    }
223
224    pub fn snapshot(&self) -> QuantizedVectorSnapshot {
225        QuantizedVectorSnapshot {
226            dim: self.dim,
227            bits: self.bits,
228            clip: self.clip,
229            rotation_seed: self.rotation_seed,
230            codes_data: self.codes.data.clone(),
231            codes_len: self.codes.len,
232            codes_bits: self.codes.bits,
233            qjl_residual: self
234                .qjl_residual
235                .as_ref()
236                .map(|residual| QjlResidualSnapshot {
237                    signs: residual.signs.clone(),
238                    dim: residual.dim,
239                    norm: residual.norm,
240                    seed: residual.seed,
241                }),
242        }
243    }
244
245    pub fn from_snapshot(snapshot: QuantizedVectorSnapshot) -> Result<Self, String> {
246        if snapshot.dim == 0 || !snapshot.dim.is_power_of_two() {
247            return Err(
248                "quantized vector snapshot dimension must be a non-zero power of two".to_owned(),
249            );
250        }
251        if snapshot.codes_len != snapshot.dim {
252            return Err("quantized vector snapshot code length must match dimension".to_owned());
253        }
254        if !(1..=8).contains(&snapshot.bits) || !(1..=8).contains(&snapshot.codes_bits) {
255            return Err("quantized vector snapshot bits must be in 1..=8".to_owned());
256        }
257        if snapshot.bits != snapshot.codes_bits {
258            return Err("quantized vector snapshot code bits mismatch".to_owned());
259        }
260        if !snapshot.clip.is_finite() || snapshot.clip <= 0.0 {
261            return Err("quantized vector snapshot clip must be finite and positive".to_owned());
262        }
263        let expected_code_bytes =
264            (snapshot.codes_len * usize::from(snapshot.codes_bits)).div_ceil(8);
265        if snapshot.codes_data.len() != expected_code_bytes {
266            return Err("quantized vector snapshot code byte length mismatch".to_owned());
267        }
268        let qjl_residual = snapshot
269            .qjl_residual
270            .map(|residual| {
271                if residual.dim != snapshot.dim {
272                    return Err(
273                        "QJL snapshot dimension must match quantized vector dimension".to_owned(),
274                    );
275                }
276                if residual.signs.len() != residual.dim.div_ceil(64) {
277                    return Err("QJL snapshot sign word length mismatch".to_owned());
278                }
279                if !residual.norm.is_finite() || residual.norm <= 0.0 {
280                    return Err("QJL snapshot norm must be finite and positive".to_owned());
281                }
282                Ok(QjlResidual {
283                    signs: residual.signs,
284                    dim: residual.dim,
285                    norm: residual.norm,
286                    seed: residual.seed,
287                })
288            })
289            .transpose()?;
290        Ok(Self {
291            dim: snapshot.dim,
292            bits: snapshot.bits,
293            clip: snapshot.clip,
294            rotation_seed: snapshot.rotation_seed,
295            codes: PackedCodes {
296                data: snapshot.codes_data,
297                len: snapshot.codes_len,
298                bits: snapshot.codes_bits,
299            },
300            qjl_residual,
301        })
302    }
303
304    pub fn decode(&self) -> Result<Vec<f32>, String> {
305        if self.codes.len() != self.dim {
306            return Err("quantized code length does not match dimension".to_owned());
307        }
308
309        let rotated = decode_scalar_codes(&self.codes, self.bits, self.clip)?;
310        let mut decoded = inverse_rotate(&rotated, self.rotation_seed)?;
311        if let Some(residual) = &self.qjl_residual {
312            let correction = residual.decode()?;
313            if correction.len() != decoded.len() {
314                return Err("QJL correction length does not match decoded vector".to_owned());
315            }
316            for (dst, corr) in decoded.iter_mut().zip(correction) {
317                *dst += corr;
318            }
319        }
320        Ok(decoded)
321    }
322
323    pub fn compressed_bits(&self) -> usize {
324        let scalar_bits = self.codes.byte_len() * 8;
325        let qjl_bits = self
326            .qjl_residual
327            .as_ref()
328            .map_or(0, |residual| residual.dim + 32);
329        scalar_bits + qjl_bits
330    }
331
332    pub fn compressed_bytes(&self) -> usize {
333        self.compressed_bits().div_ceil(8)
334    }
335
336    pub fn scalar_dot_rotated_query(&self, rotated_query: &[f32]) -> Result<f32, String> {
337        if rotated_query.len() != self.dim {
338            return Err("rotated query dimension does not match quantized vector".to_owned());
339        }
340        let rotated_query_sum = rotated_query.iter().sum();
341        self.scalar_dot_rotated_query_with_sum(rotated_query, rotated_query_sum)
342    }
343
344    fn scalar_dot_rotated_query_with_sum(
345        &self,
346        rotated_query: &[f32],
347        rotated_query_sum: f32,
348    ) -> Result<f32, String> {
349        if rotated_query.len() != self.dim {
350            return Err("rotated query dimension does not match quantized vector".to_owned());
351        }
352        self.codes
353            .decoded_dot(rotated_query, rotated_query_sum, self.clip)
354    }
355
356    fn approximate_dot_4bit_lookup(
357        &self,
358        lookup: &[[f32; 256]],
359        rotated_query_sum: f32,
360        rotated_qjl_query: Option<&[f32]>,
361    ) -> Result<f32, String> {
362        if self.bits != 4 {
363            return Err("4-bit lookup scoring requires a 4-bit quantized vector".to_owned());
364        }
365        let scale = 2.0 * self.clip / 15.0;
366        let scalar_score = (-self.clip * rotated_query_sum)
367            + scale * self.codes.weighted_code_sum_4_lookup(lookup)?;
368        Ok(scalar_score + self.residual_dot_rotated_query(rotated_qjl_query)?)
369    }
370
371    fn residual_dot_rotated_query(&self, rotated_qjl_query: Option<&[f32]>) -> Result<f32, String> {
372        match (&self.qjl_residual, rotated_qjl_query) {
373            (Some(residual), Some(query)) => residual.dot_rotated_query(query),
374            (Some(_), None) => Err("QJL residual scoring requires a rotated QJL query".to_owned()),
375            (None, _) => Ok(0.0),
376        }
377    }
378
379    pub fn scalar_dot_query(&self, query: &[f32]) -> Result<f32, String> {
380        let rotated_query = rotate(query, self.rotation_seed)?;
381        self.scalar_dot_rotated_query(&rotated_query)
382    }
383
384    pub fn approximate_dot_query(&self, query: &[f32]) -> Result<f32, String> {
385        let rotated_query = rotate(query, self.rotation_seed)?;
386        let rotated_query_sum = rotated_query.iter().sum();
387        let mut score =
388            self.scalar_dot_rotated_query_with_sum(&rotated_query, rotated_query_sum)?;
389        if let Some(residual) = &self.qjl_residual {
390            let rotated_residual_query = rotate(query, residual.seed)?;
391            score += residual.dot_rotated_query(&rotated_residual_query)?;
392        }
393        Ok(score)
394    }
395
396    pub fn approximate_dot_prepared_query(
397        &self,
398        query: &PreparedQuantizedQuery,
399    ) -> Result<f32, String> {
400        if query.dim != self.dim {
401            return Err("prepared query dimension does not match quantized vector".to_owned());
402        }
403        if query.rotation_seed != self.rotation_seed {
404            return Err("prepared query rotation seed does not match quantized vector".to_owned());
405        }
406        let rotated_qjl = if let Some(residual) = &self.qjl_residual {
407            if query.qjl_seed != Some(residual.seed) {
408                return Err("prepared query QJL seed does not match residual".to_owned());
409            }
410            Some(
411                query
412                    .rotated_qjl
413                    .as_deref()
414                    .ok_or_else(|| "prepared query is missing QJL rotation".to_owned())?,
415            )
416        } else {
417            None
418        };
419        if self.bits == 4 {
420            return self.approximate_dot_4bit_lookup(
421                &query.scalar_4bit_lookup,
422                query.rotated_sum,
423                rotated_qjl,
424            );
425        }
426        let mut score =
427            self.scalar_dot_rotated_query_with_sum(&query.rotated, query.rotated_sum)?;
428        if let Some(residual) = &self.qjl_residual {
429            let Some(rotated_qjl) = rotated_qjl else {
430                return Err("prepared query is missing QJL rotation".to_owned());
431            };
432            score += residual.dot_rotated_query(rotated_qjl)?;
433        }
434        Ok(score)
435    }
436}
437
438impl PackedCodes {
439    fn encode(values: &[u8], bits: u8) -> Result<Self, String> {
440        if !(1..=8).contains(&bits) {
441            return Err("packed code bits must be in 1..=8".to_owned());
442        }
443        let mask = (1_u16 << bits) - 1;
444        let mut data = vec![0_u8; (values.len() * usize::from(bits)).div_ceil(8)];
445        for (idx, code) in values.iter().enumerate() {
446            if u16::from(*code) > mask {
447                return Err("scalar code exceeds bit width".to_owned());
448            }
449            let bit_pos = idx * usize::from(bits);
450            let byte_idx = bit_pos / 8;
451            let bit_offset = bit_pos % 8;
452            let shifted = u16::from(*code) << bit_offset;
453            data[byte_idx] |= shifted as u8;
454            if bit_offset + usize::from(bits) > 8 {
455                data[byte_idx + 1] |= (shifted >> 8) as u8;
456            }
457        }
458
459        Ok(Self {
460            data,
461            len: values.len(),
462            bits,
463        })
464    }
465
466    fn len(&self) -> usize {
467        self.len
468    }
469
470    fn byte_len(&self) -> usize {
471        self.data.len()
472    }
473
474    fn get(&self, idx: usize) -> Result<u8, String> {
475        if idx >= self.len {
476            return Err("packed code index out of bounds".to_owned());
477        }
478        let bit_pos = idx * usize::from(self.bits);
479        let byte_idx = bit_pos / 8;
480        let bit_offset = bit_pos % 8;
481        let mut value = u16::from(self.data[byte_idx] >> bit_offset);
482        if bit_offset + usize::from(self.bits) > 8 {
483            value |= u16::from(self.data[byte_idx + 1]) << (8 - bit_offset);
484        }
485        Ok((value & ((1_u16 << self.bits) - 1)) as u8)
486    }
487
488    fn decoded_dot(
489        &self,
490        rotated_query: &[f32],
491        rotated_query_sum: f32,
492        clip: f32,
493    ) -> Result<f32, String> {
494        if rotated_query.len() != self.len {
495            return Err("rotated query length does not match packed code length".to_owned());
496        }
497        if !clip.is_finite() || clip <= 0.0 {
498            return Err("clip must be finite and positive".to_owned());
499        }
500
501        let levels = (1_u16 << self.bits) - 1;
502        let scale = 2.0 * clip / levels as f32;
503        let offset_sum = -clip * rotated_query_sum;
504        Ok(offset_sum + scale * self.weighted_code_sum(rotated_query)?)
505    }
506
507    fn weighted_code_sum(&self, rotated_query: &[f32]) -> Result<f32, String> {
508        match self.bits {
509            2 => Ok(self.weighted_code_sum_2(rotated_query)),
510            4 => Ok(self.weighted_code_sum_4(rotated_query)),
511            8 => Ok(self.weighted_code_sum_8(rotated_query)),
512            _ => self.weighted_code_sum_generic(rotated_query),
513        }
514    }
515
516    fn weighted_code_sum_2(&self, rotated_query: &[f32]) -> f32 {
517        let mut sum = 0.0_f32;
518        for (byte_idx, byte) in self.data.iter().enumerate() {
519            let idx = byte_idx * 4;
520            if idx >= self.len {
521                break;
522            }
523            sum += f32::from(byte & 0b0000_0011) * rotated_query[idx];
524            if idx + 1 < self.len {
525                sum += f32::from((byte >> 2) & 0b0000_0011) * rotated_query[idx + 1];
526            }
527            if idx + 2 < self.len {
528                sum += f32::from((byte >> 4) & 0b0000_0011) * rotated_query[idx + 2];
529            }
530            if idx + 3 < self.len {
531                sum += f32::from(byte >> 6) * rotated_query[idx + 3];
532            }
533        }
534        sum
535    }
536
537    fn weighted_code_sum_4(&self, rotated_query: &[f32]) -> f32 {
538        let mut sum = 0.0_f32;
539        for (byte_idx, byte) in self.data.iter().enumerate() {
540            let idx = byte_idx * 2;
541            if idx >= self.len {
542                break;
543            }
544            sum += f32::from(byte & 0x0f) * rotated_query[idx];
545            if idx + 1 < self.len {
546                sum += f32::from(byte >> 4) * rotated_query[idx + 1];
547            }
548        }
549        sum
550    }
551
552    fn weighted_code_sum_4_lookup(&self, lookup: &[[f32; 256]]) -> Result<f32, String> {
553        if lookup.len() != self.data.len() {
554            return Err("4-bit lookup table length does not match packed code bytes".to_owned());
555        }
556        let mut sum = 0.0_f32;
557        for (idx, byte) in self.data.iter().enumerate() {
558            sum += lookup[idx][usize::from(*byte)];
559        }
560        Ok(sum)
561    }
562
563    fn weighted_code_sum_8(&self, rotated_query: &[f32]) -> f32 {
564        let mut sum = 0.0_f32;
565        for (code, query_value) in self.data.iter().zip(rotated_query) {
566            sum += f32::from(*code) * query_value;
567        }
568        sum
569    }
570
571    fn weighted_code_sum_generic(&self, rotated_query: &[f32]) -> Result<f32, String> {
572        let mut sum = 0.0_f32;
573        let mask = (1_u32 << self.bits) - 1;
574        let mut byte_idx = 0;
575        let mut bit_buffer = 0_u32;
576        let mut bits_in_buffer = 0_u8;
577        for query_value in rotated_query {
578            while bits_in_buffer < self.bits {
579                if byte_idx >= self.data.len() {
580                    return Err("packed code buffer ended early".to_owned());
581                }
582                bit_buffer |= u32::from(self.data[byte_idx]) << bits_in_buffer;
583                bits_in_buffer += 8;
584                byte_idx += 1;
585            }
586
587            let code = bit_buffer & mask;
588            bit_buffer >>= self.bits;
589            bits_in_buffer -= self.bits;
590            sum += code as f32 * query_value;
591        }
592        Ok(sum)
593    }
594
595    fn decode_all(&self, clip: f32) -> Result<Vec<f32>, String> {
596        if !clip.is_finite() || clip <= 0.0 {
597            return Err("clip must be finite and positive".to_owned());
598        }
599        let levels = (1_u16 << self.bits) - 1;
600        let scale = 2.0 * clip / levels as f32;
601        let offset = -clip;
602        let mut decoded = Vec::with_capacity(self.len);
603        for idx in 0..self.len {
604            decoded.push(offset + f32::from(self.get(idx)?) * scale);
605        }
606        Ok(decoded)
607    }
608}
609
610impl ReferenceIndexConfig {
611    pub fn validate(&self) -> Result<(), String> {
612        if self.window_len == 0 {
613            return Err("window_len must be non-zero".to_owned());
614        }
615        if self.stride == 0 {
616            return Err("stride must be non-zero".to_owned());
617        }
618        if self.k == 0 || self.k > 31 {
619            return Err("k must be in 1..=31".to_owned());
620        }
621        self.quantizer.validate(self.dim)
622    }
623}
624
625impl ReferenceWindowIndex {
626    pub fn build(reference: &[u8], config: ReferenceIndexConfig) -> Result<Self, String> {
627        config.validate()?;
628        if reference.len() < config.window_len {
629            return Err("reference is shorter than the configured window length".to_owned());
630        }
631
632        let window_count = ((reference.len() - config.window_len) / config.stride) + 1;
633        let mut windows = Vec::with_capacity(window_count);
634        for start in (0..=reference.len() - config.window_len).step_by(config.stride) {
635            let end = start + config.window_len;
636            let sketch = dna_kmer_sketch(&reference[start..end], config.k, config.dim)?;
637            let sketch = config.quantizer.encode(&sketch)?;
638            windows.push(ReferenceWindow {
639                target_name: "reference".to_owned(),
640                target_start: start,
641                target_end: end,
642                start,
643                end,
644                sketch,
645            });
646        }
647
648        Ok(Self { config, windows })
649    }
650
651    pub fn build_records(
652        records: &[SequenceRecord],
653        config: ReferenceIndexConfig,
654    ) -> Result<Self, String> {
655        config.validate()?;
656        if records.is_empty() {
657            return Err("at least one sequence record is required".to_owned());
658        }
659
660        let mut windows = Vec::new();
661        let mut linear_offset = 0_usize;
662        for (idx, record) in records.iter().enumerate() {
663            if record.bases.len() >= config.window_len {
664                let window_count = ((record.bases.len() - config.window_len) / config.stride) + 1;
665                windows.reserve(window_count);
666                for target_start in
667                    (0..=record.bases.len() - config.window_len).step_by(config.stride)
668                {
669                    let target_end = target_start + config.window_len;
670                    let sketch = dna_kmer_sketch(
671                        &record.bases[target_start..target_end],
672                        config.k,
673                        config.dim,
674                    )?;
675                    let sketch = config.quantizer.encode(&sketch)?;
676                    let start = linear_offset + target_start;
677                    let end = linear_offset + target_end;
678                    windows.push(ReferenceWindow {
679                        target_name: record.name.clone(),
680                        target_start,
681                        target_end,
682                        start,
683                        end,
684                        sketch,
685                    });
686                }
687            }
688            linear_offset += record.bases.len();
689            if idx + 1 < records.len() {
690                linear_offset += 1;
691            }
692        }
693
694        if windows.is_empty() {
695            return Err("no reference record is long enough for the configured window".to_owned());
696        }
697
698        Ok(Self { config, windows })
699    }
700
701    pub fn search_sequence(&self, query: &[u8], top_k: usize) -> Result<Vec<SearchHit>, String> {
702        let sketch = dna_kmer_sketch(query, self.config.k, self.config.dim)?;
703        self.search_sketch(&sketch, top_k)
704    }
705
706    pub fn search_sketch(&self, query: &[f32], top_k: usize) -> Result<Vec<SearchHit>, String> {
707        if top_k == 0 {
708            return Ok(Vec::new());
709        }
710        if query.len() != self.config.dim {
711            return Err("query sketch dimension does not match index dimension".to_owned());
712        }
713
714        let rotated_query = rotate(query, self.config.quantizer.rotation_seed)?;
715        let rotated_query_sum = rotated_query.iter().sum();
716        let rotated_qjl_query = if self.config.quantizer.use_qjl_residual {
717            Some(rotate(query, self.config.quantizer.qjl_seed)?)
718        } else {
719            None
720        };
721        if self.config.quantizer.bits == 4 {
722            return self.search_rotated_4bit_lookup(
723                &rotated_query,
724                rotated_query_sum,
725                rotated_qjl_query.as_deref(),
726                top_k,
727            );
728        }
729
730        let mut top = Vec::with_capacity(top_k.min(self.windows.len()));
731        for window in &self.windows {
732            let scalar_score = window
733                .sketch
734                .scalar_dot_rotated_query_with_sum(&rotated_query, rotated_query_sum)?;
735            let score = scalar_score
736                + window
737                    .sketch
738                    .residual_dot_rotated_query(rotated_qjl_query.as_deref())?;
739            push_top_hit(
740                &mut top,
741                top_k,
742                SearchHit {
743                    target_name: window.target_name.clone(),
744                    target_start: window.target_start,
745                    target_end: window.target_end,
746                    start: window.start,
747                    end: window.end,
748                    score,
749                },
750            );
751        }
752        top.sort_by(|left, right| right.score.total_cmp(&left.score));
753        Ok(top)
754    }
755
756    fn search_rotated_4bit_lookup(
757        &self,
758        rotated_query: &[f32],
759        rotated_query_sum: f32,
760        rotated_qjl_query: Option<&[f32]>,
761        top_k: usize,
762    ) -> Result<Vec<SearchHit>, String> {
763        let lookup = build_4bit_lookup(rotated_query);
764        let mut top = Vec::with_capacity(top_k.min(self.windows.len()));
765        for window in &self.windows {
766            let score = window.sketch.approximate_dot_4bit_lookup(
767                &lookup,
768                rotated_query_sum,
769                rotated_qjl_query,
770            )?;
771            push_top_hit(
772                &mut top,
773                top_k,
774                SearchHit {
775                    target_name: window.target_name.clone(),
776                    target_start: window.target_start,
777                    target_end: window.target_end,
778                    start: window.start,
779                    end: window.end,
780                    score,
781                },
782            );
783        }
784        top.sort_by(|left, right| right.score.total_cmp(&left.score));
785        Ok(top)
786    }
787
788    pub fn window_count(&self) -> usize {
789        self.windows.len()
790    }
791
792    pub fn compressed_bytes(&self) -> usize {
793        self.windows
794            .iter()
795            .map(|window| window.sketch.compressed_bytes())
796            .sum()
797    }
798
799    pub fn raw_vector_bytes(&self) -> usize {
800        self.windows.len() * self.config.dim * std::mem::size_of::<f32>()
801    }
802
803    pub fn compression_ratio(&self) -> f32 {
804        if self.compressed_bytes() == 0 {
805            return 0.0;
806        }
807        self.raw_vector_bytes() as f32 / self.compressed_bytes() as f32
808    }
809
810    pub fn config(&self) -> ReferenceIndexConfig {
811        self.config
812    }
813}
814
815impl QjlResidual {
816    fn encode(residual: &[f32], seed: u64) -> Result<Option<Self>, String> {
817        let norm = l2_norm(residual);
818        if norm <= f32::EPSILON {
819            return Ok(None);
820        }
821
822        let projected = rotate(residual, seed)?;
823        let mut signs = vec![0_u64; projected.len().div_ceil(64)];
824        for (idx, value) in projected.iter().enumerate() {
825            if *value >= 0.0 {
826                signs[idx / 64] |= 1_u64 << (idx % 64);
827            }
828        }
829
830        Ok(Some(Self {
831            signs,
832            dim: projected.len(),
833            norm,
834            seed,
835        }))
836    }
837
838    fn decode(&self) -> Result<Vec<f32>, String> {
839        if self.dim == 0 || !self.dim.is_power_of_two() {
840            return Err("QJL residual dimension must be a non-zero power of two".to_owned());
841        }
842        let scale = (PI / 2.0).sqrt() * self.norm / self.dim as f32;
843        let mut projected = vec![0.0_f32; self.dim];
844        for (idx, slot) in projected.iter_mut().enumerate() {
845            let bit = (self.signs[idx / 64] >> (idx % 64)) & 1;
846            *slot = if bit == 1 { scale } else { -scale };
847        }
848        inverse_rotate(&projected, self.seed)
849    }
850
851    fn dot_rotated_query(&self, rotated_query: &[f32]) -> Result<f32, String> {
852        if rotated_query.len() != self.dim {
853            return Err("QJL rotated query dimension does not match residual dimension".to_owned());
854        }
855        let scale = (PI / 2.0).sqrt() * self.norm / self.dim as f32;
856        let mut sum = 0.0_f32;
857        for (idx, query_value) in rotated_query.iter().enumerate() {
858            let bit = (self.signs[idx / 64] >> (idx % 64)) & 1;
859            let sign = if bit == 1 { 1.0 } else { -1.0 };
860            sum += sign * query_value;
861        }
862        Ok(scale * sum)
863    }
864}
865
866pub fn rotate(vector: &[f32], seed: u64) -> Result<Vec<f32>, String> {
867    validate_hadamard_dim(vector.len())?;
868    let mut rotated = Vec::with_capacity(vector.len());
869    for (idx, value) in vector.iter().enumerate() {
870        rotated.push(*value * coordinate_sign(seed, idx));
871    }
872    hadamard_in_place(&mut rotated);
873    let scale = 1.0 / (vector.len() as f32).sqrt();
874    for value in &mut rotated {
875        *value *= scale;
876    }
877    Ok(rotated)
878}
879
880pub fn inverse_rotate(rotated: &[f32], seed: u64) -> Result<Vec<f32>, String> {
881    validate_hadamard_dim(rotated.len())?;
882    let mut vector = rotated.to_vec();
883    hadamard_in_place(&mut vector);
884    let scale = 1.0 / (rotated.len() as f32).sqrt();
885    for (idx, value) in vector.iter_mut().enumerate() {
886        *value *= scale * coordinate_sign(seed, idx);
887    }
888    Ok(vector)
889}
890
891pub fn dna_kmer_sketch(seq: &[u8], k: usize, dim: usize) -> Result<Vec<f32>, String> {
892    if k == 0 || k > 31 {
893        return Err("k must be in 1..=31".to_owned());
894    }
895    if dim == 0 {
896        return Err("sketch dimension must be non-zero".to_owned());
897    }
898    if seq.len() < k {
899        return Ok(vec![0.0; dim]);
900    }
901
902    let mut sketch = vec![0.0_f32; dim];
903    for window in seq.windows(k) {
904        if let Some(code) = canonical_kmer_code(window) {
905            let hash = mix_u64(code ^ ((k as u64) << 56));
906            let bucket = (hash as usize) % dim;
907            let sign = if (hash >> 63) == 0 { 1.0 } else { -1.0 };
908            sketch[bucket] += sign;
909        }
910    }
911    l2_normalize(&mut sketch);
912    Ok(sketch)
913}
914
915pub fn protein_kmer_sketch(seq: &[u8], k: usize, dim: usize) -> Result<Vec<f32>, String> {
916    if k == 0 || k > 16 {
917        return Err("protein k must be in 1..=16".to_owned());
918    }
919    if dim == 0 {
920        return Err("sketch dimension must be non-zero".to_owned());
921    }
922    if seq.len() < k {
923        return Ok(vec![0.0; dim]);
924    }
925
926    let mut sketch = vec![0.0_f32; dim];
927    for window in seq.windows(k) {
928        if let Some(code) = protein_kmer_code(window) {
929            let hash = mix_u64(code ^ ((k as u64) << 56));
930            let bucket = (hash as usize) % dim;
931            let sign = if (hash >> 63) == 0 { 1.0 } else { -1.0 };
932            sketch[bucket] += sign;
933        }
934    }
935    l2_normalize(&mut sketch);
936    Ok(sketch)
937}
938
939#[cfg(test)]
940fn parse_fasta_bytes(input: &[u8]) -> Result<Vec<SequenceRecord>, String> {
941    let mut records = Vec::new();
942    dino_seq::visit_fasta_bytes(input, |record: dino_seq::FastaVisitRecord<'_>| {
943        let mut bases = Vec::with_capacity(record.seq().len());
944        append_sequence_line(record.seq(), &mut bases).map_err(dino_seq_format_error)?;
945        records.push(SequenceRecord {
946            name: parse_record_name(record.name_without_gt()).map_err(dino_seq_format_error)?,
947            bases,
948        });
949        Ok(())
950    })
951    .map_err(|err| format!("failed to parse FASTA bytes with dino-seq: {err}"))?;
952    if records.is_empty() {
953        return Err("FASTA input did not contain any records".to_owned());
954    }
955    Ok(records)
956}
957
958pub fn read_fasta_file(path: impl AsRef<Path>) -> Result<Vec<SequenceRecord>, String> {
959    let path = path.as_ref();
960    let mut reader = dino_seq::open_fasta_for_reference(path).map_err(|err| {
961        format!(
962            "failed to open FASTA with dino-seq {}: {err}",
963            path.display()
964        )
965    })?;
966    let mut records = Vec::new();
967    reader
968        .visit_records(|record| {
969            let mut bases = Vec::with_capacity(record.seq().len());
970            append_sequence_line(record.seq(), &mut bases).map_err(dino_seq_format_error)?;
971            records.push(SequenceRecord {
972                name: parse_record_name(record.name_without_gt()).map_err(dino_seq_format_error)?,
973                bases,
974            });
975            Ok(())
976        })
977        .map_err(|err| {
978            format!(
979                "failed to parse FASTA with dino-seq {}: {err}",
980                path.display()
981            )
982        })?;
983    if records.is_empty() {
984        return Err("FASTA input did not contain any records".to_owned());
985    }
986    Ok(records)
987}
988
989pub fn read_protein_fasta_file(path: impl AsRef<Path>) -> Result<Vec<SequenceRecord>, String> {
990    let path = path.as_ref();
991    let mut reader = dino_seq::open_fasta(path).map_err(|err| {
992        format!(
993            "failed to open protein FASTA with dino-seq {}: {err}",
994            path.display()
995        )
996    })?;
997    let mut records = Vec::new();
998    reader
999        .visit_records(|record| {
1000            let mut bases = Vec::with_capacity(record.seq().len());
1001            append_protein_line(record.seq(), &mut bases).map_err(dino_seq_format_error)?;
1002            records.push(SequenceRecord {
1003                name: parse_record_name(record.name_without_gt()).map_err(dino_seq_format_error)?,
1004                bases,
1005            });
1006            Ok(())
1007        })
1008        .map_err(|err| {
1009            format!(
1010                "failed to parse protein FASTA with dino-seq {}: {err}",
1011                path.display()
1012            )
1013        })?;
1014    if records.is_empty() {
1015        return Err("protein FASTA input did not contain any records".to_owned());
1016    }
1017    Ok(records)
1018}
1019
1020#[cfg(test)]
1021fn parse_fastq_bytes(input: &[u8]) -> Result<Vec<SequenceRecord>, String> {
1022    let mut records = Vec::new();
1023    dino_seq::visit_fastq_bytes(
1024        input,
1025        dino_seq::FastqConfig::default(),
1026        |record: dino_seq::FastqVisitRecord<'_>| {
1027            let header = record.name();
1028            let name = header.strip_prefix(b"@").unwrap_or(header);
1029            let mut bases = Vec::with_capacity(record.seq().len());
1030            append_sequence_line(record.seq(), &mut bases).map_err(dino_seq_format_error)?;
1031            records.push(SequenceRecord {
1032                name: parse_record_name(name).map_err(dino_seq_format_error)?,
1033                bases,
1034            });
1035            Ok(())
1036        },
1037    )
1038    .map_err(|err| format!("failed to parse FASTQ bytes with dino-seq: {err}"))?;
1039    if records.is_empty() {
1040        return Err("FASTQ input did not contain any records".to_owned());
1041    }
1042    Ok(records)
1043}
1044
1045pub fn visit_fastq_slices_file(
1046    path: impl AsRef<Path>,
1047    visitor: impl FnMut(SequenceRecordRef<'_>) -> Result<(), String>,
1048) -> Result<(), String> {
1049    visit_fastq_slices_file_limit(path, None, visitor)
1050}
1051
1052pub fn visit_fastq_slices_file_limit(
1053    path: impl AsRef<Path>,
1054    max_records: Option<usize>,
1055    visitor: impl FnMut(SequenceRecordRef<'_>) -> Result<(), String>,
1056) -> Result<(), String> {
1057    let path = path.as_ref();
1058    let mut reader = dino_seq::open_fastq(path).map_err(|err| {
1059        format!(
1060            "failed to open FASTQ with dino-seq {}: {err}",
1061            path.display()
1062        )
1063    })?;
1064    visit_fastq_slices_with_reader(&mut reader, max_records, visitor).map_err(|err| {
1065        format!(
1066            "failed to parse FASTQ with dino-seq {}: {err}",
1067            path.display()
1068        )
1069    })
1070}
1071
1072fn visit_fastq_slices_with_reader<R: Read>(
1073    reader: &mut dino_seq::FastqReader<R>,
1074    max_records: Option<usize>,
1075    mut visitor: impl FnMut(SequenceRecordRef<'_>) -> Result<(), String>,
1076) -> Result<(), String> {
1077    let chunk_bytes = std::env::var("DINO_QUANT_FASTQ_CHUNK_BYTES")
1078        .ok()
1079        .and_then(|value| value.parse::<u64>().ok())
1080        .filter(|&value| value > 0)
1081        .unwrap_or(4 * 1024 * 1024);
1082    let chunk_config = dino_seq::FastqChunkConfig::new(chunk_bytes).min_records(1);
1083    let mut records = 0_usize;
1084    while max_records.is_none_or(|limit| records < limit)
1085        && reader
1086            .next_chunk_with_sink(chunk_config, &mut |record: dino_seq::FastqVisitRecord<
1087                '_,
1088            >| {
1089                if max_records.is_some_and(|limit| records >= limit) {
1090                    return Ok(());
1091                }
1092                let header = record.name();
1093                let name = header.strip_prefix(b"@").unwrap_or(header);
1094                let name = parse_record_name_bytes(name).map_err(dino_seq_format_error)?;
1095                validate_sequence_bases(record.seq()).map_err(dino_seq_format_error)?;
1096                visitor(SequenceRecordRef {
1097                    name,
1098                    bases: record.seq(),
1099                })
1100                .map_err(dino_seq_format_error)?;
1101                records += 1;
1102                Ok(())
1103            })
1104            .map_err(|err| err.to_string())?
1105            .is_some()
1106    {}
1107    if records == 0 {
1108        return Err("FASTQ input did not contain any records".to_owned());
1109    }
1110    Ok(())
1111}
1112
1113pub fn concatenate_records(records: &[SequenceRecord]) -> Result<Vec<u8>, String> {
1114    if records.is_empty() {
1115        return Err("at least one sequence record is required".to_owned());
1116    }
1117    let total_bases = records
1118        .iter()
1119        .map(|record| record.bases.len())
1120        .sum::<usize>();
1121    let mut concatenated = Vec::with_capacity(total_bases + records.len().saturating_sub(1));
1122    for (idx, record) in records.iter().enumerate() {
1123        if idx > 0 {
1124            concatenated.push(CONTIG_SEPARATOR);
1125        }
1126        concatenated.extend_from_slice(&record.bases);
1127    }
1128    Ok(concatenated)
1129}
1130
1131pub fn reconstruction_metrics(
1132    original: &[f32],
1133    decoded: &[f32],
1134) -> Result<ReconstructionMetrics, String> {
1135    Ok(ReconstructionMetrics {
1136        mse: mse(original, decoded)?,
1137        cosine: cosine_similarity(original, decoded)?,
1138        dot_error: (dot(original, original)? - dot(original, decoded)?).abs(),
1139    })
1140}
1141
1142pub fn mutate_dna(seq: &[u8], every: usize) -> Vec<u8> {
1143    if every == 0 {
1144        return seq.to_vec();
1145    }
1146    let mut mutated = seq.to_vec();
1147    for idx in (every - 1..mutated.len()).step_by(every) {
1148        mutated[idx] = match mutated[idx].to_ascii_uppercase() {
1149            b'A' => b'C',
1150            b'C' => b'G',
1151            b'G' => b'T',
1152            b'T' => b'A',
1153            other => other,
1154        };
1155    }
1156    mutated
1157}
1158
1159pub fn synthetic_dna(len: usize, seed: u64) -> Vec<u8> {
1160    let mut rng = SplitMix64::new(seed);
1161    let mut seq = Vec::with_capacity(len);
1162    for _ in 0..len {
1163        let base = match rng.next_u64() & 3 {
1164            0 => b'A',
1165            1 => b'C',
1166            2 => b'G',
1167            _ => b'T',
1168        };
1169        seq.push(base);
1170    }
1171    seq
1172}
1173
1174pub fn intervals_overlap(
1175    left_start: usize,
1176    left_end: usize,
1177    right_start: usize,
1178    right_end: usize,
1179) -> bool {
1180    left_start < right_end && right_start < left_end
1181}
1182
1183pub fn dot(left: &[f32], right: &[f32]) -> Result<f32, String> {
1184    if left.len() != right.len() {
1185        return Err("vectors must have equal length".to_owned());
1186    }
1187    Ok(left.iter().zip(right).map(|(a, b)| a * b).sum())
1188}
1189
1190pub fn mse(left: &[f32], right: &[f32]) -> Result<f32, String> {
1191    if left.len() != right.len() {
1192        return Err("vectors must have equal length".to_owned());
1193    }
1194    if left.is_empty() {
1195        return Err("vectors must be non-empty".to_owned());
1196    }
1197    let sum: f32 = left
1198        .iter()
1199        .zip(right)
1200        .map(|(a, b)| {
1201            let delta = a - b;
1202            delta * delta
1203        })
1204        .sum();
1205    Ok(sum / left.len() as f32)
1206}
1207
1208pub fn cosine_similarity(left: &[f32], right: &[f32]) -> Result<f32, String> {
1209    let denom = l2_norm(left) * l2_norm(right);
1210    if denom <= f32::EPSILON {
1211        return Err("cosine similarity is undefined for zero vectors".to_owned());
1212    }
1213    Ok(dot(left, right)? / denom)
1214}
1215
1216pub fn l2_normalize(vector: &mut [f32]) {
1217    let norm = l2_norm(vector);
1218    if norm <= f32::EPSILON {
1219        return;
1220    }
1221    for value in vector {
1222        *value /= norm;
1223    }
1224}
1225
1226fn encode_scalar_codes(rotated: &[f32], bits: u8, clip: f32) -> Result<PackedCodes, String> {
1227    let levels = (1_u16 << bits) - 1;
1228    let inv_width = 1.0 / (2.0 * clip);
1229    let mut codes = Vec::with_capacity(rotated.len());
1230    for value in rotated {
1231        let clipped = value.clamp(-clip, clip);
1232        let normalized = (clipped + clip) * inv_width;
1233        codes.push((normalized * levels as f32).round() as u8);
1234    }
1235    PackedCodes::encode(&codes, bits)
1236}
1237
1238fn decode_scalar_codes(codes: &PackedCodes, bits: u8, clip: f32) -> Result<Vec<f32>, String> {
1239    if !(1..=8).contains(&bits) {
1240        return Err("quantizer bits must be in 1..=8".to_owned());
1241    }
1242    if !clip.is_finite() || clip <= 0.0 {
1243        return Err("clip must be finite and positive".to_owned());
1244    }
1245    if codes.bits != bits {
1246        return Err("packed code bit width does not match quantizer bit width".to_owned());
1247    }
1248    codes.decode_all(clip)
1249}
1250
1251fn subtract(left: &[f32], right: &[f32]) -> Result<Vec<f32>, String> {
1252    if left.len() != right.len() {
1253        return Err("vectors must have equal length".to_owned());
1254    }
1255    Ok(left.iter().zip(right).map(|(a, b)| a - b).collect())
1256}
1257
1258fn push_top_hit(top: &mut Vec<SearchHit>, top_k: usize, hit: SearchHit) {
1259    if top.len() < top_k {
1260        top.push(hit);
1261        return;
1262    }
1263
1264    let mut worst_idx = 0;
1265    let mut worst_score = top[0].score;
1266    for (idx, current) in top.iter().enumerate().skip(1) {
1267        if current.score < worst_score {
1268            worst_idx = idx;
1269            worst_score = current.score;
1270        }
1271    }
1272
1273    if hit.score > worst_score {
1274        top[worst_idx] = hit;
1275    }
1276}
1277
1278fn build_4bit_lookup(rotated_query: &[f32]) -> Vec<[f32; 256]> {
1279    let mut lookup = Vec::with_capacity(rotated_query.len().div_ceil(2));
1280    for chunk in rotated_query.chunks(2) {
1281        let first = chunk[0];
1282        let second = if chunk.len() == 2 { chunk[1] } else { 0.0 };
1283        let mut table = [0.0_f32; 256];
1284        for byte in 0_u16..=255 {
1285            let low = f32::from((byte & 0x0f) as u8);
1286            let high = f32::from((byte >> 4) as u8);
1287            table[usize::from(byte)] = low * first + high * second;
1288        }
1289        lookup.push(table);
1290    }
1291    lookup
1292}
1293
1294fn trim_ascii(mut bytes: &[u8]) -> &[u8] {
1295    while matches!(bytes.first(), Some(b' ' | b'\t' | b'\r' | b'\n')) {
1296        bytes = &bytes[1..];
1297    }
1298    while matches!(bytes.last(), Some(b' ' | b'\t' | b'\r' | b'\n')) {
1299        bytes = &bytes[..bytes.len() - 1];
1300    }
1301    bytes
1302}
1303
1304fn parse_record_name(header: &[u8]) -> Result<String, String> {
1305    let name = parse_record_name_bytes(header)?;
1306    String::from_utf8(name.to_vec()).map_err(|_| "sequence record name must be UTF-8".to_owned())
1307}
1308
1309fn parse_record_name_bytes(header: &[u8]) -> Result<&[u8], String> {
1310    let trimmed = trim_ascii(header);
1311    if trimmed.is_empty() {
1312        return Err("sequence record header must contain a name".to_owned());
1313    }
1314    trimmed
1315        .split(|byte| matches!(*byte, b' ' | b'\t'))
1316        .next()
1317        .filter(|name| !name.is_empty())
1318        .ok_or_else(|| "sequence record header must contain a name".to_owned())
1319}
1320
1321fn append_sequence_line(line: &[u8], dst: &mut Vec<u8>) -> Result<(), String> {
1322    for base in line {
1323        match base.to_ascii_uppercase() {
1324            b'A' | b'C' | b'G' | b'T' | b'N' => dst.push(base.to_ascii_uppercase()),
1325            b' ' | b'\t' | b'\r' => {}
1326            other => {
1327                return Err(format!(
1328                    "unsupported sequence character '{}' in input",
1329                    char::from(other)
1330                ));
1331            }
1332        }
1333    }
1334    Ok(())
1335}
1336
1337fn append_protein_line(line: &[u8], dst: &mut Vec<u8>) -> Result<(), String> {
1338    for residue in line {
1339        let upper = residue.to_ascii_uppercase();
1340        match upper {
1341            b'A'..=b'Z' | b'*' | b'-' => dst.push(upper),
1342            b' ' | b'\t' | b'\r' => {}
1343            other => {
1344                return Err(format!(
1345                    "unsupported protein sequence character '{}' in input",
1346                    char::from(other)
1347                ));
1348            }
1349        }
1350    }
1351    Ok(())
1352}
1353
1354fn validate_sequence_bases(line: &[u8]) -> Result<(), String> {
1355    for base in line {
1356        match base.to_ascii_uppercase() {
1357            b'A' | b'C' | b'G' | b'T' | b'N' => {}
1358            other => {
1359                return Err(format!(
1360                    "unsupported sequence character '{}' in input",
1361                    char::from(other)
1362                ));
1363            }
1364        }
1365    }
1366    Ok(())
1367}
1368
1369fn protein_kmer_code(window: &[u8]) -> Option<u64> {
1370    let mut code = 0_u64;
1371    for residue in window {
1372        code = code.checked_mul(23)?;
1373        code = code.checked_add(u64::from(protein_code(*residue)?))?;
1374    }
1375    Some(code)
1376}
1377
1378fn protein_code(residue: u8) -> Option<u8> {
1379    match residue.to_ascii_uppercase() {
1380        b'A' => Some(0),
1381        b'C' => Some(1),
1382        b'D' => Some(2),
1383        b'E' => Some(3),
1384        b'F' => Some(4),
1385        b'G' => Some(5),
1386        b'H' => Some(6),
1387        b'I' => Some(7),
1388        b'K' => Some(8),
1389        b'L' => Some(9),
1390        b'M' => Some(10),
1391        b'N' => Some(11),
1392        b'P' => Some(12),
1393        b'Q' => Some(13),
1394        b'R' => Some(14),
1395        b'S' => Some(15),
1396        b'T' => Some(16),
1397        b'V' => Some(17),
1398        b'W' => Some(18),
1399        b'Y' => Some(19),
1400        b'B' | b'J' | b'O' | b'U' | b'X' | b'Z' | b'*' | b'-' => None,
1401        _ => None,
1402    }
1403}
1404
1405fn dino_seq_format_error(message: String) -> dino_seq::FastqError {
1406    dino_seq::FastqError::Format(message)
1407}
1408
1409fn l2_norm(vector: &[f32]) -> f32 {
1410    vector.iter().map(|value| value * value).sum::<f32>().sqrt()
1411}
1412
1413fn validate_hadamard_dim(dim: usize) -> Result<(), String> {
1414    if dim == 0 {
1415        return Err("Hadamard dimension must be non-zero".to_owned());
1416    }
1417    if !dim.is_power_of_two() {
1418        return Err("Hadamard dimension must be a power of two".to_owned());
1419    }
1420    Ok(())
1421}
1422
1423fn hadamard_in_place(values: &mut [f32]) {
1424    let mut stride = 1;
1425    while stride < values.len() {
1426        let step = stride * 2;
1427        for start in (0..values.len()).step_by(step) {
1428            for idx in start..start + stride {
1429                let left = values[idx];
1430                let right = values[idx + stride];
1431                values[idx] = left + right;
1432                values[idx + stride] = left - right;
1433            }
1434        }
1435        stride = step;
1436    }
1437}
1438
1439fn coordinate_sign(seed: u64, idx: usize) -> f32 {
1440    if mix_u64(seed ^ idx as u64) & 1 == 0 {
1441        1.0
1442    } else {
1443        -1.0
1444    }
1445}
1446
1447fn canonical_kmer_code(window: &[u8]) -> Option<u64> {
1448    let mut forward = 0_u64;
1449    let mut reverse = 0_u64;
1450    for (idx, base) in window.iter().enumerate() {
1451        let code = base_code(*base)?;
1452        forward = (forward << 2) | u64::from(code);
1453        let rc = u64::from(BASE_T - code);
1454        reverse |= rc << (idx * 2);
1455    }
1456    Some(forward.min(reverse))
1457}
1458
1459fn base_code(base: u8) -> Option<u8> {
1460    match base.to_ascii_uppercase() {
1461        b'A' => Some(BASE_A),
1462        b'C' => Some(BASE_C),
1463        b'G' => Some(BASE_G),
1464        b'T' => Some(BASE_T),
1465        _ => None,
1466    }
1467}
1468
1469fn mix_u64(mut value: u64) -> u64 {
1470    value = (value ^ (value >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
1471    value = (value ^ (value >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
1472    value ^ (value >> 31)
1473}
1474
1475#[cfg(test)]
1476mod tests {
1477    use super::*;
1478
1479    #[test]
1480    fn rotation_round_trips() -> Result<(), String> {
1481        let mut vector = (0..128)
1482            .map(|idx| ((idx as f32 + 1.0) * 0.17).sin())
1483            .collect::<Vec<_>>();
1484        l2_normalize(&mut vector);
1485
1486        let rotated = rotate(&vector, 17)?;
1487        let decoded = inverse_rotate(&rotated, 17)?;
1488        let err = mse(&vector, &decoded)?;
1489        assert!(err < 1.0e-12, "round-trip MSE was {err}");
1490        Ok(())
1491    }
1492
1493    #[test]
1494    fn more_scalar_bits_reduce_error() -> Result<(), String> {
1495        let mut vector = (0..256)
1496            .map(|idx| ((idx as f32 + 3.0) * 0.11).cos())
1497            .collect::<Vec<_>>();
1498        l2_normalize(&mut vector);
1499
1500        let low = QuantizerConfig {
1501            bits: 2,
1502            use_qjl_residual: false,
1503            ..QuantizerConfig::default()
1504        }
1505        .encode(&vector)?
1506        .decode()?;
1507        let high = QuantizerConfig {
1508            bits: 5,
1509            use_qjl_residual: false,
1510            ..QuantizerConfig::default()
1511        }
1512        .encode(&vector)?
1513        .decode()?;
1514
1515        assert!(mse(&vector, &high)? < mse(&vector, &low)?);
1516        Ok(())
1517    }
1518
1519    #[test]
1520    fn packed_codes_round_trip_values() -> Result<(), String> {
1521        let values = (0..97).map(|idx| (idx % 7) as u8).collect::<Vec<_>>();
1522        let packed = PackedCodes::encode(&values, 3)?;
1523
1524        assert!(packed.byte_len() < values.len());
1525        for (idx, expected) in values.iter().enumerate() {
1526            assert_eq!(packed.get(idx)?, *expected);
1527        }
1528        Ok(())
1529    }
1530
1531    #[test]
1532    fn dna_sketches_are_normalized() -> Result<(), String> {
1533        let seq = synthetic_dna(512, 9);
1534        let sketch = dna_kmer_sketch(&seq, 15, 128)?;
1535        let norm = l2_norm(&sketch);
1536        assert!((norm - 1.0).abs() < 1.0e-5, "norm was {norm}");
1537        Ok(())
1538    }
1539
1540    #[test]
1541    fn protein_sketches_are_normalized_and_skip_ambiguous_residues() -> Result<(), String> {
1542        let sketch = protein_kmer_sketch(b"MKRISTXTTITTTITITTGNGAG", 3, 128)?;
1543        let norm = l2_norm(&sketch);
1544        assert!((norm - 1.0).abs() < 1.0e-5, "norm was {norm}");
1545        Ok(())
1546    }
1547
1548    #[test]
1549    fn parses_fasta_records_and_concatenates_with_separator() -> Result<(), String> {
1550        let records = parse_fasta_bytes(b">chr1 description\nacgt\nNN\n>chr2\nTTA\n")?;
1551        assert_eq!(records.len(), 2);
1552        assert_eq!(records[0].name, "chr1");
1553        assert_eq!(records[0].bases, b"ACGTNN");
1554        assert_eq!(records[1].name, "chr2");
1555        assert_eq!(concatenate_records(&records)?, b"ACGTNNNTTA");
1556        Ok(())
1557    }
1558
1559    #[test]
1560    fn parses_fastq_records() -> Result<(), String> {
1561        let records = parse_fastq_bytes(b"@read1 comment\nacgtn\n+\nIIIII\n@read2\nTTA\n+\n###\n")?;
1562        assert_eq!(records.len(), 2);
1563        assert_eq!(records[0].name, "read1");
1564        assert_eq!(records[0].bases, b"ACGTN");
1565        assert_eq!(records[1].name, "read2");
1566        assert_eq!(records[1].bases, b"TTA");
1567        Ok(())
1568    }
1569
1570    #[test]
1571    fn fastq_slice_visitor_can_stop_after_limit() -> Result<(), String> {
1572        let input = b"@read1\nACGT\n+\nIIII\n@read2\nTTAA\n+\n####\n";
1573        let mut reader = dino_seq::FastqReader::new(input.as_slice());
1574        let mut names = Vec::new();
1575        visit_fastq_slices_with_reader(&mut reader, Some(1), |record| {
1576            names.push(String::from_utf8(record.name.to_vec()).map_err(|err| err.to_string())?);
1577            Ok(())
1578        })?;
1579        assert_eq!(names, vec!["read1"]);
1580        Ok(())
1581    }
1582
1583    #[test]
1584    fn parses_protein_fasta_records() -> Result<(), String> {
1585        let path =
1586            std::env::temp_dir().join(format!("dino_quant_protein_{}.faa", std::process::id()));
1587        std::fs::write(&path, b">p1 protein\nmkristX*\n").map_err(|err| err.to_string())?;
1588        let records = read_protein_fasta_file(&path)?;
1589        std::fs::remove_file(path).map_err(|err| err.to_string())?;
1590        assert_eq!(records.len(), 1);
1591        assert_eq!(records[0].name, "p1");
1592        assert_eq!(records[0].bases, b"MKRISTX*");
1593        Ok(())
1594    }
1595
1596    #[test]
1597    fn rejects_invalid_sequence_character() {
1598        let err = parse_fasta_bytes(b">chr1\nACGTX\n").expect_err("invalid base should fail");
1599        assert!(err.contains("unsupported sequence character"));
1600    }
1601
1602    #[test]
1603    fn qjl_residual_decodes_finite_vector() -> Result<(), String> {
1604        let seq = synthetic_dna(2048, 42);
1605        let vector = dna_kmer_sketch(&seq, 17, 256)?;
1606        let quantized = QuantizerConfig {
1607            bits: 3,
1608            use_qjl_residual: true,
1609            ..QuantizerConfig::default()
1610        }
1611        .encode(&vector)?;
1612        let decoded = quantized.decode()?;
1613
1614        assert_eq!(decoded.len(), vector.len());
1615        assert!(decoded.iter().all(|value| value.is_finite()));
1616        assert!(quantized.compressed_bits() < vector.len() * 32);
1617        Ok(())
1618    }
1619
1620    #[test]
1621    fn approximate_dot_matches_decoded_qjl_dot() -> Result<(), String> {
1622        let reference = dna_kmer_sketch(&synthetic_dna(2048, 42), 17, 256)?;
1623        let query = dna_kmer_sketch(&mutate_dna(&synthetic_dna(2048, 77), 19), 17, 256)?;
1624        let quantized = QuantizerConfig {
1625            bits: 4,
1626            use_qjl_residual: true,
1627            ..QuantizerConfig::default()
1628        }
1629        .encode(&reference)?;
1630
1631        let decoded = quantized.decode()?;
1632        let decoded_dot = dot(&decoded, &query)?;
1633        let approximate_dot = quantized.approximate_dot_query(&query)?;
1634
1635        assert!(
1636            (decoded_dot - approximate_dot).abs() < 1.0e-5,
1637            "decoded_dot={decoded_dot} approximate_dot={approximate_dot}"
1638        );
1639        Ok(())
1640    }
1641
1642    #[test]
1643    fn quantized_vector_snapshot_round_trips_scoring() -> Result<(), String> {
1644        let reference = dna_kmer_sketch(&synthetic_dna(512, 19), 11, 128)?;
1645        let query = dna_kmer_sketch(&mutate_dna(&synthetic_dna(512, 19), 7), 11, 128)?;
1646        let quantized = QuantizerConfig {
1647            bits: 4,
1648            use_qjl_residual: true,
1649            ..QuantizerConfig::default()
1650        }
1651        .encode(&reference)?;
1652        let restored = QuantizedVector::from_snapshot(quantized.snapshot())?;
1653
1654        let original_score = quantized.approximate_dot_query(&query)?;
1655        let restored_score = restored.approximate_dot_query(&query)?;
1656        assert_eq!(original_score.to_bits(), restored_score.to_bits());
1657        Ok(())
1658    }
1659
1660    #[test]
1661    fn prepared_query_matches_direct_scoring() -> Result<(), String> {
1662        let reference = dna_kmer_sketch(&synthetic_dna(512, 23), 11, 128)?;
1663        let query = dna_kmer_sketch(&mutate_dna(&synthetic_dna(512, 23), 5), 11, 128)?;
1664        let config = QuantizerConfig {
1665            bits: 4,
1666            use_qjl_residual: true,
1667            ..QuantizerConfig::default()
1668        };
1669        let quantized = config.encode(&reference)?;
1670        let prepared = QuantizedVector::prepare_approximate_query(
1671            &query,
1672            config.rotation_seed,
1673            Some(config.qjl_seed),
1674        )?;
1675
1676        let direct = quantized.approximate_dot_query(&query)?;
1677        let reused = quantized.approximate_dot_prepared_query(&prepared)?;
1678        assert!(
1679            (direct - reused).abs() < 1.0e-5,
1680            "direct={direct} reused={reused}"
1681        );
1682        Ok(())
1683    }
1684
1685    #[test]
1686    fn compressed_index_finds_mutated_source_window() -> Result<(), String> {
1687        let reference = synthetic_dna(8192, 0xabc);
1688        let config = ReferenceIndexConfig {
1689            k: 15,
1690            dim: 256,
1691            window_len: 512,
1692            stride: 128,
1693            quantizer: QuantizerConfig {
1694                bits: 4,
1695                use_qjl_residual: false,
1696                ..QuantizerConfig::default()
1697            },
1698        };
1699        let index = ReferenceWindowIndex::build(&reference, config)?;
1700        let source_start = 2816;
1701        let source_end = source_start + config.window_len;
1702        let query = mutate_dna(&reference[source_start..source_end], 41);
1703        let hits = index.search_sequence(&query, 5)?;
1704
1705        assert!(
1706            hits.iter()
1707                .any(|hit| { intervals_overlap(hit.start, hit.end, source_start, source_end) })
1708        );
1709        assert!(index.compression_ratio() > 6.0);
1710        Ok(())
1711    }
1712
1713    #[test]
1714    fn record_index_reports_target_coordinates() -> Result<(), String> {
1715        let records = vec![
1716            SequenceRecord {
1717                name: "chr1".to_owned(),
1718                bases: synthetic_dna(1024, 1),
1719            },
1720            SequenceRecord {
1721                name: "chr2".to_owned(),
1722                bases: synthetic_dna(1024, 2),
1723            },
1724        ];
1725        let config = ReferenceIndexConfig {
1726            k: 11,
1727            dim: 128,
1728            window_len: 256,
1729            stride: 128,
1730            quantizer: QuantizerConfig {
1731                bits: 4,
1732                use_qjl_residual: false,
1733                ..QuantizerConfig::default()
1734            },
1735        };
1736        let index = ReferenceWindowIndex::build_records(&records, config)?;
1737        let query = records[1].bases[256..512].to_vec();
1738        let hits = index.search_sequence(&query, 3)?;
1739        assert!(hits.iter().any(|hit| {
1740            hit.target_name == "chr2"
1741                && intervals_overlap(hit.target_start, hit.target_end, 256, 512)
1742        }));
1743        Ok(())
1744    }
1745
1746    #[test]
1747    fn canonical_kmers_match_reverse_complements() -> Result<(), String> {
1748        let left = canonical_kmer_code(b"ACGTT").ok_or("left k-mer should be valid")?;
1749        let right = canonical_kmer_code(b"AACGT").ok_or("right k-mer should be valid")?;
1750        assert_eq!(left, right);
1751        Ok(())
1752    }
1753}