From 28e9cf76153b0ef3a8db4da19654219ce7bf3a17 Mon Sep 17 00:00:00 2001 From: Emma Lien Date: Wed, 8 Jul 2026 04:38:30 +0000 Subject: [PATCH] Fix: Gemma4 checkpoint conversion failure when enable_nnx=True --- src/maxtext/layers/nnx_decoders.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/maxtext/layers/nnx_decoders.py b/src/maxtext/layers/nnx_decoders.py index 1b14803d94..a7ce5822f9 100644 --- a/src/maxtext/layers/nnx_decoders.py +++ b/src/maxtext/layers/nnx_decoders.py @@ -653,7 +653,7 @@ def _init_scanned_gemma4(self, decoder_block_classes, rngs, mesh): RemattedGemma4Block = gemma4.Gemma4ScannableBlock if scan_length > 0: - self.layers = self._create_scanned_layers( + self.scanned_blocks = self._create_scanned_layers( RemattedGemma4Block, length=scan_length, metadata_axis_name="layers", @@ -2030,8 +2030,8 @@ def _apply_gemma4_scanned_blocks( grouped_kv_caches = maxtext_utils.prepare_kv_caches_for_scan( kv_caches, scan_length, attention_pattern_length, stack=False ) - y, self.layers, _ = self._apply_layers_sequentially( - self.layers, y, *layer_args, length=scan_length, kv_caches_stacked=grouped_kv_caches, **layer_kwargs + y, self.scanned_blocks, _ = self._apply_layers_sequentially( + self.scanned_blocks, y, *layer_args, length=scan_length, kv_caches_stacked=grouped_kv_caches, **layer_kwargs ) maxtext_utils.update_kv_caches_after_scan( kv_caches, grouped_kv_caches, scan_length, attention_pattern_length, stacked=False