Skip to main content

ferrox_models/
kimi_validate.rs

1//! Validate a Kimi K3 checkpoint directory without loading all shards.
2//!
3//! Reads `model.safetensors.index.json` (or a partial weight_map) and
4//! checks that expected layer prefixes and tensor shapes are present.
5//! Used by `ferrox inspect-kimi` / CI gates before a full ~1.56 TB run.
6
7use std::collections::HashMap;
8use std::fs;
9use std::path::Path;
10
11use serde::Deserialize;
12use thiserror::Error;
13
14#[derive(Debug, Error)]
15pub enum KimiValidateError {
16    #[error("io: {0}")]
17    Io(#[from] std::io::Error),
18    #[error("json: {0}")]
19    Json(#[from] serde_json::Error),
20    #[error("{0}")]
21    Message(String),
22}
23
24#[derive(Debug, Deserialize)]
25struct IndexFile {
26    weight_map: HashMap<String, String>,
27}
28
29#[derive(Debug, Clone)]
30pub struct KimiCheckpointReport {
31    pub n_tensors: usize,
32    pub n_shards_referenced: usize,
33    pub layers_seen: Vec<usize>,
34    pub has_embed: bool,
35    pub has_lm_head: bool,
36    pub missing_required: Vec<String>,
37}
38
39/// Lightweight validation: index present, embed/lm_head named, at least
40/// one layer prefix, shard files referenced exist when `check_files`.
41pub fn validate_kimi_checkpoint_dir(
42    dir: &Path,
43    check_files: bool,
44) -> Result<KimiCheckpointReport, KimiValidateError> {
45    let index_path = dir.join("model.safetensors.index.json");
46    if !index_path.is_file() {
47        return Err(KimiValidateError::Message(format!(
48            "missing {} — not a Kimi safetensors directory",
49            index_path.display()
50        )));
51    }
52    let index: IndexFile = serde_json::from_str(&fs::read_to_string(&index_path)?)?;
53    let n_tensors = index.weight_map.len();
54    let mut shards: HashMap<String, ()> = HashMap::new();
55    let mut layers = std::collections::BTreeSet::new();
56    let mut has_embed = false;
57    let mut has_lm_head = false;
58    for (name, shard) in &index.weight_map {
59        shards.insert(shard.clone(), ());
60        if name.contains("embed_tokens") || name.ends_with("tok_embeddings.weight") {
61            has_embed = true;
62        }
63        if name.contains("lm_head") || name.contains("output.weight") {
64            has_lm_head = true;
65        }
66        if let Some(rest) = name.strip_prefix("language_model.model.layers.") {
67            if let Some(num) = rest.split('.').next().and_then(|s| s.parse::<usize>().ok()) {
68                layers.insert(num);
69            }
70        }
71    }
72    let mut missing_required = Vec::new();
73    if !has_embed {
74        missing_required.push("embed_tokens / tok_embeddings".into());
75    }
76    if layers.is_empty() {
77        missing_required.push("language_model.model.layers.*".into());
78    }
79    if check_files {
80        for shard in shards.keys() {
81            let p = dir.join(shard);
82            if !p.is_file() {
83                missing_required.push(format!("shard file missing: {shard}"));
84            }
85        }
86    }
87    Ok(KimiCheckpointReport {
88        n_tensors,
89        n_shards_referenced: shards.len(),
90        layers_seen: layers.into_iter().collect(),
91        has_embed,
92        has_lm_head,
93        missing_required,
94    })
95}
96
97#[cfg(test)]
98mod tests {
99    use super::*;
100
101    #[test]
102    fn validates_minimal_index_json() {
103        let dir = std::env::temp_dir().join(format!("ferrox_kimi_validate_{}", std::process::id()));
104        let _ = fs::remove_dir_all(&dir);
105        fs::create_dir_all(&dir).unwrap();
106        let index = r#"{
107          "weight_map": {
108            "language_model.model.embed_tokens.weight": "model-00001-of-000096.safetensors",
109            "language_model.model.layers.0.input_layernorm.weight": "model-00001-of-000096.safetensors",
110            "language_model.lm_head.weight": "model-00001-of-000096.safetensors"
111          }
112        }"#;
113        fs::write(dir.join("model.safetensors.index.json"), index).unwrap();
114        let report = validate_kimi_checkpoint_dir(&dir, false).unwrap();
115        assert_eq!(report.n_tensors, 3);
116        assert_eq!(report.layers_seen, vec![0]);
117        assert!(report.has_embed);
118        assert!(report.has_lm_head);
119        assert!(report.missing_required.is_empty());
120        let _ = fs::remove_dir_all(&dir);
121    }
122}