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>,
)
| 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, |
nothing calls this directly
no test coverage detected