From 9ce0b5c8024cfe34cebe3a178e1ce73f70ee374c Mon Sep 17 00:00:00 2001 From: Yizhuo Li <143763751+Dodojordi@users.noreply.github.com> Date: Sun, 9 Aug 2026 07:44:00 +0000 Subject: [PATCH] fix: load critic from policy checkpoints without value head --- slime/backends/megatron_utils/model.py | 38 +++++++++++++--- tests/test_critic_value_head_load.py | 62 ++++++++++++++++++++++++++ 2 files changed, 93 insertions(+), 7 deletions(-) create mode 100644 tests/test_critic_value_head_load.py diff --git a/slime/backends/megatron_utils/model.py b/slime/backends/megatron_utils/model.py index 1ad6cd7957..1410ad4b9e 100644 --- a/slime/backends/megatron_utils/model.py +++ b/slime/backends/megatron_utils/model.py @@ -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. @@ -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: diff --git a/tests/test_critic_value_head_load.py b/tests/test_critic_value_head_load.py new file mode 100644 index 0000000000..a3f220a0c3 --- /dev/null +++ b/tests/test_critic_value_head_load.py @@ -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