Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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"])
Comment thread
yushengsu-thu marked this conversation as resolved.

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