use std::collections::HashMap;
use std::path::Path;
use serde::Deserialize;
use crate::error::YomitokiError;
#[derive(Deserialize)]
struct FragmentRecord {
radius: u32,
fragment_hash: u64,
occurrence_count: u64,
}
#[derive(Deserialize)]
struct FrequencyTableFile {
total_molecules_processed: u64,
fragments: Vec<FragmentRecord>,
}
#[derive(Deserialize)]
struct CorpusDomainFile {
source_name: String,
domain: String,
synthesis_focused: bool,
description: String,
}
#[derive(Deserialize)]
struct ManifestFile {
artifact_sha256: String,
reference_distribution: Vec<f64>,
fragment_definition_version: String,
reference_distribution_version: String,
corpus_domain: CorpusDomainFile,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct CorpusDomain {
pub(crate) source_name: String,
pub(crate) domain: String,
pub(crate) synthesis_focused: bool,
pub(crate) description: String,
}
#[derive(Debug, Clone, PartialEq)]
pub struct FragmentCorpus {
pub(crate) radius: u32,
pub(crate) total_molecules_processed: u64,
pub(crate) frequency: HashMap<u64, u64>,
reference_distribution: Vec<f64>,
version: String,
pub(crate) domain: CorpusDomain,
pub(crate) fragment_definition_version: String,
pub(crate) reference_distribution_version: String,
}
impl FragmentCorpus {
pub fn load_dir(dir: impl AsRef<Path>) -> Result<FragmentCorpus, YomitokiError> {
let dir = dir.as_ref();
let table: FrequencyTableFile = read_json(&dir.join("fragment_frequencies.json"))?;
let manifest: ManifestFile = read_json(&dir.join("manifest.json"))?;
let mut radius = None;
let mut frequency = HashMap::with_capacity(table.fragments.len());
for record in table.fragments {
match radius {
None => radius = Some(record.radius),
Some(r) if r != record.radius => {
return Err(YomitokiError::ModelLoadError(format!(
"corpus references multiple radii ({r} and {}) — rebuild with a \
single --radius value",
record.radius
)));
}
_ => {}
}
frequency.insert(record.fragment_hash, record.occurrence_count);
}
let radius = radius
.ok_or_else(|| YomitokiError::ModelLoadError("corpus has no fragments".to_string()))?;
if manifest.reference_distribution.len() < 2 {
return Err(YomitokiError::ModelLoadError(
"corpus manifest has no reference_distribution — rebuild with \
tools/build-fragment-corpus (round 17 or later)"
.to_string(),
));
}
Ok(FragmentCorpus {
radius,
total_molecules_processed: table.total_molecules_processed,
frequency,
reference_distribution: manifest.reference_distribution,
version: manifest.artifact_sha256,
domain: CorpusDomain {
source_name: manifest.corpus_domain.source_name,
domain: manifest.corpus_domain.domain,
synthesis_focused: manifest.corpus_domain.synthesis_focused,
description: manifest.corpus_domain.description,
},
fragment_definition_version: manifest.fragment_definition_version,
reference_distribution_version: manifest.reference_distribution_version,
})
}
pub fn version(&self) -> &str {
&self.version
}
pub(crate) fn percentile_rank(&self, mean_document_frequency: f64) -> f64 {
let grid = &self.reference_distribution;
match grid.binary_search_by(|v| v.partial_cmp(&mean_document_frequency).unwrap()) {
Ok(i) => i as f64 / (grid.len() - 1) as f64,
Err(0) => 0.0,
Err(i) if i >= grid.len() => 1.0,
Err(i) => {
let (lo, hi) = (grid[i - 1], grid[i]);
let frac = if hi > lo {
(mean_document_frequency - lo) / (hi - lo)
} else {
0.0
};
((i - 1) as f64 + frac) / (grid.len() - 1) as f64
}
}
.clamp(0.0, 1.0)
}
}
fn read_json<T: serde::de::DeserializeOwned>(path: &Path) -> Result<T, YomitokiError> {
let bytes = std::fs::read(path)
.map_err(|e| YomitokiError::ModelLoadError(format!("could not read {path:?}: {e}")))?;
serde_json::from_slice(&bytes)
.map_err(|e| YomitokiError::ModelLoadError(format!("could not parse {path:?}: {e}")))
}