| 41 | |
| 42 | impl EncoderBlock { |
| 43 | fn load(vb: VarBuilder, cfg: &EncoderConfig, backend: Arc<dyn ComputeBackend>) -> Result<Self> { |
| 44 | let h = cfg.hidden_size; |
| 45 | let i = cfg.intermediate_size; |
| 46 | let size_q = cfg.num_heads * cfg.head_dim; |
| 47 | let size_kv = cfg.num_kv_heads * cfg.head_dim; |
| 48 | |
| 49 | let attn = vb.pp("self_attn"); |
| 50 | let q_proj_weight = attn.pp("q_proj").get((size_q, h), "weight")?; |
| 51 | let k_proj_weight = attn.pp("k_proj").get((size_kv, h), "weight")?; |
| 52 | let v_proj_weight = attn.pp("v_proj").get((size_kv, h), "weight")?; |
| 53 | let o_proj_weight = attn.pp("o_proj").get((h, size_q), "weight")?; |
| 54 | let q_norm_weight = attn.pp("q_norm").get(cfg.head_dim, "weight")?; |
| 55 | let k_norm_weight = attn.pp("k_norm").get(cfg.head_dim, "weight")?; |
| 56 | let qk_norm_eps = cfg.rms_norm_eps as f32; |
| 57 | |
| 58 | let mlp = vb.pp("mlp"); |
| 59 | let gate_proj_weight = mlp.pp("gate_proj").get((i, h), "weight")?; |
| 60 | let up_proj_weight = mlp.pp("up_proj").get((i, h), "weight")?; |
| 61 | let down_proj_weight = mlp.pp("down_proj").get((h, i), "weight")?; |
| 62 | |
| 63 | let rms_1_weight = vb.pp("input_layernorm").get(h, "weight")?; |
| 64 | let rms_2_weight = vb.pp("post_attention_layernorm").get(h, "weight")?; |
| 65 | let rms_eps = cfg.rms_norm_eps as f32; |
| 66 | |
| 67 | Ok(Self { |
| 68 | rms_1_weight, |
| 69 | rms_2_weight, |
| 70 | rms_eps, |
| 71 | q_proj_weight, |
| 72 | k_proj_weight, |
| 73 | v_proj_weight, |
| 74 | o_proj_weight, |
| 75 | q_norm_weight, |
| 76 | k_norm_weight, |
| 77 | qk_norm_eps, |
| 78 | num_heads: cfg.num_heads, |
| 79 | num_kv_heads: cfg.num_kv_heads, |
| 80 | head_dim: cfg.head_dim, |
| 81 | gate_proj_weight, |
| 82 | up_proj_weight, |
| 83 | down_proj_weight, |
| 84 | backend, |
| 85 | }) |
| 86 | } |
| 87 | |
| 88 | /// attn_mask: optional (1, seq) tensor with 1 for real tokens, 0 for padding |
| 89 | fn forward_with_mask(&self, x: &Tensor, attn_mask: Option<&Tensor>) -> Result<Tensor> { |