zsfm_flowstate/
tensor_map.rs1pub fn map_tensor_name(hf_name: &str) -> Option<String> {
16 match hf_name {
18 "embed.embed.weight" => return Some("embed.weight".into()),
19 "embed.embed.bias" => return Some("embed.bias".into()),
20 "decoder.lin.weight" => return Some("decoder.weight".into()),
21 "decoder.lin.bias" => return Some("decoder.bias".into()),
22 _ => {}
23 }
24
25 let rest = hf_name.strip_prefix("encoder.layers.")?;
27 let dot = rest.find('.')?;
28 let n: usize = rest[..dot].parse().ok()?;
29 let suffix = &rest[dot + 1..];
30
31 let name = match suffix {
32 "ssm.log_Lambda_real" => format!("blk.{n}.ssm.log_lambda_real"),
33 "ssm.Lambda_imag" => format!("blk.{n}.ssm.lambda_imag"),
34 "ssm.B_tilde_r" => format!("blk.{n}.ssm.b_r"),
35 "ssm.B_tilde_i" => format!("blk.{n}.ssm.b_i"),
36 "ssm.C_tilde_r" => format!("blk.{n}.ssm.c_r"),
37 "ssm.C_tilde_i" => format!("blk.{n}.ssm.c_i"),
38 "ssm.D" => format!("blk.{n}.ssm.d"),
39 "ssm.log_Delta" => format!("blk.{n}.ssm.log_delta"),
40 "out.weight" => format!("blk.{n}.out.weight"),
41 "out.bias" => format!("blk.{n}.out.bias"),
42 "norm.weight" => format!("blk.{n}.norm.weight"),
43 "norm.bias" => format!("blk.{n}.norm.bias"),
44 _ => return None,
45 };
46 Some(name)
47}
48
49#[cfg(test)]
50mod tests {
51 use super::*;
52
53 #[test]
54 fn test_embed() {
55 assert_eq!(map_tensor_name("embed.embed.weight"), Some("embed.weight".into()));
56 assert_eq!(map_tensor_name("embed.embed.bias"), Some("embed.bias".into()));
57 }
58
59 #[test]
60 fn test_decoder() {
61 assert_eq!(map_tensor_name("decoder.lin.weight"), Some("decoder.weight".into()));
62 assert_eq!(map_tensor_name("decoder.lin.bias"), Some("decoder.bias".into()));
63 }
64
65 #[test]
66 fn test_ssm_params() {
67 assert_eq!(
68 map_tensor_name("encoder.layers.0.ssm.log_Lambda_real"),
69 Some("blk.0.ssm.log_lambda_real".into())
70 );
71 assert_eq!(
72 map_tensor_name("encoder.layers.3.ssm.B_tilde_r"),
73 Some("blk.3.ssm.b_r".into())
74 );
75 assert_eq!(
76 map_tensor_name("encoder.layers.5.ssm.C_tilde_i"),
77 Some("blk.5.ssm.c_i".into())
78 );
79 assert_eq!(
80 map_tensor_name("encoder.layers.2.ssm.log_Delta"),
81 Some("blk.2.ssm.log_delta".into())
82 );
83 }
84
85 #[test]
86 fn test_layer_mlp_and_norm() {
87 assert_eq!(
88 map_tensor_name("encoder.layers.1.out.weight"),
89 Some("blk.1.out.weight".into())
90 );
91 assert_eq!(
92 map_tensor_name("encoder.layers.4.norm.bias"),
93 Some("blk.4.norm.bias".into())
94 );
95 }
96
97 #[test]
98 fn test_unknown_returns_none() {
99 assert_eq!(map_tensor_name("something.unknown"), None);
100 }
101}