Skip to main content

zsfm_moirai2/
tensor_map.rs

1/// Map Moirai-2.0-R-small safetensors tensor names to GGUF naming convention.
2pub fn map_tensor_name(name: &str) -> Option<String> {
3    match name {
4        // ResidualBlock in_proj
5        "in_proj.hidden_layer.weight"   => return Some("in_proj.hidden.weight".into()),
6        "in_proj.hidden_layer.bias"     => return Some("in_proj.hidden.bias".into()),
7        "in_proj.output_layer.weight"   => return Some("in_proj.output.weight".into()),
8        "in_proj.output_layer.bias"     => return Some("in_proj.output.bias".into()),
9        "in_proj.residual_layer.weight" => return Some("in_proj.residual.weight".into()),
10        "in_proj.residual_layer.bias"   => return Some("in_proj.residual.bias".into()),
11        // ResidualBlock out_proj
12        "out_proj.hidden_layer.weight"   => return Some("out_proj.hidden.weight".into()),
13        "out_proj.hidden_layer.bias"     => return Some("out_proj.hidden.bias".into()),
14        "out_proj.output_layer.weight"   => return Some("out_proj.output.weight".into()),
15        "out_proj.output_layer.bias"     => return Some("out_proj.output.bias".into()),
16        "out_proj.residual_layer.weight" => return Some("out_proj.residual.weight".into()),
17        "out_proj.residual_layer.bias"   => return Some("out_proj.residual.bias".into()),
18        // Final norm
19        "encoder.norm.weight" => return Some("norm_f.weight".into()),
20        _ => {}
21    }
22
23    if let Some(rest) = name.strip_prefix("encoder.layers.") {
24        let (n_str, rest) = rest.split_once('.')?;
25        let n: u32 = n_str.parse().ok()?;
26
27        let suffix = match rest {
28            "norm1.weight"                       => "norm1.weight",
29            "norm2.weight"                       => "norm2.weight",
30            "self_attn.q_proj.weight"            => "attn_q.weight",
31            "self_attn.k_proj.weight"            => "attn_k.weight",
32            "self_attn.v_proj.weight"            => "attn_v.weight",
33            "self_attn.out_proj.weight"          => "attn_o.weight",
34            "self_attn.q_norm.weight"            => "attn_qn.weight",
35            "self_attn.k_norm.weight"            => "attn_kn.weight",
36            "self_attn.var_attn_bias.emb.weight" => "attn_vbias.weight",
37            "ffn.fc1.weight"                     => "ffn_fc1.weight",
38            "ffn.fc2.weight"                     => "ffn_fc2.weight",
39            "ffn.fc_gate.weight"                 => "ffn_gate.weight",
40            _ => return None,
41        };
42        return Some(format!("blk.{n}.{suffix}"));
43    }
44
45    None
46}
47
48#[cfg(test)]
49mod tests {
50    use super::*;
51
52    #[test]
53    fn in_proj() {
54        assert_eq!(
55            map_tensor_name("in_proj.hidden_layer.weight"),
56            Some("in_proj.hidden.weight".into())
57        );
58    }
59
60    #[test]
61    fn encoder_block() {
62        assert_eq!(
63            map_tensor_name("encoder.layers.0.self_attn.q_proj.weight"),
64            Some("blk.0.attn_q.weight".into())
65        );
66        assert_eq!(
67            map_tensor_name("encoder.layers.5.ffn.fc_gate.weight"),
68            Some("blk.5.ffn_gate.weight".into())
69        );
70    }
71
72    #[test]
73    fn out_proj() {
74        assert_eq!(
75            map_tensor_name("out_proj.output_layer.bias"),
76            Some("out_proj.output.bias".into())
77        );
78    }
79}