use std::fmt;
use serde::{Deserialize, Serialize};
use crate::superfile::vector::distance::{Metric, SQ8_RESIDUAL_DIVISOR};
const LOW_DIM_RERANK_FLOOR_THRESHOLD: usize = 384;
const FP32_LOW_DIM_RERANK_FLOOR: usize = 20;
const FP32_HIGH_DIM_RERANK_FLOOR: usize = 50;
const SQ8_LOW_DIM_RERANK_FLOOR: usize = 50;
const SQ8_HIGH_DIM_RERANK_FLOOR: usize = 100;
pub(crate) const SQ8_FIXED_OFFSET: f32 = -1.0;
pub(crate) const SQ8_FIXED_SCALE: f32 = 2.0 / 255.0;
pub(crate) const SQ8_FIXED_RESIDUAL_DIVISOR: f32 = 256.0;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RerankCodec {
Fp32,
Sq8Residual,
Sq8FixedResidual,
RabitqOnly,
}
impl Default for RerankCodec {
fn default() -> Self {
Self::Sq8FixedResidual
}
}
impl RerankCodec {
#[inline]
pub const fn codec_id(self) -> u8 {
match self {
Self::Fp32 => 0,
Self::Sq8Residual => 1,
Self::RabitqOnly => 2,
Self::Sq8FixedResidual => 3,
}
}
#[inline]
pub const fn from_codec_id(id: u8) -> Option<Self> {
match id {
0 => Some(Self::Fp32),
1 => Some(Self::Sq8Residual),
2 => Some(Self::RabitqOnly),
3 => Some(Self::Sq8FixedResidual),
_ => None,
}
}
#[inline]
pub const fn name(self) -> &'static str {
match self {
Self::Fp32 => "fp32",
Self::Sq8Residual => "sq8_residual",
Self::RabitqOnly => "rabitq_only",
Self::Sq8FixedResidual => "sq8_fixed_residual",
}
}
#[inline]
pub const fn per_vector_bytes(self, dim: usize) -> usize {
match self {
Self::Fp32 => dim * 4,
Self::Sq8Residual | Self::Sq8FixedResidual => dim * 2,
Self::RabitqOnly => 0,
}
}
#[inline]
pub const fn writes_full(self) -> bool {
!matches!(self, Self::RabitqOnly)
}
#[inline]
pub const fn is_implemented(self) -> bool {
matches!(
self,
Self::Fp32 | Self::Sq8Residual | Self::Sq8FixedResidual | Self::RabitqOnly
)
}
#[inline]
pub const fn is_sq8_residual_family(self) -> bool {
matches!(self, Self::Sq8Residual | Self::Sq8FixedResidual)
}
#[inline]
pub const fn residual_divisor(self) -> Option<f32> {
match self {
Self::Sq8Residual => Some(SQ8_RESIDUAL_DIVISOR),
Self::Sq8FixedResidual => Some(SQ8_FIXED_RESIDUAL_DIVISOR),
Self::Fp32 | Self::RabitqOnly => None,
}
}
#[inline]
pub const fn uses_fixed_quantizer(self) -> bool {
matches!(self, Self::Sq8FixedResidual)
}
#[inline]
pub const fn supports_metric(self, metric: Metric) -> bool {
!matches!(self, Self::Sq8FixedResidual) || matches!(metric, Metric::Cosine)
}
#[inline]
pub const fn recommended_rerank_mult_floor(self, dim: usize) -> Option<usize> {
let high_dim = dim > LOW_DIM_RERANK_FLOOR_THRESHOLD;
match self {
Self::Fp32 => Some(if high_dim {
FP32_HIGH_DIM_RERANK_FLOOR
} else {
FP32_LOW_DIM_RERANK_FLOOR
}),
Self::Sq8Residual => Some(if high_dim {
SQ8_HIGH_DIM_RERANK_FLOOR
} else {
SQ8_LOW_DIM_RERANK_FLOOR
}),
Self::Sq8FixedResidual => Some(if high_dim {
SQ8_HIGH_DIM_RERANK_FLOOR
} else {
SQ8_LOW_DIM_RERANK_FLOOR
}),
Self::RabitqOnly => None,
}
}
#[inline]
pub const fn codec_meta_bytes(
self,
dim: usize,
n_docs: usize,
n_cent: usize,
metric: Metric,
) -> usize {
match self {
Self::Fp32 | Self::RabitqOnly => 0,
Self::Sq8Residual | Self::Sq8FixedResidual => {
let scale_offset_bytes = 2 * n_cent * dim * 4;
let norms_bytes = match metric {
Metric::L2Sq | Metric::Cosine => n_docs * 4,
Metric::NegDot => 0,
};
scale_offset_bytes + norms_bytes
}
}
}
}
impl fmt::Display for RerankCodec {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.name())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_is_sq8_fixed_residual() {
assert_eq!(RerankCodec::default(), RerankCodec::Sq8FixedResidual);
}
#[test]
fn fp32_codec_id_is_zero() {
assert_eq!(RerankCodec::Fp32.codec_id(), 0u8);
}
#[test]
fn codec_id_roundtrips_every_variant() {
for c in [
RerankCodec::Fp32,
RerankCodec::Sq8Residual,
RerankCodec::Sq8FixedResidual,
RerankCodec::RabitqOnly,
] {
assert_eq!(
RerankCodec::from_codec_id(c.codec_id()),
Some(c),
"round-trip mismatch for {c:?}"
);
}
}
#[test]
fn unknown_codec_id_is_none() {
for id in [4u8, 5, 16, 200, 255] {
assert_eq!(
RerankCodec::from_codec_id(id),
None,
"unknown id {id} must not map to a codec"
);
}
}
#[test]
fn per_vector_bytes_matches_spec() {
assert_eq!(RerankCodec::Fp32.per_vector_bytes(384), 1536);
assert_eq!(RerankCodec::Sq8Residual.per_vector_bytes(384), 768);
assert_eq!(RerankCodec::Sq8FixedResidual.per_vector_bytes(384), 768);
assert_eq!(RerankCodec::RabitqOnly.per_vector_bytes(384), 0);
}
#[test]
fn writes_full_matches_per_vector_bytes() {
for c in [
RerankCodec::Fp32,
RerankCodec::Sq8Residual,
RerankCodec::Sq8FixedResidual,
RerankCodec::RabitqOnly,
] {
assert_eq!(
c.writes_full(),
c.per_vector_bytes(384) > 0,
"writes_full disagrees with per_vector_bytes for {c:?}"
);
}
}
#[test]
fn all_codecs_implemented() {
assert!(RerankCodec::Fp32.is_implemented());
assert!(RerankCodec::Sq8Residual.is_implemented());
assert!(RerankCodec::Sq8FixedResidual.is_implemented());
assert!(RerankCodec::RabitqOnly.is_implemented());
}
#[test]
fn recommended_rerank_mult_floor_matches_calibration_table() {
assert_eq!(
RerankCodec::Fp32.recommended_rerank_mult_floor(384),
Some(20)
);
assert_eq!(
RerankCodec::Sq8Residual.recommended_rerank_mult_floor(384),
Some(50)
);
assert_eq!(
RerankCodec::Sq8FixedResidual.recommended_rerank_mult_floor(384),
Some(50)
);
assert_eq!(
RerankCodec::RabitqOnly.recommended_rerank_mult_floor(384),
None
);
assert_eq!(
RerankCodec::Fp32.recommended_rerank_mult_floor(1024),
Some(50)
);
assert_eq!(
RerankCodec::Sq8Residual.recommended_rerank_mult_floor(1024),
Some(100)
);
assert_eq!(
RerankCodec::Sq8FixedResidual.recommended_rerank_mult_floor(1024),
Some(100)
);
assert_eq!(
RerankCodec::RabitqOnly.recommended_rerank_mult_floor(1024),
None
);
assert_eq!(
RerankCodec::Sq8Residual.recommended_rerank_mult_floor(385),
Some(100)
);
}
#[test]
fn display_renders_stable_name() {
assert_eq!(RerankCodec::Fp32.to_string(), "fp32");
assert_eq!(RerankCodec::Sq8Residual.to_string(), "sq8_residual");
assert_eq!(
RerankCodec::Sq8FixedResidual.to_string(),
"sq8_fixed_residual"
);
assert_eq!(RerankCodec::RabitqOnly.to_string(), "rabitq_only");
for c in [
RerankCodec::Fp32,
RerankCodec::Sq8Residual,
RerankCodec::Sq8FixedResidual,
RerankCodec::RabitqOnly,
] {
assert_eq!(c.to_string(), c.name());
}
}
#[test]
fn codec_meta_bytes_matches_layout_spec() {
for c in [RerankCodec::Fp32, RerankCodec::RabitqOnly] {
for m in [Metric::L2Sq, Metric::Cosine, Metric::NegDot] {
assert_eq!(
c.codec_meta_bytes(384, 1_000_000, 1024, m),
0,
"{c:?} / {m:?}"
);
}
}
let so_bytes = 2 * 1024 * 384 * 4;
assert_eq!(
RerankCodec::Sq8Residual.codec_meta_bytes(384, 1_000_000, 1024, Metric::NegDot),
so_bytes
);
assert_eq!(
RerankCodec::Sq8Residual.codec_meta_bytes(384, 1_000_000, 1024, Metric::Cosine),
so_bytes + 1_000_000 * 4
);
assert_eq!(
RerankCodec::Sq8FixedResidual.codec_meta_bytes(384, 1_000_000, 1024, Metric::Cosine),
so_bytes + 1_000_000 * 4
);
assert_eq!(
RerankCodec::Sq8Residual.codec_meta_bytes(384, 1_000_000, 1024, Metric::L2Sq),
so_bytes + 1_000_000 * 4
);
assert_eq!(
RerankCodec::Sq8Residual.codec_meta_bytes(384, 1_000_000, 1024, Metric::NegDot),
so_bytes
);
}
#[test]
fn fixed_residual_contract_is_cosine_only() {
assert!(RerankCodec::Sq8FixedResidual.supports_metric(Metric::Cosine));
assert!(!RerankCodec::Sq8FixedResidual.supports_metric(Metric::L2Sq));
assert!(!RerankCodec::Sq8FixedResidual.supports_metric(Metric::NegDot));
assert_eq!(
RerankCodec::Sq8FixedResidual.residual_divisor(),
Some(SQ8_FIXED_RESIDUAL_DIVISOR)
);
assert!(RerankCodec::Sq8FixedResidual.uses_fixed_quantizer());
assert!(RerankCodec::Sq8FixedResidual.is_sq8_residual_family());
}
}