Skip to content

Commit ebd4a78

Browse files
Merge pull request #4378 from CIeNET-International:emma/fix-gemma4-cp-convert
PiperOrigin-RevId: 944551277
2 parents 989abc6 + 28e9cf7 commit ebd4a78

1 file changed

Lines changed: 3 additions & 3 deletions

File tree

src/maxtext/layers/nnx_decoders.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -653,7 +653,7 @@ def _init_scanned_gemma4(self, decoder_block_classes, rngs, mesh):
653653
RemattedGemma4Block = gemma4.Gemma4ScannableBlock
654654

655655
if scan_length > 0:
656-
self.layers = self._create_scanned_layers(
656+
self.scanned_blocks = self._create_scanned_layers(
657657
RemattedGemma4Block,
658658
length=scan_length,
659659
metadata_axis_name="layers",
@@ -2030,8 +2030,8 @@ def _apply_gemma4_scanned_blocks(
20302030
grouped_kv_caches = maxtext_utils.prepare_kv_caches_for_scan(
20312031
kv_caches, scan_length, attention_pattern_length, stack=False
20322032
)
2033-
y, self.layers, _ = self._apply_layers_sequentially(
2034-
self.layers, y, *layer_args, length=scan_length, kv_caches_stacked=grouped_kv_caches, **layer_kwargs
2033+
y, self.scanned_blocks, _ = self._apply_layers_sequentially(
2034+
self.scanned_blocks, y, *layer_args, length=scan_length, kv_caches_stacked=grouped_kv_caches, **layer_kwargs
20352035
)
20362036
maxtext_utils.update_kv_caches_after_scan(
20372037
kv_caches, grouped_kv_caches, scan_length, attention_pattern_length, stacked=False

0 commit comments

Comments
 (0)