Skip to content
Open
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
38 changes: 31 additions & 7 deletions slime/backends/megatron_utils/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -179,6 +179,25 @@ def _reinitialize_critic_output_layer(args: Namespace, model: Sequence[DDP]) ->
output_layer.bias.data.zero_()


@contextmanager
def _hide_critic_output_layer_during_policy_checkpoint_load(model: Sequence[DDP], enabled: bool):
"""Omit a policy-incompatible critic value head from a checkpoint load request."""
hidden_parameters = []
if enabled:
for _chunk_id, output_layer in _iter_critic_output_layers(model):
for name in ("weight", "bias"):
parameter = getattr(output_layer, name, None)
if parameter is None:
continue
hidden_parameters.append((output_layer, name, parameter))
output_layer.register_parameter(name, None)
try:
yield
finally:
for output_layer, name, parameter in hidden_parameters:
output_layer.register_parameter(name, parameter)


def get_optimizer_param_scheduler(args: Namespace, optimizer: MegatronOptimizer) -> OptimizerParamScheduler:
"""Create and configure the optimizer learning-rate/weight-decay scheduler.

Expand Down Expand Up @@ -991,13 +1010,18 @@ def initialize_model_and_optimizer(
model[0].role = role
reinit_critic_output_layer = _critic_output_layer_needs_reinit(args, model, role)
clear_memory()
iteration, _ = load_checkpoint(
model,
optimizer,
opt_param_scheduler,
checkpointing_context={},
skip_load_to_model_and_opt=False,
)
# Policy checkpoints can omit an independent output weight when embeddings
# are tied, and their language-model head is incompatible with the critic's
# scalar value head in either case. Keep strict loading for every shared
# tensor while excluding only the value head that will be initialized below.
with _hide_critic_output_layer_during_policy_checkpoint_load(model, reinit_critic_output_layer):
iteration, _ = load_checkpoint(
model,
optimizer,
opt_param_scheduler,
checkpointing_context={},
skip_load_to_model_and_opt=False,
)
if reinit_critic_output_layer:
_reinitialize_critic_output_layer(args, model)
if (args.fp16 or args.bf16) and optimizer is not None:
Expand Down
62 changes: 62 additions & 0 deletions tests/test_critic_value_head_load.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
import ast
from collections.abc import Sequence
from contextlib import contextmanager
from pathlib import Path

import pytest
import torch

NUM_GPUS = 0
MODEL_PATH = Path(__file__).resolve().parents[1] / "slime/backends/megatron_utils/model.py"


def _load_hide_output_layer_function(output_layer):
tree = ast.parse(MODEL_PATH.read_text())
function_node = next(
node
for node in tree.body
if isinstance(node, ast.FunctionDef) and node.name == "_hide_critic_output_layer_during_policy_checkpoint_load"
)
function_node.decorator_list = []
module = ast.fix_missing_locations(ast.Module(body=[function_node], type_ignores=[]))
namespace = {
"DDP": object,
"Sequence": Sequence,
"_iter_critic_output_layers": lambda _model: [(0, output_layer)],
}
exec(compile(module, str(MODEL_PATH), "exec"), namespace)
return contextmanager(namespace["_hide_critic_output_layer_during_policy_checkpoint_load"])


@pytest.mark.unit
def test_policy_checkpoint_load_temporarily_hides_full_critic_output_layer():
output_layer = torch.nn.Linear(4, 1, bias=True)
original_weight = output_layer.weight
original_bias = output_layer.bias
optimizer = torch.optim.SGD(output_layer.parameters(), lr=0.1)
hide_output_layer = _load_hide_output_layer_function(output_layer)

with hide_output_layer([object()], enabled=True):
assert output_layer.weight is None
assert output_layer.bias is None
assert dict(output_layer.named_parameters()) == {}

assert output_layer.weight is original_weight
assert output_layer.bias is original_bias
assert optimizer.param_groups[0]["params"][0] is original_weight
assert optimizer.param_groups[0]["params"][1] is original_bias


@pytest.mark.unit
def test_policy_checkpoint_load_restores_critic_output_layer_after_failure():
output_layer = torch.nn.Linear(4, 1, bias=True)
original_weight = output_layer.weight
original_bias = output_layer.bias
hide_output_layer = _load_hide_output_layer_function(output_layer)

with pytest.raises(RuntimeError, match="checkpoint load failed"):
with hide_output_layer([object()], enabled=True):
raise RuntimeError("checkpoint load failed")

assert output_layer.weight is original_weight
assert output_layer.bias is original_bias
Loading