@@ -392,7 +392,7 @@ def get_decoder_layers(self):
392392 case DecoderBlockType .GEMMA2 :
393393 return [gemma2 .Gemma2DecoderLayerToLinen ]
394394 case DecoderBlockType .GEMMA3 :
395- return [gemma3 .Gemma3DecoderLayer ]
395+ return [gemma3 .Gemma3DecoderLayerToLinen ]
396396 case DecoderBlockType .GPT3 :
397397 return [gpt3 .Gpt3DecoderLayer ]
398398 case DecoderBlockType .GPT_OSS :
@@ -485,7 +485,7 @@ def scan_decoder_layers(self, cfg, decoder_layer, length, metadata_axis_name, me
485485 length = length ,
486486 metadata_params = {nn .PARTITION_NAME : metadata_axis_name },
487487 )
488- return scan_fn (config = cfg , mesh = mesh , name = metadata_axis_name , quant = self .quant , ** kwargs )
488+ return scan_fn (config = cfg , mesh = mesh , name = metadata_axis_name , quant = self .quant , ** kwargs ) # pytype: disable=wrong-keyword-args
489489
490490 def get_pipeline_stage_module (self , decoder_blocks ):
491491 """get pipeline stage module"""
@@ -880,7 +880,7 @@ def _apply_gemma3_scanned_blocks(
880880 scan_length = cfg .num_decoder_layers // attention_pattern_length
881881
882882 policy = self .get_remat_policy ()
883- RemattedGemma3Block = self .set_remat_policy ([gemma3 .Gemma3ScannableBlock ], policy )[0 ]
883+ RemattedGemma3Block = self .set_remat_policy ([gemma3 .Gemma3ScannableBlockToLinen ], policy )[0 ]
884884
885885 layer_call_kwargs = {"bidirectional_mask" : bidirectional_mask }
886886 layer_kwargs = {"num_of_layers" : attention_pattern_length }
@@ -909,6 +909,7 @@ def _apply_gemma3_scanned_blocks(
909909 if num_remaining_layers > 0 :
910910 # We name the remainder block with a 'remainder' suffix to avoid parameter name collisions
911911 rem_layer_kwargs = {"num_of_layers" : num_remaining_layers }
912+ # pytype: disable=wrong-keyword-args
912913 layer = RemattedGemma3Block (
913914 config = cfg , mesh = mesh , quant = self .quant , model_mode = self .model_mode , name = "layers_remainder" , ** rem_layer_kwargs
914915 )
0 commit comments