diff --git a/docs/advanced/lora.md b/docs/advanced/lora.md index a91e32afe3b..654e1c17e9b 100644 --- a/docs/advanced/lora.md +++ b/docs/advanced/lora.md @@ -182,7 +182,7 @@ alternative aligned-expert path. Bridge mode. - **Checkpoints.** miles saves native per-rank adapter shards and optimizer/scheduler state. Exact resume expects the same TP/PP topology. It - also attempts a best-effort HF PEFT `adapter_model.bin` plus + also attempts a best-effort HF PEFT `adapter_model.safetensors` plus `adapter_config.json` export for external serving and warns if that export fails. Direct HF PEFT-to-Bridge resume is not implemented yet; native Inkling supplies a model-specific HF adapter importer. diff --git a/miles/backends/megatron_utils/lora_utils.py b/miles/backends/megatron_utils/lora_utils.py index 7fe3804e7a7..8cfcf14a3da 100644 --- a/miles/backends/megatron_utils/lora_utils.py +++ b/miles/backends/megatron_utils/lora_utils.py @@ -4,13 +4,16 @@ import os from argparse import Namespace from collections.abc import Sequence +from contextlib import ExitStack from pathlib import Path +from tempfile import TemporaryDirectory from typing import Any import torch import torch.distributed as dist from miles.backends.training_utils.parallel import get_parallel_state +from miles.utils.distributed_utils import get_gloo_group from miles.utils.lora import is_lora_enabled, lora_rollout_enabled # noqa: F401 (re-exported) logger = logging.getLogger(__name__) @@ -416,9 +419,8 @@ def save_lora_checkpoint( """Save LoRA adapter checkpoint to disk. Saves in two formats: - 1. **HF PEFT format** (``adapter_model.bin`` + ``adapter_config.json``) for - external tool compatibility. Uses Megatron-Bridge's ``export_adapter_weights`` - which correctly handles fused QKV / gate-up weight splitting and TP gathering. + 1. **HF PEFT format** (``adapter_model.safetensors`` + ``adapter_config.json``) + through Megatron-Bridge. 2. **Megatron-native format** (``adapter_megatron_rank{global_rank}.pt``) for fast checkpoint resume without name/weight conversion. Each TP/PP rank saves its own shard with original parameter names. @@ -427,20 +429,14 @@ def save_lora_checkpoint( also saved per-rank for checkpoint resume. Base model weights are frozen and never change, so they are not saved. - This function is collective: **all ranks must call it** because the bridge - export performs TP all-gather internally. Only ``dp_rank == 0`` writes files. + This function is collective: every rank writes its native shard and participates + in the Bridge export; Bridge rank 0 writes the HF files. """ - import json - from megatron.bridge import AutoBridge from miles.utils import megatron_bridge_utils save_path = Path(save_dir) - parallel_state = get_parallel_state() - is_dp_cp_rank_0 = parallel_state.effective_dp.rank == 0 and parallel_state.cp.rank == 0 - tp_rank = parallel_state.tp.rank - pp_rank = parallel_state.pp.rank save_path.mkdir(parents=True, exist_ok=True) if dist.is_initialized(): @@ -457,44 +453,77 @@ def save_lora_checkpoint( torch.save(adapter_state, native_path) logger.info(f"Saved {len(adapter_state)} adapter tensors (native) to {native_path}") - # ---- HF PEFT format (uses bridge for correct name/weight conversion) ---- - # Bridge export is collective: all TP ranks participate in the all-gather, - # so every rank must call export_adapter_weights. + # ---- HF PEFT format ---- + hf_export_err: Exception | None = None try: - bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True) - - lora_state_dict: dict[str, torch.Tensor] = {} - with megatron_bridge_utils.patch_megatron_model(model): - for hf_name, weight, _megatron_name in bridge.export_adapter_weights( - model, - cpu=True, - show_progress=False, - ): - lora_state_dict[hf_name] = weight - - if is_dp_cp_rank_0 and tp_rank == 0 and pp_rank == 0: - torch.save(lora_state_dict, save_path / "adapter_model.bin") - - target_modules_hf = ( - convert_target_modules_to_hf(list(args.target_modules)) - if args.target_modules - else ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"] - ) - config = { - "peft_type": "LORA", - "r": args.lora_rank, - "lora_alpha": args.lora_alpha, - "target_modules": target_modules_hf, - "lora_dropout": args.lora_dropout, - "bias": "none", - "task_type": "CAUSAL_LM", - } - with open(save_path / "adapter_config.json", "w") as f: - json.dump(config, f, indent=2) - - os.sync() - logger.info(f"Saved HF PEFT adapter to {save_path} with {len(lora_state_dict)} tensors") - except Exception as hf_export_err: + with ExitStack() as stack: + peft_export_path = None + try: + # Staging must be set up inside the guarded block: a rank that fails here + # would otherwise skip the consensus below and hang its peers. + peft_export_path = Path(stack.enter_context(TemporaryDirectory(prefix=".peft-export-", dir=save_path))) + bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True) + peft_config = create_lora_instance(args) + stack.enter_context(megatron_bridge_utils.patch_megatron_model(model)) + except Exception as setup_err: + hf_export_err = setup_err + + if dist.is_initialized(): + # No rank may enter Bridge's initial barrier unless every rank can make the call. + # Report over gloo: the failure may be a poisoned GPU communicator. + group = get_gloo_group() + setup_errors: list[str | None] = [None] * dist.get_world_size(group=group) + dist.all_gather_object( + setup_errors, + repr(hf_export_err) if hf_export_err is not None else None, + group=group, + ) + first_setup_err = next((error for error in setup_errors if error is not None), None) + if hf_export_err is None and first_setup_err is not None: + hf_export_err = RuntimeError(f"HF PEFT export setup failed on another rank: {first_setup_err}") + + if hf_export_err is None: + try: + bridge.save_hf_adapter( + model, + path=peft_export_path, + peft_config=peft_config, + base_model_name_or_path=args.hf_checkpoint, + show_progress=False, + ) + except Exception as save_err: + hf_export_err = save_err + if dist.is_initialized(): + # Match Bridge's final barrier on ranks where the collective save did not raise. + dist.barrier() + else: + try: + if not dist.is_initialized() or dist.get_rank() == 0: + os.sync() + (peft_export_path / "adapter_model.safetensors").replace( + save_path / "adapter_model.safetensors" + ) + (peft_export_path / "adapter_config.json").replace(save_path / "adapter_config.json") + except Exception as promotion_err: + hf_export_err = promotion_err + if dist.is_initialized(): + group = get_gloo_group() + save_errors: list[str | None] = [None] * dist.get_world_size(group=group) + dist.all_gather_object( + save_errors, + repr(hf_export_err) if hf_export_err is not None else None, + group=group, + ) + first_save_err = next((error for error in save_errors if error is not None), None) + if hf_export_err is None and first_save_err is not None: + hf_export_err = RuntimeError(f"HF PEFT export failed on another rank: {first_save_err}") + except Exception as cleanup_err: + if hf_export_err is None: + hf_export_err = cleanup_err + + if hf_export_err is None: + logger.info(f"Saved HF PEFT adapter to {save_path}") + else: logger.warning( f"HF PEFT adapter export skipped ({hf_export_err}); the per-rank native " f"shards + training state are sufficient for training resume." @@ -530,7 +559,7 @@ def load_lora_adapter( Attempts to load from Megatron-native format first (per-rank ``.pt`` files), which preserves the exact TP/PP sharding and requires no name conversion. - Falls back to HF PEFT ``adapter_model.bin`` if native files are not found + Falls back to HF PEFT ``adapter_model.safetensors``/``.bin`` if native files are not found (not yet implemented for HF PEFT format). When ``optimizer`` is provided, also restores training state (optimizer + @@ -577,8 +606,15 @@ def load_lora_adapter( return True, iteration # ---- HF PEFT format (future work) ---- - hf_path = adapter_dir / "adapter_model.bin" - if hf_path.exists(): + hf_path = next( + ( + path + for path in (adapter_dir / "adapter_model.safetensors", adapter_dir / "adapter_model.bin") + if path.exists() + ), + None, + ) + if hf_path is not None: logger.warning( f"Found HF PEFT adapter at {hf_path} but direct HF PEFT loading into " f"Megatron is not yet supported. Please save using Megatron-native format " diff --git a/miles/backends/megatron_utils/multi_lora_utils.py b/miles/backends/megatron_utils/multi_lora_utils.py index 4d6d7689098..857bfa2b1ed 100644 --- a/miles/backends/megatron_utils/multi_lora_utils.py +++ b/miles/backends/megatron_utils/multi_lora_utils.py @@ -20,6 +20,21 @@ _shard_topology: tuple[bool, tuple[tuple[int, int, int], ...]] | None = None +def _raise_if_any_rank_failed(local_error: Exception | None, operation: str) -> None: + message = None if local_error is None else f"{type(local_error).__name__}: {local_error}" + if dist.is_initialized(): + group = get_gloo_group() + messages: list[str | None] = [None] * dist.get_world_size(group=group) + dist.all_gather_object(messages, message, group=group) + message = next((item for item in messages if item is not None), None) + + if message is not None: + error = RuntimeError(f"{operation} failed on at least one rank: {message}") + if local_error is not None: + raise error from local_error + raise error + + def create_multi_lora_instance(args: Namespace): """Create a MultiLoRA instance from training args.""" from megatron.bridge.peft.multi_lora import MultiLoRA @@ -181,6 +196,43 @@ def slice_lora_to_rank(hf_name: str, tensor: torch.Tensor, adapter_rank: int) -> return tensor +def _build_multi_lora_peft_export( + adapter_weights, + *, + rank: int, + alpha: int, + dropout: float, + base_model_name_or_path: str, +) -> tuple[dict[str, torch.Tensor], dict[str, object]]: + """Slice Bridge adapter weights to this adapter's rank, then convert to PEFT state + config.""" + from megatron.bridge.models.conversion.peft_bridge import ( + build_adapter_config_dict, + convert_adapter_weights_to_peft_state, + infer_rank_pattern_from_adapter_weights, + infer_target_modules_from_adapter_weights, + ) + + # fp32 like Bridge's own save_hf_adapter, so both export paths write the same dtype. + sliced_weights = [ + item._replace(weight=slice_lora_to_rank(item.param_name, item.weight, rank).clone().float()) + for item in adapter_weights + ] + if not sliced_weights: + # Bridge guards its own export path; this one must too, or an empty + # safetensors plus target_modules: [] gets promoted as a complete checkpoint. + raise RuntimeError("No adapter weights were exported; ensure the adapter slot is exposed before export.") + + state_dict, module_weight_names, target_parameters = convert_adapter_weights_to_peft_state(sliced_weights) + config = build_adapter_config_dict( + Namespace(dim=rank, alpha=alpha, dropout=dropout), + target_modules=infer_target_modules_from_adapter_weights(module_weight_names), + target_parameters=target_parameters, + base_model_name_or_path=base_model_name_or_path, + rank_pattern=infer_rank_pattern_from_adapter_weights(sliced_weights, default_rank=rank), + ) + return state_dict, config + + def save_multi_lora_checkpoints( args, model, @@ -200,7 +252,6 @@ def save_multi_lora_checkpoints( from megatron.bridge.peft.multi_lora_layers import expose_adapter_slot from safetensors.torch import save_file as save_safetensors - from miles.backends.megatron_utils.lora_utils import convert_target_modules_to_hf from miles.utils import megatron_bridge_utils parallel_state = get_parallel_state() @@ -212,13 +263,13 @@ def save_multi_lora_checkpoints( is_shard_writer, _ = adapter_shard_topology() is_global_writer = is_shard_writer and tp_rank == 0 and pp_rank == 0 and ep_rank == 0 - target_modules_hf = ( - convert_target_modules_to_hf(list(args.target_modules)) - if args.target_modules - else ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"] - ) - - bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True) + bridge = None + setup_error = None + try: + bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True) + except Exception as error: + setup_error = error + _raise_if_any_rank_failed(setup_error, "Multi-LoRA checkpoint export setup") for adapter_name, adapter in adapters.items(): config = adapter.config @@ -231,69 +282,73 @@ def save_multi_lora_checkpoints( final_dir = config.save / "checkpoints" / f"step_{iteration}" tmp_dir = config.save / "checkpoints" / f"_tmp_step_{iteration}" - if is_shard_writer: - tmp_dir.mkdir(parents=True, exist_ok=True) - if dist.is_initialized(): - dist.barrier() - - with expose_adapter_slot(model, adapter.slot): - # Megatron checkpoints + native_error = None + try: if is_shard_writer: - shard: dict[str, torch.Tensor] = { - name: param.data.cpu() - for batch in model - for name, param in batch.named_parameters() - if ".adapter." in name - } + tmp_dir.mkdir(parents=True, exist_ok=True) + # The slot must stay exposed while walking named_parameters(): outside the + # context the parameters are named ``.adapters.{slot}.`` and the shard would + # be empty, so resume would silently restart from a fresh adapter. + with expose_adapter_slot(model, adapter.slot): + shard: dict[str, torch.Tensor] = { + name: param.data.cpu() + for batch in model + for name, param in batch.named_parameters() + if ".adapter." in name + } native_path = tmp_dir / megatron_shard_name(tp_rank, pp_rank, ep_rank, ep_size) torch.save(shard, native_path) logger.info(f"{log_prefix} saved Megatron shard " f"({len(shard)} tensors) to {native_path}") - - hf_state: dict[str, torch.Tensor] = {} - with megatron_bridge_utils.patch_megatron_model(model): - for hf_name, weight, _megatron_name in bridge.export_adapter_weights( - model, - cpu=True, - show_progress=False, - ): - # Slice from the shared --lora-rank down to this adapter's real rank to - # match adapter_config's r; clone() since safetensors rejects aliased views. - hf_state[hf_name] = slice_lora_to_rank(hf_name, weight, config.rank).clone() - - if is_global_writer: - save_safetensors( - hf_state, - str(tmp_dir / "adapter_model.safetensors"), - metadata={"format": "pt"}, - ) - adapter_config_json = { - "peft_type": "LORA", - "r": config.rank, - "lora_alpha": config.alpha, - "target_modules": target_modules_hf, - "lora_dropout": getattr(args, "lora_dropout", 0.0), - "bias": "none", - "task_type": "CAUSAL_LM", - } - with open(tmp_dir / "adapter_config.json", "w") as f: - json.dump(adapter_config_json, f, indent=2) - os.sync() - logger.info(f"{log_prefix} saved HF PEFT to {tmp_dir} " f"({len(hf_state)} tensors)") - - if dist.is_initialized(): - dist.barrier() + except Exception as error: + native_error = error + _raise_if_any_rank_failed(native_error, f"{log_prefix} native checkpoint save") + + hf_state = None + adapter_config_json = None + export_error = None + try: + with expose_adapter_slot(model, adapter.slot), megatron_bridge_utils.patch_megatron_model(model): + hf_state, adapter_config_json = _build_multi_lora_peft_export( + bridge.export_adapter_weights(model, cpu=True, show_progress=False), + rank=config.rank, + alpha=config.alpha, + dropout=getattr(args, "lora_dropout", 0.0), + base_model_name_or_path=args.hf_checkpoint, + ) + except Exception as error: + export_error = error + _raise_if_any_rank_failed(export_error, f"{log_prefix} PEFT conversion") + + peft_write_error = None + try: + if is_global_writer: + save_safetensors( + hf_state, + str(tmp_dir / "adapter_model.safetensors"), + metadata={"format": "pt"}, + ) + with open(tmp_dir / "adapter_config.json", "w") as f: + json.dump(adapter_config_json, f, indent=2) + os.sync() + logger.info(f"{log_prefix} saved HF PEFT to {tmp_dir} " f"({len(hf_state)} tensors)") + except Exception as error: + peft_write_error = error + _raise_if_any_rank_failed(peft_write_error, f"{log_prefix} PEFT checkpoint write") # Write to a temp dir and move into place so readers never see a # partially written checkpoint. - if is_global_writer: - if final_dir.exists(): - import shutil - - shutil.rmtree(final_dir) - os.replace(tmp_dir, final_dir) - logger.info(f"{log_prefix} promoted checkpoint to {final_dir}") - if dist.is_initialized(): - dist.barrier() + promotion_error = None + try: + if is_global_writer: + if final_dir.exists(): + import shutil + + shutil.rmtree(final_dir) + os.replace(tmp_dir, final_dir) + logger.info(f"{log_prefix} promoted checkpoint to {final_dir}") + except Exception as error: + promotion_error = error + _raise_if_any_rank_failed(promotion_error, f"{log_prefix} checkpoint promotion") def _register_adapter(adapter: AdapterRun, model) -> int: diff --git a/tests/e2e/lora/test_lora_qwen2.5_0.5B.py b/tests/e2e/lora/test_lora_qwen2.5_0.5B.py index 1619e20237f..e77dd465a5a 100644 --- a/tests/e2e/lora/test_lora_qwen2.5_0.5B.py +++ b/tests/e2e/lora/test_lora_qwen2.5_0.5B.py @@ -11,8 +11,12 @@ Triggered by label: run-ci-lora """ +import glob +import json import os +import torch + from tests.ci.ci_register import register_cuda_ci, register_rocm_ci import miles.utils.external_utils.command_utils as U @@ -26,12 +30,39 @@ MODEL_NAME = "Qwen2.5-0.5B-Instruct" MODEL_TYPE = "qwen2.5-0.5B" NUM_GPUS = 4 +SAVE_DIR = "/root/checkpoints/lora-qwen2.5-0.5B-ci" def prepare(): U.exec_command_cpu("mkdir -p /root/models /root/datasets") U.exec_command_cpu(f"hf download Qwen/{MODEL_NAME} --local-dir /root/models/{MODEL_NAME}") U.exec_command_cpu("hf download --repo-type dataset zhuzilin/gsm8k --local-dir /root/datasets/gsm8k") + U.exec_command_cpu(f"rm -rf {SAVE_DIR}") + + +def _assert_peft_export(): + from peft import PeftModel + from safetensors.torch import load_file + from transformers import AutoModelForCausalLM + + adapter_dirs = sorted(glob.glob(f"{SAVE_DIR}/iter_*/adapter")) + assert adapter_dirs + adapter_dir = adapter_dirs[-1] + on_disk = load_file(f"{adapter_dir}/adapter_model.safetensors") + assert on_disk + base = AutoModelForCausalLM.from_pretrained(f"/root/models/{MODEL_NAME}", dtype=torch.float32) + loaded = PeftModel.from_pretrained(base, adapter_dir).state_dict() + + expected_names = {name.replace(".weight", ".default.weight") for name in on_disk} + loaded_names = {name for name in loaded if ".lora_A." in name or ".lora_B." in name} + assert loaded_names == expected_names + for name, expected in on_disk.items(): + loaded_name = name.replace(".weight", ".default.weight") + assert torch.equal(loaded[loaded_name], expected), name + + with open(f"{adapter_dir}/adapter_config.json") as config_file: + config = json.load(config_file) + assert MODEL_NAME in config["base_model_name_or_path"] def execute(): @@ -96,7 +127,7 @@ def execute(): ci_args = "--ci-test " - save_args = "--save-interval 2 " "--save /root/checkpoints/lora-qwen2.5-0.5B-ci " + save_args = f"--save-interval 2 --save {SAVE_DIR} " misc_args = ( "--attention-dropout 0.0 " @@ -131,6 +162,7 @@ def execute(): num_gpus_per_node=NUM_GPUS, megatron_model_type=MODEL_TYPE, ) + _assert_peft_export() if __name__ == "__main__": diff --git a/tests/fast/backends/megatron_utils/test_lora_utils.py b/tests/fast/backends/megatron_utils/test_lora_utils.py index ca902420b40..a0280f3f670 100644 --- a/tests/fast/backends/megatron_utils/test_lora_utils.py +++ b/tests/fast/backends/megatron_utils/test_lora_utils.py @@ -4,7 +4,11 @@ exclude-module parsing, and LoRA sync config building — all without GPU. """ +import sys +import types from argparse import Namespace +from contextlib import nullcontext +from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -348,6 +352,219 @@ def test_canonical_target_modules(self): assert config["r"] == 8 +# --------------------------------------------------------------------------- +# HF PEFT export +# --------------------------------------------------------------------------- + + +class TestHfPeftExport: + """Single-adapter export should delegate the PEFT format to Bridge.""" + + @staticmethod + def _install_bridge_stub(monkeypatch, bridge): + class LoRA: + def __init__(self, *, dim, alpha, dropout, **_kwargs): + self.dim = dim + self.alpha = alpha + self.dropout = dropout + + class CanonicalLoRA(LoRA): + pass + + bridge_module = types.ModuleType("megatron.bridge") + bridge_module.__path__ = [] + bridge_module.AutoBridge = SimpleNamespace(from_hf_pretrained=lambda *args, **kwargs: bridge) + peft_module = types.ModuleType("megatron.bridge.peft") + peft_module.__path__ = [] + lora_module = types.ModuleType("megatron.bridge.peft.lora") + lora_module.LoRA = LoRA + canonical_module = types.ModuleType("megatron.bridge.peft.canonical_lora") + canonical_module.CanonicalLoRA = CanonicalLoRA + for module in (bridge_module, peft_module, lora_module, canonical_module): + monkeypatch.setitem(sys.modules, module.__name__, module) + + def test_save_delegates_to_bridge(self, tmp_path, monkeypatch): + import miles.backends.megatron_utils.lora_utils as lora_utils + + bridge = MagicMock() + self._install_bridge_stub(monkeypatch, bridge) + monkeypatch.setattr("miles.utils.megatron_bridge_utils.patch_megatron_model", lambda _: nullcontext()) + model = MagicMock() + model.named_parameters.return_value = [] + + def save_hf_adapter(*_args, **kwargs): + path = kwargs["path"] + (path / "adapter_config.json").write_text("{}") + (path / "adapter_model.safetensors").write_bytes(b"weights") + + bridge.save_hf_adapter.side_effect = save_hf_adapter + lora_utils.save_lora_checkpoint( + [model], + Namespace( + hf_checkpoint="/models/Qwen2.5-0.5B-Instruct", + target_modules=["linear_q", "linear_fc2"], + lora_rank=4, + lora_alpha=4, + lora_dropout=0.0, + ), + str(tmp_path), + ) + + bridge.save_hf_adapter.assert_called_once() + call = bridge.save_hf_adapter.call_args + assert call.args[0] == [model] + assert call.kwargs["path"].parent == tmp_path + assert call.kwargs["base_model_name_or_path"] == "/models/Qwen2.5-0.5B-Instruct" + assert call.kwargs["show_progress"] is False + assert call.kwargs["peft_config"].dim == 4 + assert call.kwargs["peft_config"].alpha == 4 + assert (tmp_path / "adapter_config.json").read_text() == "{}" + assert (tmp_path / "adapter_model.safetensors").read_bytes() == b"weights" + assert not list(tmp_path.glob(".peft-export-*")) + + def test_rank0_write_failure_recovers_before_outer_barrier(self, tmp_path, monkeypatch): + import miles.backends.megatron_utils.lora_utils as lora_utils + + args = Namespace( + hf_checkpoint="/models/Qwen2.5-0.5B-Instruct", + target_modules=["linear_q"], + lora_rank=4, + lora_alpha=4, + lora_dropout=0.0, + ) + + def run_rank(rank, *, fail_during_rank0_write): + bridge = MagicMock() + self._install_bridge_stub(monkeypatch, bridge) + monkeypatch.setattr("miles.utils.megatron_bridge_utils.patch_megatron_model", lambda _: nullcontext()) + + events = [] + collectives = [] + monkeypatch.setattr(lora_utils.dist, "is_initialized", lambda: True) + monkeypatch.setattr(lora_utils.dist, "get_rank", lambda: rank) + monkeypatch.setattr(lora_utils.dist, "get_world_size", lambda group=None: 2) + monkeypatch.setattr(lora_utils, "get_gloo_group", lambda: None) + + def barrier(): + events.append("barrier") + collectives.append("barrier") + + consensus_count = 0 + + def all_gather_object(output, value, group=None): + nonlocal consensus_count + collectives.append("all_gather_object") + if consensus_count == 0: + assert value is None + events.append("setup-consensus") + output[:] = [None, None] + else: + assert (value == "OSError('disk full')") is (rank == 0) + events.append("save-consensus") + output[:] = ["OSError('disk full')", None] + consensus_count += 1 + + monkeypatch.setattr(lora_utils.dist, "barrier", barrier) + monkeypatch.setattr(lora_utils.dist, "all_gather_object", all_gather_object) + + export_path = tmp_path / f"rank{rank}" + + def save_hf_adapter(*_args, **kwargs): + lora_utils.dist.barrier() # Bridge initial barrier. + if rank == 0: + (kwargs["path"] / "adapter_config.json").write_text("{}") + if fail_during_rank0_write: + events.append("rank-0-write-failed") + raise OSError("disk full") + lora_utils.dist.barrier() # Bridge final barrier. + + bridge.save_hf_adapter.side_effect = save_hf_adapter + model = MagicMock() + model.named_parameters.return_value = [] + optimizer = MagicMock() + optimizer.state_dict.side_effect = lambda: events.append("training-state") or {} + + lora_utils.save_lora_checkpoint( + [model], + args, + str(export_path), + optimizer=optimizer, + ) + assert not (export_path / "adapter_config.json").exists() + assert not (export_path / "adapter_model.safetensors").exists() + assert not list(export_path.glob(".peft-export-*")) + return events, collectives + + rank0_events, rank0_collectives = run_rank(0, fail_during_rank0_write=True) + peer_events, peer_collectives = run_rank(1, fail_during_rank0_write=False) + + # Deadlock freedom: both ranks issue the same collectives in the same order, + # even though only rank 0 failed inside Bridge's collective save. + assert rank0_collectives == peer_collectives + assert rank0_events == [ + "barrier", + "setup-consensus", + "barrier", + "rank-0-write-failed", + "barrier", + "save-consensus", + "training-state", + "barrier", + ] + assert peer_events == [ + "barrier", + "setup-consensus", + "barrier", + "barrier", + "save-consensus", + "training-state", + "barrier", + ] + + def test_setup_failure_is_shared_before_collective_save(self, tmp_path, monkeypatch): + import miles.backends.megatron_utils.lora_utils as lora_utils + + args = Namespace( + hf_checkpoint="/models/Qwen2.5-0.5B-Instruct", + target_modules=["linear_q"], + lora_rank=4, + lora_alpha=4, + lora_dropout=0.0, + ) + rank0_error = "OSError('rank 0 setup failed')" + + def run_rank(rank): + bridge = MagicMock() + self._install_bridge_stub(monkeypatch, bridge) + if rank == 0: + sys.modules["megatron.bridge"].AutoBridge.from_hf_pretrained = MagicMock( + side_effect=OSError("rank 0 setup failed") + ) + monkeypatch.setattr("miles.utils.megatron_bridge_utils.patch_megatron_model", lambda _: nullcontext()) + + collectives = [] + monkeypatch.setattr(lora_utils.dist, "is_initialized", lambda: True) + monkeypatch.setattr(lora_utils.dist, "get_rank", lambda: rank) + monkeypatch.setattr(lora_utils.dist, "get_world_size", lambda group=None: 2) + monkeypatch.setattr(lora_utils, "get_gloo_group", lambda: None) + monkeypatch.setattr(lora_utils.dist, "barrier", lambda: collectives.append("barrier")) + + def all_gather_object(output, value, group=None): + assert (value == rank0_error) is (rank == 0) + collectives.append("all_gather_object") + output[:] = [rank0_error, None] + + monkeypatch.setattr(lora_utils.dist, "all_gather_object", all_gather_object) + model = MagicMock() + model.named_parameters.return_value = [] + + lora_utils.save_lora_checkpoint([model], args, str(tmp_path / f"setup-rank{rank}")) + bridge.save_hf_adapter.assert_not_called() + return collectives + + assert run_rank(0) == run_rank(1) == ["barrier", "all_gather_object", "barrier"] + + # --------------------------------------------------------------------------- # LORA_ADAPTER_NAME constant # --------------------------------------------------------------------------- diff --git a/tests/fast/backends/megatron_utils/test_multi_lora_export_errors.py b/tests/fast/backends/megatron_utils/test_multi_lora_export_errors.py new file mode 100644 index 00000000000..5a5997eaaad --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_multi_lora_export_errors.py @@ -0,0 +1,19 @@ +from unittest.mock import MagicMock + +import pytest + +import miles.backends.megatron_utils.multi_lora_utils as multi_lora_utils + + +def test_checkpoint_stage_failure_is_reported_to_peer_ranks(monkeypatch): + monkeypatch.setattr(multi_lora_utils.dist, "is_initialized", lambda: True) + monkeypatch.setattr(multi_lora_utils.dist, "get_world_size", lambda group: 2) + monkeypatch.setattr(multi_lora_utils, "get_gloo_group", MagicMock(return_value=object())) + + def gather(messages, _local_message, group): + messages[:] = ["OSError: disk full", None] + + monkeypatch.setattr(multi_lora_utils.dist, "all_gather_object", gather) + + with pytest.raises(RuntimeError, match="PEFT checkpoint write.*disk full"): + multi_lora_utils._raise_if_any_rank_failed(None, "PEFT checkpoint write") diff --git a/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py b/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py index 2df14f54b1f..ebdb5136128 100644 --- a/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py +++ b/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py @@ -1,10 +1,19 @@ """slice_lora_to_rank trims max-rank-padded LoRA tensors to the adapter's real rank; used by weight-sync and HF PEFT export (PEFT rejects tensors padded past the declared rank).""" +from typing import NamedTuple +from unittest.mock import patch + import pytest import torch -from miles.backends.megatron_utils.multi_lora_utils import slice_lora_to_rank +from miles.backends.megatron_utils.multi_lora_utils import _build_multi_lora_peft_export, slice_lora_to_rank + + +class _Weight(NamedTuple): + param_name: str + weight: torch.Tensor + megatron_param_name: str | None = None def _padded(shape, live_rows=None, live_cols=None): @@ -81,3 +90,34 @@ def test_packed_expert_fewer_experts_than_rank_is_not_confused(): tensor[:, :4] = 1.0 out = slice_lora_to_rank("x.experts.gate_proj.lora_A.weight", tensor, 4) assert out.shape == (2, 4, 8) + + +def test_multi_lora_slices_before_bridge_conversion(): + tensor = torch.zeros(2, 8, 4, dtype=torch.bfloat16) + tensor[:, :4] = 1 + weights = [_Weight("x.experts.gate_proj.lora_A.weight", tensor)] + + from megatron.bridge.models.conversion import peft_bridge + + with patch.object( + peft_bridge, + "convert_adapter_weights_to_peft_state", + return_value=({"weight": torch.ones(1)}, ["x.experts.gate_proj"], []), + ) as convert: + _, config = _build_multi_lora_peft_export( + iter(weights), + rank=4, + alpha=8, + dropout=0.0, + base_model_name_or_path="/base", + ) + + converted_weights = convert.call_args.args[0] + assert converted_weights[0].weight.shape == (2, 4, 4) + assert converted_weights[0].weight.dtype == torch.float32 + assert config["r"] == 4 + + +def test_multi_lora_export_rejects_an_empty_adapter(): + with pytest.raises(RuntimeError, match="No adapter weights"): + _build_multi_lora_peft_export(iter([]), rank=4, alpha=8, dropout=0.0, base_model_name_or_path="/base")