Skip to content
86 changes: 85 additions & 1 deletion miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from miles.utils.ft_utils.indep_dp import IndepDPInfo
from miles.utils.hf_config import load_hf_config
from miles.utils.memory_utils import clear_memory, print_memory
from miles.utils.multi_lora import is_multi_lora_enabled
from miles.utils.processing_utils import load_tokenizer
from miles.utils.ray_utils import Box
from miles.utils.reloadable_process_group import destroy_process_groups, monkey_patch_torch_dist, reload_process_groups
Expand Down Expand Up @@ -235,6 +236,12 @@ def init(
is_lora=is_lora_enabled(args),
)

# Adapters currently loaded into Megatron slots on this rank.
self.loaded_adapters: dict[str, object] = {}
# Adapters with stale engine-side weights (newly loaded or just trained);
# consumed by the next update_weights. Identical on every rank.
self._multi_lora_pending_push: set[str] = set()

# empty cache after initialization
clear_memory()

Expand Down Expand Up @@ -513,12 +520,66 @@ def train_actor(
logger.info(f"Updating ref model at rollout_id {rollout_id}")
self.weights_backuper.backup("ref")

if train_step_outcome == TrainStepOutcome.NORMAL and is_multi_lora_enabled(self.args):
from miles.backends.megatron_utils.multi_lora_utils import commit_trained_batch

commit_trained_batch(rollout_data, rollout_id, self._multi_lora_pending_push)

log_perf_data(rollout_id, self.args, extra_metrics=self.weight_updater.pop_metrics())

self._heartbeat.bump()
return train_step_outcome

@with_logs
@timer
def reconcile_adapters(self) -> None:
"""Load adapters the controller wants served; retire deregistered ones, dropping their untrained tail."""
if not is_multi_lora_enabled(self.args):
return
from miles.backends.megatron_utils.multi_lora_utils import cleanup_adapters as _cleanup_adapters
from miles.backends.megatron_utils.multi_lora_utils import load_adapters as _load_adapters
from miles.ray.multi_lora.controller import get_multi_lora_controller

broadcast_buffer = [None]
if is_first_replica_megatron_main_rank():
controller = get_multi_lora_controller()
ray.get(controller.retire_adapters.remote())
broadcast_buffer[0] = ray.get(controller.snapshot.remote())
if dist.is_initialized():
dist.broadcast_object_list(broadcast_buffer, src=0, group=get_gloo_group())
snapshot = broadcast_buffer[0]
should_be_loaded = {**snapshot["active"], **snapshot["pending"], **snapshot["retiring"]}
cleanup_names = set(snapshot["cleanup"])

loaded_names = set(self.loaded_adapters)
# Sorted so per-adapter collectives (checkpoint export) run in the same
# order on every rank; set iteration order is process-specific.
adapters_to_load = sorted(
(adapter for name, adapter in should_be_loaded.items() if name not in loaded_names),
key=lambda adapter: adapter.name,
)
adapters_to_clean_up = sorted(
(self.loaded_adapters[n] for n in loaded_names if n in cleanup_names or n not in should_be_loaded),
key=lambda adapter: adapter.name,
)
if adapters_to_load:
_load_adapters(self.args, self.model, self.optimizer, adapters_to_load)
for adapter in adapters_to_load:
self.loaded_adapters[adapter.name] = adapter
self._multi_lora_pending_push.add(adapter.name)
self.weights_backuper.backup("actor")
if adapters_to_clean_up:
_cleanup_adapters(self.args, self.model, self.optimizer, adapters_to_clean_up)
for adapter in adapters_to_clean_up:
self.loaded_adapters.pop(adapter.name, None)
self._multi_lora_pending_push.discard(adapter.name)
self.weights_backuper.backup("actor")

# Deregistered before ever being loaded: nothing to save or clear.
if is_first_replica_megatron_main_rank():
for name in cleanup_names - loaded_names:
ray.get(get_multi_lora_controller().free_slot.remote(name))

@timer
def save_model(self, rollout_id: int, force_sync: bool = False) -> None:
self._heartbeat.bump()
Expand All @@ -534,7 +595,15 @@ def save_model(self, rollout_id: int, force_sync: bool = False) -> None:

maybe_finalize_async_save(blocking=True)

save(rollout_id, self.model, self.optimizer, self.opt_param_scheduler)
if is_multi_lora_enabled(self.args):
from miles.backends.megatron_utils.multi_lora_utils import save_due_adapter_checkpoints

if not save_due_adapter_checkpoints(self.args, self.model):
if self.args.offload_train:
destroy_process_groups()
return
Comment thread
yushengsu-thu marked this conversation as resolved.
else:
save(rollout_id, self.model, self.optimizer, self.opt_param_scheduler)

if force_sync and self.args.async_save:
maybe_finalize_async_save(blocking=True)
Expand All @@ -549,6 +618,7 @@ def save_model(self, rollout_id: int, force_sync: bool = False) -> None:
maybe_finalize_async_save(blocking=True)

from megatron.training.checkpointing import get_checkpoint_name

from miles.utils.misc import load_function

checkpoint_dir = get_checkpoint_name(self.args.save, rollout_id, return_base_dir=True)
Expand Down Expand Up @@ -598,11 +668,25 @@ def update_weights(self, info: "EnginesAndLock") -> None:
destroy_process_groups()
return

version_update_names: list[str] = []
if is_multi_lora_enabled(self.args):
from miles.backends.megatron_utils.multi_lora_utils import select_adapters_to_push

self.weight_updater.multi_lora_adapters, version_update_names = select_adapters_to_push(
self.loaded_adapters, self._multi_lora_pending_push, has_new_engines
)

with torch_memory_saver.disable() if self.args.offload_train else nullcontext():
print_memory("before update_weights")
self.weight_updater.update_weights()
print_memory("after update_weights")

if is_multi_lora_enabled(self.args):
from miles.backends.megatron_utils.multi_lora_utils import commit_weight_push

self._multi_lora_pending_push.clear()
commit_weight_push(version_update_names, self._is_first_replica_megatron_main_rank)

if self.args.ci_test and len(rollout_engines) > 0 and not is_lora_enabled(self.args):
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
Expand Down
3 changes: 3 additions & 0 deletions miles/backends/megatron_utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
16 changes: 14 additions & 2 deletions miles/backends/megatron_utils/bridge_lora_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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,
Expand Down
101 changes: 77 additions & 24 deletions miles/backends/megatron_utils/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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()

Expand Down
Loading
Loading