ferrox_models/
kimi_validate.rs1use 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
39pub 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}