Skip to content

Commit 08f9d41

Browse files
committed
feat: [gemma4] externalize per_layer_token_embd weights as bin file
- Export per_layer_token_embd as external bin to avoid 35x duplicate storage in bmodel - Pre-multiply per_layer_embed_scale into bin file, remove MulConstOp - Replace input_ids (INT32) input with pre-computed per_layer_embeds (BF16/F16) from CPU runtime - Fix loader not passed to base class __init__ - Fix tie_word_embeddings to read from config instead of hardcoded True - Change default audio_length from 200 to 750 - Add traceback print on gen_mlir failure in LlmConverter Change-Id: I4aae592f1c29b97840e70d6cc24eed61571ba183
1 parent ccc90aa commit 08f9d41

2 files changed

Lines changed: 76 additions & 47 deletions

File tree

python/llm/Gemma4Converter.py

Lines changed: 74 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -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"

python/llm/LlmConverter.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -341,10 +341,12 @@ def gen_all_mlir(self):
341341
# This will raise exceptions if any occurred during thread execution
342342
future.result()
343343
except Exception as e:
344+
import traceback
344345
for future in futures:
345346
if not future.done():
346347
future.cancel()
347348
logger.error("gen mlir failed: %s", e)
349+
traceback.print_exc()
348350
sys.exit(1)
349351

350352
def load_pretrained(self, config):

0 commit comments

Comments
 (0)