Skip to main content

zsfm_flowstate/
tensor_map.rs

1/// Map a FlowState HuggingFace tensor name to the GGUF canonical name.
2/// Returns `None` for tensors unused at inference time (caller skips them).
3///
4/// HF prefix in safetensors:  (none — model is saved as FlowStateForPrediction,
5/// which wraps FlowStateModel in `self.model`, but the checkpoint stores
6/// FlowStateModel weights directly without the `.model.` prefix.)
7///
8/// Verified tensor names from model.safetensors header:
9///   embed.embed.weight / embed.embed.bias
10///   encoder.layers.{N}.ssm.{log_Lambda_real, Lambda_imag, B_tilde_r, B_tilde_i,
11///                                C_tilde_r, C_tilde_i, D, log_Delta}
12///   encoder.layers.{N}.out.{weight, bias}
13///   encoder.layers.{N}.norm.{weight, bias}
14///   decoder.lin.{weight, bias}
15pub fn map_tensor_name(hf_name: &str) -> Option<String> {
16    // Embedding
17    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    // encoder.layers.{N}.ssm.*  and  encoder.layers.{N}.{out,norm}.*
26    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}