MCPcopy Create free account
hub / github.com/evilsocket/cake / load_custom

Method load_custom

cake-core/src/models/common/attention.rs:76–149  ·  view source on GitHub ↗

Custom load with explicit per-layer options (used by Gemma3 for interleaved local/global).

(
        vb: VarBuilder,
        cfg: &super::Config,
        use_qk_norm: bool,
        sliding_window: Option<usize>,
        use_rope: bool,
        backend: Arc<dyn ComputeBackend>,
    )

Source from the content-addressed store, hash-verified

74
75 /// Custom load with explicit per-layer options (used by Gemma3 for interleaved local/global).
76 pub fn load_custom(
77 vb: VarBuilder,
78 cfg: &super::Config,
79 use_qk_norm: bool,
80 sliding_window: Option<usize>,
81 use_rope: bool,
82 backend: Arc<dyn ComputeBackend>,
83 ) -> Result<Self> {
84 let size_in = cfg.hidden_size;
85 let head_dim = cfg.head_dim.unwrap_or(cfg.hidden_size / cfg.num_attention_heads);
86 let rotary_dim = (head_dim as f32 * cfg.partial_rotary_factor) as usize;
87 let size_q = head_dim * cfg.num_attention_heads;
88 let size_kv = head_dim * cfg.num_key_value_heads;
89
90 let (qkv_proj_weight, qkv_proj_bias) = if cfg.fused_qkv_proj {
91 // Phi-3/4 style: weights already fused as a single 'qkv_proj' tensor.
92 let w = vb.pp("qkv_proj").get((size_q + 2 * size_kv, size_in), "weight")?;
93 let w = backend.preprocess_linear_weight(&w)?;
94 (w, None)
95 } else if cfg.use_qkv_bias {
96 let q_w = vb.pp("q_proj").get((size_q, size_in), "weight")?;
97 let k_w = vb.pp("k_proj").get((size_kv, size_in), "weight")?;
98 let v_w = vb.pp("v_proj").get((size_kv, size_in), "weight")?;
99 let fused_w = Tensor::cat(&[&q_w, &k_w, &v_w], 0)?;
100 let fused_w = backend.preprocess_linear_weight(&fused_w)?;
101
102 let q_b = vb.pp("q_proj").get(size_q, "bias")?;
103 let k_b = vb.pp("k_proj").get(size_kv, "bias")?;
104 let v_b = vb.pp("v_proj").get(size_kv, "bias")?;
105 let fused_b = Tensor::cat(&[&q_b, &k_b, &v_b], 0)?;
106
107 (fused_w, Some(fused_b))
108 } else {
109 let q_w = vb.pp("q_proj").get((size_q, size_in), "weight")?;
110 let k_w = vb.pp("k_proj").get((size_kv, size_in), "weight")?;
111 let v_w = vb.pp("v_proj").get((size_kv, size_in), "weight")?;
112 let fused_w = Tensor::cat(&[&q_w, &k_w, &v_w], 0)?;
113 let fused_w = backend.preprocess_linear_weight(&fused_w)?;
114 (fused_w, None)
115 };
116
117 let o_w = vb.pp("o_proj").get((size_in, size_q), "weight")?;
118 let o_proj_weight = backend.preprocess_linear_weight(&o_w)?;
119
120 let (q_norm_weight, k_norm_weight) = if use_qk_norm {
121 let norm_dim = if cfg.pre_reshape_qk_norm { size_q } else { head_dim };
122 let norm_kv_dim = if cfg.pre_reshape_qk_norm { size_kv } else { head_dim };
123 let residual = cfg.residual_rms_norm;
124 let qn = load_rms_norm_weight(norm_dim, residual, vb.pp("q_norm"))?;
125 let kn = load_rms_norm_weight(norm_kv_dim, residual, vb.pp("k_norm"))?;
126 (Some(qn), Some(kn))
127 } else {
128 (None, None)
129 };
130
131 Ok(Self {
132 qkv_proj_weight,
133 qkv_proj_bias,

Callers

nothing calls this directly

Calls 3

load_rms_norm_weightFunction · 0.85
getMethod · 0.45

Tested by

no test coverage detected