4848#include " gemma/attention.h" // includes highway.h
4949#include " gemma/tiled_attention.h"
5050#include " gemma/gemma-inl.h"
51+ #include " ops/fast_ops-inl.h"
5152#include " ops/ops-inl.h"
5253
5354HWY_BEFORE_NAMESPACE ();
@@ -239,6 +240,8 @@ struct Gemma4MoE {
239240 activations.s_expert_in .Notify (layer.layer_idx , tmp_in, env.ctx , 0 ,
240241 cluster_idx, parallelism);
241242
243+ const size_t expert_ff_hidden_dim =
244+ layer.moe_gating_einsum_w1 [expert_idx].Rows ();
242245 const ActivationType activation = layer.layer_config .activation ;
243246 MatPtrT<BF16 >& C1 = per_cluster.moe_C1 ;
244247 C1 .OverrideRows (expert_size);
@@ -275,19 +278,20 @@ struct Gemma4MoE {
275278 cluster_idx, parallelism);
276279
277280 // Hidden layer -> output layer.
278- size_t expert_ff_hidden_dim = layer.moe_gating_einsum_w1 [expert_idx].Rows ();
279- MatStorageT<BF16 > C1_narrow (" C1_n" ,
280- Extents2D (expert_size, expert_ff_hidden_dim),
281- env.ctx .allocator , MatPadding::kOdd );
282- for (size_t i = 0 ; i < expert_size; ++i) {
283- memcpy (C1_narrow.Row (i), C1 .Row (i), expert_ff_hidden_dim * sizeof (BF16 ));
281+ {
282+ MatStorageT<BF16 > C1_narrow (" C1_n" ,
283+ Extents2D (expert_size, expert_ff_hidden_dim),
284+ env.ctx .allocator , MatPadding::kOdd );
285+ for (size_t i = 0 ; i < expert_size; ++i) {
286+ memcpy (C1_narrow.Row (i), C1 .Row (i),
287+ expert_ff_hidden_dim * sizeof (BF16 ));
288+ }
289+ CallMatMul (C1_narrow, layer.moe_linear_w [expert_idx],
290+ /* add=*/ nullptr , env, expert_out, options);
284291 }
285292
286-
287- CallMatMul (C1_narrow, layer.moe_linear_w [expert_idx],
288- /* add=*/ nullptr , env, expert_out, options);
289-
290- if (layer.p_expert_sc .HasPtr ()) {
293+ if (layer.p_expert_sc .HasPtr ()
294+ ) {
291295 const MatPtrT<BF16 > sc_mat (layer.p_expert_sc );
292296 const BF16 * sc_data = sc_mat.Row (0 );
293297 const float expert_scale =
@@ -466,9 +470,12 @@ void Gemma4MoETransformerLayer(size_t num_tokens, size_t layer_idx,
466470 /* is_attention=*/ true , env.ctx );
467471
468472 // Dual-Path FFW
469- pre_norm (
470- layer.pre_ffw2_ns .HasPtr () ? layer.pre_ffw2_ns : layer.pre_ffw_norm_scale ,
471- activations.pre_ffw_rms_out );
473+ const MatPtr& shared_norm = layer.pre_ffw2_ns .HasPtr () ? layer.pre_ffw2_ns : layer.pre_ffw_norm_scale ;
474+ const MatPtr& moe_norm = layer.pre_ffw_norm_scale ;
475+ const MatPtr& shared_post_norm = layer.post_ffw2_ns ;
476+ const MatPtr& moe_post_norm = layer.post_ffw1_ns ;
477+
478+ pre_norm (shared_norm, activations.pre_ffw_rms_out );
472479
473480 // Shared MLP Path
474481 FFWNoVit (layer, activations, env); // writes to activations.ffw_out
@@ -482,17 +489,17 @@ void Gemma4MoETransformerLayer(size_t num_tokens, size_t layer_idx,
482489 }
483490 }
484491
485- if (layer. post_ffw2_ns .HasPtr ()) {
486- rms_norm_inplace (layer. post_ffw2_ns , activations.attention .att_sums );
492+ if (shared_post_norm .HasPtr ()) {
493+ rms_norm_inplace (shared_post_norm , activations.attention .att_sums );
487494 }
488495
489496 // MoE Path
490- pre_norm (layer. pre_ffw_norm_scale , activations.pre_ffw_rms_out );
497+ pre_norm (moe_norm , activations.pre_ffw_rms_out );
491498
492499 Gemma4MoE::MoEFFW (layer, activations, env); // writes to activations.ffw_out
493500
494- if (layer. post_ffw1_ns .HasPtr ()) {
495- rms_norm_inplace (layer. post_ffw1_ns , activations.ffw_out );
501+ if (moe_post_norm .HasPtr ()) {
502+ rms_norm_inplace (moe_post_norm , activations.ffw_out );
496503 }
497504
498505 // Combine & Final Norm (Fix for dual-path combination)
0 commit comments