Skip to content

Commit b9bbc5f

Browse files
Merge pull request #4050 from AI-Hypercomputer:jackyf/lora-ckpt-converter
PiperOrigin-RevId: 938170156
2 parents df07bd6 + 81c0c16 commit b9bbc5f

4 files changed

Lines changed: 336 additions & 17 deletions

File tree

src/maxtext/checkpoint_conversion/to_huggingface.py

Lines changed: 53 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -116,12 +116,35 @@ def _get_lora_delta(key, lora_state_dict, lora_scaling):
116116
a_key, b_key = key[7:] + "_lora_a", key[7:] + "_lora_b"
117117

118118
if a_key in lora_state_dict and b_key in lora_state_dict:
119-
data_a, data_b = jnp.asarray(lora_state_dict[a_key], dtype=jnp.float32), jnp.asarray(
120-
lora_state_dict[b_key], dtype=jnp.float32
121-
)
122-
if data_a.ndim > 2:
119+
data_a = jnp.asarray(lora_state_dict[a_key], dtype=jnp.float32)
120+
data_b = jnp.asarray(lora_state_dict[b_key], dtype=jnp.float32)
121+
122+
is_attention = "attention" in key.lower() or "attn" in key.lower()
123+
124+
if is_attention and data_a.ndim > 2:
125+
if data_a.ndim == 4:
126+
# Scanned attention projection: [num_layers, input_dim, heads, rank] & [num_layers, rank, heads, output_dim]
127+
return jnp.einsum("lipr,lrpo->lipo", data_a, data_b) * lora_scaling
128+
# Unscanned attention projection: [input_dim, heads, rank] & [rank, heads, output_dim]
123129
return jnp.einsum("ipr,rpo->ipo", data_a, data_b) * lora_scaling
124-
return jnp.matmul(data_a, data_b) * lora_scaling
130+
else:
131+
if data_a.ndim == 3:
132+
# Scanned standard linear projection: can be [num_layers, input_dim, rank] or [input_dim, num_layers, rank]
133+
rank = data_a.shape[2]
134+
if rank == data_b.shape[1] and rank != data_b.shape[0]:
135+
# Case A: [num_layers, input_dim, rank] & [num_layers, rank, output_dim]
136+
return jnp.einsum("lir,lro->lio", data_a, data_b) * lora_scaling
137+
elif rank == data_b.shape[0] and rank != data_b.shape[1]:
138+
# Case B: [input_dim, num_layers, rank] & [rank, num_layers, output_dim]
139+
return jnp.einsum("ilr,rlo->ilo", data_a, data_b) * lora_scaling
140+
else:
141+
# Disambiguate using key names (Case B is typically 'wo' or 'out-kernel' / 'out_proj')
142+
if any(term in key for term in ["wo", "out-kernel", "out_proj"]):
143+
return jnp.einsum("ilr,rlo->ilo", data_a, data_b) * lora_scaling
144+
else:
145+
return jnp.einsum("lir,lro->lio", data_a, data_b) * lora_scaling
146+
# Unscanned standard linear projection
147+
return jnp.matmul(data_a, data_b) * lora_scaling
125148
return None
126149

127150

@@ -312,19 +335,38 @@ def _transform_weights_to_adapter(param_map, state_dict):
312335
if a_key in state_dict and b_key in state_dict:
313336
data_a, data_b = state_dict[a_key], state_dict[b_key]
314337
hf_paths = [hf_paths] if not isinstance(hf_paths, list) else hf_paths
315-
for i in range(min(data_a.shape[1] if data_a.ndim > 2 else 1, len(hf_paths))):
316-
found_hf_modules.add(hf_paths[i].split(".")[-2])
317-
name = hf_paths[i].replace(".weight", "")
338+
for i, hf_path in enumerate(hf_paths):
339+
found_hf_modules.add(hf_path.split(".")[-2])
340+
name = hf_path.replace(".weight", "")
341+
342+
if data_a.ndim > 2:
343+
if data_a.shape[0] == len(hf_paths):
344+
# Case A: layer dimension is axis 0
345+
layer_a = data_a[i, ...]
346+
layer_b = data_b[i, ...]
347+
else:
348+
# Case B: layer dimension is axis 1
349+
layer_a = data_a[:, i, ...]
350+
layer_b = data_b[:, i, ...]
351+
else:
352+
layer_a = data_a
353+
layer_b = data_b
354+
355+
if layer_a.ndim > 2:
356+
layer_a = layer_a[:, 0, :]
357+
if layer_b.ndim > 2:
358+
layer_b = layer_b[:, 0, :]
359+
318360
processed_params_list.append(
319361
(
320362
f"base_model.model.{name}.lora_A.weight",
321-
jax.numpy.asarray((data_a[:, i, :] if data_a.ndim > 2 else data_a).T),
363+
jax.numpy.asarray(layer_a.T),
322364
)
323365
)
324366
processed_params_list.append(
325367
(
326368
f"base_model.model.{name}.lora_B.weight",
327-
jax.numpy.asarray((data_b[:, i, :] if data_b.ndim > 2 else data_b).T),
369+
jax.numpy.asarray(layer_b.T),
328370
)
329371
)
330372
return dict(processed_params_list), found_hf_modules
@@ -454,9 +496,7 @@ def main(argv: Sequence[str]) -> None:
454496
maxtext_state_dict = detect_and_extract_checkpoint(checkpoint_dict)
455497

456498
# Validate that checkpoint keys match the parameter mapping
457-
state_keys = set(maxtext_state_dict) | {
458-
k.replace("_lora_a", "").replace("_lora_b", "") for k in maxtext_state_dict if "_lora_" in k
459-
}
499+
state_keys = {k.replace("_lora_a", "").replace("_lora_b", "") for k in maxtext_state_dict}
460500
filtered_map_keys = validate_and_filter_param_map_keys(param_map, state_keys)
461501

462502
# When not converting a multimodal model, skip vision encoder weights even if

src/maxtext/checkpoint_conversion/utils/utils.py

Lines changed: 18 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -817,6 +817,16 @@ def format_meter(
817817
return super().format_meter(n=n, total=total, elapsed=elapsed, postfix=postfix, **extra_kwargs)
818818

819819

820+
def _recursive_update(d: dict, u: dict) -> dict:
821+
"""Recursively updates dictionary d with dictionary u in place."""
822+
for k, v in u.items():
823+
if isinstance(v, dict) and isinstance(d.get(k), dict):
824+
_recursive_update(d[k], v)
825+
else:
826+
d[k] = v
827+
return d
828+
829+
820830
def load_orbax_checkpoint(config) -> dict:
821831
"""Loads Orbax checkpoints from Base and/or LoRA paths in config.
822832
@@ -852,15 +862,21 @@ def create_restore_args(tree_metadata):
852862
paths = [p for p in [config.load_parameters_path, lora_path] if p]
853863

854864
merged_dict = {}
855-
for path in paths:
865+
for i, path in enumerate(paths):
856866
checkpoint_path = epath.Path(path)
857867
metadata = ckptr.metadata(checkpoint_path)
858868
restore_args = jax.tree_util.tree_map(
859869
lambda x: create_restore_args(x) if hasattr(x, "shape") else None,
860870
metadata.item_metadata.tree,
861871
is_leaf=lambda x: hasattr(x, "shape"),
862872
)
863-
merged_dict.update(ckptr.restore(checkpoint_path, restore_args=restore_args))
873+
restored = ckptr.restore(checkpoint_path, restore_args=restore_args)
874+
875+
if i == 0:
876+
merged_dict = restored
877+
else:
878+
# Recursively update base checkpoint with LoRA adapter checkpoint keys to avoid overwriting
879+
_recursive_update(merged_dict, restored)
864880

865881
return merged_dict
866882

src/maxtext/configs/post_train/lora_module_path.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -19,8 +19,8 @@ llama3.1: "decoder/layers/.*(self_attention/(query|key|value|out)|mlp/(wi_0|wi_1
1919
qwen3: "decoder/layers/self_attention/(query|key|value|out)|decoder/layers/mlp/(wi_0|wi_1|wo)"
2020
mistral: "decoder/layers/.*(attention/(query|key|value|out)|mlp/(wi_0|wi_1|wo))"
2121
deepseek2: "decoder/(dense_layers|moe_stack)/self_attention/(query|out|wkv_a|wkv_b)|decoder/(dense_layers|moe_stack)/(mlp|shared_experts)/(wi_0|wi_1|wo)"
22-
gemma2: "decoder/layers/(self_attention_local|self_attention_global)/(query|key|value|out)|decoder/layers/(mlp_local|mlp_global)/(wi_0|wi_1|wo)"
23-
gemma3: "decoder/layers/.*(self_attention/(query|key|value|out)|mlp/(wi_0|wi_1|wo|gate|up|down))"
22+
gemma2: "decoder/(scanned_blocks|layers_remainder|layers)/(self_attention_local|self_attention_global)/(query|key|value|out)|decoder/(scanned_blocks|layers_remainder|layers)/(mlp_local|mlp_global)/(wi_0|wi_1|wo)"
23+
gemma3: "decoder/(scanned_blocks|layers_remainder|layers)/.*(self_attention/(query|key|value|out)|mlp/(wi_0|wi_1|wo|gate|up|down))"
2424
gemma4: "decoder/(scanned_blocks|layers_remainder)/layers.*/.*(self_attention/(query|key|value|out)|mlp/.*(wi_0|wi_1|wo|shared_experts/(wi_0|wi_1|wo)))"
2525
olmo3: "decoder/layers/.*(attention/(query|key|value|out)|mlp/(wi_0|wi_1|wo))"
2626
gpt3: "decoder/layers/(self_attention/(qkv_proj|out)|mlp/(wi|wo))"

0 commit comments

Comments
 (0)