Skip to main content

zsfm_tabdpt/
tensor_map.rs

1/// Map TabDPT safetensors tensor names to GGUF naming convention. The per-layer `kappa`/
2/// `max_len_f`/`n0` registered buffers are skipped — they're constants derived from
3/// `base_len`/`max_len` (identical across every layer), recomputed from config instead of
4/// round-tripped through GGUF.
5pub fn map_tensor_name(name: &str) -> Option<String> {
6    match name {
7        "encoder.weight" => return Some("encoder.weight".into()),
8        "encoder.bias" => return Some("encoder.bias".into()),
9        "head.0.weight" => return Some("head_fc1.weight".into()),
10        "head.0.bias" => return Some("head_fc1.bias".into()),
11        "head.2.weight" => return Some("head_fc2.weight".into()),
12        "head.2.bias" => return Some("head_fc2.bias".into()),
13        "thinking_embed" => return Some("thinking_embed".into()),
14        _ => {}
15    }
16
17    if let Some(rest) = name.strip_prefix("transformer_encoder.") {
18        let (n_str, rest) = rest.split_once('.')?;
19        let n: u32 = n_str.parse().ok()?;
20        let suffix = match rest {
21            "attn_norm.weight" => "attn_norm.weight",
22            "attn_norm.bias" => "attn_norm.bias",
23            "ff_norm.weight" => "ff_norm.weight",
24            "ff_norm.bias" => "ff_norm.bias",
25            "q_proj.weight" => "q_proj.weight",
26            "k_proj.weight" => "k_proj.weight",
27            "v_proj.weight" => "v_proj.weight",
28            "out_proj.weight" => "out_proj.weight",
29            "q_gate.weight" => "q_gate.weight",
30            "q_norm.weight" => "q_norm.weight",
31            "k_norm.weight" => "k_norm.weight",
32            "ff.up.weight" => "ff_up.weight",
33            "ff.down.weight" => "ff_down.weight",
34            "kappa" | "max_len_f" | "n0" => return None,
35            _ => return None,
36        };
37        return Some(format!("blk.{n}.{suffix}"));
38    }
39
40    if let Some(rest) = name.strip_prefix("y_encoders.") {
41        let (n_str, rest) = rest.split_once('.')?;
42        let n: u32 = n_str.parse().ok()?;
43        let suffix = match rest {
44            "0.weight" => "fc1.weight",
45            "0.bias" => "fc1.bias",
46            "2.weight" => "fc2.weight",
47            "2.bias" => "fc2.bias",
48            _ => return None,
49        };
50        return Some(format!("y_enc.{n}.{suffix}"));
51    }
52
53    None
54}
55
56#[cfg(test)]
57mod tests {
58    use super::*;
59
60    #[test]
61    fn top_level() {
62        assert_eq!(map_tensor_name("encoder.weight"), Some("encoder.weight".into()));
63        assert_eq!(map_tensor_name("head.2.bias"), Some("head_fc2.bias".into()));
64        assert_eq!(map_tensor_name("thinking_embed"), Some("thinking_embed".into()));
65    }
66
67    #[test]
68    fn block_and_y_encoder() {
69        assert_eq!(map_tensor_name("transformer_encoder.0.q_proj.weight"), Some("blk.0.q_proj.weight".into()));
70        assert_eq!(map_tensor_name("transformer_encoder.31.ff.down.weight"), Some("blk.31.ff_down.weight".into()));
71        assert_eq!(map_tensor_name("y_encoders.5.2.weight"), Some("y_enc.5.fc2.weight".into()));
72    }
73
74    #[test]
75    fn skips_scalar_buffers() {
76        assert_eq!(map_tensor_name("transformer_encoder.0.kappa"), None);
77        assert_eq!(map_tensor_name("transformer_encoder.0.max_len_f"), None);
78        assert_eq!(map_tensor_name("transformer_encoder.0.n0"), None);
79    }
80}