Skip to content

Commit a95a3da

Browse files
feat(peft): expand_shared_outer export flag for shared-outer MoE LoRA (#5588)
Signed-off-by: Hollow Man <hollowman@opensuse.org> Co-authored-by: Kamran Jafari <kjafarisadeg@nvidia.com>
1 parent f54f47b commit a95a3da

6 files changed

Lines changed: 361 additions & 6 deletions

File tree

examples/conversion/adapter/export_adapter.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -113,6 +113,15 @@ def parse_args() -> argparse.Namespace:
113113
"Can be specified multiple times; e.g. `mtp.layers` excludes MTP adapters."
114114
),
115115
)
116+
parser.add_argument(
117+
"--expand-shared-outer",
118+
action="store_true",
119+
help=(
120+
"For shared-outer MoE LoRA, replicate the shared factor across experts under "
121+
"per-expert 2D names (vLLM `pack_moe`) instead of a single shared `[1, ...]` tensor "
122+
"(SGLang). Multiplies the shared factor's on-disk size by the expert count."
123+
),
124+
)
116125
parser.add_argument("--tp", type=int, default=1, help="Tensor parallel size for distributed GPU export.")
117126
parser.add_argument("--pp", type=int, default=1, help="Pipeline parallel size for distributed GPU export.")
118127
parser.add_argument("--ep", type=int, default=1, help="Expert parallel size for distributed GPU export.")
@@ -231,6 +240,7 @@ def _export_adapter_distributed(args: argparse.Namespace) -> None:
231240
peft_config=lora,
232241
base_model_name_or_path=args.hf_model_path,
233242
exclude_adapter_base_prefixes=tuple(args.exclude_adapter_base_prefix),
243+
expand_shared_outer=args.expand_shared_outer,
234244
)
235245
finally:
236246
if parallel_state.is_initialized():
@@ -253,6 +263,7 @@ def main() -> None:
253263
peft_checkpoint=args.lora_checkpoint,
254264
output_path=args.output,
255265
exclude_adapter_base_prefixes=tuple(args.exclude_adapter_base_prefix),
266+
expand_shared_outer=args.expand_shared_outer,
256267
)
257268

258269

src/megatron/bridge/models/conversion/auto_bridge.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -858,6 +858,7 @@ def export_adapter_weights(
858858
cpu: bool = True,
859859
show_progress: bool = True,
860860
exclude_adapter_base_prefixes: Iterable[str] | None = None,
861+
expand_shared_outer: bool = False,
861862
) -> Iterable["HFWeightTuple"]:
862863
"""
863864
Export only adapter weights from a Megatron model without merging them into base tensors.
@@ -871,16 +872,26 @@ def export_adapter_weights(
871872
show_progress: Display progress bar during export
872873
exclude_adapter_base_prefixes: Megatron adapter base prefixes to
873874
skip before resolving HuggingFace parameter mappings.
875+
expand_shared_outer: Replicate the shared factor across experts under per-expert
876+
names (vLLM 2D ``pack_moe``) instead of a shared ``[1, ...]`` tensor (SGLang).
877+
Default ``False``; no effect for non-shared-outer adapters.
874878
875879
Yields:
876880
HFWeightTuple: Named tuples of (param_name, weight_tensor) for adapter parameters
881+
882+
Note:
883+
With ``expand_shared_outer``, the per-expert copies of the shared factor alias one
884+
storage rather than being cloned. ``safetensors.torch.save_file`` rejects tensors
885+
that share memory, so callers serializing these tensors directly must clone them
886+
first — :meth:`save_hf_adapter` already does.
877887
"""
878888
bridge = self._model_bridge
879889
return bridge.stream_adapter_weights_megatron_to_hf(
880890
model,
881891
cpu=cpu,
882892
show_progress=show_progress,
883893
exclude_adapter_base_prefixes=exclude_adapter_base_prefixes,
894+
expand_shared_outer=expand_shared_outer,
884895
)
885896

886897
def save_hf_adapter(
@@ -891,6 +902,7 @@ def save_hf_adapter(
891902
base_model_name_or_path: Optional[str] = None,
892903
show_progress: bool = True,
893904
exclude_adapter_base_prefixes: Iterable[str] | None = None,
905+
expand_shared_outer: bool = False,
894906
) -> None:
895907
"""Save LoRA adapter weights as a HuggingFace PEFT-compatible directory.
896908
@@ -909,6 +921,8 @@ def save_hf_adapter(
909921
show_progress: Display progress bar during export.
910922
exclude_adapter_base_prefixes: Megatron adapter base prefixes to
911923
skip before resolving HuggingFace parameter mappings.
924+
expand_shared_outer: Replicate the shared factor across experts under per-expert
925+
names (vLLM 2D ``pack_moe``). Default ``False`` keeps the PEFT shared ``[1, ...]`` layout.
912926
913927
Example:
914928
>>> bridge.save_hf_adapter(
@@ -949,6 +963,7 @@ def save_hf_adapter(
949963
cpu=True,
950964
show_progress=show_progress,
951965
exclude_adapter_base_prefixes=exclude_adapter_base_prefixes,
966+
expand_shared_outer=expand_shared_outer,
952967
)
953968
]
954969
if not raw_adapter_weights:
@@ -1561,6 +1576,7 @@ def export_adapter_ckpt(
15611576
output_path: str | Path,
15621577
show_progress: bool = True,
15631578
exclude_adapter_base_prefixes: Iterable[str] | None = None,
1579+
expand_shared_outer: bool = False,
15641580
) -> None:
15651581
"""Export LoRA adapter weights from a Megatron PEFT checkpoint to HuggingFace PEFT format.
15661582
@@ -1582,6 +1598,10 @@ def export_adapter_ckpt(
15821598
show_progress: Display progress bar during export.
15831599
exclude_adapter_base_prefixes: Megatron adapter base prefixes to
15841600
skip before resolving HuggingFace parameter mappings.
1601+
expand_shared_outer: Replicate the shared factor across experts under per-expert
1602+
names (vLLM 2D ``pack_moe``). Default ``False`` keeps the PEFT shared ``[1, ...]``
1603+
layout. Enabling it multiplies the shared factor's on-disk size by the expert
1604+
count; no effect for non-shared-outer adapters.
15851605
15861606
Example:
15871607
>>> bridge = AutoBridge.from_hf_pretrained("meta-llama/Llama-3.2-1B")
@@ -1689,6 +1709,7 @@ def _load_and_export_adapter(model):
16891709
base_model_name_or_path=base_model_name,
16901710
show_progress=show_progress,
16911711
exclude_adapter_base_prefixes=exclude_adapter_base_prefixes,
1712+
expand_shared_outer=expand_shared_outer,
16921713
)
16931714

16941715
model_context = (

src/megatron/bridge/models/conversion/model_bridge.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2370,6 +2370,7 @@ def stream_adapter_weights_megatron_to_hf(
23702370
cpu: bool = True,
23712371
show_progress: bool = True,
23722372
exclude_adapter_base_prefixes: Optional[Iterable[str]] = None,
2373+
expand_shared_outer: bool = False,
23732374
) -> Iterable[HFWeightTuple]:
23742375
"""Bridge only adapter weights from Megatron to HuggingFace format."""
23752376
...
@@ -2462,13 +2463,15 @@ def _adapter_stream_registered_impl(
24622463
cpu: bool = True,
24632464
show_progress: bool = True,
24642465
exclude_adapter_base_prefixes: Optional[Iterable[str]] = None,
2466+
expand_shared_outer: bool = False,
24652467
) -> Iterable[HFWeightTuple]:
24662468
bridge = bridge_class()
24672469
return bridge.stream_adapter_weights_megatron_to_hf(
24682470
megatron_model,
24692471
cpu=cpu,
24702472
show_progress=show_progress,
24712473
exclude_adapter_base_prefixes=exclude_adapter_base_prefixes,
2474+
expand_shared_outer=expand_shared_outer,
24722475
)
24732476

24742477
# Set meaningful names for debugging

src/megatron/bridge/models/conversion/peft_bridge.py

Lines changed: 27 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -853,6 +853,7 @@ def stream_adapter_weights_megatron_to_hf(
853853
cpu: bool = True,
854854
show_progress: bool = True,
855855
exclude_adapter_base_prefixes: Iterable[str] | None = None,
856+
expand_shared_outer: bool = False,
856857
) -> Iterable["HFWeightTuple"]:
857858
"""Stream only adapter weights without merging them into base tensors.
858859
@@ -895,6 +896,7 @@ def stream_adapter_weights_megatron_to_hf(
895896
linear_out_tensor,
896897
num_moe_experts,
897898
cpu,
899+
expand_shared_outer=expand_shared_outer,
898900
)
899901
continue
900902

@@ -1034,14 +1036,15 @@ def _stream_shared_outer_adapter_weights(
10341036
linear_out_tensor: torch.Tensor,
10351037
num_moe_experts: int,
10361038
cpu: bool,
1039+
expand_shared_outer: bool,
10371040
) -> Iterable["HFWeightTuple"]:
10381041
"""Stream a shared-outer grouped-expert LoRA adapter (SGLang PR #21466).
10391042
1040-
One side is a 2D LoRA matrix replicated across local experts; the other
1041-
is a per-expert 3D pack. The shared side is emitted once as a ``[1, ...]``
1042-
tensor under the expert-agnostic HF name (so the serving loader takes its
1043-
3D-shared branch); the per-expert side is gathered across EP ranks and
1044-
emitted once per global expert.
1043+
One side is a 2D LoRA matrix shared across experts; the other is a
1044+
per-expert 3D pack. By default the shared side is emitted once as a
1045+
``[1, ...]`` tensor under the expert-agnostic HF name. With
1046+
``expand_shared_outer``, it is replicated under per-expert 2D names
1047+
(vLLM 2D ``pack_moe`` contract); the training-side parameter stays shared.
10451048
"""
10461049

10471050
from megatron.bridge.models.conversion.model_bridge import HFWeightTuple
@@ -1051,7 +1054,7 @@ def _stream_shared_outer_adapter_weights(
10511054
(linear_in_tensor, ".linear_in.weight"),
10521055
(linear_out_tensor, ".linear_out.weight"),
10531056
):
1054-
if side_tensor.ndim == 2:
1057+
if side_tensor.ndim == 2 and not expand_shared_outer:
10551058
# Shared side: emit one [1, out, in] tensor. A shared linear_in
10561059
# feeding a fused gate/up FC1 maps to two HF names, so the same
10571060
# tensor is emitted for each projection.
@@ -1066,6 +1069,24 @@ def _stream_shared_outer_adapter_weights(
10661069
yield HFWeightTuple(hf_name, current)
10671070
continue
10681071

1072+
if side_tensor.ndim == 2 and expand_shared_outer:
1073+
# Expand the shared factor under per-expert 2D names (vLLM pack_moe).
1074+
# The tensor is reused across experts, not cloned.
1075+
shared_current = side_tensor.cpu() if cpu else side_tensor
1076+
for expert_idx in range(num_moe_experts):
1077+
base_hf_weight_names = self._get_base_hf_param_names_for_adapter(
1078+
mapping_registry,
1079+
adapter_task.global_base_prefix,
1080+
adapter_task.adapter_key,
1081+
f".weight{expert_idx}",
1082+
)
1083+
for base_name in base_hf_weight_names:
1084+
hf_name = self._make_lora_param_name(base_name, side_suffix)
1085+
if hf_name is None:
1086+
continue
1087+
yield HFWeightTuple(hf_name, shared_current)
1088+
continue
1089+
10691090
# Per-expert side: emit one slice per global expert. A fused FC1
10701091
# linear_out (gate+up) is split per HF projection name; otherwise the
10711092
# single projection is emitted directly.

0 commit comments

Comments
 (0)