diff --git a/miles/backends/megatron_utils/arguments.py b/miles/backends/megatron_utils/arguments.py index 668fb6b9ac..06e198b661 100644 --- a/miles/backends/megatron_utils/arguments.py +++ b/miles/backends/megatron_utils/arguments.py @@ -15,6 +15,9 @@ def set_default_megatron_args(args): args.use_distributed_optimizer = (args.optimizer is None or args.optimizer.lower() == "adam") and not getattr( args, "debug_disable_optimizer", False ) + # Multi-LoRA: per-slot LayerWise optimizers require plain DDP all-reduce. + if getattr(args, "multi_lora_n_adapters", 0) > 0: + args.use_distributed_optimizer = False # TODO: maybe change this after megatron has good fp8 support args.bf16 = not args.fp16 # placeholders diff --git a/miles/backends/megatron_utils/bridge_lora_helpers.py b/miles/backends/megatron_utils/bridge_lora_helpers.py index c56b516c7c..9114a14c5b 100644 --- a/miles/backends/megatron_utils/bridge_lora_helpers.py +++ b/miles/backends/megatron_utils/bridge_lora_helpers.py @@ -12,8 +12,9 @@ from megatron.core.utils import get_attr_wrapped_model from miles.utils.hf_config import load_hf_config +from miles.utils.multi_lora import is_multi_lora_enabled -from .lora_utils import create_lora_instance, patch_param_grad_buffer_for_colocate_mode_lora +from .lora_utils import patch_param_grad_buffer_for_colocate_mode_lora @dataclass @@ -111,7 +112,14 @@ def _setup_lora_model_via_bridge(args: Namespace) -> list: provider.dsa_attention_backend = getattr(args, "dsa_attention_backend", "megatron") provider.finalize() - lora = create_lora_instance(args) + if is_multi_lora_enabled(args): + from miles.backends.megatron_utils.multi_lora_utils import create_multi_lora_instance + + lora = create_multi_lora_instance(args) + else: + from .lora_utils import create_lora_instance + + lora = create_lora_instance(args) def apply_lora_hook(model_chunks): transformed = lora(model_chunks, training=True) @@ -129,6 +137,10 @@ def apply_lora_hook(model_chunks): provider.register_pre_wrap_hook(_make_value_model_hook(hidden_size, provider.sequence_parallel)) use_distributed_optimizer = "muon" not in (args.optimizer or "").lower() + if is_multi_lora_enabled(args): + # Per-slot LayerWise optimizers: plain DDP all-reduce keeps full grads on + # every rank (whole-param sharding + retained-gradient idempotency). + use_distributed_optimizer = False ddp_config = DistributedDataParallelConfig( use_distributed_optimizer=use_distributed_optimizer, grad_reduce_in_fp32=args.accumulate_allreduce_grads_in_fp32, diff --git a/miles/backends/megatron_utils/model.py b/miles/backends/megatron_utils/model.py index 61a49fe395..30e70661b9 100644 --- a/miles/backends/megatron_utils/model.py +++ b/miles/backends/megatron_utils/model.py @@ -6,6 +6,7 @@ import math from argparse import Namespace from collections.abc import Callable, Sequence +from contextlib import nullcontext from functools import partial from pathlib import Path @@ -31,6 +32,7 @@ from miles.utils.audit_utils.witness.module import witness_dump_and_clear_stale from miles.utils.dumper_utils import DumperMegatronUtil, DumperPhase from miles.utils.memory_utils import clear_memory +from miles.utils.multi_lora import is_multi_lora_enabled from miles.utils.test_utils.ft_test_actions import FTTestActionActorExecutor from miles.utils.tracking_utils.structured_log import log_structured @@ -137,7 +139,11 @@ def setup_model_and_optimizer( assert not args.moe_use_upcycling assert args.load is not None or args.pretrained_checkpoint is not None - if is_lora_enabled(args) and role == "actor" and args.megatron_to_hf_mode == "bridge": + # Multi-LoRA and single-LoRA (actor, bridge) both build via the bridge helper, + # which picks the adapter type internally. + if is_multi_lora_enabled(args) or ( + is_lora_enabled(args) and role == "actor" and args.megatron_to_hf_mode == "bridge" + ): model = _setup_lora_model_via_bridge(args) else: model = get_model(get_model_provider_func(args, role), ModelType.encoder_or_decoder) @@ -165,6 +171,10 @@ def setup_model_and_optimizer( use_gloo_process_groups=args.enable_gloo_process_groups, layer_wise_distributed_optimizer="dist" in config.optimizer.lower(), ) + elif is_multi_lora_enabled(args): + from miles.backends.megatron_utils.multi_lora_optimizer import build_multi_lora_optimizer + + optimizer = build_multi_lora_optimizer(args, config, model) else: optimizer = get_megatron_optimizer( config=config, @@ -291,6 +301,11 @@ def forward_step( total_lengths = batch["total_lengths"] response_lengths = batch["response_lengths"] + if "adapter_token_counts" in batch: + from megatron.bridge.peft.multi_lora_layers import set_tokens_per_adapter_slot + + set_tokens_per_adapter_slot(model, batch["adapter_token_counts"]) + output_tensor = model( input_ids=tokens, position_ids=None, @@ -355,6 +370,13 @@ def forward_step( return rollout_data +def _zero_grads(model: Sequence[DDP], optimizer: MegatronOptimizer | None, disable_optimizer: bool) -> None: + for model_chunk in model: + model_chunk.zero_grad_buffer() + if not disable_optimizer: + optimizer.zero_grad() + + def train_one_step( args: Namespace, rollout_id: int, @@ -373,6 +395,10 @@ def train_one_step( Runs forward/backward over ``num_microbatches``, applies optimizer step and one scheduler step when gradients are valid. + Multi-LoRA: gradients are retained across train calls (per-adapter + gradient accumulation); only the slots in the batch's ``step_slots`` step, + and only their gradients are zeroed. + Args: args: Runtime arguments. rollout_id: Rollout identifier. @@ -390,12 +416,16 @@ def train_one_step( parallel_state = get_parallel_state() dumper_phase_util = DumperMegatronUtil(args, model, DumperPhase.FWD_BWD, rollout_id=rollout_id) disable_optimizer = args.debug_disable_optimizer or optimizer is None + multi_lora = is_multi_lora_enabled(args) - # Set grad to zero. - for model_chunk in model: - model_chunk.zero_grad_buffer() - if not disable_optimizer: - optimizer.zero_grad() + if multi_lora: + from miles.backends.megatron_utils.multi_lora_optimizer import reset_grad_metadata_keep_grads + + # Retain accumulated per-adapter gradients; reset only the per-iteration + # DDP bookkeeping. Slot grads are zeroed selectively at step time. + reset_grad_metadata_keep_grads(model) + else: + _zero_grads(model, optimizer, disable_optimizer) if args.custom_megatron_before_train_step_hook_path: from miles.utils.misc import load_function @@ -445,6 +475,11 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p allgather_cp=args.allgather_cp, ) + if "adapter_token_counts" in batch: + from megatron.bridge.peft.multi_lora_layers import set_tokens_per_adapter_slot + + set_tokens_per_adapter_slot(model, batch["adapter_token_counts"]) + from miles.utils.replay_base import all_replay_managers old_stages = [m.stage for m in all_replay_managers] @@ -517,7 +552,7 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p outcome = TrainStepOutcome.DISCARDED_SHOULD_RETRY valid_step = False - if (not disable_optimizer) and (not getattr(args, "check_for_nan_in_loss_and_grad", True)): + if (not disable_optimizer) and (not multi_lora) and (not getattr(args, "check_for_nan_in_loss_and_grad", True)): found_inf_flag = optimizer.prepare_grads() if found_inf_flag: valid_step = False @@ -542,18 +577,24 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p dumper_phase_util.finalize(model) if not disable_optimizer and valid_step: - # Update parameters. - update_successful, grad_norm, num_zeros_in_grad = optimizer.step() + if multi_lora: + from miles.backends.megatron_utils.multi_lora_utils import step_stepped_adapter_slots - # Update learning rate. - assert update_successful - opt_param_scheduler.step(increment=args.global_batch_size) + grad_norm = step_stepped_adapter_slots( + args, model, optimizer, data_iterator[0].rollout_data, rollout_id, step_id + ) + else: + # Update parameters. + update_successful, grad_norm, num_zeros_in_grad = optimizer.step() - # release grad - for model_chunk in model: - model_chunk.zero_grad_buffer() - if not disable_optimizer: - optimizer.zero_grad() + # Update learning rate. + assert update_successful + opt_param_scheduler.step(increment=args.global_batch_size) + + # release grad (multi-LoRA retains accumulated grads; stepped slots were + # zeroed selectively inside step_adapter_slots) + if not multi_lora: + _zero_grads(model, optimizer, disable_optimizer) log_structured( logger.info, @@ -909,13 +950,25 @@ def initialize_model_and_optimizer( model, optimizer, opt_param_scheduler = setup_model_and_optimizer(args, role) model[0].role = role clear_memory() - iteration, _ = load_checkpoint( - model, - optimizer, - opt_param_scheduler, - checkpointing_context=checkpointing_context, - skip_load_to_model_and_opt=False, - ) + + multi_lora = is_multi_lora_enabled(args) + if multi_lora: + # Hide adapter params so the bridge's conversion-task walk doesn't see them + # while loading the base checkpoint. + from megatron.bridge.peft.multi_lora_layers import hide_adapters + + load_ctx = hide_adapters(model) + else: + load_ctx = nullcontext() + + with load_ctx: + iteration, _ = load_checkpoint( + model, + optimizer, + opt_param_scheduler, + checkpointing_context=checkpointing_context, + skip_load_to_model_and_opt=False, + ) check_peak_gpu_memory_after_load(args) clear_memory() diff --git a/miles/backends/megatron_utils/multi_lora_optimizer.py b/miles/backends/megatron_utils/multi_lora_optimizer.py new file mode 100644 index 0000000000..1397bd7569 --- /dev/null +++ b/miles/backends/megatron_utils/multi_lora_optimizer.py @@ -0,0 +1,195 @@ +"""Per-slot decoupled Adam optimizers for multi-LoRA, chained under Megatron's LayerWiseDistributedOptimizer; +requires plain DDP all-reduce (use_distributed_optimizer OFF) so cross-batch gradient retention stays idempotent.""" + +import logging +from argparse import Namespace +from collections.abc import Sequence +from contextlib import contextmanager + +import torch +from megatron.core.optimizer import get_megatron_optimizer +from megatron.core.optimizer.clip_grads import clip_grad_by_total_norm_fp32, get_grad_norm_fp32 +from megatron.core.optimizer.layer_wise_optimizer import LayerWiseDistributedOptimizer +from megatron.core.optimizer.optimizer import MegatronOptimizer +from megatron.core.optimizer.optimizer_config import OptimizerConfig +from megatron.core.process_groups_config import ProcessGroupCollection + +logger = logging.getLogger(__name__) + + +def adapter_slot_parameters(model, slot: int) -> list[torch.nn.Parameter]: + """All parameters belonging to one adapter slot, across model chunks.""" + from megatron.bridge.peft.multi_lora_layers import MultiLoRALinear + + parameters = [] + seen = set() + model_chunks = model if isinstance(model, (list, tuple)) else [model] + for model_chunk in model_chunks: + for module in model_chunk.modules(): + if not isinstance(module, MultiLoRALinear): + continue + for param in module.adapters[slot].parameters(): + if id(param) not in seen: + parameters.append(param) + seen.add(id(param)) + return parameters + + +def _adam_init_state_fn(opt, config=None): + for group in opt.param_groups: + for p in group["params"]: + if len(opt.state[p]) == 0: + opt.state[p]["exp_avg"] = torch.zeros_like(p.data) + opt.state[p]["exp_avg_sq"] = torch.zeros_like(p.data) + + +@contextmanager +def _only_slot_trainable(model_chunks, slot_params: list[torch.nn.Parameter]): + """Temporarily freeze every trainable param outside ``slot_params`` so the + stock param-group builder sees exactly one slot (the Muon construction + pattern from megatron's ``get_megatron_muon_optimizer``).""" + slot_ids = {id(p) for p in slot_params} + frozen = [] + for model_chunk in model_chunks: + for param in model_chunk.parameters(): + if param.requires_grad and id(param) not in slot_ids: + param.requires_grad = False + frozen.append(param) + try: + yield + finally: + for param in frozen: + param.requires_grad = True + + +def build_multi_lora_optimizer( + args: Namespace, + config: OptimizerConfig, + model_chunks: Sequence, +) -> MegatronOptimizer: + """Build one Float16-wrapped Adam per adapter slot under a LayerWiseDistributedOptimizer (ChainedOptimizer); + each child's param groups are tagged with ``miles_multi_lora_slot`` and narrowed to this rank's shard.""" + assert not config.use_distributed_optimizer, ( + "multi-LoRA per-slot optimizers require use_distributed_optimizer=False: " + "gradient retention relies on all-reduce idempotency, and LayerWise " + "sharding replaces byte-level ZeRO" + ) + assert not config.fp16, "multi-LoRA per-slot optimizers require bf16 (no dynamic loss scaler)" + assert (config.optimizer or "").lower() == "adam", ( + "multi-LoRA per-slot optimizers only implement Adam semantics (state init, " + f"slot retirement cleanup, step clocks); got optimizer={config.optimizer!r}" + ) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups() + + # Defer bf16 master-weight creation into LayerWise (post-sharding) so fp32 masters exist only for owned params. + reset_bf16 = config.bf16 + config.bf16 = False + + base_optimizers: list = [] + init_fns: list = [] + slot_child_indices: dict[int, list[int]] = {} + try: + for slot in range(args.multi_lora_n_adapters): + slot_params = adapter_slot_parameters(model_chunks, slot) + assert slot_params, f"adapter slot {slot} has no parameters; is this a multi-LoRA model?" + with _only_slot_trainable(model_chunks, slot_params): + chained = get_megatron_optimizer( + config, + list(model_chunks), + use_gloo_process_groups=args.enable_gloo_process_groups, + ) + children = [ + child + for child in chained.chained_optimizers + if getattr(child, "optimizer", None) is not None and child.get_parameters() + ] + assert children, f"adapter slot {slot} produced no optimizer children" + slot_child_indices[slot] = list(range(len(base_optimizers), len(base_optimizers) + len(children))) + for child in children: + for group in child.param_groups: + group["miles_multi_lora_slot"] = slot + base_optimizers.append(child) + init_fns.append(_adam_init_state_fn) + finally: + config.bf16 = reset_bf16 + + optimizer = LayerWiseDistributedOptimizer(base_optimizers, config, pg_collection, init_state_fn_list=init_fns) + + # Params are scattered whole across DP ranks, so per-child norm/clip reductions must span the world. + for child in optimizer.chained_optimizers: + child.grad_stats_parallel_group = None + + optimizer.miles_slot_child_indices = slot_child_indices + logger.info( + f"Built multi-LoRA LayerWise optimizer: {args.multi_lora_n_adapters} slots, " + f"{len(optimizer.chained_optimizers)} chained children" + ) + return optimizer + + +def _slot_children(optimizer, slot: int): + return [optimizer.chained_optimizers[i] for i in optimizer.miles_slot_child_indices[slot]] + + +def reset_grad_metadata_keep_grads(model_chunks) -> None: + """Reset DDP per-iteration grad bookkeeping WITHOUT zeroing grad buffers, so per-adapter accumulation + survives across train batches (replaces ``DistributedDataParallel.zero_grad_buffer``).""" + for model_chunk in model_chunks: + if getattr(model_chunk.config, "cuda_graph_impl", "none") != "transformer_engine": + for param in model_chunk.params_with_grad: + param.grad_added_to_main_grad = False + for bucket_group in model_chunk.bucket_groups + model_chunk.expert_parallel_bucket_groups: + bucket_group.reset() + + +def zero_adapter_slot_grads(model, slot: int) -> None: + """Zero one slot's gradients everywhere they live: the DDP ``main_grad`` buffer views + and any lingering ``grad``/``main_param.grad`` references.""" + for param in adapter_slot_parameters(model, slot): + if (main_grad := getattr(param, "main_grad", None)) is not None: + main_grad.zero_() + param.grad = None + if (main_param := getattr(param, "main_param", None)) is not None: + main_param.grad = None + + +def step_adapter_slots( + optimizer, + model, + step_batch_sizes: dict[int, int], + clip_grad: float, +) -> dict[int, float]: + """Step exactly the slots in ``step_batch_sizes`` (slot -> batch size), retaining all other slots' gradients; + scales each slot's accumulated grad sum by 1/batch_size and returns the grad norm per stepped slot.""" + grad_norms: dict[int, float] = {} + + for slot, batch_size in step_batch_sizes.items(): + children = _slot_children(optimizer, slot) + # Copy accumulated main_grads into the owned masters' grads, then scale the sum to the adapter-batch mean. + for child in children: + child.prepare_grads() + for main_param in child.get_parameters(): + if main_param.grad is not None: + main_param.grad.mul_(1.0 / batch_size) + + # Per-slot grad norm over the slot's children, reduced across the whole world (whole-param DP scatter). + grads_for_norm = [] + slot_params = [] + for child in children: + grads_for_norm += child.get_main_grads_for_grad_norm() + slot_params += child.get_parameters() + slot_norm = get_grad_norm_fp32(grads_for_norm, grad_stats_parallel_group=None) + if clip_grad > 0.0 and slot_params: + clip_grad_by_total_norm_fp32(slot_params, clip_grad, slot_norm, False) + grad_norms[slot] = float(slot_norm) + + for child in children: + child.step_with_ready_grads() + + zero_adapter_slot_grads(model, slot) + + if step_batch_sizes: + optimizer.allgather_params() + + return grad_norms diff --git a/miles/backends/megatron_utils/multi_lora_scheduler.py b/miles/backends/megatron_utils/multi_lora_scheduler.py new file mode 100644 index 0000000000..4f148eb599 --- /dev/null +++ b/miles/backends/megatron_utils/multi_lora_scheduler.py @@ -0,0 +1,88 @@ +"""Per-adapter LR/WD schedules for multi-LoRA: one ``OptimizerParamScheduler`` per adapter slot, positioned by +the adapter's own trained samples. Adapters without a known ``num_step`` warm up, then hold ``--lr`` constant.""" + +import logging +from argparse import Namespace + +from megatron.core.optimizer_param_scheduler import OptimizerParamScheduler + +logger = logging.getLogger(__name__) + + +class _SlotParamGroups: + """Minimal optimizer facade: the scheduler only reads ``.param_groups``.""" + + def __init__(self, param_groups: list[dict]): + self.param_groups = param_groups + + +def build_slot_scheduler(args: Namespace, optimizer, adapter, resume_step: int) -> OptimizerParamScheduler: + """Build the slot's scheduler and position it at the adapter's committed + samples. Rebuilt on every adapter load, so slot reuse starts fresh.""" + from miles.backends.megatron_utils.multi_lora_optimizer import _slot_children + + groups = [group for child in _slot_children(optimizer, adapter.slot) for group in child.param_groups] + samples_per_step = adapter.config.adapter_global_batch_size + num_step = adapter.config.num_step + + decay_steps = num_step * samples_per_step if num_step is not None else None + if args.lr_warmup_fraction is not None and decay_steps is not None: + lr_warmup_steps = args.lr_warmup_fraction * decay_steps + else: + lr_warmup_steps = args.lr_warmup_iters * samples_per_step + if decay_steps is None: + # No horizon: warm up, then hold constant. The decay steps only need + # to satisfy the scheduler's warmup < decay invariant. + lr_decay_style = "constant" + decay_steps = int(lr_warmup_steps) + 1 + else: + lr_decay_style = args.lr_decay_style + + scheduler = OptimizerParamScheduler( + _SlotParamGroups(groups), + init_lr=args.lr_warmup_init, + max_lr=args.lr, + min_lr=args.min_lr, + lr_warmup_steps=lr_warmup_steps, + lr_decay_steps=decay_steps, + lr_decay_style=lr_decay_style, + start_wd=args.start_weight_decay, + end_wd=args.end_weight_decay, + wd_incr_steps=decay_steps, + wd_incr_style=args.weight_decay_incr_style, + use_checkpoint_opt_param_scheduler=False, + override_opt_param_scheduler=False, + wsd_decay_steps=( + args.lr_wsd_decay_iters * samples_per_step + if lr_decay_style == "WSD" and args.lr_wsd_decay_iters is not None + else None + ), + lr_wsd_decay_style=args.lr_wsd_decay_style, + ) + if resume_step: + scheduler.step(increment=resume_step * samples_per_step) + return scheduler + + +def install_slot_scheduler(args: Namespace, optimizer, adapter, resume_step: int) -> None: + """Attach the adapter's scheduler to the optimizer, keyed by slot.""" + if not hasattr(optimizer, "miles_slot_schedulers"): + optimizer.miles_slot_schedulers = {} + optimizer.miles_slot_schedulers[adapter.slot] = build_slot_scheduler(args, optimizer, adapter, resume_step) + + +def drop_slot_scheduler(optimizer, slot: int) -> None: + """Detach a retired slot's scheduler (the next tenant installs its own).""" + getattr(optimizer, "miles_slot_schedulers", {}).pop(slot, None) + + +def step_slot_schedulers(optimizer, step_batch_sizes: dict[int, int]) -> dict[int, float]: + """Advance exactly the stepped slots' schedules by their batch samples. + Returns slot -> new learning rate, for logging.""" + lr_by_slot: dict[int, float] = {} + for slot, batch_size in step_batch_sizes.items(): + scheduler = optimizer.miles_slot_schedulers[slot] + scheduler.step(increment=batch_size) + if scheduler.optimizer.param_groups: # empty on ranks owning none of the slot's params + lr_by_slot[slot] = scheduler.optimizer.param_groups[0]["lr"] + return lr_by_slot diff --git a/miles/backends/megatron_utils/multi_lora_utils.py b/miles/backends/megatron_utils/multi_lora_utils.py new file mode 100644 index 0000000000..707541fa4e --- /dev/null +++ b/miles/backends/megatron_utils/multi_lora_utils.py @@ -0,0 +1,441 @@ +import json +import logging +import os +from argparse import Namespace +from collections.abc import Mapping +from pathlib import Path + +import ray +import torch +import torch.distributed as dist + +from miles.backends.training_utils.parallel import get_parallel_state +from miles.ray.multi_lora.controller import get_multi_lora_controller +from miles.utils.adapter_config import AdapterRun + +logger = logging.getLogger(__name__) + + +def create_multi_lora_instance(args: Namespace): + """Create a MultiLoRA instance from training args.""" + from megatron.bridge.peft.multi_lora import MultiLoRA + + from miles.backends.megatron_utils.lora_utils import convert_target_modules_to_megatron + + lora_type_name = getattr(args, "lora_type", "lora").lower() + if lora_type_name == "canonical_lora": + from megatron.bridge.peft.canonical_lora import CanonicalLoRA + + lora_cls = CanonicalLoRA + else: + from megatron.bridge.peft.lora import LoRA + + lora_cls = LoRA + + return MultiLoRA( + target_modules=convert_target_modules_to_megatron(args.target_modules, lora_type=lora_cls), + n_adapters=args.multi_lora_n_adapters, + dim=args.lora_rank, + alpha=args.lora_alpha, + dropout=getattr(args, "lora_dropout", 0.0), + lora_A_init_method=getattr(args, "lora_A_init_method", "xavier"), + lora_B_init_method=getattr(args, "lora_B_init_method", "zero"), + ) + + +def all_megatron_checkpoints_exist(step_dir: Path, tp_size, pp_size) -> bool: + return all( + (step_dir / f"adapter_megatron_tp{tp}_pp{pp}.pt").exists() for tp in range(tp_size) for pp in range(pp_size) + ) + + +def find_latest_checkpoint(ckpt_dir: Path) -> tuple[Path | None, int]: + if not ckpt_dir.exists(): + return None, 0 + + parallel_state = get_parallel_state() + tp_size = parallel_state.tp.size + pp_size = parallel_state.pp.size + tp_rank = parallel_state.tp.rank + pp_rank = parallel_state.pp.rank + + def get_step(d): + return int(d.name.split("_")[1]) + + step_dirs = sorted( + [d for d in ckpt_dir.iterdir() if d.is_dir() and d.name.startswith("step_")], + key=get_step, + reverse=True, + ) + for step_dir in step_dirs: + step = get_step(step_dir) + if all_megatron_checkpoints_exist(step_dir, tp_size, pp_size): + return step_dir / f"adapter_megatron_tp{tp_rank}_pp{pp_rank}.pt", step + + return None, 0 + + +def zero_optimizer_state_for_adapter(optimizer, model, idx: int) -> None: + from megatron.bridge.peft.multi_lora_layers import MultiLoRALinear, _iter_multi_lora_modules + + target_main_params = set() + for module in _iter_multi_lora_modules(model): + if not isinstance(module, MultiLoRALinear): + continue + adapter = module.adapters[idx] + for param in adapter.parameters(): + main = getattr(param, "main_param", None) + target_main_params.add(id(main if main is not None else param)) + + chained = getattr(optimizer, "chained_optimizers", [optimizer]) + for chained_optimizer in chained: + inner = getattr(chained_optimizer, "optimizer", chained_optimizer) + if inner is None: + continue + # TE/apex FusedAdam tracks the Adam step per param GROUP, not per param; + # reset the retired slot's groups so the next tenant restarts bias correction. + for group in inner.param_groups: + if group.get("miles_multi_lora_slot") == idx and "step" in group: + if isinstance(group["step"], torch.Tensor): + group["step"].zero_() + else: + group["step"] = 0 + for param, state in inner.state.items(): + if id(param) not in target_main_params: + continue + if "exp_avg" in state: + state["exp_avg"].zero_() + if "exp_avg_sq" in state: + state["exp_avg_sq"].zero_() + # Bias correction restarts for the slot's next tenant. + if "step" in state: + if isinstance(state["step"], torch.Tensor): + state["step"].zero_() + else: + state["step"] = 0 + + +def slice_lora_to_rank(hf_name: str, tensor: torch.Tensor, adapter_rank: int) -> torch.Tensor: + if "lora_A" in hf_name and adapter_rank < tensor.shape[0]: + remainder = tensor[adapter_rank:] + assert remainder.abs().max() == 0, ( + f"lora_A padded dims are non-zero: {hf_name}, " + f"max={remainder.abs().max().item():.6e}, shape={tensor.shape}, rank={adapter_rank}" + ) + return tensor[:adapter_rank] + if "lora_B" in hf_name and adapter_rank < tensor.shape[1]: + remainder = tensor[:, adapter_rank:] + assert remainder.abs().max() == 0, ( + f"lora_B padded dims are non-zero: {hf_name}, " + f"max={remainder.abs().max().item():.6e}, shape={tensor.shape}, rank={adapter_rank}" + ) + return tensor[:, :adapter_rank] + return tensor + + +def save_multi_lora_checkpoints( + args, + model, + adapter_steps: Mapping[str, int], + adapters: Mapping[str, AdapterRun], +): + """Save per-adapter checkpoints in two formats per adapter. + + Layout (per adapter):: + + {adapter.save}/checkpoints/step_{iteration}/ + ├── adapter_megatron_tp{tp}_pp{pp}.pt ← per-rank shard, fast resume + ├── adapter_model.safetensors ← gathered HF, inference / external + └── adapter_config.json ← HF PEFT metadata (r, alpha, ...) + """ + from megatron.bridge import AutoBridge + 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() + tp_rank = parallel_state.tp.rank + pp_rank = parallel_state.pp.rank + # One writer per (tp, pp) shard: LoRA params are replicated across DP AND + # CP, so gate on the combined dp×cp group. Gating on intra_dp alone left + # every CP rank writing the same shard file and racing the os.replace. + is_dp_cp_rank_0 = parallel_state.intra_dp_cp.rank == 0 + is_global_writer = is_dp_cp_rank_0 and tp_rank == 0 and pp_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) + + for adapter_name, adapter in adapters.items(): + config = adapter.config + log_prefix = f"[multilora] ({adapter_name})" + iteration = adapter_steps[adapter_name] + + if config.save is None: + logger.info(f"{log_prefix} skipping checkpoint (no save dir configured)") + continue + + final_dir = config.save / "checkpoints" / f"step_{iteration}" + tmp_dir = config.save / "checkpoints" / f"_tmp_step_{iteration}" + if is_dp_cp_rank_0: + tmp_dir.mkdir(parents=True, exist_ok=True) + if dist.is_initialized(): + dist.barrier() + + with expose_adapter_slot(model, adapter.slot): + # Megatron checkpoints + if is_dp_cp_rank_0: + 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 / f"adapter_megatron_tp{tp_rank}_pp{pp_rank}.pt" + 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() + + # 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() + + +def _register_adapter(adapter: AdapterRun, model) -> int: + """Install one adapter on this rank's local model shard. Returns the step + of the checkpoint it resumed from (0 for a fresh adapter).""" + from megatron.bridge.peft.multi_lora_layers import init_adapter_slot, load_adapter + + name = adapter.name + config = adapter.config + slot = adapter.slot + log_prefix = f"[multilora] ({name})" + + step = 0 + if config.save is not None: + ckpt_root = config.save / "checkpoints" + ckpt, step = find_latest_checkpoint(ckpt_root) + else: + ckpt = None + + if ckpt is None: + logger.info(f"{log_prefix} no checkpoint, starting from random init") + step = 0 + else: + state_dict = torch.load(ckpt, map_location="cpu", weights_only=True) + loaded = load_adapter(model, slot, state_dict) + assert loaded > 0, ( + f"{log_prefix} loaded 0 tensors from {ckpt} " + f"(state_dict has {len(state_dict)} entries) — name mismatch?" + ) + logger.info(f"{log_prefix} loaded from {ckpt} ({loaded} tensors)") + + init_adapter_slot(model, slot, rank=config.rank, alpha=config.alpha) + logger.info(f"{log_prefix} installed at slot {slot}") + return step + + +def _deregister_adapter(adapter: AdapterRun, args, model, optimizer) -> None: + """Model-side cleanup for one adapter.""" + from megatron.bridge.peft.multi_lora_layers import clear_adapter_slot + + name = adapter.name + slot = adapter.slot + log_prefix = f"[multilora] ({name})" + + if args.save_interval is not None: + # The controller still holds the step count until free_slot runs. + step = ray.get(get_multi_lora_controller().adapter_step.remote(name)) + save_multi_lora_checkpoints(args, model, {name: step}, {name: adapter}) + logger.info(f"{log_prefix} saved final checkpoint at step {step}") + else: + logger.info(f"{log_prefix} save_interval unset; skipping final checkpoint") + + clear_adapter_slot(model, slot) + logger.info(f"{log_prefix} cleared adapter slot {slot}") + + # Prevent future slot tenants from inheriting optimizer momentum or the + # previous tenant's partially accumulated gradients. + from miles.backends.megatron_utils.multi_lora_optimizer import zero_adapter_slot_grads + + from miles.backends.megatron_utils.multi_lora_scheduler import drop_slot_scheduler + + zero_optimizer_state_for_adapter(optimizer, model, slot) + zero_adapter_slot_grads(model, slot) + drop_slot_scheduler(optimizer, slot) + optimizer.reload_model_params() + logger.info(f"{log_prefix} cleared optimizer state and retained grads for slot {slot}") + + +def load_adapters(args, model, optimizer, adapters) -> int: + """Load adapters into Megatron slots; resumes step counts from checkpoints.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + from miles.utils.distributed_utils import get_gloo_group + + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + if not adapters: + return 0 + from miles.backends.megatron_utils.multi_lora_scheduler import install_slot_scheduler + + resume_steps: dict[str, int] = {} + for adapter in adapters: + resume_steps[adapter.name] = _register_adapter(adapter, model) + # Per-adapter LR/WD schedule, positioned at the resumed step count. + install_slot_scheduler(args, optimizer, adapter, resume_steps[adapter.name]) + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + optimizer.reload_model_params() + if is_first_replica_megatron_main_rank(): + for name, step in resume_steps.items(): + if step > 0: + ray.get(get_multi_lora_controller().set_adapter_step.remote(name, step)) + return len(adapters) + + +def cleanup_adapters(args, model, optimizer, adapters) -> int: + """Save final ckpt + clear Megatron slot, then free_slot on the controller.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + from miles.utils.distributed_utils import get_gloo_group + + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + if not adapters: + return 0 + for adapter in adapters: + _deregister_adapter(adapter, args, model, optimizer) + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + if is_first_replica_megatron_main_rank(): + for adapter in adapters: + ray.get(get_multi_lora_controller().free_slot.remote(adapter.name)) + return len(adapters) + + +def step_stepped_adapter_slots(args, model, optimizer, rollout_data, rollout_id: int, step_id: int) -> float: + """Optimizer-step the slots whose adapter batch completes with this train batch and advance + their per-adapter LR/WD schedules. Returns the max grad norm across stepped slots (0.0 if none).""" + from miles.backends.megatron_utils.multi_lora_optimizer import step_adapter_slots + from miles.backends.megatron_utils.multi_lora_scheduler import step_slot_schedulers + from miles.utils.tracking_utils.structured_log import log_structured + + # slot -> adapter_global_batch_size for adapter batches completing now. + step_batch_sizes = dict(rollout_data.get("step_adapter_batch_sizes", {})) + grad_norms_by_slot = step_adapter_slots( + optimizer, + model, + step_batch_sizes, + clip_grad=args.clip_grad, + ) + + if lr_by_slot := step_slot_schedulers(optimizer, step_batch_sizes): + log_structured( + logger.info, + op="adapter_lr", + rollout=rollout_id, + step=step_id, + **{f"slot_{slot}": lr for slot, lr in lr_by_slot.items()}, + ) + return max(grad_norms_by_slot.values(), default=0.0) + + +def commit_trained_batch(rollout_data, rollout_id: int, pending_push: set) -> None: + """A train call landed: schedule the stepped adapters' engine push and + commit the batch on the controller (main rank only). The stepped set ships + with the train data, identical on all ranks.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + + pending_push.update(rollout_data.get("step_adapter_names", [])) + if is_first_replica_megatron_main_rank(): + ray.get(get_multi_lora_controller().mark_batch_trained.remote(rollout_id)) + + +def save_due_adapter_checkpoints(args, model) -> bool: + """Save per-adapter checkpoints for adapters at a save-interval multiple + without a checkpoint on disk. Rank 0 picks and broadcasts, so the + collective export lines up. Returns False when nothing is due.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + from miles.utils.distributed_utils import get_gloo_group + + due_buffer = [None] + if is_first_replica_megatron_main_rank() and args.save_interval is not None: + snapshot = ray.get(get_multi_lora_controller().snapshot.remote()) + adapters = {**snapshot["active"], **snapshot["retiring"]} + due_buffer[0] = { + name: adapter + for name, adapter in adapters.items() + if adapter.step > 0 + and adapter.step % args.save_interval == 0 + and adapter.config.save is not None + and not (Path(adapter.config.save) / "checkpoints" / f"step_{adapter.step}").exists() + } + if dist.is_initialized(): + dist.broadcast_object_list(due_buffer, src=0, group=get_gloo_group()) + due_adapters = due_buffer[0] + if not due_adapters: + return False + adapter_steps = {name: adapter.step for name, adapter in due_adapters.items()} + save_multi_lora_checkpoints(args, model, adapter_steps, due_adapters) + return True + + +def select_adapters_to_push(loaded_adapters: dict, pending_push: set, has_new_engines: bool) -> tuple[dict, list]: + """Pick the stale adapters to push (all loaded adapters when engines are new). Returns + (adapters to push keyed by name, names to version-bump — only those whose weights changed).""" + pending = pending_push & set(loaded_adapters) + push_names = set(loaded_adapters) if has_new_engines else pending + return {name: loaded_adapters[name] for name in sorted(push_names)}, sorted(pending) + + +def commit_weight_push(version_update_names: list, is_main_rank: bool) -> None: + """A weight push landed: bump the pushed adapters' slot versions on the + controller (promotes PENDING adapters to ACTIVE).""" + if version_update_names and is_main_rank: + ray.get(get_multi_lora_controller().record_weight_update.remote(version_update_names)) diff --git a/miles/backends/training_utils/data.py b/miles/backends/training_utils/data.py index b5dccb8061..92468eb18b 100644 --- a/miles/backends/training_utils/data.py +++ b/miles/backends/training_utils/data.py @@ -128,7 +128,7 @@ def get_batch( Steps: - Fetch raw fields via iterator. - Save original token tensors under "unconcat_tokens". - - Slice tokens into two chunks for Context Parallelism (CP), concatenate, and pad to a configurable multiple. + - Slice tokens into two batches for Context Parallelism (CP), concatenate, and pad to a configurable multiple. - Build cu_seqlens and `PackedSeqParams` with T-H-D layout (T: sequence length, H: attention heads, D: head dimension). Args: @@ -146,6 +146,10 @@ def get_batch( parallel_state = get_parallel_state() assert "tokens" in keys + # get_batch consumes adapter_slots itself (per-adapter token counts below); + # fetch it here so callers don't have to know. None for non-multi-LoRA runs. + if "adapter_slots" not in keys: + keys = [*keys, "adapter_slots"] batch = data_iterator.get_next(keys) if "dynamic_global_batch_size" in data_iterator.rollout_data: @@ -167,7 +171,9 @@ def get_batch( if qkv_format == "bshd": max_seqlen = batch["max_seq_lens"][0] assert max([t.size(0) for t in tokens]) <= max_seqlen + if allgather_cp: + assert batch.get("adapter_slots") is None, "allgather CP is currently not supported with multi-LoRA: " assert max_seqlen % cp_size == 0, f"max_seqlen {max_seqlen} not divisible by cp_size {cp_size}" local_len = max_seqlen // cp_size start = parallel_state.cp.rank * local_len @@ -176,21 +182,23 @@ def get_batch( ] else: tokens = [slice_with_cp(t, pad_token_id, qkv_format, max_seqlen) for t in tokens] + sample_token_lengths = [t.size(0) for t in tokens] tokens = torch.stack(tokens) elif qkv_format == "thd": cp_rank = parallel_state.cp.rank if allgather_cp: + assert batch.get("adapter_slots") is None, "allgather CP is currently not supported with multi-LoRA: " # DSA mode: concatenate all sequences first, then slice once with CP. - # We also pad the *global* concatenated stream to make per-rank chunks equal. + # We also pad the *global* concatenated stream to make per-rank batches equal. cu_seqlens_list: list[int] = [0] for t in tokens: cu_seqlens_list.append(cu_seqlens_list[-1] + t.size(0)) tokens = torch.cat(tokens, dim=0) - # Pad global stream so (1) divisible by cp_size (equal chunks), + # Pad global stream so (1) divisible by cp_size (equal batches), # (2) divisible by pad_size (reduce fragmentation). global_pad_size = cp_size * pad_size pad = (global_pad_size - tokens.size(0) % global_pad_size) % global_pad_size @@ -202,6 +210,7 @@ def get_batch( tokens = tokens.chunk(cp_size, dim=0)[cp_rank] else: tokens = [slice_with_cp(t, pad_token_id, qkv_format) for t in tokens] + sample_token_lengths = [t.size(0) for t in tokens] cu_seqlens = [0] for t in tokens: @@ -227,6 +236,21 @@ def get_batch( else: raise ValueError(f"Unsupported qkv_format: {qkv_format}") + # Multi-LoRA: compute per-adapter token counts from post-CP per-sample lengths. + # NOTE: allgather CP is currently not supported + adapter_slots = batch.get("adapter_slots") + if adapter_slots is not None: + assert all( + adapter_slots[i] <= adapter_slots[i + 1] for i in range(len(adapter_slots) - 1) + ), f"adapter_slots not sorted in micro-batch: {adapter_slots}" + n_adapters = data_iterator.rollout_data["n_adapters"] + total_tokens = tokens.numel() + counts = torch.zeros(n_adapters, dtype=torch.int32, device=torch.cuda.current_device()) + for slot, length in zip(adapter_slots, sample_token_lengths, strict=True): + counts[slot] += length + counts[adapter_slots[-1]] += total_tokens - counts.sum().item() + batch["adapter_token_counts"] = counts + batch["tokens"] = tokens def _compute_transform_like_token_ids(ids_list: list): @@ -346,7 +370,7 @@ def get_next(self, keys: Sequence[str]) -> dict[str, list[object] | None]: - If `micro_batch_indices` is provided, selects rows according to the current index list for each requested key. - - Otherwise, slices a contiguous window of size `micro_batch_size` starting + - Otherwise, slices a contiguous adapter batch of size `micro_batch_size` starting at the current offset. Returns a dict mapping each key to a list subset (or None if absent). @@ -426,6 +450,12 @@ def _generate_data_iterator(rollout_data, micro_batch_size, micro_batch_indices= return data_iterator if not args.use_dynamic_batch_size: + if "adapter_slots" in rollout_data and num_local_gbs % args.micro_batch_size != 0: + raise ValueError( + "A multi-LoRA local batch must be divisible by --micro-batch-size; " + f"got local_batch_size={num_local_gbs}, micro_batch_size={args.micro_batch_size}. " + "Use --use-dynamic-batch-size or choose compatible adapter batch shapes." + ) num_microbatches = [num_local_gbs // args.micro_batch_size for _ in range(num_steps_per_rollout)] data_iterator = _generate_data_iterator(rollout_data, args.micro_batch_size) else: @@ -463,6 +493,10 @@ def _generate_data_iterator(rollout_data, micro_batch_size, micro_batch_indices= for j in range(num_mbs): for k in range(len(partitions[j])): partitions[j][k] += start + # Multi-LoRA: microbatches must be contiguous-by-slot for the + # grouped GEMM's per-adapter token-count math. + if "adapter_slots" in rollout_data: + partitions[j].sort(key=lambda index: rollout_data["adapter_slots"][index]) micro_batch_indices.extend(partitions) assert len(set(sum(micro_batch_indices, []))) == num_local_samples diff --git a/miles/backends/training_utils/log_utils.py b/miles/backends/training_utils/log_utils.py index 45b073ce9f..a4aa866ed7 100644 --- a/miles/backends/training_utils/log_utils.py +++ b/miles/backends/training_utils/log_utils.py @@ -137,6 +137,12 @@ def log_rollout_data(rollout_id: int, args: Namespace, rollout_data: RolloutBatc "witness_ids", "weight_versions", "metadata", + "n_adapters", + "adapter_slots", + "step_slots", + "step_adapter_names", + "step_adapter_batch_sizes", + "prompt_group_sizes", ]: continue # Upload per sample mean for each rollout value diff --git a/miles/backends/training_utils/loss.py b/miles/backends/training_utils/loss.py index 270fcde3e4..d817d93bf1 100644 --- a/miles/backends/training_utils/loss.py +++ b/miles/backends/training_utils/loss.py @@ -12,6 +12,7 @@ from miles.backends.training_utils.parallel import get_parallel_state from miles.utils.audit_utils.event_logger.logger import get_event_logger, is_event_logger_initialized from miles.utils.audit_utils.event_logger.models import TrainAdvantageComputationEvent +from miles.utils.multi_lora import is_multi_lora_enabled from miles.utils.types import RolloutBatch @@ -154,6 +155,11 @@ def loss_function( # Here we need to divide by cp_size because to cancel the multiply in Megatron. assert args.use_dynamic_global_batch_size == ("dynamic_global_batch_size" in batch) global_batch_size = batch.get("dynamic_global_batch_size", args.global_batch_size) + # Multi-LoRA: samples enter the gradient buffers with weight 1; per-adapter + # normalization (1/adapter_global_batch_size, a constant known in advance) + # is applied to the accumulated slot gradient at optimizer-step time. + if is_multi_lora_enabled(args): + global_batch_size = 1 if not args.calculate_per_token_loss: if apply_megatron_loss_scaling: loss_parallel_size = ( diff --git a/tests/fast/backends/megatron_utils/test_multi_lora_scheduler.py b/tests/fast/backends/megatron_utils/test_multi_lora_scheduler.py new file mode 100644 index 0000000000..be48d7118e --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_multi_lora_scheduler.py @@ -0,0 +1,114 @@ +"""Per-adapter LR schedules: parameters come from the global args, position is per adapter. +Pins two fixes: late loads don't inherit the decayed position; resume rebuilds position from committed steps.""" + +from types import SimpleNamespace + +import pytest + +from miles.backends.megatron_utils.multi_lora_scheduler import install_slot_scheduler, step_slot_schedulers + +LR = 2e-5 + + +def make_args(**overrides) -> SimpleNamespace: + args = SimpleNamespace( + lr=LR, + min_lr=0.0, + lr_warmup_init=0.0, + lr_warmup_fraction=None, + lr_warmup_iters=0, + lr_decay_style="cosine", + start_weight_decay=0.1, + end_weight_decay=0.1, + weight_decay_incr_style="constant", + lr_wsd_decay_iters=None, + lr_wsd_decay_style=None, + ) + for key, value in overrides.items(): + setattr(args, key, value) + return args + + +def make_optimizer(n_slots: int = 2) -> SimpleNamespace: + children = [SimpleNamespace(param_groups=[{"lr": 0.0, "weight_decay": 0.0}]) for _ in range(n_slots)] + return SimpleNamespace( + chained_optimizers=children, + miles_slot_child_indices={slot: [slot] for slot in range(n_slots)}, + ) + + +def make_adapter(slot: int, num_step: int | None, samples_per_step: int = 64) -> SimpleNamespace: + config = SimpleNamespace(num_step=num_step, adapter_global_batch_size=samples_per_step) + return SimpleNamespace(slot=slot, name=f"a{slot}", config=config) + + +def slot_lr(optimizer, slot: int) -> float: + return optimizer.chained_optimizers[slot].param_groups[0]["lr"] + + +def test_decaying_adapter_walks_its_own_cosine_schedule(): + optimizer = make_optimizer() + adapter = make_adapter(slot=0, num_step=10) + install_slot_scheduler(make_args(), optimizer, adapter, resume_step=0) + + assert slot_lr(optimizer, 0) == pytest.approx(LR) # fresh: top of the schedule + + step_slot_schedulers(optimizer, {0: 5 * 64}) # half the horizon + assert slot_lr(optimizer, 0) == pytest.approx(LR / 2) + + step_slot_schedulers(optimizer, {0: 100 * 64}) # far past the horizon + assert slot_lr(optimizer, 0) == pytest.approx(0.0) # clamped at min_lr + + +def test_adapter_without_num_step_holds_constant(): + optimizer = make_optimizer() + install_slot_scheduler(make_args(), optimizer, make_adapter(slot=0, num_step=None), resume_step=0) + + step_slot_schedulers(optimizer, {0: 12345 * 64}) + assert slot_lr(optimizer, 0) == pytest.approx(LR) # no horizon: never decays + + +def test_resume_position_is_deterministic_from_committed_steps(): + stepped = make_optimizer() + install_slot_scheduler(make_args(), stepped, make_adapter(slot=0, num_step=10), resume_step=0) + step_slot_schedulers(stepped, {0: 5 * 64}) + + resumed = make_optimizer() + install_slot_scheduler(make_args(), resumed, make_adapter(slot=0, num_step=10), resume_step=5) + + assert slot_lr(resumed, 0) == pytest.approx(slot_lr(stepped, 0)) + + +def test_only_stepped_slots_advance(): + optimizer = make_optimizer() + args = make_args() + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=10), resume_step=0) + install_slot_scheduler(args, optimizer, make_adapter(slot=1, num_step=10), resume_step=0) + + lr_by_slot = step_slot_schedulers(optimizer, {0: 5 * 64}) + + assert set(lr_by_slot) == {0} + assert slot_lr(optimizer, 0) == pytest.approx(LR / 2) + assert slot_lr(optimizer, 1) == pytest.approx(LR) # co-tenant untouched + + +def test_slot_reuse_installs_a_fresh_schedule(): + optimizer = make_optimizer() + args = make_args() + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=10), resume_step=0) + step_slot_schedulers(optimizer, {0: 5 * 64}) + + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=20), resume_step=0) + assert slot_lr(optimizer, 0) == pytest.approx(LR) # next tenant starts at the top + + +def test_warmup_ramps_from_init_lr(): + optimizer = make_optimizer() + args = make_args(lr_warmup_iters=2) # 2 adapter steps of warmup + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=10), resume_step=0) + + assert slot_lr(optimizer, 0) == pytest.approx(0.0) # init_lr + step_slot_schedulers(optimizer, {0: 64}) + assert slot_lr(optimizer, 0) == pytest.approx(LR / 2) # mid-warmup + step_slot_schedulers(optimizer, {0: 64}) + assert slot_lr(optimizer, 0) == pytest.approx(LR) # warmed up diff --git a/tests/fast/backends/megatron_utils/test_multi_lora_slot_cleanup.py b/tests/fast/backends/megatron_utils/test_multi_lora_slot_cleanup.py new file mode 100644 index 0000000000..93fcd3bd93 --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_multi_lora_slot_cleanup.py @@ -0,0 +1,91 @@ +"""zero_optimizer_state_for_adapter must reset a retired slot's Adam moments and step clock +(group-level FusedAdam or per-param torch AdamW) while leaving co-tenant slots untouched.""" + +import sys +import types +from types import SimpleNamespace + +import pytest +import torch + +from miles.backends.megatron_utils.multi_lora_utils import zero_optimizer_state_for_adapter + +MLL_MODULE = "megatron.bridge.peft.multi_lora_layers" + + +class FakeAdapter: + def __init__(self, params): + self._params = list(params) + + def parameters(self): + return self._params + + +class FakeMultiLoRALinear: + def __init__(self, adapters): + self.adapters = adapters + + +@pytest.fixture() +def rig(monkeypatch): + # Stub the lazily imported bridge module so the test needs no bridge build that ships multi-LoRA. + p0 = torch.nn.Parameter(torch.ones(4)) + p1 = torch.nn.Parameter(torch.ones(4)) + module = FakeMultiLoRALinear({0: FakeAdapter([p0]), 1: FakeAdapter([p1])}) + stub = types.ModuleType(MLL_MODULE) + stub.MultiLoRALinear = FakeMultiLoRALinear + stub._iter_multi_lora_modules = lambda model: [module] + monkeypatch.setitem(sys.modules, MLL_MODULE, stub) + return SimpleNamespace(p0=p0, p1=p1, model=object()) + + +def make_optimizer(groups, state): + inner = SimpleNamespace(param_groups=groups, state=state) + return inner, SimpleNamespace(chained_optimizers=[SimpleNamespace(optimizer=inner)]) + + +def test_group_level_fused_adam_clock_resets_only_for_the_retired_slot(rig): + inner, optimizer = make_optimizer( + groups=[ + {"params": [rig.p0], "miles_multi_lora_slot": 0, "step": 50}, + {"params": [rig.p1], "miles_multi_lora_slot": 1, "step": 50}, + ], + state={ + rig.p0: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4)}, + rig.p1: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4)}, + }, + ) + + zero_optimizer_state_for_adapter(optimizer, rig.model, 0) + + assert inner.param_groups[0]["step"] == 0 + assert inner.param_groups[1]["step"] == 50 # co-tenant slot untouched + assert float(inner.state[rig.p0]["exp_avg"].abs().sum()) == 0.0 + assert float(inner.state[rig.p0]["exp_avg_sq"].abs().sum()) == 0.0 + assert float(inner.state[rig.p1]["exp_avg"].abs().sum()) == 4.0 + + +def test_tensor_valued_group_clock_resets_in_place(rig): + step = torch.tensor(50) + inner, optimizer = make_optimizer( + groups=[{"params": [rig.p0], "miles_multi_lora_slot": 0, "step": step}], + state={rig.p0: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4)}}, + ) + + zero_optimizer_state_for_adapter(optimizer, rig.model, 0) + + assert int(step) == 0 # zeroed in place, no rebinding needed + + +def test_per_param_adamw_fallback_clock_resets(rig): + # torch.optim.AdamW keeps the clock per param; groups carry no "step". + inner, optimizer = make_optimizer( + groups=[{"params": [rig.p0], "miles_multi_lora_slot": 0}], + state={ + rig.p0: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4), "step": torch.tensor(50.0)}, + }, + ) + + zero_optimizer_state_for_adapter(optimizer, rig.model, 0) + + assert float(inner.state[rig.p0]["step"]) == 0.0 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 new file mode 100644 index 0000000000..6b22665766 --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py @@ -0,0 +1,48 @@ +"""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).""" + +import pytest +import torch + +from miles.backends.megatron_utils.multi_lora_utils import slice_lora_to_rank + + +def _padded(shape, live_rows=None, live_cols=None): + t = torch.zeros(shape) + if live_rows is not None: + t[:live_rows] = 1.0 + if live_cols is not None: + t[:, :live_cols] = 1.0 + return t + + +def test_lora_a_is_sliced_on_the_rank_dim(): + tensor = _padded((32, 8), live_rows=16) + out = slice_lora_to_rank("base_model.q_proj.lora_A.weight", tensor, 16) + assert out.shape == (16, 8) + assert torch.equal(out, tensor[:16]) + + +def test_lora_b_is_sliced_on_the_rank_dim(): + tensor = _padded((8, 32), live_cols=16) + out = slice_lora_to_rank("base_model.q_proj.lora_B.weight", tensor, 16) + assert out.shape == (8, 16) + assert torch.equal(out, tensor[:, :16]) + + +def test_nonzero_padding_is_rejected(): + # Live values beyond the adapter's rank mean the pad rows were trained — + # slicing would silently drop signal, so it must hard-fail instead. + tensor = torch.ones(32, 8) + with pytest.raises(AssertionError, match="padded dims are non-zero"): + slice_lora_to_rank("x.lora_A.weight", tensor, 16) + + +def test_full_rank_tensor_passes_through(): + tensor = torch.ones(16, 8) + assert slice_lora_to_rank("x.lora_A.weight", tensor, 16) is tensor + + +def test_non_lora_names_pass_through(): + tensor = torch.ones(32, 8) + assert slice_lora_to_rank("x.some_other.weight", tensor, 16) is tensor diff --git a/tests/fast/backends/training_utils/test_get_batch_multi_lora_cp.py b/tests/fast/backends/training_utils/test_get_batch_multi_lora_cp.py new file mode 100644 index 0000000000..cdc4e07ca1 --- /dev/null +++ b/tests/fast/backends/training_utils/test_get_batch_multi_lora_cp.py @@ -0,0 +1,134 @@ +"""CP=2 tests for get_batch's multi-LoRA per-adapter token counts; CUDA is stubbed to run on CPU.""" + +from types import SimpleNamespace + +import pytest +import torch + +import miles.backends.training_utils.cp_utils as cp_utils_mod +import miles.backends.training_utils.data as data_mod +from miles.backends.training_utils.cp_utils import slice_with_cp +from miles.backends.training_utils.data import get_batch + + +def _parallel_state(cp_rank: int, cp_size: int, tp_size: int = 1) -> SimpleNamespace: + return SimpleNamespace( + cp=SimpleNamespace(rank=cp_rank, size=cp_size), + tp=SimpleNamespace(size=tp_size), + ) + + +class _FakeIterator: + def __init__(self, batch: dict, n_adapters: int): + self._batch = batch + self.rollout_data = {"n_adapters": n_adapters} + + def get_next(self, keys): + return {key: self._batch[key] for key in keys} + + +KEYS = ["tokens", "loss_masks", "total_lengths", "response_lengths", "adapter_slots"] + +# 5 samples over 3 of 4 slots (sorted), ragged/odd lengths to exercise the +# per-sample zigzag padding (2 * cp_size chunking). +LENGTHS = [7, 13, 5, 9, 4] +SLOTS = [0, 0, 1, 2, 2] +N_ADAPTERS = 4 + + +def _make_batch(max_seqlen: int | None = None) -> dict: + # Token values start at 1 so zigzag pad (value 0) is distinguishable. + tokens = [torch.arange(1, length + 1, dtype=torch.long) for length in LENGTHS] + response_lengths = [max(1, length // 2) for length in LENGTHS] + batch = { + "tokens": tokens, + "loss_masks": [torch.ones(r, dtype=torch.int) for r in response_lengths], + "total_lengths": list(LENGTHS), + "response_lengths": response_lengths, + "adapter_slots": list(SLOTS), + } + if max_seqlen is not None: + batch["max_seq_lens"] = [max_seqlen] * len(LENGTHS) + return batch + + +@pytest.fixture(autouse=True) +def _stub_cuda(monkeypatch): + monkeypatch.setattr(torch.Tensor, "cuda", lambda self, *args, **kwargs: self, raising=False) + monkeypatch.setattr(torch.cuda, "current_device", lambda: "cpu", raising=False) + + +def _patch_state(monkeypatch, state: SimpleNamespace) -> None: + # data.get_batch and cp_utils.slice_with_cp each resolve the state themselves. + monkeypatch.setattr(data_mod, "get_parallel_state", lambda: state) + monkeypatch.setattr(cp_utils_mod, "get_parallel_state", lambda: state) + + +def _expected_thd_counts(state: SimpleNamespace, local_total: int) -> torch.Tensor: + sliced_lengths = [ + slice_with_cp(t, 0, "thd", parallel_state=state).numel() + for t in (torch.arange(1, length + 1, dtype=torch.long) for length in LENGTHS) + ] + expected = torch.zeros(N_ADAPTERS, dtype=torch.int32) + for slot, sliced_length in zip(SLOTS, sliced_lengths, strict=True): + expected[slot] += sliced_length + stream_pad = local_total - int(sum(sliced_lengths)) + assert stream_pad >= 0, "rank-local stream shorter than the sliced samples" + expected[SLOTS[-1]] += stream_pad + return expected + + +@pytest.mark.parametrize("cp_rank", [0, 1]) +def test_thd_cp2_adapter_token_counts(monkeypatch, cp_rank): + state = _parallel_state(cp_rank, cp_size=2) + _patch_state(monkeypatch, state) + + out = get_batch(_FakeIterator(_make_batch(), N_ADAPTERS), KEYS, pad_multiplier=8, qkv_format="thd") + + counts = out["adapter_token_counts"] + local_total = out["tokens"].numel() + assert counts.dtype == torch.int32 + assert counts.tolist() == _expected_thd_counts(state, local_total).tolist() + assert int(counts.sum()) == local_total, "counts must cover every rank-local token incl. padding" + + +def test_thd_cp2_counts_identical_across_ranks(monkeypatch): + # Zigzag gives each rank the same padded share of every sample, so the + # grouped-GEMM routing counts must not depend on the CP rank. + per_rank = [] + for cp_rank in (0, 1): + state = _parallel_state(cp_rank, cp_size=2) + _patch_state(monkeypatch, state) + out = get_batch(_FakeIterator(_make_batch(), N_ADAPTERS), KEYS, pad_multiplier=8, qkv_format="thd") + per_rank.append(out["adapter_token_counts"].tolist()) + assert per_rank[0] == per_rank[1] + + +def test_unsorted_adapter_slots_rejected(monkeypatch): + _patch_state(monkeypatch, _parallel_state(cp_rank=0, cp_size=2)) + batch = _make_batch() + batch["adapter_slots"] = [2, 0, 1, 0, 2] + with pytest.raises(AssertionError, match="not sorted"): + get_batch(_FakeIterator(batch, N_ADAPTERS), KEYS, pad_multiplier=8, qkv_format="thd") + + +def test_bshd_cp2_tokens_are_single_sliced(monkeypatch): + # Regression: the bshd path used to slice tokens twice under CP>1, zeroing the back half of each rank's tokens. + max_seqlen = 16 # divisible by 2 * cp_size + state = _parallel_state(cp_rank=0, cp_size=2) + _patch_state(monkeypatch, state) + + out = get_batch( + _FakeIterator(_make_batch(max_seqlen=max_seqlen), N_ADAPTERS), + KEYS + ["max_seq_lens"], + pad_multiplier=8, + qkv_format="bshd", + ) + + expected = torch.stack( + [ + slice_with_cp(torch.arange(1, length + 1, dtype=torch.long), 0, "bshd", max_seqlen, parallel_state=state) + for length in LENGTHS + ] + ) + assert torch.equal(out["tokens"], expected)