(self, spec, module, quant_type=common_spec.Quantization.CT2)
| 2262 | # Gemma4 uses output * gamma (ones-initialized), not output * (1 + gamma) |
| 2263 | |
| 2264 | def set_decoder(self, spec, module, quant_type=common_spec.Quantization.CT2): |
| 2265 | spec.scale_embeddings = True |
| 2266 | spec.start_from_zero_embedding = False |
| 2267 | self.set_embeddings(spec.embeddings, module.embed_tokens) |
| 2268 | self.set_layer_norm(spec.layer_norm, module.norm) |
| 2269 | |
| 2270 | attention_k_eq_v = getattr(self, "_attention_k_eq_v", False) |
| 2271 | ghd = getattr(self, "_global_head_dim", None) |
| 2272 | grd = getattr(self, "_global_rotary_dim", None) |
| 2273 | # HF's proportional partial-RoPE pairs channels [0:R/2]↔[HD/2:HD/2+R/2]; |
| 2274 | # CT2's RotaryEmbeddings pairs [0:R/2]↔[R/2:R]. Permute Q/K accordingly. |
| 2275 | partial_perm = None |
| 2276 | if ghd and grd and 0 < grd < ghd: |
| 2277 | partial_perm = ( |
| 2278 | list(range(0, grd // 2)) |
| 2279 | + list(range(ghd // 2, ghd // 2 + grd // 2)) |
| 2280 | + list(range(grd // 2, ghd // 2)) |
| 2281 | + list(range(ghd // 2 + grd // 2, ghd)) |
| 2282 | ) |
| 2283 | |
| 2284 | for layer_spec, layer in zip(spec.layer, module.layers): |
| 2285 | self.set_layer_norm(layer_spec.input_layer_norm, layer.input_layernorm) |
| 2286 | self.set_layer_norm( |
| 2287 | layer_spec.post_attention_layer_norm, layer.post_attention_layernorm |
| 2288 | ) |
| 2289 | self.set_layer_norm( |
| 2290 | layer_spec.pre_feedforward_layer_norm, layer.pre_feedforward_layernorm |
| 2291 | ) |
| 2292 | self.set_layer_norm( |
| 2293 | layer_spec.post_feedforward_layer_norm, layer.post_feedforward_layernorm |
| 2294 | ) |
| 2295 | self.set_layer_norm( |
| 2296 | layer_spec.self_attention.q_norm, layer.self_attn.q_norm |
| 2297 | ) |
| 2298 | self.set_layer_norm( |
| 2299 | layer_spec.self_attention.k_norm, layer.self_attn.k_norm |
| 2300 | ) |
| 2301 | |
| 2302 | # v_norm has no learnable scale; supply all-ones gamma (pure RMS norm) |
| 2303 | layer_spec.self_attention.v_norm.gamma = ( |
| 2304 | torch.ones_like(layer.self_attn.k_norm.weight).float().numpy() |
| 2305 | ) |
| 2306 | |
| 2307 | # When attention_k_eq_v is set, full-attention layers have no v_proj — |
| 2308 | # values are the same as keys, so we reuse k_proj weights. |
| 2309 | is_full_attn = layer.self_attn.layer_type == "full_attention" |
| 2310 | use_k_as_v = attention_k_eq_v and is_full_attn |
| 2311 | |
| 2312 | split_layers = [common_spec.LinearSpec() for _ in range(3)] |
| 2313 | self.set_linear( |
| 2314 | split_layers[0], layer.self_attn.q_proj, quant_type=quant_type |
| 2315 | ) |
| 2316 | self.set_linear( |
| 2317 | split_layers[1], layer.self_attn.k_proj, quant_type=quant_type |
| 2318 | ) |
| 2319 | if use_k_as_v: |
| 2320 | self.set_linear( |
| 2321 | split_layers[2], layer.self_attn.k_proj, quant_type=quant_type |
no test coverage detected