Skip to content
Draft
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
47 changes: 43 additions & 4 deletions miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,11 @@
from ...utils.profile_utils import TrainProfiler
from ...utils.tensor_backper import TensorBackuper
from ..training_utils.data import DataIterator, get_data_iterator, get_rollout_data, sync_actor_critic_data
from ..training_utils.higgs_policy import (
is_higgs_policy_enabled,
validate_higgs_logprob_parity,
validate_higgs_weight_versions,
)
from ..training_utils.log_utils import log_cpu_memory, log_perf_data, log_rollout_data
from ..training_utils.loss import compute_advantages_and_returns, get_log_probs_and_entropy, get_values
from ..training_utils.parallel import get_parallel_state
Expand All @@ -39,9 +44,6 @@
from .parallel import verify_megatron_parallel_state
from .replay_utils import register_replay_list_moe
from .update_weight.common import named_params_and_buffers
from .update_weight.update_weight_from_distributed.broadcast import UpdateWeightFromDistributed
from .update_weight.update_weight_from_distributed.p2p import UpdateWeightP2P
from .update_weight.update_weight_from_tensor import UpdateWeightFromTensor

if TYPE_CHECKING:
from miles.ray.rollout.rollout_manager import EnginesAndLock
Expand Down Expand Up @@ -173,16 +175,22 @@ def init(
self.args.vocab_size = self.tokenizer.vocab_size

if self.args.colocate:
from .update_weight.update_weight_from_tensor import UpdateWeightFromTensor

update_weight_cls = UpdateWeightFromTensor
else:
if self.args.update_weight_transfer_mode == "broadcast":
from .update_weight.update_weight_from_distributed.broadcast import UpdateWeightFromDistributed

update_weight_cls = UpdateWeightFromDistributed
elif self.args.update_weight_transfer_mode == "disk-delta":
# Lazy import: keeps the delta deps (numpy/zstandard/xxhash) off the other paths.
from .update_weight.update_weight_from_distributed.delta import UpdateWeightFromDiskDelta

update_weight_cls = UpdateWeightFromDiskDelta
else:
from .update_weight.update_weight_from_distributed.p2p import UpdateWeightP2P

update_weight_cls = UpdateWeightP2P
self.weight_updater = update_weight_cls(
self.args,
Expand Down Expand Up @@ -322,6 +330,12 @@ def _use_rollout_replay(self, m) -> bool:
return getattr(self.args, f"use_rollout_{m.name}_replay", False)

def train_actor(self, rollout_id: int, rollout_data: RolloutBatch) -> None:
higgs_policy = is_higgs_policy_enabled(self.args)
if higgs_policy:
validate_higgs_weight_versions(
rollout_data.get("weight_versions"),
trainer_weight_version=self.weight_updater.weight_version,
)
# Create data iterator for log_probs and train.
data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data)

Expand Down Expand Up @@ -364,7 +378,7 @@ def train_actor(self, rollout_id: int, rollout_data: RolloutBatch) -> None:
)
)
self._switch_model("old_actor" if self.args.keep_old_actor else "actor")
if not self.args.use_rollout_logprobs or self.args.get_mismatch_metrics:
if higgs_policy or not self.args.use_rollout_logprobs or self.args.get_mismatch_metrics:
for m in all_replay_managers:
if m.enabled:
if self._use_rollout_replay(m):
Expand All @@ -381,6 +395,25 @@ def train_actor(self, rollout_id: int, rollout_data: RolloutBatch) -> None:
for m in all_replay_managers:
if self._use_rollout_replay(m):
m.clear_all_forward()
if higgs_policy:
parity = validate_higgs_logprob_parity(
rollout_data["action_traces"],
rollout_data["log_probs"],
atol=self.args.higgs_logprob_parity_atol,
)
log = logger.info if parity["within_tolerance"] else logger.warning
log(
"Higgs pre-optimizer joint-logprob parity %s: "
"max_abs_diff=%.6g mean_abs_diff=%.6g warning_atol=%.6g",
(
"within measured tolerance"
if parity["within_tolerance"]
else "exceeded measured tolerance"
),
parity["max_abs_diff"],
parity["mean_abs_diff"],
self.args.higgs_logprob_parity_atol,
)

if self.args.use_critic:
sync_actor_critic_data(
Expand Down Expand Up @@ -466,6 +499,12 @@ def save_model(self, rollout_id: int, force_sync: bool = False) -> None:
if self.args.offload_train:
destroy_process_groups()

@timer
def disconnect_rollout_engines(self) -> None:
disconnect = getattr(self.weight_updater, "disconnect_rollout_engines", None)
if disconnect is not None:
disconnect()

@timer
def update_weights(self, info: "EnginesAndLock") -> None:
if self.args.debug_train_only or self.args.debug_rollout_only:
Expand Down
119 changes: 107 additions & 12 deletions miles/backends/megatron_utils/checkpoint.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
import logging
import os
import re
from collections import Counter
from pathlib import Path

import torch
import torch.distributed as dist

# TODO: may need to copy those 2 functions and do refactoring.
Expand Down Expand Up @@ -97,23 +99,94 @@ def _init_from_local_shards_and_global_metadata( # type: ignore[override]
__all__ = ["save_checkpoint", "save_checkpoint_with_lora", "load_checkpoint"]


def _normalize_torch_optimizer_steps_for_checkpoint_load(optimizer) -> None:
"""Make disposable native-Adam load templates internally consistent."""

states = []
for wrapped_optimizer in getattr(optimizer, "chained_optimizers", (optimizer,)):
torch_optimizer = getattr(wrapped_optimizer, "optimizer", None)
if torch_optimizer is None:
continue
serialized_state = torch_optimizer.state_dict().get("state", {})
if not serialized_state:
initialize_states = getattr(wrapped_optimizer, "_init_optimizer_states_with_dummy_values", None)
if initialize_states is not None:
logger.info("Initializing temporary native optimizer state for Higgs checkpoint load")
initialize_states()
serialized_state = torch_optimizer.state_dict().get("state", {})
states.extend(state for state in serialized_state.values() if "step" in state)

step_counts = Counter(float(state["step"].item()) for state in states)
if len(step_counts) <= 1:
return

logger.warning(
"Normalizing divergent temporary Torch optimizer steps before Higgs checkpoint load: %s",
dict(sorted(step_counts.items())),
)
for state in states:
step = state["step"]
if isinstance(step, torch.Tensor):
step.zero_()
else:
state["step"] = 0


def _load_higgs_megatron_checkpoint_with_consistent_optimizer_steps(checkpoint_optimizer, **load_kwargs):
"""Normalize native-Adam template steps at the point Megatron materializes them."""

patched_state_dicts = []
for wrapped_optimizer in getattr(checkpoint_optimizer, "chained_optimizers", (checkpoint_optimizer,)):
original_state_dict = wrapped_optimizer.state_dict

def state_dict(_optimizer=wrapped_optimizer, _original=original_state_dict):
_normalize_torch_optimizer_steps_for_checkpoint_load(_optimizer)
return _original()

patched_state_dicts.append((wrapped_optimizer, original_state_dict))
wrapped_optimizer.state_dict = state_dict
try:
return _load_checkpoint_megatron(**load_kwargs)
finally:
for wrapped_optimizer, original_state_dict in patched_state_dicts:
wrapped_optimizer.state_dict = original_state_dict


def load_checkpoint(ddp_model, optimizer, opt_param_scheduler, checkpointing_context, skip_load_to_model_and_opt):
# ref: how megatron `load_checkpoint` gets directory
args = get_args()
load_path = args.load

from miles.backends.training_utils.higgs_policy import is_higgs_policy_enabled

if is_higgs_policy_enabled(args):
from .higgs_checkpoint import resolve_higgs_checkpoint_path

load_path = str(resolve_higgs_checkpoint_path(load_path))
args.load = load_path

assert Path(load_path).exists() and _is_dir_nonempty(
load_path
), f"{args.load=} does not exist or is an empty directory. Did you specify the wrong folder?"

if _is_megatron_checkpoint(load_path):
result = _load_checkpoint_megatron(
ddp_model=ddp_model,
optimizer=optimizer,
opt_param_scheduler=opt_param_scheduler,
checkpointing_context=checkpointing_context,
skip_load_to_model_and_opt=skip_load_to_model_and_opt,
)
if is_higgs_policy_enabled(args) and optimizer is not None:
result = _load_higgs_megatron_checkpoint_with_consistent_optimizer_steps(
optimizer,
ddp_model=ddp_model,
optimizer=optimizer,
opt_param_scheduler=opt_param_scheduler,
checkpointing_context=checkpointing_context,
skip_load_to_model_and_opt=skip_load_to_model_and_opt,
)
else:
result = _load_checkpoint_megatron(
ddp_model=ddp_model,
optimizer=optimizer,
opt_param_scheduler=opt_param_scheduler,
checkpointing_context=checkpointing_context,
skip_load_to_model_and_opt=skip_load_to_model_and_opt,
)
else:
result = _load_checkpoint_hf(
ddp_model=ddp_model,
Expand Down Expand Up @@ -172,14 +245,36 @@ def _is_megatron_checkpoint(path: str | Path) -> bool:


def _load_checkpoint_hf(ddp_model, optimizer, args, load_path: str):
assert args.megatron_to_hf_mode == "bridge", "Only bridge mode is supported for loading HF checkpoint"
from megatron.bridge import AutoBridge
from miles.backends.training_utils.higgs_policy import is_higgs_policy_enabled, validate_higgs_single_device_config

logger.info(f"Load checkpoint from HuggingFace model into Megatron (path={load_path})")

with megatron_bridge_utils.patch_megatron_model(ddp_model):
bridge = AutoBridge.from_hf_pretrained(load_path, trust_remote_code=True)
bridge.load_hf_weights(ddp_model)
if is_higgs_policy_enabled(args):
from megatron.core.utils import unwrap_model

from .higgs_checkpoint import HIGGS_TEXT_VOCAB_SIZE, load_higgs_policy_checkpoint

validate_higgs_single_device_config(args)
if args.megatron_to_hf_mode != "raw":
raise ValueError("Higgs HF loading requires megatron_to_hf_mode='raw'")
if args.vocab_size != HIGGS_TEXT_VOCAB_SIZE or args.padded_vocab_size != HIGGS_TEXT_VOCAB_SIZE:
raise ValueError(
"Higgs HF loading requires vocab_size=padded_vocab_size="
f"{HIGGS_TEXT_VOCAB_SIZE}, got vocab_size={args.vocab_size!r} "
f"and padded_vocab_size={args.padded_vocab_size!r}"
)
unwrapped_model = unwrap_model(ddp_model)
if len(unwrapped_model) != 1:
raise ValueError("the initial Higgs raw loader requires exactly one Megatron model chunk")
load_higgs_policy_checkpoint(unwrapped_model[0], load_path)
else:
if args.megatron_to_hf_mode != "bridge":
raise ValueError("only bridge mode is supported for loading a generic HF checkpoint")
from megatron.bridge import AutoBridge

with megatron_bridge_utils.patch_megatron_model(ddp_model):
bridge = AutoBridge.from_hf_pretrained(load_path, trust_remote_code=True)
bridge.load_hf_weights(ddp_model)

# Copied from Megatron-core :: load_checkpoint (with simplifications)
if (args.fp16 or args.bf16) and optimizer is not None:
Expand Down
Loading