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}