@@ -14,15 +14,15 @@ class Gemma4Converter(LlmConverter):
1414
1515 def __init__ (self , args , config , loader = None ):
1616 self .rmsnorm_type = WeightType .RMSNORM
17- super ().__init__ (args , config , loader = None )
17+ super ().__init__ (args , config , loader = loader )
1818 # Override values set by base class __init__
1919 self .do_vit = True
2020 self .vit_f16_out_bf16 = True # Gemma4 vit is f16, but we force output to bf16
2121 self .do_audio = True
2222 if self .do_audio :
2323 if args .audio_length <= 0 :
24- self .audio_length = 200
25- print ("audio_length not specified, using default 200 " )
24+ self .audio_length = 750
25+ print ("audio_length not specified, using default 750 " )
2626 elif args .audio_length > 750 :
2727 self .audio_length = 750
2828 print (f"audio_length { args .audio_length } exceeds max 750, capped to 750" )
@@ -34,6 +34,10 @@ def __init__(self, args, config, loader=None):
3434 # which already set self.cos_sliding, self.sin_sliding, self.cos_full, self.sin_full
3535 # and self.cos, self.sin
3636
37+ # Export per_layer_token_embd as external bin file to avoid 35x duplicate storage in bmodel
38+ if self .hidden_size_per_layer_input > 0 :
39+ self .all_gen_mlirs .append (self .gen_per_layer_embedding_bin )
40+
3741 @override
3842 def load_pretrained (self , config ):
3943 super ().load_pretrained (config )
@@ -44,7 +48,7 @@ def load_pretrained(self, config):
4448 @override
4549 def init_config (self ):
4650 super ().init_config ()
47- self .tie_word_embeddings = True
51+ self .tie_word_embeddings = getattr ( self . llm_config , 'tie_word_embeddings' , True )
4852 self .do_lmhead_merge = self .tie_word_embeddings and not self .embedding_disk and self .num_device < 2
4953 # Gemma4 specific config
5054 self .layer_types = getattr (self .llm_config , 'layer_types' , None )
@@ -119,6 +123,41 @@ def rotary_embedding(self):
119123 # Return sliding cos/sin for base class compatibility
120124 return cos_sliding , sin_sliding
121125
126+ def gen_per_layer_embedding_bin (self ):
127+ """Export per_layer_token_embd weights as external bin file.
128+
129+ The full [vocab, N*D] weight matrix is pre-multiplied by per_layer_embed_scale
130+ and saved as raw BF16/F16 bytes. At runtime, the CPU loads this file, does a
131+ single gather on the full matrix, then slices per-layer results for each block.
132+ This avoids storing 35 separate [vocab, D] slices inside the bmodel (~4.7 GB saved).
133+ """
134+ bin_file = os .path .join (self .config_dir , 'per_layer_token_embd.bin' )
135+ if os .path .exists (bin_file ):
136+ logger .info ("%s already exists. Skipping export." , bin_file )
137+ return
138+
139+ embed_per_layer = self .model_info .weights [LlmList .EMBEDING_PER_LAYER ]
140+ full_emb = self .model .read (embed_per_layer + ".weight" ) # [vocab, N*D]
141+
142+ # Pre-multiply scale so runtime doesn't need MulConstOp
143+ per_layer_embed_scale = self .hidden_size_per_layer_input ** 0.5
144+ full_emb_scaled = full_emb * per_layer_embed_scale
145+
146+ import ctypes
147+ weight = torch .from_numpy (full_emb_scaled )
148+ if self .half_precision_quantize == 'bf16' :
149+ tensor_data = weight .to (torch .bfloat16 )
150+ elif self .half_precision_quantize == 'f16' :
151+ tensor_data = weight .to (torch .float16 )
152+ else :
153+ raise NotImplementedError (
154+ f"per_layer_embedding_bin not supported for quantize={ self .quantize } " )
155+ data_ptr = tensor_data .untyped_storage ().data_ptr ()
156+ buffer = (ctypes .c_byte * (tensor_data .numel () * 2 )).from_address (data_ptr )
157+ with open (bin_file , 'wb' ) as f :
158+ f .write (buffer )
159+ tqdm .write (f"exported { bin_file } ({ os .path .getsize (bin_file ) / (1024 ** 2 ):.1f} MB)" )
160+
122161 def _compute_layer_params (self , idx ):
123162 """Compute layer-specific parameters based on layer type and KV sharing."""
124163 layer_type = self .layer_types [idx ]
@@ -456,10 +495,7 @@ def gen_block_mlir(self, idx):
456495 model_projection = self .model_info .weights [LlmList .PER_LAYER_MODEL_PROJECTION ]
457496 projection_norm = self .model_info .weights [LlmList .PER_LAYER_PROJECTION_NORM ]
458497 per_layer_dim = self .hidden_size_per_layer_input
459- # Slice embed_tokens_per_layer.weight: [vocab, N*D] → [vocab, D] for layer idx
460- full_emb = self .model .read (embed_per_layer + ".weight" )
461- emb_slice = full_emb [:, idx * per_layer_dim :(idx + 1 ) * per_layer_dim ]
462- weight_dict [embed_per_layer + f".weight.{ idx } " ] = emb_slice .copy ()
498+ # per_layer_token_embd is stored externally as bin file (see gen_per_layer_embedding_bin)
463499 # Slice per_layer_model_projection.weight: [N*D, hidden] → transpose → [hidden, N*D]
464500 # then slice columns → [hidden, D] for layer idx
465501 full_proj = self .model .read (model_projection + ".weight" )
@@ -505,8 +541,8 @@ def gen_mlp(mlir_gen, input_shape, in_op):
505541 ip = ip ).output
506542 return new_op
507543
508- def gen_per_layer_input (mlir_gen , input_shape , hidden_states , residual_op , ids_op ,
509- embeds_op ):
544+ def gen_per_layer_input (mlir_gen , input_shape , hidden_states , residual_op , embeds_op ,
545+ per_layer_embeds_op ):
510546 """Generate per_layer_input subgraph with per-layer sliced weights.
511547 Each block uses its own slice of embed_tokens_per_layer and per_layer_model_projection,
512548 avoiding the need to compute the full N*D tensor and then slice."""
@@ -516,27 +552,13 @@ def gen_per_layer_input(mlir_gen, input_shape, hidden_states, residual_op, ids_o
516552 len = input_shape [1 ]
517553
518554 # Scale constants from HF source
519- per_layer_embed_scale = per_layer_dim ** 0.5
520555 model_projection_scale = self .hidden_size ** - 0.5
521556 per_layer_input_scale = 2.0 ** - 0.5
522-
523- # Path A: embed_tokens_per_layer(ids) * per_layer_embed_scale (sliced for layer idx)
524557 per_layer_slice_shape = [batch , len , per_layer_dim ]
525- emb_per_layer_weight = mlir_gen .create_weight_op (
526- embed_per_layer + f".weight.{ idx } " ,
527- [self .vocab_size_per_layer_input , per_layer_dim ])
528- gather_op = top .GatherOp (mlir_gen .get_tensor_type (per_layer_slice_shape ),
529- emb_per_layer_weight ,
530- ids_op ,
531- axis = 0 ,
532- loc = self .get_loc ("embed_tokens_per_layer.gather" , mlir_gen ),
533- ip = ip ).output
534- per_layer_inputs_op = top .MulConstOp (mlir_gen .get_tensor_type (per_layer_slice_shape ),
535- gather_op ,
536- const_val = per_layer_embed_scale ,
537- loc = self .get_loc ("embed_tokens_per_layer.scale" ,
538- mlir_gen ),
539- ip = ip ).output
558+
559+ # Path A: per_layer_embeds (pre-computed by CPU from external bin file)
560+ # The bin file already has scale pre-multiplied, so no MulConstOp needed here.
561+ per_layer_inputs_op = per_layer_embeds_op
540562
541563 # Path B: per_layer_model_projection(embeds) * model_projection_scale (sliced for layer idx)
542564 proj_op = self .linear (mlir_gen , model_projection + f".{ idx } " , embeds_op ,
@@ -591,8 +613,10 @@ def gen_block_by_length(name, input_len):
591613 return_ops_list = []
592614
593615 if has_per_layer_input :
594- input_shapes .extend ([id_shape , input_shape ])
595- input_types .extend (["INT32" , "F32" ])
616+ embed_dtype = self .half_precision_quantize .upper ()
617+ per_layer_embeds_shape = [1 , input_len , self .hidden_size_per_layer_input ]
618+ input_shapes .extend ([per_layer_embeds_shape , input_shape ])
619+ input_types .extend ([embed_dtype , "F32" ])
596620
597621 if is_shared :
598622 # Shared layer: receives shared_k and shared_v as inputs, outputs only hidden_states
@@ -624,10 +648,10 @@ def L(name):
624648 in1_op = block_mlir .create_input_op (L ("position_ids" ), 1 )
625649 in2_op = block_mlir .create_input_op (L ("attention_mask" ), 2 )
626650 input_idx = 3
627- ids_op = None
628651 embeds_op = None
652+ per_layer_embeds_op = None
629653 if has_per_layer_input :
630- ids_op = block_mlir .create_input_op (L ("input_ids " ), input_idx )
654+ per_layer_embeds_op = block_mlir .create_input_op (L ("per_layer_embeds " ), input_idx )
631655 embeds_op = block_mlir .create_input_op (L ("inputs_embeds" ), input_idx + 1 )
632656 input_idx += 2
633657
@@ -715,8 +739,8 @@ def L(name):
715739
716740 # per_layer_input
717741 if has_per_layer_input :
718- new_op = gen_per_layer_input (block_mlir , input_shape , new_op , residual_mlp , ids_op ,
719- embeds_op )
742+ new_op = gen_per_layer_input (block_mlir , input_shape , new_op , residual_mlp ,
743+ embeds_op , per_layer_embeds_op )
720744
721745 # layer_scalar
722746 layer_scalar_loc = "layer_scalar" if do_norm else "output_states"
@@ -753,9 +777,10 @@ def gen_block_cache():
753777 input_shapes = [input_shape , id_shape , mask_shape ]
754778
755779 if has_per_layer_input :
756- input_ids_shape = [self .batch , 1 ]
757- input_shapes .extend ([input_ids_shape , input_shape ])
758- input_types .extend (["INT32" , "F32" ])
780+ embed_dtype = self .half_precision_quantize .upper ()
781+ per_layer_embeds_shape = [self .batch , 1 , self .hidden_size_per_layer_input ]
782+ input_shapes .extend ([per_layer_embeds_shape , input_shape ])
783+ input_types .extend ([embed_dtype , "F32" ])
759784
760785 if is_shared :
761786 # Shared KV layer in decode mode
@@ -789,10 +814,10 @@ def L(name):
789814 in1_op = block_mlir .create_input_op (L ("position_ids" ), 1 )
790815 in2_op = block_mlir .create_input_op (L ("attention_mask" ), 2 )
791816 input_idx = 3
792- ids_op = None
793817 embeds_op = None
818+ per_layer_embeds_op = None
794819 if has_per_layer_input :
795- ids_op = block_mlir .create_input_op (L ("input_ids " ), input_idx )
820+ per_layer_embeds_op = block_mlir .create_input_op (L ("per_layer_embeds " ), input_idx )
796821 embeds_op = block_mlir .create_input_op (L ("inputs_embeds" ), input_idx + 1 )
797822 input_idx += 2
798823
@@ -908,8 +933,8 @@ def L(name):
908933
909934 # per_layer_input
910935 if has_per_layer_input :
911- new_op = gen_per_layer_input (block_mlir , input_shape , new_op , residual_mlp , ids_op ,
912- embeds_op )
936+ new_op = gen_per_layer_input (block_mlir , input_shape , new_op , residual_mlp ,
937+ embeds_op , per_layer_embeds_op )
913938
914939 # layer_scalar
915940 layer_scalar_loc = "layer_scalar" if do_norm else "output_states"
@@ -947,8 +972,10 @@ def gen_block_with_kv():
947972 input_shapes = [input_shape , id_shape , mask_shape ]
948973
949974 if has_per_layer_input :
950- input_shapes .extend ([id_shape , input_shape ])
951- input_types .extend (["INT32" , "F32" ])
975+ embed_dtype = self .half_precision_quantize .upper ()
976+ per_layer_embeds_shape = [1 , input_len , self .hidden_size_per_layer_input ]
977+ input_shapes .extend ([per_layer_embeds_shape , input_shape ])
978+ input_types .extend ([embed_dtype , "F32" ])
952979
953980 if is_shared :
954981 # Shared KV layer: receives shared_k (full length), shared_v (full length)
@@ -981,10 +1008,10 @@ def L(name):
9811008 in1_op = block_mlir .create_input_op (L ("position_ids" ), 1 )
9821009 in2_op = block_mlir .create_input_op (L ("attention_mask" ), 2 )
9831010 input_idx = 3
984- ids_op = None
9851011 embeds_op = None
1012+ per_layer_embeds_op = None
9861013 if has_per_layer_input :
987- ids_op = block_mlir .create_input_op (L ("input_ids " ), input_idx )
1014+ per_layer_embeds_op = block_mlir .create_input_op (L ("per_layer_embeds " ), input_idx )
9881015 embeds_op = block_mlir .create_input_op (L ("inputs_embeds" ), input_idx + 1 )
9891016 input_idx += 2
9901017
@@ -1104,8 +1131,8 @@ def L(name):
11041131
11051132 # per_layer_input
11061133 if has_per_layer_input :
1107- new_op = gen_per_layer_input (block_mlir , input_shape , new_op , residual_mlp , ids_op ,
1108- embeds_op )
1134+ new_op = gen_per_layer_input (block_mlir , input_shape , new_op , residual_mlp ,
1135+ embeds_op , per_layer_embeds_op )
11091136
11101137 # layer_scalar
11111138 layer_scalar_loc = "layer_scalar" if do_norm else "output_states"
0 commit comments