zsfm_tabdpt/
tensor_map.rs1pub 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}