diff --git a/tests/special_sanity/check_device_api_usage.py b/tests/special_sanity/check_device_api_usage.py index 0f924a8fa..471eba8f3 100644 --- a/tests/special_sanity/check_device_api_usage.py +++ b/tests/special_sanity/check_device_api_usage.py @@ -25,6 +25,7 @@ # directory or file path must contain keyword ".cuda" or "cuda" CUDA_KEYWORD_CHECK_WHITELIST = [ "verl_omni/workers/engine/fsdp/diffusers_impl.py", # appear in default device_name + "verl_omni/workers/engine/fsdp/distillation_impl.py", # device=[...] registry declaration "verl_omni/trainer/diffusion/ray_diffusion_trainer.py", # appear in default device_name "verl_omni/workers/engine/fsdp/omni_impl.py", # device=[...] registry declaration "verl_omni/workers/engine/veomni/diffusion_impl.py", # device=[...] registry declaration diff --git a/tests/trainer/diffusion/test_distillation_checkpoint_on_cpu.py b/tests/trainer/diffusion/test_distillation_checkpoint_on_cpu.py new file mode 100644 index 000000000..1ed705085 --- /dev/null +++ b/tests/trainer/diffusion/test_distillation_checkpoint_on_cpu.py @@ -0,0 +1,205 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""CPU tests for atomic multi-role distillation checkpoint orchestration.""" + +import os +import random +from dataclasses import replace + +import numpy as np +import pytest +import torch +from omegaconf import OmegaConf + +from verl_omni.trainer.diffusion.distillation.contracts import PhaseRequest +from verl_omni.trainer.diffusion.distillation.controller import ( + DistillationTrainerController, + FakeBatchProvider, + FakeDistillationHooks, + FakePhaseExecutor, +) +from verl_omni.trainer.diffusion.distillation.ray_trainer import DistillationBatchProvider, DistillationRayTrainer +from verl_omni.trainer.diffusion.distillation.recipes import build_plan + + +class CheckpointExecutor(FakePhaseExecutor): + def __init__(self, *, fail_save=False): + super().__init__() + self.role_state = { + "student": 1.0, + "fake_score": 2.0, + "student_ema": 1.5, + "student_optimizer": 3, + "fake_optimizer": 4, + "student_scheduler": 5, + "fake_scheduler": 6, + } + self.fail_save = fail_save + self.loaded_path = None + + def save_checkpoint(self, local_path, global_step): + os.makedirs(local_path, exist_ok=True) + torch.save({"global_step": global_step, "role_state": self.role_state}, os.path.join(local_path, "roles.pt")) + if self.fail_save: + raise RuntimeError("injected save failure") + + def load_checkpoint(self, local_path): + state = torch.load(os.path.join(local_path, "roles.pt"), weights_only=False) + self.role_state = state["role_state"] + self.loaded_path = local_path + + +class StatefulLoader: + def __init__(self): + self.position = 0 + + def state_dict(self): + return {"position": self.position} + + def load_state_dict(self, state): + self.position = state["position"] + + +class TestDistillationBatchProvider: + def test_reuse_student_returns_cached_batch_without_advancing(self): + batches = [ + {"values": torch.tensor([[1.0]]), "responses": torch.tensor([[9.0]])}, + {"values": torch.tensor([[2.0]]), "responses": torch.tensor([[8.0]])}, + ] + provider = DistillationBatchProvider(batches) + student = PhaseRequest("student", 0, 0, "fresh", ("student",), True) + reused = PhaseRequest("fake_score", 0, 0, "reuse_student", ("fake_score",), False) + fresh = PhaseRequest("fake_score", 0, 1, "fresh", ("fake_score",), False) + student_batch = provider.next(student) + reused_batch = provider.next(reused) + fresh_batch = provider.next(fresh) + torch.testing.assert_close(student_batch["values"], reused_batch["values"]) + torch.testing.assert_close(fresh_batch["values"], torch.tensor([[2.0]])) + assert "responses" not in student_batch + + def test_reuse_before_student_fails(self): + provider = DistillationBatchProvider([{"values": torch.tensor([[1.0]])}]) + request = PhaseRequest("fake_score", 0, 0, "reuse_student", ("fake_score",), False) + with pytest.raises(RuntimeError, match="before a student batch"): + provider.next(request) + + +class TestDistillationCheckpoint: + @staticmethod + def make_trainer(tmp_path, *, fail_save=False): + plan = build_plan( + "dmd2", + {"model_path": "/m", "fake_update_ratio": 1}, + frozenset({"distribution_matching"}), + ) + executor = CheckpointExecutor(fail_save=fail_save) + hooks = FakeDistillationHooks() + controller = DistillationTrainerController( + plan=plan, + executor=executor, + batch_provider=FakeBatchProvider(num_batches=10), + hooks=hooks, + ) + controller.run_cycle() + + trainer = DistillationRayTrainer( + plan=plan, executor=executor, batch_provider=FakeBatchProvider(10), hooks=hooks + ) + trainer.controller_instance = controller + trainer._production = True + trainer.global_steps = controller.counters.global_step + trainer.train_dataloader = StatefulLoader() + trainer.train_dataloader.position = 7 + trainer.config = OmegaConf.create( + { + "trainer": { + "default_local_dir": str(tmp_path), + "default_hdfs_dir": None, + "resume_mode": "auto", + "resume_from_path": None, + } + } + ) + return trainer, executor + + def test_round_trip_restores_roles_counters_dataloader_and_rng(self, tmp_path): + trainer, executor = self.make_trainer(tmp_path) + random.seed(11) + np.random.seed(12) + torch.manual_seed(13) + trainer._save_checkpoint() + expected_random = random.random() + expected_numpy = float(np.random.random()) + expected_torch = float(torch.rand(())) + + executor.role_state = {"corrupt": True} + trainer.controller.counters.global_step = 99 + trainer.controller.counters.optimizer_steps = {"student": 99} + trainer.train_dataloader.position = 99 + random.seed(101) + np.random.seed(102) + torch.manual_seed(103) + + restored_step = trainer._load_checkpoint() + assert restored_step == 1 + assert trainer.controller.counters.global_step == 1 + assert trainer.controller.counters.optimizer_steps == {"student": 1, "fake_score": 1} + assert trainer.controller.counters.completed_cycles == 1 + assert trainer.train_dataloader.position == 7 + assert executor.role_state["student_optimizer"] == 3 + assert executor.role_state["fake_scheduler"] == 6 + assert random.random() == expected_random + assert float(np.random.random()) == expected_numpy + assert float(torch.rand(())) == expected_torch + + checkpoint = tmp_path / "global_step_1" + assert (checkpoint / "manifest.json").is_file() + assert (checkpoint / "trainer_state.pt").is_file() + assert (checkpoint / "data.pt").is_file() + assert (checkpoint / "rng.pt").is_file() + assert executor.loaded_path == str(checkpoint / "workers") + + def test_failed_save_never_publishes_a_checkpoint(self, tmp_path): + trainer, _ = self.make_trainer(tmp_path, fail_save=True) + with pytest.raises(RuntimeError, match="injected save failure"): + trainer._save_checkpoint() + assert not (tmp_path / "global_step_1").exists() + assert not list(tmp_path.glob(".global_step_1_*")) + assert not (tmp_path / "latest_checkpointed_iteration.txt").exists() + + def test_equivalent_plan_mappings_have_identical_fingerprints(self, tmp_path): + trainer, _ = self.make_trainer(tmp_path) + fingerprint = trainer.checkpoint_fingerprint() + trainer.plan = replace(trainer.plan, objective=dict(reversed(list(trainer.plan.objective.items())))) + assert trainer.checkpoint_fingerprint() == fingerprint + + def test_changed_plan_is_rejected_before_worker_restore(self, tmp_path): + trainer, executor = self.make_trainer(tmp_path) + trainer._save_checkpoint() + trainer.plan = build_plan( + "dmd2", + {"model_path": "/m", "fake_update_ratio": 2}, + frozenset({"distribution_matching"}), + ) + with pytest.raises(ValueError, match="does not match the active run"): + trainer._load_checkpoint() + assert executor.loaded_path is None + + def test_incomplete_checkpoint_is_rejected(self, tmp_path): + trainer, _ = self.make_trainer(tmp_path) + incomplete = tmp_path / "global_step_1" + incomplete.mkdir() + (tmp_path / "latest_checkpointed_iteration.txt").write_text("1") + with pytest.raises(FileNotFoundError, match="Incomplete distillation checkpoint"): + trainer._load_checkpoint() diff --git a/tests/trainer/diffusion/test_distillation_config_on_cpu.py b/tests/trainer/diffusion/test_distillation_config_on_cpu.py index a1d4788c9..89f607f29 100644 --- a/tests/trainer/diffusion/test_distillation_config_on_cpu.py +++ b/tests/trainer/diffusion/test_distillation_config_on_cpu.py @@ -36,6 +36,12 @@ def test_defaults_select_dmd2_without_enabling_opd(self): assert config.distribution_matching.recipe == "dmd2" assert config.distribution_matching.profile is None assert config.distribution_matching.fake_update_ratio is None + assert config.distribution_matching.role_storage == "shared_base_adapters" + assert config.distribution_matching.student_micro_batch_size_per_gpu == 1 + assert config.distribution_matching.fake_score_micro_batch_size_per_gpu == 1 + assert config.distribution_matching.ema_decay == pytest.approx(0.999) + assert config.distribution_matching.ema_start_step == 0 + assert config.distribution_matching.fake_score_optim.lr == pytest.approx(2e-5) @pytest.mark.parametrize( "kwargs,error", @@ -47,6 +53,12 @@ def test_defaults_select_dmd2_without_enabling_opd(self): ({"rollout_strategy": "typo"}, "Invalid rollout_strategy"), ({"data_mode": "typo"}, "Invalid data_mode"), ({"export_role": "teacher_score"}, "Invalid export_role"), + ({"role_storage": "remote"}, "Invalid role_storage"), + ({"student_micro_batch_size_per_gpu": 0}, "greater than 0"), + ({"fake_score_micro_batch_size_per_gpu": 0}, "greater than 0"), + ({"ema_decay": -0.1}, "ema_decay"), + ({"ema_decay": 1.1}, "ema_decay"), + ({"ema_start_step": -1}, "non-negative"), ], ) def test_invalid_values_fail_closed(self, kwargs, error): @@ -219,6 +231,7 @@ def test_cli_distribution_matching_overrides_do_not_enable_opd(self): cfg = self._compose( [ "algorithm.trainer_type=distillation", + "algorithm.sample_source=offline", "distillation.distribution_matching.recipe=dmd2", "distillation.distribution_matching.fake_update_ratio=2", "distillation.distribution_matching.rollout_strategy=consistency_renoise", @@ -228,6 +241,7 @@ def test_cli_distribution_matching_overrides_do_not_enable_opd(self): assert config.enabled is False assert config.distribution_matching.fake_update_ratio == 2 assert config.distribution_matching.rollout_strategy == "consistency_renoise" + assert config.distribution_matching.fake_score_optim.lr == pytest.approx(2e-5) def test_composed_config_builds_validated_plan(self): from verl_omni.trainer.diffusion.distillation.recipes import build_plan_from_config @@ -235,6 +249,7 @@ def test_composed_config_builds_validated_plan(self): cfg = self._compose( [ "algorithm.trainer_type=distillation", + "algorithm.sample_source=offline", "actor_rollout_ref.model.path=/m", "distillation.distribution_matching.fake_update_ratio=2", ] @@ -250,6 +265,7 @@ def test_null_overrides_use_each_recipe_default(self): cfg = self._compose( [ "algorithm.trainer_type=distillation", + "algorithm.sample_source=offline", "actor_rollout_ref.model.path=/m", "distillation.distribution_matching.recipe=dmd", ] diff --git a/tests/trainer/diffusion/test_distillation_contracts_on_cpu.py b/tests/trainer/diffusion/test_distillation_contracts_on_cpu.py index 8dc61d40f..851305754 100644 --- a/tests/trainer/diffusion/test_distillation_contracts_on_cpu.py +++ b/tests/trainer/diffusion/test_distillation_contracts_on_cpu.py @@ -280,6 +280,16 @@ def test_causal_recipes_use_separate_causal_and_bidirectional_groups(self): assert role_groups["student"] == role_groups["student_ema"] == "causal_base" assert role_groups["teacher_score"] == role_groups["fake_score"] == "bidirectional_base" + def test_colocated_independent_materializes_one_group_per_role(self): + plan = build_plan( + "dmd2", + {"model_path": "/m", "role_storage": "colocated_independent"}, + ALL_CAPS, + ) + assert {group.storage for group in plan.role_layout.groups} == {"independent_module"} + assert len(plan.role_layout.groups) == len(plan.role_layout.bindings) == 4 + assert all(binding.group == f"{binding.role}_model" for binding in plan.role_layout.bindings) + @pytest.mark.parametrize( "name,config,error", [ @@ -291,6 +301,7 @@ def test_causal_recipes_use_separate_causal_and_bidirectional_groups(self): ("dmd2", {"fake_update_ratio": 1.5, "model_path": "/m"}, "integer"), ("dmd2", {"fake_update_ratio": True, "model_path": "/m"}, "integer"), ("dmd2", {"fake_warmup_cycles": 1.5, "model_path": "/m"}, "integer"), + ("dmd2", {"role_storage": "remote", "model_path": "/m"}, "role_storage"), ], ) def test_invalid_recipe_values_fail_closed(self, name, config, error): @@ -368,6 +379,20 @@ def test_warmup_requires_fake_only_phases(self): warmup_cycles=1, ) + def test_warmup_cannot_reuse_a_missing_student_batch(self): + with pytest.raises(ValueError, match="cannot reuse a student batch"): + UpdateSchedule( + phases=(self.student_phase(), self.fake_phase()), + warmup_phases=( + UpdatePhaseSpec( + kind="fake_score", + batch_policy="reuse_student", + trainable_roles=("fake_score",), + ), + ), + warmup_cycles=1, + ) + def test_warmup_cycles_require_warmup_phases(self): with pytest.raises(ValueError, match="requires at least one warmup phase"): UpdateSchedule(phases=(self.student_phase(), self.fake_phase()), warmup_cycles=1) diff --git a/tests/trainer/diffusion/test_distillation_controller_on_cpu.py b/tests/trainer/diffusion/test_distillation_controller_on_cpu.py index ae954d2ba..adc500e90 100644 --- a/tests/trainer/diffusion/test_distillation_controller_on_cpu.py +++ b/tests/trainer/diffusion/test_distillation_controller_on_cpu.py @@ -217,9 +217,33 @@ def test_reset_clears_healthy_driver_state(self): assert controller.counters.optimizer_steps == {} assert controller.metrics == {} - def test_failed_driver_cannot_be_reset(self): + def test_state_dict_round_trip_restores_completed_counters(self): + controller, _, _ = make_controller(make_plan(fake_repeats=2)) + controller.run(3) + state = controller.state_dict() + + restored, _, _ = make_controller(make_plan(fake_repeats=2)) + restored.load_state_dict(state) + assert restored.counters.global_step == 3 + assert restored.counters.completed_cycles == 3 + assert restored.counters.optimizer_steps == {"student": 3, "fake_score": 6} + + def test_invalid_checkpoint_state_is_rejected(self): + controller, _, _ = make_controller(make_plan()) + with pytest.raises(ValueError, match="exactly"): + controller.load_state_dict({"global_step": 1}) + with pytest.raises(ValueError, match="non-negative integer"): + controller.load_state_dict({"global_step": -1, "optimizer_steps": {}, "completed_cycles": 0}) + with pytest.raises(ValueError, match="unknown optimizer roles"): + controller.load_state_dict({"global_step": 0, "optimizer_steps": {"unknown": 1}, "completed_cycles": 1}) + with pytest.raises(ValueError, match="must equal global_step"): + controller.load_state_dict({"global_step": 2, "optimizer_steps": {"student": 1}, "completed_cycles": 2}) + + def test_failed_driver_cannot_be_checkpointed_or_reset(self): controller, _, _ = make_controller(make_plan(), executor=FakePhaseExecutor(fail_on="student")) with pytest.raises(RuntimeError): controller.run_cycle() + with pytest.raises(RuntimeError, match="Cannot checkpoint"): + controller.state_dict() with pytest.raises(RuntimeError, match="cannot be reset"): controller.reset() diff --git a/tests/trainer/diffusion/test_distillation_trainer_routing_on_cpu.py b/tests/trainer/diffusion/test_distillation_trainer_routing_on_cpu.py index fb1d8f272..ecc780b40 100644 --- a/tests/trainer/diffusion/test_distillation_trainer_routing_on_cpu.py +++ b/tests/trainer/diffusion/test_distillation_trainer_routing_on_cpu.py @@ -18,17 +18,21 @@ ``distillation.enabled=true`` together with ``trainer_type=policy_gradient``. """ +from types import SimpleNamespace +from unittest.mock import Mock + import pytest from omegaconf import OmegaConf from verl_omni.trainer.config.algorithm import DiffusionAlgoConfig from verl_omni.trainer.diffusion.distillation.ray_trainer import DistillationRayTrainer -from verl_omni.trainer.main_diffusion import _get_trainer_cls +from verl_omni.trainer.main_diffusion import TaskRunner, _get_trainer_cls class FakeAlgorithmConfig: def __init__(self, trainer_type): self.trainer_type = trainer_type + self.sample_source = "offline" if trainer_type == "distillation" else "online" class FakeTrainerConfig: @@ -36,6 +40,15 @@ def __init__(self, trainer_type): self.algorithm = FakeAlgorithmConfig(trainer_type) +def initialize_fake_base_trainer(self, **kwargs): + for name, value in kwargs.items(): + setattr(self, name, value) + self.config = kwargs["config"] + self.total_training_steps = 4 + self.train_dataloader = [] + self.device_name = "cuda" + + class TestTrainerRouting: def test_distillation_routes_to_distillation_trainer(self): assert _get_trainer_cls(FakeTrainerConfig("distillation")) is DistillationRayTrainer @@ -54,12 +67,35 @@ def test_unknown_trainer_type_lists_distillation(self): with pytest.raises(ValueError, match="distillation"): _get_trainer_cls(FakeTrainerConfig("bogus")) + def test_task_runner_rejects_external_rollout_for_distillation(self): + config = FakeTrainerConfig("distillation") + config.algorithm.sample_source = "online" + with pytest.raises(ValueError, match="sample_source=offline"): + TaskRunner().add_actor_rollout_worker(config) + + def test_task_runner_registers_the_distillation_worker(self, monkeypatch): + import ray + from verl.trainer.ppo.ray_trainer import Role + + from verl_omni.workers.diffusion_distillation_worker import DiffusionDistillationWorker + + monkeypatch.setattr(ray, "remote", lambda worker_cls: worker_cls) + runner = TaskRunner() + worker_cls, _ = runner.add_actor_rollout_worker(FakeTrainerConfig("distillation")) + assert worker_cls is DiffusionDistillationWorker + assert runner.role_worker_mapping[Role.Actor] is DiffusionDistillationWorker + assert runner.mapping[Role.Actor] == "global_pool" + class TestAlgorithmConfig: def test_distillation_is_a_valid_trainer_type(self): - config = DiffusionAlgoConfig(trainer_type="distillation") + config = DiffusionAlgoConfig(trainer_type="distillation", sample_source="offline") assert config.trainer_type == "distillation" + def test_distillation_rejects_external_rollout_sampling(self): + with pytest.raises(ValueError, match="sample_source='offline'"): + DiffusionAlgoConfig(trainer_type="distillation", sample_source="online") + def test_existing_trainer_types_still_valid(self): assert DiffusionAlgoConfig(trainer_type="policy_gradient").trainer_type == "policy_gradient" assert DiffusionAlgoConfig(trainer_type="direct_preference").trainer_type == "direct_preference" @@ -75,7 +111,7 @@ def test_invalid_trainer_type_raises(self): def runtime_config(): return OmegaConf.create( { - "algorithm": {"trainer_type": "distillation"}, + "algorithm": {"trainer_type": "distillation", "sample_source": "offline"}, "actor_rollout_ref": { "actor": { "diffusion_loss": {"loss_mode": "flow_grpo"}, @@ -99,7 +135,39 @@ def runtime_config(): ) -class TestPR1DataPlaneBoundary: +class TestRuntimeValidation: + @staticmethod + def trainer_config(*, role_storage="shared_base_adapters", strategy="fsdp2", lora_rank=8, use_orig=True): + return OmegaConf.create( + { + "distillation": {"distribution_matching": {"role_storage": role_storage}}, + "actor_rollout_ref": { + "model": {"lora_rank": lora_rank}, + "actor": {"strategy": strategy, "fsdp_config": {"use_orig_params": use_orig}}, + }, + } + ) + + def test_shared_base_requires_lora(self): + from verl_omni.trainer.diffusion.distillation.recipes import build_plan + + trainer = object.__new__(DistillationRayTrainer) + trainer.plan = build_plan("dmd2", {"model_path": "/m"}, frozenset({"distribution_matching"})) + trainer.config = self.trainer_config(lora_rank=0) + with pytest.raises(ValueError, match="lora_rank > 0"): + trainer.validate_runtime_config() + + def test_fsdp1_shared_base_requires_orig_params(self): + from verl_omni.trainer.diffusion.distillation.recipes import build_plan + + trainer = object.__new__(DistillationRayTrainer) + trainer.plan = build_plan("dmd2", {"model_path": "/m"}, frozenset({"distribution_matching"})) + trainer.config = self.trainer_config(strategy="fsdp", use_orig=False) + with pytest.raises(ValueError, match="use_orig_params=true"): + trainer.validate_runtime_config() + + +class TestDataPlaneBoundary: def test_production_constructor_reaches_explicit_pr2_boundary(self): config = runtime_config() trainer = DistillationRayTrainer( @@ -115,7 +183,7 @@ def test_production_constructor_reaches_explicit_pr2_boundary(self): train_sampler=object(), ) assert trainer.config is config - with pytest.raises(NotImplementedError, match="PR 2"): + with pytest.raises(NotImplementedError, match="composed diffusion trainer config"): trainer.init_workers() def test_constructor_rejects_opd_switch(self): @@ -136,9 +204,56 @@ def test_fit_without_executor_reports_pr2_boundary(self): plan = build_plan("dmd2", {"model_path": "/m"}, frozenset({"distribution_matching"})) trainer = DistillationRayTrainer(plan=plan) - with pytest.raises(NotImplementedError, match="PR 2"): + with pytest.raises(NotImplementedError, match="composed diffusion trainer config"): trainer.fit(num_cycles=1) + def test_production_constructor_and_worker_lifecycle_are_wired(self, monkeypatch): + from verl.trainer.ppo.ray_trainer import Role + + from verl_omni.trainer.diffusion.distillation.recipes import build_plan + from verl_omni.trainer.diffusion.ray_diffusion_trainer import BaseRayDiffusionTrainer + from verl_omni.workers.diffusion_distillation_worker import DiffusionDistillationWorkerGroup + + config = runtime_config() + config.trainer = { + "device": "cuda", + "ray_master_port_range": None, + "n_gpus_per_node": 1, + "nnodes": 1, + } + config.data = {"train_batch_size": 1} + config.actor_rollout_ref.model.lora_rank = 8 + config.actor_rollout_ref.actor.strategy = "fsdp2" + config.actor_rollout_ref.actor.fsdp_config = { + "use_orig_params": False, + "ulysses_sequence_parallel_size": 1, + } + config.distillation.distribution_matching.role_storage = "shared_base_adapters" + config.distillation.distribution_matching.fake_score_optim = {"total_training_steps": -1} + plan = build_plan( + "dmd2", + {"model_path": "/m", "fake_update_ratio": 2}, + frozenset({"distribution_matching"}), + ) + + monkeypatch.setattr(BaseRayDiffusionTrainer, "__init__", initialize_fake_base_trainer) + resource_manager = SimpleNamespace(create_resource_pool=Mock(), get_resource_pool=Mock(return_value="pool")) + worker_group = SimpleNamespace(init_model=Mock()) + worker_group_factory = Mock(return_value=worker_group) + trainer = DistillationRayTrainer( + config=config, + plan=plan, + role_worker_mapping={Role.Actor: object}, + resource_pool_manager=resource_manager, + ray_worker_group_cls=worker_group_factory, + ) + trainer.init_workers() + resource_manager.create_resource_pool.assert_called_once() + resource_manager.get_resource_pool.assert_called_once_with(Role.Actor) + worker_group.init_model.assert_called_once() + assert isinstance(trainer.executor, DiffusionDistillationWorkerGroup) + assert config.distillation.distribution_matching.fake_score_optim.total_training_steps == 8 + def test_controller_binds_when_collaborators_are_supplied(self): from verl_omni.trainer.diffusion.distillation.controller import ( FakeBatchProvider, diff --git a/tests/workers/test_diffusers_fsdp_lora_adapter.py b/tests/workers/test_diffusers_fsdp_lora_adapter.py index 5d815d34c..2334152d7 100644 --- a/tests/workers/test_diffusers_fsdp_lora_adapter.py +++ b/tests/workers/test_diffusers_fsdp_lora_adapter.py @@ -16,15 +16,20 @@ import os import shutil import tempfile +from copy import deepcopy from functools import partial import pytest import ray import torch +from verl.single_controller.base import Worker from verl.single_controller.base.decorator import Dispatch, register from verl.single_controller.ray import RayClassWithInitArgs, RayResourcePool, RayWorkerGroup from verl.utils import tensordict_utils as tu +from verl.utils.distributed import initialize_global_process_group_ray, set_numa_affinity +from verl_omni.trainer.diffusion.distillation.contracts import RoleBinding, RoleGroupSpec +from verl_omni.workers.engine.fsdp.distillation_impl import DistillationRoleGroupEngine from verl_omni.workers.engine_workers import TrainingWorker from verl_omni.workers.utils.losses import diffusion_loss from verl_omni.workers.utils.padding import embeds_padding_2_no_padding @@ -53,6 +58,89 @@ def _require_model_path() -> str: return _DEFAULT_MODEL_PATH +class DistillationLoRAFSDPTestWorker(Worker): + """Tiny-model worker exercising the real multi-role FSDP engine.""" + + def __init__(self, training_config, model_path): + Worker.__init__(self) + initialize_global_process_group_ray(timeout_second=None) + set_numa_affinity() + model_config = deepcopy(training_config.model_config) + engine_config = deepcopy(training_config.engine_config) + optimizer_config = deepcopy(training_config.optimizer_config) + checkpoint_config = deepcopy(training_config.checkpoint_config) + object.__setattr__(model_config, "model_type", "diffusion_distillation_model") + object.__setattr__(engine_config, "use_orig_params", True) + group = RoleGroupSpec( + name="base", + model_ref=model_path, + storage="shared_base_adapters", + placement="colocated", + ) + bindings = ( + RoleBinding("student", "base", "student", True, "student_optim"), + RoleBinding("teacher_score", "base", None, False, None), + RoleBinding("fake_score", "base", "fake_score", True, "fake_score_optim"), + RoleBinding("student_ema", "base", "student_ema", False, None), + ) + fake_optimizer_config = deepcopy(optimizer_config) + object.__setattr__(fake_optimizer_config, "lr", optimizer_config.lr * 0.5) + self.engine = DistillationRoleGroupEngine( + model_config=model_config, + engine_config=engine_config, + optimizer_config=optimizer_config, + checkpoint_config=checkpoint_config, + role_group=group, + role_bindings=bindings, + optimizer_configs={"student": optimizer_config, "fake_score": fake_optimizer_config}, + ) + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def init_model(self): + self.engine.initialize() + + def collect_adapter(self, role): + params, config = self.engine.get_per_tensor_param(adapter_name=role, base_sync_done=True) + return {name: tensor.detach().cpu() for name, tensor in params}, config + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def collect_roles(self): + student, student_config = self.collect_adapter("student") + fake, fake_config = self.collect_adapter("fake_score") + ema, ema_config = self.collect_adapter("student_ema") + return { + "student": student, + "fake_score": fake, + "student_ema": ema, + "student_config": student_config, + "fake_config": fake_config, + "ema_config": ema_config, + } + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def run_student_step_across_role_switches(self): + self.engine.optimizer_zero_grad("student") + with self.engine.use_role("student"): + student_parameters = self.engine.parameters_for_role("student") + loss = sum(parameter.float().square().mean() for parameter in student_parameters if parameter.numel()) + with self.engine.use_role("teacher_score"): + assert not torch.is_grad_enabled() + with self.engine.use_role("fake_score", grad_enabled=False): + assert not torch.is_grad_enabled() + self.engine.backward_role("student", loss) + stepped, grad_norm = self.engine.optimizer_step("student") + self.engine.update_role_ema("student", "student_ema", decay=0.5) + return {"stepped": stepped, "grad_norm": grad_norm} + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def save_role_group(self, path, step): + self.engine.save_role_group_checkpoint(path, step) + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def load_role_group(self, path): + self.engine.load_role_group_checkpoint(path) + + class LoRAFSDPTestWorker(TrainingWorker): @register(dispatch_mode=Dispatch.ONE_TO_ALL) def report_fsdp_topology(self): @@ -319,3 +407,62 @@ def test_diffusers_fsdp_lora_adapter_copy_ema(strategy): if not torch.cuda.is_available(): pytest.skip("CUDA is required for FSDP LoRA adapter tests.") _run_copy_ema_adapter_test(strategy) + + +@pytest.mark.parametrize("strategy", ["fsdp", "fsdp2"]) +def test_distillation_role_isolation_ema_and_resume(strategy): + if not torch.cuda.is_available(): + pytest.skip("CUDA is required for distillation role-group tests.") + base_model_path = _require_model_path() + device_count = _resolve_lora_test_device_count(strategy) + + ray.init() + tmp_dir = tempfile.mkdtemp(prefix="qwen_image_distillation_roles_") + try: + sp_enabled = device_count > 1 and _diffusers_sp_supported() + if sp_enabled: + model_path = _create_sp_compatible_model(tmp_dir, base_model_path, num_attention_heads=2) + else: + model_path = base_model_path + training_config, _ = create_training_config( + model_type="diffusion_distillation_model", + strategy=strategy, + device_count=device_count, + model=model_path, + policy_state_adapters=("default", "student", "fake_score", "student_ema"), + ) + ray_cls_with_init = RayClassWithInitArgs( + cls=ray.remote(DistillationLoRAFSDPTestWorker), + training_config=training_config, + model_path=model_path, + ) + resource_pool = RayResourcePool(process_on_nodes=[device_count]) + wg = RayWorkerGroup(resource_pool=resource_pool, ray_cls_with_init=ray_cls_with_init) + wg.init_model() + + initial = wg.collect_roles()[0] + _lora_params_close(initial["student"], initial["student_ema"]) + _lora_params_close(initial["student"], initial["fake_score"]) + assert initial["student_config"] == initial["fake_config"] == initial["ema_config"] + + step_result = wg.run_student_step_across_role_switches()[0] + assert step_result["stepped"] + assert step_result["grad_norm"] > 0 + after_step = wg.collect_roles()[0] + _lora_params_differ(after_step["student"], initial["student"]) + _lora_params_close(after_step["fake_score"], initial["fake_score"]) + _assert_ema_blend(after_step["student_ema"], initial["student_ema"], after_step["student"], decay=0.5) + + checkpoint_path = os.path.join(tmp_dir, "checkpoint") + wg.save_role_group(checkpoint_path, 1) + wg.run_student_step_across_role_switches() + changed = wg.collect_roles()[0] + _lora_params_differ(changed["student"], after_step["student"]) + wg.load_role_group(checkpoint_path) + restored = wg.collect_roles()[0] + _lora_params_close(restored["student"], after_step["student"]) + _lora_params_close(restored["fake_score"], after_step["fake_score"]) + _lora_params_close(restored["student_ema"], after_step["student_ema"]) + finally: + ray.shutdown() + shutil.rmtree(tmp_dir, ignore_errors=True) diff --git a/tests/workers/test_diffusers_fsdp_merged_lora_on_cpu.py b/tests/workers/test_diffusers_fsdp_merged_lora_on_cpu.py index e77169511..603d15ad7 100644 --- a/tests/workers/test_diffusers_fsdp_merged_lora_on_cpu.py +++ b/tests/workers/test_diffusers_fsdp_merged_lora_on_cpu.py @@ -32,6 +32,10 @@ def __init__(self): # Carrying a ``peft_config`` is all ``get_per_tensor_param`` needs to # take the LoRA branch. self.peft_config = {"default": SimpleNamespace(to_dict=lambda: {"r": 8})} + self.active_adapter = "default" + + def set_adapter(self, name): + self.active_adapter = name def _make_engine(module, lora_config: dict) -> PPODiffusersFSDPEngine: diff --git a/tests/workers/test_diffusion_distillation_lora_on_cpu.py b/tests/workers/test_diffusion_distillation_lora_on_cpu.py new file mode 100644 index 000000000..38fcf3ef6 --- /dev/null +++ b/tests/workers/test_diffusion_distillation_lora_on_cpu.py @@ -0,0 +1,132 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""CPU regressions for role-aware LoRA switching and export.""" + +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch + +from verl_omni.workers.engine.fsdp.diffusers_impl import DiffusersFSDPEngine +from verl_omni.workers.engine.lora_adapter_mixin import LoRAAdapterMixin + + +class PeftConfig: + def __init__(self, name): + self.name = name + + def to_dict(self): + return {"name": self.name} + + +class PeftModule(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor(1.0)) + self.peft_config = { + "default": PeftConfig("default"), + "student": PeftConfig("student"), + "student_ema": PeftConfig("student_ema"), + } + self.active_adapter = "student" + self.adapters_enabled = True + + @property + def active_adapters(self): + return [self.active_adapter] + + def set_adapter(self, name): + self.active_adapter = name[0] if isinstance(name, list) else name + + def disable_adapters(self): + self.adapters_enabled = False + + def enable_adapters(self): + self.adapters_enabled = True + + +class MixinHarness(LoRAAdapterMixin): + def __init__(self): + self.module = PeftModule() + self._is_offload_param = False + + +class TestAdapterContext: + def test_nested_named_adapter_context_restores_exact_selection(self): + harness = MixinHarness() + with harness.use_adapter("student_ema"): + assert harness.module.active_adapter == "student_ema" + with harness.use_adapter("default"): + assert harness.module.active_adapter == "default" + assert harness.module.active_adapter == "student_ema" + assert harness.module.active_adapter == "student" + + def test_exception_restores_previous_adapter(self): + harness = MixinHarness() + with pytest.raises(RuntimeError, match="boom"): + with harness.use_adapter("student_ema"): + raise RuntimeError("boom") + assert harness.module.active_adapter == "student" + + def test_reference_context_reenables_prior_named_adapter(self): + harness = MixinHarness() + with harness.use_adapter("reference"): + assert not harness.module.adapters_enabled + assert harness.module.active_adapter == "student" + assert harness.module.adapters_enabled + assert harness.module.active_adapter == "student" + + +class TestAdapterAwareExport: + @pytest.fixture(autouse=True) + def mock_gpu_memory_logging(self, monkeypatch): + monkeypatch.setattr("verl_omni.workers.engine.fsdp.diffusers_impl.log_gpu_memory_usage", Mock()) + + @pytest.mark.parametrize("adapter_name", [None, "default", "student", "student_ema"]) + def test_selected_adapter_exports_its_own_peft_config(self, monkeypatch, adapter_name): + harness = MixinHarness() + harness._uses_fsdp2_cpu_offload_policy = True + harness.model_config = SimpleNamespace(lora={"merge": False}, fsdp_layer_prefixes=["transformer_blocks."]) + collect_lora_params = Mock(return_value={"adapter.weight": torch.tensor([2.0])}) + + monkeypatch.setattr( + "verl_omni.workers.engine.fsdp.diffusers_impl.collect_lora_params", + collect_lora_params, + ) + monkeypatch.setattr( + "verl_omni.workers.engine.fsdp.diffusers_impl.convert_weight_keys", + lambda params, module: params, + ) + params, peft_config = DiffusersFSDPEngine.get_per_tensor_param( + harness, + base_sync_done=True, + adapter_name=adapter_name, + ) + assert dict(params) == {"transformer.adapter.weight": torch.tensor([2.0])} + assert peft_config == {"name": adapter_name or "default"} + collect_lora_params.assert_called_once() + assert collect_lora_params.call_args.kwargs["adapter_name"] == (adapter_name or "default") + assert harness.module.active_adapter == "student" + + def test_unknown_adapter_fails_before_export(self): + harness = MixinHarness() + harness._uses_fsdp2_cpu_offload_policy = True + harness.model_config = SimpleNamespace(lora={"merge": False}, fsdp_layer_prefixes=[]) + with pytest.raises(ValueError, match="unknown LoRA adapter"): + DiffusersFSDPEngine.get_per_tensor_param( + harness, + base_sync_done=True, + adapter_name="missing", + ) diff --git a/tests/workers/test_diffusion_distillation_runtime_on_cpu.py b/tests/workers/test_diffusion_distillation_runtime_on_cpu.py new file mode 100644 index 000000000..92d38dd33 --- /dev/null +++ b/tests/workers/test_diffusion_distillation_runtime_on_cpu.py @@ -0,0 +1,565 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""CPU tests for the generic multi-role distillation data plane.""" + +from contextlib import contextmanager, nullcontext +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch +from omegaconf import OmegaConf +from verl.protocol import DataProtoFuture +from verl.utils import tensordict_utils as tu + +from verl_omni.trainer.diffusion.distillation.contracts import PhaseRequest +from verl_omni.trainer.diffusion.distillation.recipes import build_plan +from verl_omni.workers.diffusion_distillation_worker import ( + DiffusionDistillationWorker, + DiffusionDistillationWorkerGroup, + DistillationPhaseComputation, + DistillationRoleRuntime, + resolve_profiler_configs, +) +from verl_omni.workers.engine.fsdp.distillation_impl import DistillationRoleGroupEngine + +_CAPABILITIES = frozenset({"distribution_matching"}) + + +def test_distillation_worker_instantiates_nested_profiler_tool_config(): + config = OmegaConf.create( + { + "_target_": "verl.utils.profiler.ProfilerConfig", + "tool": "torch", + "enable": False, + "all_ranks": False, + "ranks": [], + "save_path": "outputs/profile", + "tool_config": { + "torch": { + "_target_": "verl.utils.profiler.config.TorchProfilerToolConfig", + "contents": [], + "discrete": False, + "name": "torch", + } + }, + } + ) + + profiler_config, tool_config = resolve_profiler_configs(config) + + assert profiler_config.tool == "torch" + assert tool_config.name == "torch" + assert tool_config.contents == [] + + +class ToyRoleEngine: + def __init__(self, roles, initial=None): + values = initial or {} + self.parameters = { + role: torch.nn.Parameter(torch.tensor(float(values.get(role, index + 1)))) + for index, role in enumerate(roles) + } + self.optimizers = { + role: torch.optim.SGD([parameter], lr=0.1) + for role, parameter in self.parameters.items() + if role in {"student", "fake_score"} + } + self.scheduler = object() + self.model_config = object() + self.active_role = None + + @contextmanager + def use_role(self, role): + previous = self.active_role + self.active_role = role + try: + yield self.parameters[role] + finally: + self.active_role = previous + + def optimizer_zero_grad(self, role=None): + optimizers = self.optimizers.values() if role is None else (self.optimizers[role],) + for optimizer in optimizers: + optimizer.zero_grad() + + def backward_role(self, role, loss, retain_graph=False): + assert self.active_role is None + loss.backward(retain_graph=retain_graph) + + def train_mode(self): + return nullcontext() + + def get_data_parallel_group(self): + return None + + def parameters_for_role(self, role): + return (self.parameters[role],) + + def optimizer_step(self, role=None): + parameter = self.parameters[role] + grad_norm = float(parameter.grad.detach().abs()) + self.optimizers[role].step() + return True, grad_norm + + def update_role_ema(self, source_role, target_role, decay): + with torch.no_grad(): + self.parameters[target_role].lerp_(self.parameters[source_role], 1.0 - decay) + + def update_module_ema_from(self, source, decay): + source_parameter = next(iter(source.parameters.values())) + target_parameter = next(iter(self.parameters.values())) + with torch.no_grad(): + target_parameter.lerp_(source_parameter, 1.0 - decay) + + +class ToyElementMeanComputer: + def compute_phase(self, request, batch, runtime): + role = request.trainable_roles[0] + parameter = runtime.engine_for_role(role).parameters[role] + count = float(batch["count"].sum()) + loss = ((parameter - batch["target"]).square() * batch["count"]).sum() / count + return DistillationPhaseComputation( + {role: loss}, {"ode/loss": float(loss.detach()), "ode/active_elements": count}, loss_normalizer=count + ) + + +class TestElementNormalizedRuntime: + @staticmethod + def make_runtime(): + plan = build_plan("dmd2", {"model_path": "/m"}, _CAPABILITIES) + engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema")) + runtime = DistillationRoleRuntime(plan, {"base": engine}, ema_decay=0.5, ema_start_step=0) + return runtime, engine, PhaseRequest("student", 0, 0, "fresh", ("student",), False) + + @pytest.mark.parametrize("micro_batch_size", [1, 2]) + def test_element_mean_is_independent_of_micro_batch_size(self, micro_batch_size): + runtime, engine, request = self.make_runtime() + parameter = engine.parameters["student"] + targets = torch.tensor([1.0, 4.0]) + counts = torch.tensor([2.0, 6.0]) + runtime.zero_grad(("student",)) + for target, count in zip(targets.split(micro_batch_size), counts.split(micro_batch_size), strict=True): + loss = ((parameter - target).square() * count).sum() / count.sum() + computation = DistillationPhaseComputation({"student": loss}, {}, loss_normalizer=float(count.sum())) + runtime.backward_micro_batch(request, computation, weight=target.numel() / targets.numel()) + runtime.step_phase(request) + torch.testing.assert_close(parameter, torch.tensor(1.45)) + torch.testing.assert_close(parameter.grad, torch.tensor(-4.5)) + + def test_denominator_is_averaged_over_the_same_dp_group_as_gradients(self, monkeypatch): + import verl_omni.workers.diffusion_distillation_worker as worker_module + + runtime, engine, request = self.make_runtime() + group = object() + monkeypatch.setattr(engine, "get_data_parallel_group", lambda: group) + monkeypatch.setattr(worker_module, "get_device_id", lambda: "cpu") + parameter = engine.parameters["student"] + computation = DistillationPhaseComputation({"student": (parameter - 3).square()}, {}, loss_normalizer=2) + runtime.backward_micro_batch(request, computation, weight=1.0) + # FSDP averages this rank's numerator gradient (-8) and its peer's (24). + parameter.grad.fill_(8.0) + + def average_counts(value, *, op, group): + assert group is engine.get_data_parallel_group() + assert op == torch.distributed.ReduceOp.AVG + torch.testing.assert_close(value, torch.tensor(2.0, dtype=value.dtype)) + value.fill_(4.0) # mean of the two ranks' counts: (2 + 6) / 2 + + monkeypatch.setattr(torch.distributed, "all_reduce", average_counts) + runtime.step_phase(request) + torch.testing.assert_close(parameter, torch.tensor(0.8)) + + @pytest.mark.parametrize("normalizer", [0, -1, float("nan"), float("inf"), True]) + def test_invalid_normalizers_fail_before_backward(self, normalizer): + runtime, engine, request = self.make_runtime() + computation = DistillationPhaseComputation( + {"student": engine.parameters["student"].square()}, {}, loss_normalizer=normalizer + ) + with pytest.raises(ValueError, match="normalizer"): + runtime.backward_micro_batch(request, computation, weight=1.0) + assert engine.parameters["student"].grad is None + + @pytest.mark.parametrize("first,second", [(None, 2), (2, None)]) + def test_mixed_reductions_fail_and_zero_grad_resets_phase(self, first, second): + runtime, engine, request = self.make_runtime() + parameter = engine.parameters["student"] + runtime.backward_micro_batch( + request, + DistillationPhaseComputation({"student": parameter.square()}, {}, loss_normalizer=first), + weight=0.5, + ) + with pytest.raises(ValueError, match="reduction"): + runtime.backward_micro_batch( + request, + DistillationPhaseComputation({"student": parameter.square()}, {}, loss_normalizer=second), + weight=0.5, + ) + runtime.zero_grad(("student",)) + runtime.backward_micro_batch( + request, + DistillationPhaseComputation({"student": parameter.square()}, {}, loss_normalizer=second), + weight=1.0, + ) + runtime.step_phase(request) + torch.testing.assert_close(parameter, torch.tensor(0.8)) + + +class TestElementMeanWorker: + @pytest.mark.parametrize("micro_batch_size", [1, 2, 3]) + def test_element_reduction_normalizes_worker_loss_metrics_and_gradients(self, monkeypatch, micro_batch_size): + plan = build_plan("dmd2", {"model_path": "/m"}, _CAPABILITIES) + engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema")) + worker = object.__new__(DiffusionDistillationWorker) + worker.runtime = DistillationRoleRuntime( + plan, + {"base": engine}, + ema_decay=0.9, + ema_start_step=0, + micro_batch_sizes={"student": micro_batch_size, "fake_score": 1}, + ) + worker.dm_computer = ToyElementMeanComputer() + device = Mock() + device.max_memory_allocated.return_value = 0 + device.max_memory_reserved.return_value = 0 + monkeypatch.setattr("verl_omni.workers.diffusion_distillation_worker.get_torch_device", lambda: device) + monkeypatch.setattr("verl_omni.workers.diffusion_distillation_worker.get_device_id", lambda: "cpu") + batch = tu.get_tensordict({"target": torch.tensor([0.0, 2.0, 6.0]), "count": torch.tensor([2.0, 3.0, 5.0])}) + tu.assign_non_tensor(batch, phase_request=PhaseRequest("student", 0, 0, "fresh", ("student",), False)) + result = worker.execute_phase(batch) + metrics = tu.get(result, "metrics") + assert metrics["student/loss"] == pytest.approx(13.0) + assert metrics["ode/loss"] == pytest.approx(13.0) + assert metrics["student/grad_norm"] == pytest.approx(5.2) + torch.testing.assert_close(engine.parameters["student"], torch.tensor(1.52)) + assert tu.get(result, "optimizer_steps") == {"student": 1} + + +class TestRoleRuntime: + def test_role_group_engines_must_match_the_plan_exactly(self): + plan = build_plan("dmd2", {"model_path": "/m"}, _CAPABILITIES) + engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema")) + with pytest.raises(ValueError, match="missing=.*base"): + DistillationRoleRuntime(plan, {}, ema_decay=0.9, ema_start_step=0) + with pytest.raises(ValueError, match="extra=.*other"): + DistillationRoleRuntime(plan, {"base": engine, "other": engine}, ema_decay=0.9, ema_start_step=0) + + def test_invalid_micro_batch_configuration_fails_closed(self): + plan = build_plan("dmd2", {"model_path": "/m"}, _CAPABILITIES) + engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema")) + with pytest.raises(ValueError, match="micro_batch_sizes"): + DistillationRoleRuntime( + plan, + {"base": engine}, + ema_decay=0.9, + ema_start_step=0, + micro_batch_sizes={"student": 1, "fake_score": 0}, + ) + + def test_student_and_fake_optimizers_are_isolated_and_ema_updates(self): + plan = build_plan( + "dmd2", + {"model_path": "/m", "fake_update_ratio": 1}, + _CAPABILITIES, + ) + engine = ToyRoleEngine( + ("student", "teacher_score", "fake_score", "student_ema"), + {"student": 1.0, "teacher_score": 7.0, "fake_score": 2.0, "student_ema": -5.0}, + ) + runtime = DistillationRoleRuntime(plan, {"base": engine}, ema_decay=0.5, ema_start_step=0) + assert engine.parameters["student_ema"].item() == pytest.approx(1.0) + + fake_before = engine.parameters["fake_score"].detach().clone() + student_request = PhaseRequest( + kind="student", + global_step=0, + repeat_index=0, + batch_policy="fresh", + trainable_roles=("student",), + update_ema=True, + ) + student_loss = (engine.parameters["student"] - 0.0).square() + steps, metrics = runtime.backward_and_step( + student_request, + DistillationPhaseComputation(losses={"student": student_loss}, metrics={}), + ) + assert steps == {"student": 1} + assert metrics["student/loss"] == pytest.approx(1.0) + assert engine.parameters["fake_score"].detach().equal(fake_before) + assert engine.parameters["student"].item() == pytest.approx(0.8) + assert engine.parameters["student_ema"].item() == pytest.approx(0.9) + + student_before = engine.parameters["student"].detach().clone() + fake_request = PhaseRequest( + kind="fake_score", + global_step=1, + repeat_index=0, + batch_policy="fresh", + trainable_roles=("fake_score",), + ) + fake_loss = (engine.parameters["fake_score"] - 0.0).square() + steps, _ = runtime.backward_and_step( + fake_request, + DistillationPhaseComputation(losses={"fake_score": fake_loss}, metrics={}), + ) + assert steps == {"fake_score": 1} + assert engine.parameters["student"].detach().equal(student_before) + assert engine.parameters["fake_score"].item() == pytest.approx(1.6) + assert engine.parameters["teacher_score"].grad is None + + def test_independent_module_ema_is_initialized_and_updated(self): + plan = build_plan( + "dmd2", + {"model_path": "/m", "role_storage": "colocated_independent"}, + _CAPABILITIES, + ) + engines = {} + for binding in plan.role_layout.bindings: + engines[binding.group] = ToyRoleEngine( + (binding.role,), {binding.role: 3.0 if binding.role == "student" else -2.0} + ) + runtime = DistillationRoleRuntime(plan, engines, ema_decay=0.25, ema_start_step=0) + student = runtime.engine_for_role("student").parameters["student"] + ema = runtime.engine_for_role("student_ema").parameters["student_ema"] + assert ema.item() == pytest.approx(student.item()) + with torch.no_grad(): + student.fill_(7.0) + runtime.update_ema() + assert ema.item() == pytest.approx(6.0) + + def test_gradient_accumulation_matches_full_batch_mean(self): + plan = build_plan("dmd2", {"model_path": "/m"}, _CAPABILITIES) + accumulated_engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema"), {"student": 1.0}) + full_engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema"), {"student": 1.0}) + accumulated = DistillationRoleRuntime(plan, {"base": accumulated_engine}, ema_decay=0.9, ema_start_step=0) + full = DistillationRoleRuntime(plan, {"base": full_engine}, ema_decay=0.9, ema_start_step=0) + request = PhaseRequest("student", 0, 0, "fresh", ("student",), False) + + accumulated.zero_grad(("student",)) + for target in (torch.tensor(0.0), torch.tensor(2.0)): + parameter = accumulated_engine.parameters["student"] + accumulated.backward_micro_batch( + request, + DistillationPhaseComputation(losses={"student": (parameter - target).square()}, metrics={}), + weight=0.5, + ) + accumulated.step_phase(request) + + full.zero_grad(("student",)) + parameter = full_engine.parameters["student"] + full_loss = torch.stack(((parameter - 0.0).square(), (parameter - 2.0).square())).mean() + full.backward_and_step( + request, + DistillationPhaseComputation(losses={"student": full_loss}, metrics={}), + ) + assert accumulated_engine.parameters["student"].item() == pytest.approx( + full_engine.parameters["student"].item() + ) + + def test_export_uses_the_plan_selected_semantic_role(self): + plan = build_plan( + "dmd2", + {"model_path": "/m", "export_role": "student_ema"}, + _CAPABILITIES, + ) + engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema")) + engine.iter_export_tensors = lambda role, base_sync_done: (role, base_sync_done) + runtime = DistillationRoleRuntime(plan, {"base": engine}, ema_decay=0.9, ema_start_step=0) + assert runtime.export_tensors(base_sync_done=True) == ("student_ema", True) + + def test_phase_losses_must_match_requested_roles(self): + plan = build_plan("dmd2", {"model_path": "/m"}, _CAPABILITIES) + engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema")) + runtime = DistillationRoleRuntime(plan, {"base": engine}, ema_decay=0.9, ema_start_step=0) + request = PhaseRequest("student", 0, 0, "fresh", ("student",), True) + with pytest.raises(ValueError, match="must match requested roles"): + runtime.backward_and_step( + request, + DistillationPhaseComputation( + losses={"fake_score": engine.parameters["fake_score"].square()}, metrics={} + ), + ) + + def test_pr2_rejects_multi_optimizer_adversarial_phase(self): + plan = build_plan("dmd2", {"model_path": "/m", "profile": "paper"}, _CAPABILITIES | {"adversarial"}) + engine = ToyRoleEngine(("student", "teacher_score", "fake_score", "student_ema", "discriminator")) + runtime = DistillationRoleRuntime(plan, {"base": engine}, ema_decay=0.9, ema_start_step=0) + request = plan.update_schedule.next_cycle(SimpleNamespace(global_step=0, completed_cycles=0)).requests[-1] + with pytest.raises(NotImplementedError, match="Multi-role optimizer phases"): + runtime.backward_and_step( + request, + DistillationPhaseComputation( + losses={ + "fake_score": engine.parameters["fake_score"].square(), + "discriminator": engine.parameters["discriminator"].square(), + }, + metrics={}, + ), + ) + + +class ToyPeftModule(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor(1.0)) + self.active_adapter = "student" + self.adapters_enabled = True + + def set_adapter(self, name): + self.active_adapter = name + + def enable_adapters(self): + self.adapters_enabled = True + + def disable_adapters(self): + self.adapters_enabled = False + + +class TestRoleEngineValidation: + @staticmethod + def uninitialized_engine(): + engine = object.__new__(DistillationRoleGroupEngine) + engine.role_group = SimpleNamespace(name="base", storage="shared_base_adapters") + engine.role_bindings = { + "student": SimpleNamespace(role="student", group="base", trainable=True), + "teacher_score": SimpleNamespace(role="teacher_score", group="base", trainable=False), + } + engine.optimizer_configs = {"student": object()} + return engine + + def test_fsdp1_shared_base_requires_orig_params(self): + engine = self.uninitialized_engine() + with pytest.raises(ValueError, match="use_orig_params=true"): + engine.validate_constructor_inputs(SimpleNamespace(strategy="fsdp", use_orig_params=False)) + + def test_optimizer_configs_must_match_trainable_roles(self): + engine = self.uninitialized_engine() + engine.optimizer_configs = {"fake_score": object()} + with pytest.raises(ValueError, match="must match trainable roles"): + engine.validate_constructor_inputs(SimpleNamespace(strategy="fsdp2", use_orig_params=False)) + + def test_gradient_leak_into_inactive_role_is_rejected(self): + engine = self.uninitialized_engine() + student = torch.nn.Parameter(torch.tensor(1.0)) + fake_score = torch.nn.Parameter(torch.tensor(2.0)) + fake_score.grad = torch.tensor(1.0) + engine._role_parameters = {"student": (student,), "fake_score": (fake_score,)} + with pytest.raises(RuntimeError, match="fake_score"): + engine.assert_gradient_isolation({"student"}) + + +class TestRoleContext: + @staticmethod + def make_engine(): + engine = object.__new__(DistillationRoleGroupEngine) + engine.role_group = SimpleNamespace(name="base", storage="shared_base_adapters") + engine.role_bindings = { + "student": SimpleNamespace(adapter="student", trainable=True), + "student_ema": SimpleNamespace(adapter="student_ema", trainable=False), + "teacher_score": SimpleNamespace(adapter=None, trainable=False), + } + engine.module = ToyPeftModule() + engine.optimizers = {} + engine.lr_schedulers = {} + engine.optimizer_configs = {} + engine._active_role = "student" + engine._primary_role = None + return engine + + def test_frozen_role_is_eval_no_grad_and_context_restores_on_error(self): + engine = self.make_engine() + engine.module.train() + with pytest.raises(RuntimeError, match="boom"): + with engine.use_role("teacher_score") as module: + assert not module.training + assert not torch.is_grad_enabled() + assert not module.adapters_enabled + raise RuntimeError("boom") + assert engine.module.training + assert engine.module.adapters_enabled + assert engine.module.active_adapter == "student" + assert engine._active_role == "student" + + def test_non_student_export_is_rejected(self): + engine = self.make_engine() + with pytest.raises(ValueError, match="Only student or student_ema"): + engine.iter_export_tensors("teacher_score", base_sync_done=False) + + def test_export_uses_the_semantic_roles_adapter(self): + engine = self.make_engine() + parameter = torch.tensor(1.0) + engine.get_per_tensor_param = Mock(return_value=(iter((("weight", parameter),)), {"adapter": "student_ema"})) + tensors, peft_config = engine.iter_export_tensors("student_ema", base_sync_done=False) + assert list(tensors) == [("weight", parameter)] + assert peft_config == {"adapter": "student_ema"} + engine.get_per_tensor_param.assert_called_once_with(adapter_name="student_ema", base_sync_done=False) + + +class TestWorkerGroupFacade: + def test_rank_failure_surfaces_before_lazy_collect_metadata(self, monkeypatch): + import ray + from verl.single_controller.base.decorator import MAGIC_ATTR + from verl.single_controller.ray.base import func_generator + + from verl_omni.workers.diffusion_distillation_worker import DiffusionDistillationWorker + + registration = getattr(DiffusionDistillationWorker.execute_phase, MAGIC_ATTR) + futures = [object(), object()] + + get_results = Mock(side_effect=ValueError("rank 0: invalid conditioning")) + collect = Mock(side_effect=AssertionError("Lazy metadata RPCs queued behind a peer collective")) + + monkeypatch.setattr(ray, "get", get_results) + execute = func_generator( + object(), + "execute_phase", + lambda group, *args, **kwargs: (args, kwargs), + collect, + lambda name, *args, **kwargs: futures, + registration["blocking"], + ) + with pytest.raises(ValueError, match="rank 0: invalid conditioning"): + execute(tu.get_tensordict({}, {})) + get_results.assert_called_once_with(futures) + collect.assert_not_called() + + def test_tensordict_result_is_converted_to_phase_result(self): + worker_group = SimpleNamespace( + execute_phase=Mock( + return_value=tu.get_tensordict( + tensor_dict={}, + non_tensor_dict={"metrics": {"loss": 1.25}, "optimizer_steps": {"student": 1}}, + ) + ) + ) + facade = DiffusionDistillationWorkerGroup(worker_group) + request = PhaseRequest("student", 0, 0, "fresh", ("student",), True) + result = facade.execute_phase(request, tu.get_tensordict({}, {})) + assert result.metrics == {"loss": 1.25} + assert result.optimizer_steps == {"student": 1} + assert tu.get(worker_group.execute_phase.call_args.args[0], "phase_request").kind == "student" + + def test_future_result_is_resolved_before_conversion(self, monkeypatch): + output = tu.get_tensordict( + tensor_dict={}, + non_tensor_dict={"metrics": {"loss": 2.5}, "optimizer_steps": {"fake_score": 1}}, + ) + monkeypatch.setattr(DataProtoFuture, "get", Mock(return_value=output)) + future = DataProtoFuture(collect_fn=None, futures=[]) + facade = DiffusionDistillationWorkerGroup(SimpleNamespace(execute_phase=Mock(return_value=future))) + request = PhaseRequest("fake_score", 0, 0, "fresh", ("fake_score",), False) + result = facade.execute_phase(request, tu.get_tensordict({}, {})) + + assert result.metrics == {"loss": 2.5} + assert result.optimizer_steps == {"fake_score": 1} diff --git a/tests/workers/test_distillation_fsdp_roles.py b/tests/workers/test_distillation_fsdp_roles.py new file mode 100644 index 000000000..7062d3cc1 --- /dev/null +++ b/tests/workers/test_distillation_fsdp_roles.py @@ -0,0 +1,269 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Small GPU tests for role switching on FSDP1 and FSDP2 LoRA modules.""" + +import os +import tempfile +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed as dist +from peft import LoraConfig, get_peft_model +from torch import nn +from torch.distributed.tensor import DTensor + +from verl_omni.trainer.diffusion.distillation.contracts import RoleBinding, RoleGroupSpec +from verl_omni.workers.engine.fsdp.distillation_impl import DistillationRoleGroupEngine + + +class TinyCheckpointManager: + def __init__(self, module, optimizer, scheduler): + self.module = module + self.optimizer = optimizer + self.scheduler = scheduler + + def save_checkpoint(self, local_path, **kwargs): + os.makedirs(local_path, exist_ok=True) + torch.save( + { + "model": self.module.state_dict(), + "optimizer": self.optimizer.state_dict(), + "scheduler": self.scheduler.state_dict(), + }, + os.path.join(local_path, "primary.pt"), + ) + + def load_checkpoint(self, local_path, **kwargs): + state = torch.load(os.path.join(local_path, "primary.pt"), weights_only=False) + self.module.load_state_dict(state["model"]) + self.optimizer.load_state_dict(state["optimizer"]) + self.scheduler.load_state_dict(state["scheduler"]) + + +class TinyModel(nn.Module): + def __init__(self): + super().__init__() + self.proj = nn.Linear(4, 4) + + def forward(self, inputs): + return self.proj(inputs) + + +def wrap_model(strategy): + model = get_peft_model( + TinyModel(), + LoraConfig(r=2, lora_alpha=2, target_modules=["proj"]), + adapter_name="student", + ).cuda() + adapter_config = model.peft_config["student"] + model.add_adapter("fake_score", adapter_config) + model.add_adapter("student_ema", adapter_config) + model.set_adapter("student") + if strategy == "fsdp": + from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + + return FSDP(model, use_orig_params=True, device_id=torch.cuda.current_device()) + from torch.distributed.fsdp import fully_shard + + fully_shard(model) + return model + + +def wrap_independent_model(strategy, role): + model = get_peft_model( + TinyModel(), + LoraConfig(r=2, lora_alpha=2, target_modules=["proj"]), + adapter_name="default", + ).cuda() + model.add_adapter(role, model.peft_config["default"]) + model.set_adapter(role) + if strategy == "fsdp": + from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + + return FSDP(model, use_orig_params=True, device_id=torch.cuda.current_device()) + from torch.distributed.fsdp import fully_shard + + fully_shard(model) + return model + + +def engine_shell(module): + engine = object.__new__(DistillationRoleGroupEngine) + engine.module = module + engine.role_group = RoleGroupSpec( + name="base", model_ref="/tiny", storage="shared_base_adapters", placement="colocated" + ) + engine.role_bindings = { + "student": RoleBinding("student", "base", "student", True, "student_optim"), + "teacher_score": RoleBinding("teacher_score", "base", None, False, None), + "fake_score": RoleBinding("fake_score", "base", "fake_score", True, "fake_score_optim"), + "student_ema": RoleBinding("student_ema", "base", "student_ema", False, None), + } + engine.optimizers = {} + engine.lr_schedulers = {} + engine.optimizer_configs = {} + engine._active_role = "student" + engine._primary_role = "student" + role_parameters = {} + for role in ("student", "fake_score"): + with engine.use_role(role): + role_parameters[role] = tuple(parameter for parameter in module.parameters() if parameter.requires_grad) + engine._role_parameters = role_parameters + engine.optimizers = {role: torch.optim.AdamW(parameters, lr=0.1) for role, parameters in role_parameters.items()} + engine.lr_schedulers = { + role: torch.optim.lr_scheduler.LambdaLR(optimizer, lambda _: 1.0) + for role, optimizer in engine.optimizers.items() + } + engine.optimizer_configs = {role: SimpleNamespace(clip_grad=1.0) for role in engine.optimizers} + engine.optimizer = engine.optimizers["student"] + engine.lr_scheduler = engine.lr_schedulers["student"] + engine.optimizer_config = engine.optimizer_configs["student"] + engine.rank = dist.get_rank() + engine._is_offload_param = False + engine._is_offload_optimizer = False + engine._uses_fsdp2_cpu_offload_policy = False + engine.checkpoint_manager = TinyCheckpointManager(engine.module, engine.optimizer, engine.lr_scheduler) + return engine + + +def independent_engine_shell(module, role, trainable): + engine = object.__new__(DistillationRoleGroupEngine) + engine.module = module + engine.role_group = SimpleNamespace(name=f"{role}_model", storage="independent_module") + engine.role_bindings = { + role: RoleBinding(role, f"{role}_model", role, trainable, f"{role}_optim" if trainable else None) + } + engine.optimizers = {} + engine.lr_schedulers = {} + engine.optimizer_configs = {} + engine._active_role = role + engine._primary_role = role if trainable else None + return engine + + +def adapter_snapshot(engine, role): + binding = engine.role_bindings[role] + with engine._adapter_state_context(), torch.no_grad(): + parameters = engine._active_adapter_trainable_params(binding.adapter) + return tuple( + (parameter.full_tensor() if isinstance(parameter, DTensor) else parameter).detach().cpu().clone() + for parameter in parameters + ) + + +def assert_tensors_equal(left, right): + assert len(left) == len(right) + for left_tensor, right_tensor in zip(left, right, strict=True): + torch.testing.assert_close(left_tensor, right_tensor, rtol=0, atol=0) + + +@pytest.mark.parametrize("strategy", ["fsdp", "fsdp2"]) +def test_distillation_role_switch_preserves_graph_ema_and_state(strategy): + if not torch.cuda.is_available(): + pytest.skip("CUDA is required for FSDP role-isolation tests.") + if dist.is_initialized(): + pytest.skip("This isolated one-rank FSDP test requires no existing process group.") + + with tempfile.TemporaryDirectory(prefix="distillation_fsdp_role_") as tmp_dir: + torch.cuda.set_device(0) + dist.init_process_group( + backend="nccl", + init_method=f"file://{os.path.join(tmp_dir, 'rendezvous')}", + rank=0, + world_size=1, + ) + try: + torch.manual_seed(7) + engine = engine_shell(wrap_model(strategy)) + engine.copy_adapter("student", "student_ema") + student_before = adapter_snapshot(engine, "student") + fake_before = adapter_snapshot(engine, "fake_score") + ema_before = adapter_snapshot(engine, "student_ema") + assert_tensors_equal(student_before, ema_before) + + inputs = torch.randn(2, 4, device="cuda") + engine.optimizer_zero_grad("student") + with engine.use_role("student") as module: + student_loss = module(inputs).float().square().mean() + with engine.use_role("teacher_score") as module: + teacher_before = module(inputs).detach().clone() + assert not torch.is_grad_enabled() + with engine.use_role("fake_score", grad_enabled=False) as module: + fake_output_before = module(inputs).detach().clone() + assert not torch.is_grad_enabled() + engine.backward_role("student", student_loss) + engine.assert_gradient_isolation({"student"}) + engine.optimizers["student"].step() + engine.update_role_ema("student", "student_ema", decay=0.5) + + student_after = adapter_snapshot(engine, "student") + fake_after = adapter_snapshot(engine, "fake_score") + ema_after = adapter_snapshot(engine, "student_ema") + assert any( + not torch.equal(before, after) for before, after in zip(student_before, student_after, strict=True) + ) + assert_tensors_equal(fake_before, fake_after) + for before, student, ema in zip(ema_before, student_after, ema_after, strict=True): + torch.testing.assert_close(ema.float(), (before.float() + student.float()) * 0.5) + with engine.use_role("teacher_score") as module: + torch.testing.assert_close(module(inputs), teacher_before, rtol=0, atol=0) + with engine.use_role("fake_score", grad_enabled=False) as module: + torch.testing.assert_close(module(inputs), fake_output_before, rtol=0, atol=0) + + engine.optimizer_zero_grad("fake_score") + with engine.use_role("fake_score") as module: + fake_loss = module(inputs).float().square().mean() + engine.backward_role("fake_score", fake_loss) + engine.optimizers["fake_score"].step() + engine.lr_schedulers["fake_score"].step() + checkpoint_student = adapter_snapshot(engine, "student") + checkpoint_fake = adapter_snapshot(engine, "fake_score") + checkpoint_ema = adapter_snapshot(engine, "student_ema") + checkpoint_path = os.path.join(tmp_dir, f"{strategy}_checkpoint") + engine.save_role_group_checkpoint(checkpoint_path, global_step=1) + + with engine.use_role("student") as module: + second_loss = module(inputs).float().square().mean() + engine.optimizer_zero_grad("student") + engine.backward_role("student", second_loss) + engine.optimizers["student"].step() + with engine._adapter_state_context(), torch.no_grad(): + for parameter in engine._active_adapter_trainable_params("fake_score"): + parameter.fill_(17.0) + engine.load_role_group_checkpoint(checkpoint_path) + assert_tensors_equal(adapter_snapshot(engine, "student"), checkpoint_student) + assert_tensors_equal(adapter_snapshot(engine, "fake_score"), checkpoint_fake) + assert_tensors_equal(adapter_snapshot(engine, "student_ema"), checkpoint_ema) + + torch.manual_seed(17) + independent_student = independent_engine_shell(wrap_independent_model(strategy, "student"), "student", True) + torch.manual_seed(19) + independent_ema = independent_engine_shell( + wrap_independent_model(strategy, "student_ema"), "student_ema", False + ) + with independent_student._adapter_state_context(), torch.no_grad(): + for parameter in independent_student._active_adapter_trainable_params("student"): + parameter.fill_(4.0) + with independent_ema._adapter_state_context(), torch.no_grad(): + for parameter in independent_ema._active_adapter_trainable_params("student_ema"): + parameter.fill_(0.0) + independent_ema.update_module_ema_from(independent_student, decay=0.25) + independent_values = adapter_snapshot(independent_ema, "student_ema") + assert independent_values + assert all( + torch.allclose(value.float(), torch.full_like(value.float(), 3.0)) for value in independent_values + ) + finally: + dist.destroy_process_group() diff --git a/verl_omni/pipelines/model_base.py b/verl_omni/pipelines/model_base.py index bd5671d98..447f0e92b 100644 --- a/verl_omni/pipelines/model_base.py +++ b/verl_omni/pipelines/model_base.py @@ -254,6 +254,28 @@ def forward( return module(**model_inputs)[0] +class DistributionMatchingModelAdapter: + """Optional capability mixin for DMD-family architecture adapters. + + The generic distillation runtime owns role placement, optimization, EMA, and + checkpointing. Architecture packages implement the differentiable phase + program and declare capabilities without adding recipe branches to the + trainer or worker. + """ + + @classmethod + def distillation_capabilities(cls) -> frozenset[str]: + """Return capabilities accepted by distillation recipe validation.""" + return frozenset({"distribution_matching"}) + + @classmethod + def build_distribution_matching_computer(cls, model_config, plan): + """Build the architecture-owned phase computation used by the worker.""" + raise NotImplementedError( + f"{cls.__name__} declares distribution-matching support but builds no distribution-matching computer." + ) + + class DiffusionI2IModelBase(DiffusionModelBase): """Base class for image-conditioned diffusion model training helpers. diff --git a/verl_omni/trainer/config/_generated_diffusion_trainer.yaml b/verl_omni/trainer/config/_generated_diffusion_trainer.yaml index 3cb671deb..1cac21b23 100644 --- a/verl_omni/trainer/config/_generated_diffusion_trainer.yaml +++ b/verl_omni/trainer/config/_generated_diffusion_trainer.yaml @@ -507,6 +507,30 @@ distillation: rollout_strategy: null data_mode: null export_role: student_ema + role_storage: shared_base_adapters + student_micro_batch_size_per_gpu: 1 + fake_score_micro_batch_size_per_gpu: 1 + fake_score_optim: + _target_: verl.workers.config.FSDPOptimizerConfig + optimizer: AdamW + optimizer_impl: torch.optim + lr: 2.0e-05 + lr_warmup_steps_ratio: 0.0 + total_training_steps: -1 + weight_decay: 0.01 + lr_warmup_steps: -1 + betas: + - 0.9 + - 0.999 + clip_grad: 1.0 + min_lr_ratio: 0.0 + num_cycles: 0.5 + lr_scheduler_type: constant + zero_indexed_step: true + warmup_style: null + override_optimizer_config: null + ema_decay: 0.999 + ema_start_step: 0 algorithm: _target_: verl_omni.trainer.config.DiffusionAlgoConfig trainer_type: policy_gradient diff --git a/verl_omni/trainer/config/_generated_diffusion_veomni_trainer.yaml b/verl_omni/trainer/config/_generated_diffusion_veomni_trainer.yaml index d2c2ac6a7..fae2f0845 100644 --- a/verl_omni/trainer/config/_generated_diffusion_veomni_trainer.yaml +++ b/verl_omni/trainer/config/_generated_diffusion_veomni_trainer.yaml @@ -548,6 +548,30 @@ distillation: rollout_strategy: null data_mode: null export_role: student_ema + role_storage: shared_base_adapters + student_micro_batch_size_per_gpu: 1 + fake_score_micro_batch_size_per_gpu: 1 + fake_score_optim: + _target_: verl.workers.config.FSDPOptimizerConfig + optimizer: AdamW + optimizer_impl: torch.optim + lr: 2.0e-05 + lr_warmup_steps_ratio: 0.0 + total_training_steps: -1 + weight_decay: 0.01 + lr_warmup_steps: -1 + betas: + - 0.9 + - 0.999 + clip_grad: 1.0 + min_lr_ratio: 0.0 + num_cycles: 0.5 + lr_scheduler_type: constant + zero_indexed_step: true + warmup_style: null + override_optimizer_config: null + ema_decay: 0.999 + ema_start_step: 0 algorithm: _target_: verl_omni.trainer.config.DiffusionAlgoConfig trainer_type: policy_gradient diff --git a/verl_omni/trainer/config/algorithm.py b/verl_omni/trainer/config/algorithm.py index e213bd410..47862d864 100644 --- a/verl_omni/trainer/config/algorithm.py +++ b/verl_omni/trainer/config/algorithm.py @@ -46,6 +46,8 @@ def __post_init__(self): valid_trainer_types = {"policy_gradient", "direct_preference", "distillation"} if self.trainer_type not in valid_trainer_types: raise ValueError(f"Invalid trainer_type: {self.trainer_type}. Must be one of {sorted(valid_trainer_types)}") + if self.trainer_type == "distillation" and self.sample_source != "offline": + raise ValueError("The distillation trainer requires sample_source='offline'; rollout runs inside FSDP.") valid_adv_modes = {"continuous", "positive_only", "negative_only", "one_only", "binary"} if self.adv_mode not in valid_adv_modes: raise ValueError(f"Invalid adv_mode: {self.adv_mode}. Must be one of {sorted(valid_adv_modes)}") diff --git a/verl_omni/trainer/config/diffusion/distillation/diffusion_distillation.yaml b/verl_omni/trainer/config/diffusion/distillation/diffusion_distillation.yaml index 367ede769..0749ecfac 100644 --- a/verl_omni/trainer/config/diffusion/distillation/diffusion_distillation.yaml +++ b/verl_omni/trainer/config/diffusion/distillation/diffusion_distillation.yaml @@ -75,3 +75,69 @@ distribution_matching: # Semantic role exported to inference replicas export_role: student_ema + + # Physical role storage: shared_base_adapters or colocated_independent + role_storage: shared_base_adapters + + # Per-device micro-batch size for student phases + student_micro_batch_size_per_gpu: 1 + + # Per-device micro-batch size for fake-score phases + fake_score_micro_batch_size_per_gpu: 1 + + # Independent fake-score optimizer and scheduler configuration + fake_score_optim: + + # Target optimizer configuration class + _target_: verl.workers.config.FSDPOptimizerConfig + + # Optimizer class name + optimizer: AdamW + + # Python module containing the optimizer class + optimizer_impl: torch.optim + + # Fake-score learning rate + lr: 2.0e-5 + + # Warmup steps ratio when lr_warmup_steps is non-positive + lr_warmup_steps_ratio: 0.0 + + # Total fake-score optimizer steps; populated by the trainer + total_training_steps: -1 + + # Weight decay + weight_decay: 0.01 + + # Explicit warmup steps; non-positive delegates to the ratio + lr_warmup_steps: -1 + + # Adam beta coefficients + betas: [0.9, 0.999] + + # Gradient clipping threshold + clip_grad: 1.0 + + # Minimum learning-rate ratio for cosine scheduling + min_lr_ratio: 0.0 + + # Number of cosine cycles + num_cycles: 0.5 + + # Learning-rate scheduler type + lr_scheduler_type: constant + + # Whether scheduler step counting starts at zero + zero_indexed_step: true + + # Deprecated scheduler alias + warmup_style: null + + # Optional optimizer implementation overrides + override_optimizer_config: null + + # EMA decay after a completed student optimizer step + ema_decay: 0.999 + + # First completed student step that updates EMA + ema_start_step: 0 diff --git a/verl_omni/trainer/diffusion/distillation/__init__.py b/verl_omni/trainer/diffusion/distillation/__init__.py index f37526bab..31fd79373 100644 --- a/verl_omni/trainer/diffusion/distillation/__init__.py +++ b/verl_omni/trainer/diffusion/distillation/__init__.py @@ -13,12 +13,14 @@ # limitations under the License. """Distribution-matching distillation runtime (DMD, DMD2, CausVid, Self-Forcing). -PR 1 provides the architecture-neutral trainer controller, immutable execution -contracts, recipe/objective/rollout registries, and pure DMD tensor utilities. It -defines no model pipeline, Ray worker, FSDP model, or GPU runtime. +The package separates immutable plans, pure tensor utilities, and the trainer +controller from the lazily imported Ray/FSDP data plane. Architecture-owned +phase runners plug into the generic runtime without adding model branches here. """ -from verl_omni.trainer.diffusion.distillation import contracts, controller, ray_trainer, recipes, utils +from importlib import import_module + +from verl_omni.trainer.diffusion.distillation import contracts, controller, recipes, utils from verl_omni.trainer.diffusion.distillation.contracts import ( CanonicalPrediction, ConditionBundle, @@ -56,9 +58,19 @@ FakeDistillationHooks, FakePhaseExecutor, ) -from verl_omni.trainer.diffusion.distillation.ray_trainer import DistillationRayTrainer from verl_omni.trainer.diffusion.distillation.recipes import build_plan, build_plan_from_config, recipe_registry + +def __getattr__(name: str): + if name == "DistillationRayTrainer": + from verl_omni.trainer.diffusion.distillation.ray_trainer import DistillationRayTrainer + + return DistillationRayTrainer + if name == "ray_trainer": + return import_module("verl_omni.trainer.diffusion.distillation.ray_trainer") + raise AttributeError(name) + + __all__ = [ # submodules "contracts", @@ -94,7 +106,7 @@ "TeacherScoreProvider", "RoleCheckpointManifest", "DistillationCheckpointState", - # controller / executor + # control plane / executor "DistillationTrainerController", "BatchProvider", "DistillationTrainerHooks", diff --git a/verl_omni/trainer/diffusion/distillation/contracts.py b/verl_omni/trainer/diffusion/distillation/contracts.py index 4f860171c..d493f85bc 100644 --- a/verl_omni/trainer/diffusion/distillation/contracts.py +++ b/verl_omni/trainer/diffusion/distillation/contracts.py @@ -364,6 +364,8 @@ def __post_init__(self) -> None: raise ValueError("warmup_phases require warmup_cycles > 0.") if any(phase.kind == "student" for phase in self.warmup_phases): raise ValueError("Warmup phases must not contain a student phase.") + if any(phase.batch_policy == "reuse_student" for phase in self.warmup_phases): + raise ValueError("Warmup phases cannot reuse a student batch before any student phase has run.") def next_cycle(self, counters: TrainerCounters) -> UpdateCycle: """Expand either the next warmup cycle or the normal static phases.""" diff --git a/verl_omni/trainer/diffusion/distillation/controller.py b/verl_omni/trainer/diffusion/distillation/controller.py index 864c615bc..c6345e57c 100644 --- a/verl_omni/trainer/diffusion/distillation/controller.py +++ b/verl_omni/trainer/diffusion/distillation/controller.py @@ -20,7 +20,7 @@ (RFC §13.1). The cycle state machine follows RFC §14. Worker allocation, role binding, and -checkpoint restore are deliberately delegated to the PR 2 executor; this module +checkpoint restore are deliberately delegated to the bound phase executor; this module only controls validated phase execution: - Optional fake/discriminator warmup cycles: emit fake-only ``UpdateCycle`` @@ -198,6 +198,50 @@ def metrics(self) -> dict[str, dict]: """Metrics recorded for the most recent phase of each kind.""" return self.phase_metrics + def state_dict(self) -> dict[str, Any]: + """Return completed-cycle driver state for atomic checkpointing.""" + if self.failed: + raise RuntimeError("Cannot checkpoint a failed controller; restore the last completed cycle instead.") + return { + "global_step": self.counters.global_step, + "optimizer_steps": dict(self.counters.optimizer_steps), + "completed_cycles": self.counters.completed_cycles, + } + + def load_state_dict(self, state: dict[str, Any]) -> None: + """Restore validated counters into a fresh controller.""" + if self.failed: + raise RuntimeError("Cannot restore into a failed controller; construct a new driver.") + required = {"global_step", "optimizer_steps", "completed_cycles"} + if set(state) != required: + raise ValueError(f"Controller state must contain exactly {sorted(required)}, got {sorted(state)}.") + global_step = state["global_step"] + completed_cycles = state["completed_cycles"] + optimizer_steps = state["optimizer_steps"] + if isinstance(global_step, bool) or not isinstance(global_step, int) or global_step < 0: + raise ValueError(f"global_step must be a non-negative integer, got {global_step!r}.") + if isinstance(completed_cycles, bool) or not isinstance(completed_cycles, int) or completed_cycles < 0: + raise ValueError(f"completed_cycles must be a non-negative integer, got {completed_cycles!r}.") + if not isinstance(optimizer_steps, dict) or any( + not isinstance(role, str) or isinstance(steps, bool) or not isinstance(steps, int) or steps < 0 + for role, steps in optimizer_steps.items() + ): + raise ValueError("optimizer_steps must map role names to non-negative integer counters.") + trainable_roles = {binding.role for binding in self.plan.role_layout.bindings if binding.trainable} + unknown_roles = set(optimizer_steps) - trainable_roles + if unknown_roles: + raise ValueError(f"Checkpoint contains unknown optimizer roles: {sorted(unknown_roles)}.") + if optimizer_steps.get("student", 0) != global_step: + raise ValueError("The student optimizer counter must equal global_step.") + if completed_cycles < global_step: + raise ValueError("completed_cycles cannot be less than global_step.") + self.counters = TrainerCounters( + global_step=global_step, + optimizer_steps=dict(optimizer_steps), + completed_cycles=completed_cycles, + ) + self.phase_metrics = {} + def reset(self) -> None: """Reset a healthy driver; failed drivers must be reconstructed from checkpoint.""" if self.failed: diff --git a/verl_omni/trainer/diffusion/distillation/ray_trainer.py b/verl_omni/trainer/diffusion/distillation/ray_trainer.py index 3f9dd9fd1..20df259c4 100644 --- a/verl_omni/trainer/diffusion/distillation/ray_trainer.py +++ b/verl_omni/trainer/diffusion/distillation/ray_trainer.py @@ -11,28 +11,97 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Ray-entrypoint-compatible shell around the pure distillation control plane. - -PR 1 deliberately stops before model allocation. The shell accepts the same -constructor protocol and lifecycle calls as the existing diffusion trainers so -``algorithm.trainer_type=distillation`` reaches an explicit PR 2 boundary rather -than failing with a Python signature error. -""" +"""Ray trainer for the architecture-neutral distillation runtime.""" from __future__ import annotations +import hashlib +import json +import os +import random +import shutil +import tempfile +import time +from collections.abc import Mapping +from dataclasses import asdict from typing import Any, Optional -from verl_omni.trainer.diffusion.diffusion_trainer_utils import validate_distillation_config -from verl_omni.trainer.diffusion.distillation.contracts import DistillationPlan +import numpy as np +import torch +from omegaconf import OmegaConf, open_dict +from tensordict import TensorDict +from tqdm import tqdm +from verl import DataProto +from verl.single_controller.ray import RayClassWithInitArgs, RayWorkerGroup +from verl.trainer.ppo.ray_trainer import Role +from verl.utils.checkpoint.checkpoint_manager import find_latest_ckpt_path +from verl.utils.config import omega_conf_to_dataclass +from verl.utils.tracking import Tracking + +from verl_omni.pipelines.model_base import DiffusionModelBase, DistributionMatchingModelAdapter +from verl_omni.trainer.diffusion.diffusion_trainer_utils import ( + _to_diffusion_worker_tensordict, + validate_distillation_config, +) +from verl_omni.trainer.diffusion.distillation.contracts import DistillationPlan, PhaseRequest, TrainerCounters from verl_omni.trainer.diffusion.distillation.controller import DistillationTrainerController from verl_omni.trainer.diffusion.distillation.recipes import build_plan_from_config +from verl_omni.trainer.diffusion.ray_diffusion_trainer import BaseRayDiffusionTrainer +from verl_omni.workers.config import DiffusionModelConfig +from verl_omni.workers.diffusion_distillation_worker import DiffusionDistillationWorkerGroup __all__ = ["DistillationRayTrainer"] -class DistillationRayTrainer: - """Production-compatible driver shell for a validated distillation plan.""" +def checkpoint_json_value(value: Any): + """Serialize immutable plan containers independently of insertion order and hash seed.""" + if isinstance(value, Mapping): + return dict(value) + if isinstance(value, set | frozenset): + return sorted(value) + raise TypeError(f"Unsupported checkpoint fingerprint value: {type(value).__name__}.") + + +class DistillationBatchProvider: + """Stateful dataloader adapter implementing phase batch reuse semantics.""" + + def __init__(self, dataloader) -> None: + self.dataloader = dataloader + self._iterator = iter(dataloader) + self._student_batch: Optional[TensorDict] = None + self.last_data_proto: Optional[DataProto] = None + + def reset_iterator(self) -> None: + """Recreate the iterator after a dataloader state restore.""" + self._iterator = iter(self.dataloader) + self._student_batch = None + self.last_data_proto = None + + def fresh_batch(self) -> TensorDict: + """Read the next batch, restarting the dataloader at epoch boundaries.""" + try: + batch_dict = next(self._iterator) + except StopIteration: + self._iterator = iter(self.dataloader) + batch_dict = next(self._iterator) + batch = DataProto.from_single_dict(batch_dict) + self.last_data_proto = batch + return _to_diffusion_worker_tensordict(batch) + + def next(self, phase: PhaseRequest) -> TensorDict: + """Return a fresh or student-reused batch according to the phase contract.""" + if phase.batch_policy == "reuse_student": + if self._student_batch is None: + raise RuntimeError("A reuse_student phase ran before a student batch was cached.") + return self._student_batch.copy() + batch = self.fresh_batch() + if phase.kind == "student": + self._student_batch = batch.copy() + return batch + + +class DistillationRayTrainer(BaseRayDiffusionTrainer): + """Ray driver over the generic control plane and multi-role worker group.""" def __init__( self, @@ -61,41 +130,159 @@ def __init__( validate_distillation_config(config) if plan is not None and config is not None and capabilities is not None: raise ValueError("Pass either an explicit plan or config+capabilities, not both.") - if plan is None and config is not None and capabilities is not None: - plan = build_plan_from_config(config, capabilities) - - self.config = config - self.tokenizer = tokenizer - self.processor = processor - self.role_worker_mapping = role_worker_mapping - self.resource_pool_manager = resource_pool_manager - self.ray_worker_group_cls = ray_worker_group_cls - self.train_dataset = train_dataset - self.val_dataset = val_dataset - self.collate_fn = collate_fn - self.train_sampler = train_sampler - self.device_name = device_name + + self._production = config is not None and OmegaConf.select(config, "trainer") is not None + if self._production: + role_worker_mapping = role_worker_mapping or {} + ray_worker_group_cls = ray_worker_group_cls or RayWorkerGroup + super().__init__( + config=config, + tokenizer=tokenizer, + processor=processor, + role_worker_mapping=role_worker_mapping, + resource_pool_manager=resource_pool_manager, + ray_worker_group_cls=ray_worker_group_cls, + train_dataset=train_dataset, + val_dataset=val_dataset, + collate_fn=collate_fn, + train_sampler=train_sampler, + device_name=device_name, + ) + else: + self.config = config + self.tokenizer = tokenizer + self.processor = processor + self.role_worker_mapping = role_worker_mapping + self.resource_pool_manager = resource_pool_manager + self.ray_worker_group_cls = ray_worker_group_cls + self.train_dataset = train_dataset + self.val_dataset = val_dataset + self.collate_fn = collate_fn + self.train_sampler = train_sampler + self.device_name = device_name + + if plan is None and config is not None: + if capabilities is None and self._production: + model_config: DiffusionModelConfig = omega_conf_to_dataclass(config.actor_rollout_ref.model) + adapter_cls = DiffusionModelBase.get_class(model_config) + if not issubclass(adapter_cls, DistributionMatchingModelAdapter): + raise TypeError( + f"{adapter_cls.__name__} must mix in DistributionMatchingModelAdapter for distillation." + ) + capabilities = adapter_cls.distillation_capabilities() + if capabilities is not None: + plan = build_plan_from_config(config, capabilities) + self.plan = plan self.capabilities = capabilities self.executor = executor self.batch_provider = batch_provider - self.hooks = hooks + self.hooks = hooks or self self.controller_instance: Optional[DistillationTrainerController] = None + self.distillation_worker_group = None + self.global_steps = 0 + self._logger = None + if self._production and self.plan is not None: + self.validate_runtime_config() - def init_workers(self) -> None: - """Validate the PR 1 boundary before PR 2 supplies role-group workers.""" - if self.executor is None or self.batch_provider is None: - raise NotImplementedError( - "The multi-role distillation workers and architecture capability binding land in PR 2. " - "PR 1 accepts the production trainer interface but does not allocate model workers." + def validate_runtime_config(self) -> None: + """Reject unsupported role storage and distributed batch layouts.""" + distribution_matching = self.config.distillation.distribution_matching + strategy = self.config.actor_rollout_ref.actor.strategy + if strategy not in {"fsdp", "fsdp2"}: + raise ValueError(f"Distillation role groups require strategy 'fsdp' or 'fsdp2', got {strategy!r}.") + if any(group.placement != "colocated" for group in self.plan.role_layout.groups): + raise NotImplementedError("The current runtime supports colocated role groups only.") + if any(binding.role == "discriminator" for binding in self.plan.role_layout.bindings): + raise NotImplementedError("The DMD2 adversarial discriminator data plane lands in PR 4.") + if distribution_matching.role_storage == "shared_base_adapters": + model_config = self.config.actor_rollout_ref.model + lora_rank = model_config.get("lora_rank", model_config.get("lora", {}).get("rank", 0)) + if lora_rank <= 0: + raise ValueError("shared_base_adapters requires actor_rollout_ref.model.lora_rank > 0.") + if strategy == "fsdp" and not self.config.actor_rollout_ref.actor.fsdp_config.use_orig_params: + raise ValueError("shared_base_adapters with FSDP1 requires actor.fsdp_config.use_orig_params=true.") + + world_size = self.config.trainer.n_gpus_per_node * self.config.trainer.nnodes + sequence_parallel_size = self.config.actor_rollout_ref.actor.fsdp_config.ulysses_sequence_parallel_size + if world_size % sequence_parallel_size != 0: + raise ValueError("Trainer world size must be divisible by the Ulysses sequence-parallel size.") + data_parallel_size = world_size // sequence_parallel_size + global_batch_size = self.config.data.get("gen_batch_size", self.config.data.train_batch_size) + if global_batch_size % data_parallel_size != 0: + raise ValueError( + f"Distillation batch size {global_batch_size} must be divisible by data-parallel size " + f"{data_parallel_size}." ) + + def configure_role_steps(self) -> None: + """Set the fake-score scheduler horizon in its own optimizer-step units.""" + distribution_matching = self.config.distillation.distribution_matching + fake_repeats = sum(phase.repeats for phase in self.plan.update_schedule.phases if phase.kind == "fake_score") + warmup_fake_repeats = sum( + phase.repeats for phase in self.plan.update_schedule.warmup_phases if phase.kind == "fake_score" + ) + fake_total_steps = self.total_training_steps * fake_repeats + fake_total_steps += self.plan.update_schedule.warmup_cycles * warmup_fake_repeats + with open_dict(distribution_matching.fake_score_optim): + distribution_matching.fake_score_optim.total_training_steps = fake_total_steps + + def init_workers(self) -> None: + """Create the colocated multi-role Ray worker group or validate injected fakes.""" + if self.executor is not None or self.batch_provider is not None: + if self.executor is None or self.batch_provider is None: + raise ValueError("executor and batch_provider must be supplied together.") + if self.plan is None: + raise ValueError("A validated DistillationPlan is required when an executor is bound.") + return + if not self._production: + raise NotImplementedError("Production multi-role workers require the composed diffusion trainer config.") if self.plan is None: - raise ValueError("A validated DistillationPlan is required when an executor is bound.") + raise ValueError("No distribution-matching architecture adapter produced a DistillationPlan.") + if Role.Actor not in self.role_worker_mapping: + raise ValueError("Distillation training requires a Role.Actor worker mapping.") + + self.configure_role_steps() + self.resource_pool_manager.create_resource_pool() + resource_pool = self.resource_pool_manager.get_resource_pool(Role.Actor) + worker_cls = RayClassWithInitArgs( + cls=self.role_worker_mapping[Role.Actor], + config=self.config, + plan=self.plan, + ) + worker_group_kwargs = {"device_name": self.device_name} + register_timeout = OmegaConf.select(self.config.trainer, "ray_wait_register_center_timeout") + if register_timeout is not None: + worker_group_kwargs["ray_wait_register_center_timeout"] = register_timeout + master_port_range = OmegaConf.select(self.config.trainer, "ray_master_port_range") + if master_port_range is not None: + worker_group_kwargs["master_port_range"] = list(master_port_range) + profile_steps = OmegaConf.select(self.config, "global_profiler.steps") + if profile_steps is not None: + worker_group_kwargs["profile_steps"] = profile_steps + if OmegaConf.select(self.config, "global_profiler.tool") == "nsys": + worker_options = OmegaConf.select( + self.config, + "global_profiler.global_tool_config.nsys.worker_nsight_options", + ) + if worker_options is None: + raise ValueError("Nsight worker options are required when global profiling uses nsys.") + worker_group_kwargs["worker_nsight_options"] = OmegaConf.to_container(worker_options) + self.distillation_worker_group = self.ray_worker_group_cls( + resource_pool=resource_pool, + ray_cls_with_init=worker_cls, + **worker_group_kwargs, + ) + self.distillation_worker_group.init_model() + self.executor = DiffusionDistillationWorkerGroup(self.distillation_worker_group) + self.batch_provider = DistillationBatchProvider(self.train_dataloader) def build_controller(self) -> DistillationTrainerController: - """Construct the pure controller from a plan and bound collaborators.""" + """Construct the pure control plane from the bound worker data plane.""" self.init_workers() assert self.plan is not None + assert self.executor is not None + assert self.batch_provider is not None self.controller_instance = DistillationTrainerController( plan=self.plan, executor=self.executor, @@ -106,11 +293,182 @@ def build_controller(self) -> DistillationTrainerController: @property def controller(self) -> DistillationTrainerController: - """Return the lazily constructed distillation trainer controller.""" if self.controller_instance is None: return self.build_controller() return self.controller_instance - def fit(self, num_cycles: int = 0) -> None: - """Drive the injected CPU controller; production data plane arrives in PR 2.""" - self.controller.run(num_cycles) + def after_completed_step(self, counters: TrainerCounters, metrics: dict, executor: Any) -> None: + """Checkpoint completed student cycles at configured boundaries.""" + self.global_steps = counters.global_step + if not self._production: + return + save_freq = self.config.trainer.save_freq + if save_freq > 0 and (self.global_steps % save_freq == 0 or self.global_steps >= self.total_training_steps): + checkpoint_start = time.perf_counter() + self._save_checkpoint() + metrics.setdefault("system", {})["perf/checkpoint_s"] = time.perf_counter() - checkpoint_start + + def checkpoint_fingerprint(self) -> str: + """Hash canonical plan and optimizer configuration for resume validation.""" + payload = {"plan": asdict(self.plan)} + if self.config is not None and OmegaConf.select(self.config, "distillation.distribution_matching") is not None: + payload["distribution_matching"] = OmegaConf.to_container( + self.config.distillation.distribution_matching, resolve=True + ) + payload["model"] = OmegaConf.to_container(self.config.actor_rollout_ref.model, resolve=True) + payload["student_optimizer"] = OmegaConf.to_container( + self.config.actor_rollout_ref.actor.optim, resolve=True + ) + return hashlib.sha256(json.dumps(payload, sort_keys=True, default=checkpoint_json_value).encode()).hexdigest() + + @staticmethod + def driver_rng_state() -> dict[str, Any]: + """Capture driver RNG state separately from worker sampling streams.""" + return { + "python": random.getstate(), + "numpy": np.random.get_state(), + "torch": torch.get_rng_state(), + } + + @staticmethod + def restore_driver_rng_state(state: dict[str, Any]) -> None: + """Restore the driver RNG streams saved at a completed cycle.""" + random.setstate(state["python"]) + np.random.set_state(state["numpy"]) + torch.set_rng_state(state["torch"]) + + def _save_checkpoint(self) -> None: + """Atomically publish worker, driver, dataloader, and RNG state.""" + if self.executor is None or not hasattr(self.executor, "save_checkpoint"): + raise RuntimeError("The bound distillation executor cannot save checkpoints.") + root = os.path.abspath(self.config.trainer.default_local_dir) + os.makedirs(root, exist_ok=True) + final_path = os.path.join(root, f"global_step_{self.global_steps}") + if os.path.exists(final_path): + raise FileExistsError(f"Refusing to overwrite existing checkpoint {final_path}.") + temporary_path = tempfile.mkdtemp(prefix=f".global_step_{self.global_steps}_", dir=root) + try: + self.executor.save_checkpoint(os.path.join(temporary_path, "workers"), self.global_steps) + torch.save(self.controller.state_dict(), os.path.join(temporary_path, "trainer_state.pt")) + torch.save(self.train_dataloader.state_dict(), os.path.join(temporary_path, "data.pt")) + torch.save(self.driver_rng_state(), os.path.join(temporary_path, "rng.pt")) + with open(os.path.join(temporary_path, "manifest.json"), "w", encoding="utf-8") as file: + json.dump( + { + "plan_name": self.plan.name, + "plan_version": self.plan.version, + "global_step": self.global_steps, + "export_role": self.plan.export.role, + "fingerprint": self.checkpoint_fingerprint(), + }, + file, + indent=2, + sort_keys=True, + ) + os.replace(temporary_path, final_path) + except Exception: + shutil.rmtree(temporary_path, ignore_errors=True) + raise + + latest_tmp = os.path.join(root, ".latest_checkpointed_iteration.tmp") + with open(latest_tmp, "w", encoding="utf-8") as file: + file.write(str(self.global_steps)) + os.replace(latest_tmp, os.path.join(root, "latest_checkpointed_iteration.txt")) + + def _load_checkpoint(self) -> int: + """Restore the last atomically completed distillation cycle.""" + if not self._production or self.config.trainer.resume_mode == "disable": + return 0 + if self.config.trainer.default_hdfs_dir is not None: + raise NotImplementedError("Distillation checkpoint restore from HDFS is not implemented.") + checkpoint_root = os.path.abspath(self.config.trainer.default_local_dir) + if self.config.trainer.resume_mode == "auto": + checkpoint_path = find_latest_ckpt_path(checkpoint_root) + if checkpoint_path is None: + return 0 + elif self.config.trainer.resume_mode == "resume_path": + checkpoint_path = os.path.abspath(self.config.trainer.resume_from_path) + else: + raise ValueError(f"Unsupported trainer.resume_mode {self.config.trainer.resume_mode!r}.") + + manifest_path = os.path.join(checkpoint_path, "manifest.json") + if not os.path.isfile(manifest_path): + raise FileNotFoundError(f"Incomplete distillation checkpoint: missing {manifest_path}.") + with open(manifest_path, encoding="utf-8") as file: + manifest = json.load(file) + if manifest.get("plan_name") != self.plan.name or manifest.get("plan_version") != self.plan.version: + raise ValueError("Checkpoint recipe identity does not match the active DistillationPlan.") + if manifest.get("fingerprint") != self.checkpoint_fingerprint(): + raise ValueError("Checkpoint distillation plan or optimizer configuration does not match the active run.") + + self.executor.load_checkpoint(os.path.join(checkpoint_path, "workers")) + trainer_state = torch.load(os.path.join(checkpoint_path, "trainer_state.pt"), weights_only=False) + self.controller.load_state_dict(trainer_state) + self.train_dataloader.load_state_dict(torch.load(os.path.join(checkpoint_path, "data.pt"), weights_only=False)) + if hasattr(self.batch_provider, "reset_iterator"): + self.batch_provider.reset_iterator() + self.restore_driver_rng_state(torch.load(os.path.join(checkpoint_path, "rng.pt"), weights_only=False)) + self.global_steps = self.controller.counters.global_step + return self.global_steps + + @staticmethod + def flatten_metrics(metrics: dict[str, dict]) -> dict[str, float]: + """Flatten phase metrics for the existing tracking backends.""" + return {key: value for phase_metrics in metrics.values() for key, value in phase_metrics.items()} + + def profile_workers(self, *, start: bool, step: int) -> None: + """Start or stop the configured distributed profiler.""" + if self.distillation_worker_group is None: + return + if start: + self.distillation_worker_group.start_profile(role="distillation", profile_step=step) + else: + self.distillation_worker_group.stop_profile() + + def fit(self, num_cycles: Optional[int] = None) -> None: + """Run injected CPU cycles or the production dataloader-backed control plane.""" + if not self._production: + if num_cycles is None: + num_cycles = 0 + self.controller.run(num_cycles) + return + + controller = self.controller + self._load_checkpoint() + self._logger = Tracking( + project_name=self.config.trainer.project_name, + experiment_name=self.config.trainer.experiment_name, + default_backend=self.config.trainer.logger, + config=OmegaConf.to_container(self.config, resolve=True), + ) + target_steps = ( + self.total_training_steps + if num_cycles is None + else min(self.total_training_steps, controller.counters.global_step + num_cycles) + ) + progress_bar = tqdm( + total=target_steps, + initial=controller.counters.global_step, + desc="Distillation Training", + ) + profile_steps = OmegaConf.select(self.config, "global_profiler.steps", default=None) + while controller.counters.global_step < target_steps: + before_global_step = controller.counters.global_step + profile_step = before_global_step + 1 + do_profile = profile_steps is not None and profile_step in profile_steps + self.profile_workers(start=do_profile, step=profile_step) + try: + controller.run_cycle() + finally: + if do_profile: + self.profile_workers(start=False, step=profile_step) + self.global_steps = controller.counters.global_step + metrics = self.flatten_metrics(controller.metrics) + metrics["training/global_step"] = float(self.global_steps) + metrics["training/completed_cycles"] = float(controller.counters.completed_cycles) + self._logger.log(data=metrics, step=self.global_steps) + if self.global_steps > before_global_step: + progress_bar.update(1) + if hasattr(self.train_dataset, "on_batch_end"): + self.train_dataset.on_batch_end(batch=self.batch_provider.last_data_proto) + progress_bar.close() diff --git a/verl_omni/trainer/diffusion/distillation/recipes.py b/verl_omni/trainer/diffusion/distillation/recipes.py index 1b6c0c90b..c26e3bf77 100644 --- a/verl_omni/trainer/diffusion/distillation/recipes.py +++ b/verl_omni/trainer/diffusion/distillation/recipes.py @@ -306,6 +306,46 @@ def causal_bidirectional_layout(causal_model_ref: str, bidirectional_model_ref: ) +def apply_role_storage(layout: RoleLayoutSpec, role_storage: str) -> RoleLayoutSpec: + """Materialize either shared groups or one colocated model per logical role.""" + require_choice( + "role_storage", + role_storage, + {"shared_base_adapters", "colocated_independent"}, + ) + if role_storage == "shared_base_adapters": + return layout + + source_groups = {group.name: group for group in layout.groups} + groups = [] + bindings = [] + for binding in layout.bindings: + source_group = source_groups[binding.group] + group_name = f"{binding.role}_model" + groups.append( + RoleGroupSpec( + name=group_name, + model_ref=source_group.model_ref, + storage="independent_module", + placement="colocated", + ) + ) + bindings.append( + RoleBinding( + role=binding.role, + group=group_name, + adapter=binding.adapter, + trainable=binding.trainable, + optimizer_key=binding.optimizer_key, + ) + ) + return RoleLayoutSpec( + groups=tuple(groups), + bindings=tuple(bindings), + score_transport=layout.score_transport, + ) + + def build_update_schedule( fake_repeats: int, fake_warmup_cycles: int = 0, with_discriminator: bool = False ) -> UpdateSchedule: @@ -380,7 +420,10 @@ def build_plan(cls, config, capabilities) -> DistillationPlan: model_ref = get_config_value(config, "model_path", "") or "" return DistillationPlan( name="dmd", - role_layout=shared_base_layout(model_ref), + role_layout=apply_role_storage( + shared_base_layout(model_ref), + get_config_or_default(config, "role_storage", "shared_base_adapters"), + ), data_requirements={"mode": data_mode}, objective={"name": "dmd", "profile": profile}, rollout={"strategy": rollout}, @@ -414,7 +457,10 @@ def build_plan(cls, config, capabilities) -> DistillationPlan: model_ref = get_config_value(config, "model_path", "") or "" return DistillationPlan( name="dmd2", - role_layout=shared_base_layout(model_ref, with_discriminator=adversarial), + role_layout=apply_role_storage( + shared_base_layout(model_ref, with_discriminator=adversarial), + get_config_or_default(config, "role_storage", "shared_base_adapters"), + ), data_requirements={"mode": data_mode}, objective={"name": "dmd2", "profile": profile, "adversarial": adversarial}, rollout={"strategy": rollout}, @@ -448,7 +494,10 @@ def build_plan(cls, config, capabilities) -> DistillationPlan: bidirectional_ref = get_config_value(config, "bidirectional_model_path", model_ref) or model_ref return DistillationPlan( name="causvid", - role_layout=causal_bidirectional_layout(causal_ref, bidirectional_ref), + role_layout=apply_role_storage( + causal_bidirectional_layout(causal_ref, bidirectional_ref), + get_config_or_default(config, "role_storage", "shared_base_adapters"), + ), data_requirements={"mode": data_mode}, objective={"name": "dmd", "profile": "distribution_only"}, rollout={"strategy": rollout}, @@ -480,7 +529,10 @@ def build_plan(cls, config, capabilities) -> DistillationPlan: bidirectional_ref = get_config_value(config, "bidirectional_model_path", model_ref) or model_ref return DistillationPlan( name="self_forcing", - role_layout=causal_bidirectional_layout(causal_ref, bidirectional_ref), + role_layout=apply_role_storage( + causal_bidirectional_layout(causal_ref, bidirectional_ref), + get_config_or_default(config, "role_storage", "shared_base_adapters"), + ), data_requirements={"mode": data_mode}, objective={"name": "dmd", "profile": "distribution_only"}, rollout={"strategy": rollout}, @@ -526,7 +578,7 @@ def build_plan_from_config(config, capabilities) -> DistillationPlan: "fake_warmup_cycles": get_config_value(distribution_matching, "fake_warmup_cycles", 0), "export_role": get_config_value(distribution_matching, "export_role", "student_ema"), } - for optional_key in ("profile", "fake_update_ratio", "rollout_strategy", "data_mode"): + for optional_key in ("profile", "fake_update_ratio", "rollout_strategy", "data_mode", "role_storage"): value = get_config_value(distribution_matching, optional_key) if value is not None: recipe_config[optional_key] = value diff --git a/verl_omni/trainer/main_diffusion.py b/verl_omni/trainer/main_diffusion.py index a0ace0966..cfcc1223e 100644 --- a/verl_omni/trainer/main_diffusion.py +++ b/verl_omni/trainer/main_diffusion.py @@ -160,6 +160,17 @@ def add_actor_rollout_worker(self, config): from verl.single_controller.ray import RayWorkerGroup from verl.trainer.ppo.ray_trainer import Role + if config.algorithm.trainer_type == "distillation": + from verl_omni.workers.diffusion_distillation_worker import DiffusionDistillationWorker + + if config.algorithm.sample_source != "offline": + raise ValueError("The distillation trainer requires algorithm.sample_source=offline.") + if not hasattr(Role, "Actor"): + raise ValueError("Distillation training requires verl Role.Actor support.") + self.role_worker_mapping[Role.Actor] = ray.remote(DiffusionDistillationWorker) + self.mapping[Role.Actor] = "global_pool" + return DiffusionDistillationWorker, RayWorkerGroup + from verl_omni.workers.engine_workers import ActorRolloutRefWorker actor_rollout_cls = ActorRolloutRefWorker diff --git a/verl_omni/workers/config/diffusion/distillation.py b/verl_omni/workers/config/diffusion/distillation.py index 3e008628d..78863465a 100644 --- a/verl_omni/workers/config/diffusion/distillation.py +++ b/verl_omni/workers/config/diffusion/distillation.py @@ -16,6 +16,12 @@ from typing import Optional from verl.base_config import BaseConfig +from verl.workers.config import FSDPOptimizerConfig + + +def default_fake_score_optimizer() -> FSDPOptimizerConfig: + return FSDPOptimizerConfig(lr=2e-5, weight_decay=0.01, clip_grad=1.0, lr_scheduler_type="constant") + __all__ = [ "DiffusionDistillationTeacherModelConfig", @@ -73,6 +79,18 @@ class DiffusionDistributionMatchingConfig(BaseConfig): data_mode: Optional[str] = None # Semantic role exported to inference replicas. export_role: str = "student_ema" + # Physical storage used by the initial colocated runtime. + role_storage: str = "shared_base_adapters" + # Per-device student phase micro-batch size. + student_micro_batch_size_per_gpu: int = 1 + # Per-device fake-score phase micro-batch size. + fake_score_micro_batch_size_per_gpu: int = 1 + # Independent fake-score optimizer and scheduler configuration. + fake_score_optim: FSDPOptimizerConfig = field(default_factory=default_fake_score_optimizer) + # EMA decay applied after successful student optimizer steps. + ema_decay: float = 0.999 + # First completed student step that updates EMA. + ema_start_step: int = 0 def __post_init__(self): valid_recipes = {"dmd", "dmd2", "causvid", "self_forcing"} @@ -103,6 +121,22 @@ def __post_init__(self): valid_export_roles = {"student", "student_ema"} if self.export_role not in valid_export_roles: raise ValueError(f"Invalid export_role: {self.export_role}. Must be one of {sorted(valid_export_roles)}") + valid_role_storage = {"shared_base_adapters", "colocated_independent"} + if self.role_storage not in valid_role_storage: + raise ValueError(f"Invalid role_storage: {self.role_storage}. Must be one of {sorted(valid_role_storage)}") + if self.student_micro_batch_size_per_gpu <= 0: + raise ValueError( + f"student_micro_batch_size_per_gpu must be greater than 0, got {self.student_micro_batch_size_per_gpu}" + ) + if self.fake_score_micro_batch_size_per_gpu <= 0: + raise ValueError( + "fake_score_micro_batch_size_per_gpu must be greater than 0, " + f"got {self.fake_score_micro_batch_size_per_gpu}" + ) + if not 0.0 <= self.ema_decay <= 1.0: + raise ValueError(f"ema_decay must be in [0, 1], got {self.ema_decay}") + if self.ema_start_step < 0: + raise ValueError(f"ema_start_step must be non-negative, got {self.ema_start_step}") @dataclass diff --git a/verl_omni/workers/diffusion_distillation_worker.py b/verl_omni/workers/diffusion_distillation_worker.py new file mode 100644 index 000000000..85fe47a9f --- /dev/null +++ b/verl_omni/workers/diffusion_distillation_worker.py @@ -0,0 +1,611 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Ray worker and local runtime for multi-role diffusion distillation.""" + +from __future__ import annotations + +import json +import math +import os +import time +from collections.abc import Mapping +from contextlib import ExitStack, contextmanager +from copy import deepcopy +from dataclasses import asdict, dataclass +from typing import Optional, Protocol, runtime_checkable + +import torch +from omegaconf import DictConfig +from tensordict import TensorDict +from verl.protocol import DataProtoFuture +from verl.single_controller.base import Worker +from verl.single_controller.base.decorator import Dispatch, make_nd_compute_dataproto_dispatch_fn, register +from verl.utils import tensordict_utils as tu +from verl.utils.config import omega_conf_to_dataclass +from verl.utils.device import get_device_id, get_device_name, get_torch_device, is_npu_available +from verl.utils.distributed import initialize_global_process_group_ray, set_numa_affinity +from verl.utils.profiler import DistProfiler, DistProfilerExtension, ProfilerConfig +from verl.workers.config import FSDPOptimizerConfig +from verl.workers.engine import EngineRegistry + +from verl_omni.pipelines.model_base import DiffusionModelBase, DistributionMatchingModelAdapter +from verl_omni.trainer.diffusion.distillation.contracts import ( + DistillationPlan, + PhaseRequest, + PhaseResult, + RoleBinding, +) +from verl_omni.utils.fs import resolve_model_local_dir +from verl_omni.workers.config import DiffusionModelConfig +from verl_omni.workers.config.diffusion import DiffusionDistributionMatchingConfig +from verl_omni.workers.engine.fsdp.distillation_impl import DistillationRoleGroupEngine + +__all__ = [ + "DistillationPhaseComputation", + "DistributionMatchingComputer", + "DistillationRoleRuntime", + "DiffusionDistillationWorker", + "DiffusionDistillationWorkerGroup", +] + + +def resolve_profiler_configs(omega_profiler_config): + """Resolve the selected profiler and its typed tool configuration.""" + profiler_config = omega_conf_to_dataclass(omega_profiler_config, dataclass_type=ProfilerConfig) + tool = omega_profiler_config.get("tool", None) + if tool in {"npu", "nsys", "torch", "torch_memory", "precision_debugger"}: + tool_config = omega_conf_to_dataclass(omega_profiler_config.get("tool_config", {}).get(tool)) + else: + tool_config = None + return profiler_config, tool_config + + +@dataclass +class DistillationPhaseComputation: + """Role means and metrics; an explicit denominator opts into global element reduction. + + Without ``loss_normalizer``, losses are sample means. Otherwise each scalar + is a numerator divided by this count, shared by its ``*/loss`` metrics. + """ + + losses: dict[str, torch.Tensor] + metrics: dict[str, float] + loss_normalizer: Optional[float] = None + + +@runtime_checkable +class DistributionMatchingComputer(Protocol): + """Architecture-owned differentiable computation for one generic phase.""" + + def compute_phase( + self, + request: PhaseRequest, + batch: TensorDict, + runtime: DistillationRoleRuntime, + ) -> DistillationPhaseComputation: + """Build scalar losses while the runtime owns role modules and optimization.""" + ... + + def state_dict(self) -> dict: + """Return architecture-owned RNG and rollout state.""" + ... + + def load_state_dict(self, state: dict) -> None: + """Restore architecture-owned RNG and rollout state.""" + ... + + +class DistillationRoleRuntime: + """Role-to-engine router shared by local tests and the Ray worker.""" + + def __init__( + self, + plan: DistillationPlan, + engines: Mapping[str, DistillationRoleGroupEngine], + *, + ema_decay: float, + ema_start_step: int, + micro_batch_sizes: Optional[Mapping[str, int]] = None, + ) -> None: + self.plan = plan + self.engines = dict(engines) + self.bindings = {binding.role: binding for binding in plan.role_layout.bindings} + self.ema_decay = ema_decay + self.ema_start_step = ema_start_step + self.micro_batch_sizes = dict(micro_batch_sizes or {"student": 1, "fake_score": 1}) + self._loss_normalizers: dict[str, Optional[float]] = {} + expected_groups = {group.name for group in plan.role_layout.groups} + missing_groups = expected_groups - set(self.engines) + extra_groups = set(self.engines) - expected_groups + if missing_groups or extra_groups: + raise ValueError( + f"Role-group engines must match the plan exactly; missing={sorted(missing_groups)}, " + f"extra={sorted(extra_groups)}." + ) + if not 0.0 <= ema_decay <= 1.0: + raise ValueError(f"EMA decay must be in [0, 1], got {ema_decay}.") + if ema_start_step < 0: + raise ValueError(f"EMA start step must be non-negative, got {ema_start_step}.") + if set(self.micro_batch_sizes) != {"student", "fake_score"} or any( + isinstance(size, bool) or not isinstance(size, int) or size <= 0 for size in self.micro_batch_sizes.values() + ): + raise ValueError("micro_batch_sizes must define positive integer student and fake_score sizes.") + self.initialize_ema() + + def engine_for_role(self, role: str) -> DistillationRoleGroupEngine: + """Resolve the physical engine backing a semantic role.""" + try: + binding = self.bindings[role] + except KeyError: + raise KeyError(f"Unknown distillation role {role!r}; bound roles: {sorted(self.bindings)}.") from None + return self.engines[binding.group] + + @contextmanager + def use_role(self, role: str, *, grad_enabled: Optional[bool] = None): + """Yield a role model with explicit gradient intent and restore it afterward.""" + engine = self.engine_for_role(role) + with engine.use_role(role, grad_enabled=grad_enabled) as module: + yield module + + def scheduler_for_role(self, role: str): + """Return the diffusion scheduler belonging to a role's physical group.""" + return self.engine_for_role(role).scheduler + + def micro_batch_size(self, phase_kind: str) -> int: + """Return the per-device micro-batch size for one phase kind.""" + try: + return self.micro_batch_sizes[phase_kind] + except KeyError: + raise ValueError(f"Unknown distillation phase kind {phase_kind!r}.") from None + + def model_config_for_role(self, role: str): + """Return the resolved model config belonging to a semantic role.""" + return self.engine_for_role(role).model_config + + def export_tensors(self, *, base_sync_done: bool): + """Export the plan-selected student or EMA role for inference sync.""" + role = self.plan.export.role + return self.engine_for_role(role).iter_export_tensors(role, base_sync_done) + + def zero_grad(self, roles: tuple[str, ...]) -> None: + """Clear gradient state for each requested trainable role.""" + for role in roles: + self.engine_for_role(role).optimizer_zero_grad(role) + self._loss_normalizers.pop(role, None) + + def validate_computation( + self, + request: PhaseRequest, + computation: DistillationPhaseComputation, + ) -> str: + """Require one graph-bearing scalar loss owned by the requested role.""" + if not isinstance(computation, DistillationPhaseComputation): + raise TypeError(f"compute_phase must return DistillationPhaseComputation, got {type(computation)}.") + expected_roles = set(request.trainable_roles) + if set(computation.losses) != expected_roles: + raise ValueError( + f"Phase computation losses must match requested roles {sorted(expected_roles)}, " + f"got {sorted(computation.losses)}." + ) + if len(request.trainable_roles) != 1: + raise NotImplementedError( + "Multi-role optimizer phases are not supported by the current distillation runtime." + ) + role = request.trainable_roles[0] + loss = computation.losses[role] + if loss.ndim != 0: + raise ValueError(f"Role loss must be scalar, got shape {tuple(loss.shape)} for {role!r}.") + if not loss.requires_grad: + raise ValueError(f"Role loss for {role!r} must retain an autograd graph.") + return role + + def backward_micro_batch( + self, + request: PhaseRequest, + computation: DistillationPhaseComputation, + *, + weight: float, + ) -> None: + """Accumulate one weighted micro-batch loss without stepping.""" + if not 0.0 < weight <= 1.0: + raise ValueError(f"Micro-batch weight must be in (0, 1], got {weight}.") + role = self.validate_computation(request, computation) + normalizer = computation.loss_normalizer + if normalizer is not None and ( + isinstance(normalizer, bool) or not math.isfinite(normalizer) or normalizer <= 0 + ): + raise ValueError(f"Loss normalizer must be finite and positive, got {normalizer}.") + previous = self._loss_normalizers.get(role) + if role in self._loss_normalizers and (previous is None) != (normalizer is None): + raise ValueError(f"Role {role!r} cannot mix sample and element loss reductions in one phase.") + self._loss_normalizers[role] = None if normalizer is None else (previous or 0.0) + normalizer + self.engine_for_role(role).backward_role( + role, computation.losses[role] * (weight if normalizer is None else normalizer) + ) + + def normalize_role_gradients(self, role: str) -> None: + """Divide accumulated numerator gradients by the matching DP-averaged count.""" + normalizer = self._loss_normalizers.pop(role, None) + if normalizer is None: + return + engine = self.engine_for_role(role) + group = engine.get_data_parallel_group() + if group is not None: + count = torch.tensor(normalizer, dtype=torch.float32, device=get_device_id()) + torch.distributed.all_reduce(count, op=torch.distributed.ReduceOp.AVG, group=group) + normalizer = float(count.item()) + for parameter in engine.parameters_for_role(role): + if parameter.grad is not None: + parameter.grad.div_(normalizer) + + def step_phase(self, request: PhaseRequest) -> tuple[dict[str, int], dict[str, float]]: + """Step the phase optimizer once after all micro-batches were accumulated.""" + if len(request.trainable_roles) != 1: + raise NotImplementedError( + "Multi-role optimizer phases are not supported by the current distillation runtime." + ) + role = request.trainable_roles[0] + for role_engine in self.engines.values(): + if hasattr(role_engine, "assert_gradient_isolation"): + role_engine.assert_gradient_isolation({role}) + engine = self.engine_for_role(role) + self.normalize_role_gradients(role) + optimizer_start = time.perf_counter() + stepped, grad_norm = engine.optimizer_step(role) + metrics = { + f"{role}/grad_norm": grad_norm, + f"perf/{role}_optimizer_s": time.perf_counter() - optimizer_start, + } + if stepped and getattr(engine, "lr_schedulers", {}).get(role) is not None: + metrics[f"{role}/lr"] = float(engine.lr_schedulers[role].get_last_lr()[0]) + if not stepped: + return {}, metrics + if request.update_ema and request.global_step + 1 >= self.ema_start_step: + ema_start = time.perf_counter() + self.update_ema() + metrics["ema/decay"] = self.ema_decay + metrics["perf/ema_update_s"] = time.perf_counter() - ema_start + return {role: 1}, metrics + + def backward_and_step( + self, + request: PhaseRequest, + computation: DistillationPhaseComputation, + ) -> tuple[dict[str, int], dict[str, float]]: + """Convenience path for a one-micro-batch phase.""" + self.backward_micro_batch(request, computation, weight=1.0) + optimizer_steps, metrics = self.step_phase(request) + role = request.trainable_roles[0] + metrics.update(computation.metrics) + metrics[f"{role}/loss"] = float(computation.losses[role].detach().float().item()) + return optimizer_steps, metrics + + def initialize_ema(self) -> None: + """Initialize the semantic EMA role exactly from the student role.""" + self.update_ema_parameters(decay=0.0) + + def update_ema(self) -> None: + """Update the semantic student EMA in shared or independent storage.""" + self.update_ema_parameters(decay=self.ema_decay) + + def update_ema_parameters(self, decay: float) -> None: + """Route EMA updates according to the physical role-group layout.""" + student_engine = self.engine_for_role("student") + ema_engine = self.engine_for_role("student_ema") + if student_engine is ema_engine: + student_engine.update_role_ema("student", "student_ema", decay) + else: + ema_engine.update_module_ema_from(student_engine, decay) + + def reduce_metrics(self, metrics: dict[str, float]) -> dict[str, float]: + """Average scalar metrics across data-parallel replicas.""" + if not metrics: + return metrics + first_engine = next(iter(self.engines.values())) + group = first_engine.get_data_parallel_group() + if group is None: + return metrics + names = sorted(metrics) + values = torch.tensor([metrics[name] for name in names], dtype=torch.float32, device=get_device_id()) + torch.distributed.all_reduce(values, op=torch.distributed.ReduceOp.AVG, group=group) + return {name: value for name, value in zip(names, values.cpu().tolist(), strict=True)} + + def group_metrics(self) -> dict[str, float]: + """Return stable placement diagnostics for logging.""" + return { + "memory/role_group_count": float(len(self.engines)), + "memory/base_model_copies": float(len(self.engines)), + } + + +class DiffusionDistillationWorker(Worker, DistProfilerExtension): + """Own colocated role-group engines and execute one phase per Ray RPC.""" + + def __init__(self, config: DictConfig, plan: DistillationPlan): + Worker.__init__(self) + if is_npu_available: + os.environ["PYTORCH_NPU_ALLOC_CONF"] = "expandable_segments:True" + initialize_global_process_group_ray(timeout_second=None) + set_numa_affinity() + + self.config = config + self.plan = plan + self.device_name = get_device_name() + profiler_config, tool_config = resolve_profiler_configs(config.actor_rollout_ref.actor.get("profiler", {})) + DistProfilerExtension.__init__( + self, + DistProfiler(rank=self.rank, config=profiler_config, tool_config=tool_config), + ) + self.runtime: Optional[DistillationRoleRuntime] = None + self.dm_computer: Optional[DistributionMatchingComputer] = None + + def build_optimizer_configs( + self, + bindings: tuple[RoleBinding, ...], + student_optimizer_config: FSDPOptimizerConfig, + distillation_config: DiffusionDistributionMatchingConfig, + ) -> dict[str, FSDPOptimizerConfig]: + """Give each trainable role its own optimizer configuration.""" + configs = {} + for binding in bindings: + if not binding.trainable: + continue + if binding.role == "student": + configs[binding.role] = deepcopy(student_optimizer_config) + elif binding.role == "fake_score": + configs[binding.role] = deepcopy(distillation_config.fake_score_optim) + else: + raise NotImplementedError(f"No optimizer config is defined for role {binding.role!r}.") + return configs + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def init_model(self) -> None: + """Allocate every physical group once and bind the architecture phase runner.""" + model_config: DiffusionModelConfig = omega_conf_to_dataclass(self.config.actor_rollout_ref.model) + # DMD losses live in the phase runner, so instantiate only the actor sub-configs the role engine consumes. + actor_config = self.config.actor_rollout_ref.actor + actor_engine_config = omega_conf_to_dataclass(actor_config.fsdp_config) + object.__setattr__(actor_engine_config, "strategy", actor_config.strategy) + student_optimizer_config: FSDPOptimizerConfig = omega_conf_to_dataclass( + actor_config.optim, + dataclass_type=FSDPOptimizerConfig, + ) + checkpoint_config = omega_conf_to_dataclass(actor_config.checkpoint) + distillation_config: DiffusionDistributionMatchingConfig = omega_conf_to_dataclass( + self.config.distillation.distribution_matching, + dataclass_type=DiffusionDistributionMatchingConfig, + ) + adapter_cls = DiffusionModelBase.get_class(model_config) + if not issubclass(adapter_cls, DistributionMatchingModelAdapter): + raise TypeError( + f"{adapter_cls.__name__} must mix in DistributionMatchingModelAdapter for distillation training." + ) + + engines = {} + resolved_model_paths = {} + for group in self.plan.role_layout.groups: + if group.placement != "colocated": + raise NotImplementedError("The current runtime implements colocated role groups only.") + bindings = tuple(binding for binding in self.plan.role_layout.bindings if binding.group == group.name) + trainable_bindings = tuple(binding for binding in bindings if binding.trainable) + group_model_config = deepcopy(model_config) + object.__setattr__(group_model_config, "model_type", "diffusion_distillation_model") + object.__setattr__(group_model_config, "path", group.model_ref) + if group.model_ref not in resolved_model_paths: + resolved_model_paths[group.model_ref] = resolve_model_local_dir( + group.model_ref, use_shm=group_model_config.use_shm + ) + object.__setattr__(group_model_config, "local_path", resolved_model_paths[group.model_ref]) + adapters = tuple(binding.adapter for binding in bindings if binding.adapter is not None) + if group.storage == "shared_base_adapters" and group_model_config.lora_rank <= 0: + raise ValueError("shared_base_adapters requires actor_rollout_ref.model.lora_rank > 0.") + if group.storage == "independent_module" and not adapters: + object.__setattr__(group_model_config, "lora_rank", 0) + object.__setattr__(group_model_config, "lora_adapter_path", None) + if group_model_config.lora_rank > 0 and adapters: + object.__setattr__( + group_model_config, + "policy_state_adapters", + tuple(dict.fromkeys(("default", *adapters, "reference"))), + ) + + group_engine_config = deepcopy(actor_engine_config) + object.__setattr__(group_engine_config, "forward_only", not trainable_bindings) + role_optimizer_configs = self.build_optimizer_configs( + bindings, student_optimizer_config, distillation_config + ) + primary_optimizer_config = next(iter(role_optimizer_configs.values()), deepcopy(student_optimizer_config)) + engine = EngineRegistry.new( + model_type="diffusion_distillation_model", + backend=group_engine_config.strategy, + model_config=group_model_config, + engine_config=group_engine_config, + optimizer_config=primary_optimizer_config, + checkpoint_config=deepcopy(checkpoint_config), + role_group=group, + role_bindings=bindings, + optimizer_configs=role_optimizer_configs, + ) + engine.initialize() + engines[group.name] = engine + + self.runtime = DistillationRoleRuntime( + self.plan, + engines, + ema_decay=distillation_config.ema_decay, + ema_start_step=distillation_config.ema_start_step, + micro_batch_sizes={ + "student": distillation_config.student_micro_batch_size_per_gpu, + "fake_score": distillation_config.fake_score_micro_batch_size_per_gpu, + }, + ) + self.dm_computer = adapter_cls.build_distribution_matching_computer(model_config, self.plan) + if not isinstance(self.dm_computer, DistributionMatchingComputer): + raise TypeError( + "build_distribution_matching_computer() must return an object implementing " + "compute_phase(), state_dict(), and load_state_dict()." + ) + first_engine = next(iter(engines.values())) + self._register_dispatch_collect_info( + mesh_name="distillation", + dp_rank=first_engine.get_data_parallel_rank(), + is_collect=first_engine.is_mp_src_rank_with_outputs(), + ) + + @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="distillation"), blocking=True) + @DistProfiler.annotate(color="red", role="distillation_phase") + def execute_phase(self, data: TensorDict) -> TensorDict: + """Resolve rank failures before lazy collect metadata RPCs can wait behind peer collectives.""" + if self.runtime is None or self.dm_computer is None: + raise RuntimeError("init_model() must be called before execute_phase().") + request = tu.pop(data, key="phase_request") + if not isinstance(request, PhaseRequest): + raise TypeError(f"phase_request must be PhaseRequest, got {type(request)}.") + self.runtime.zero_grad(request.trainable_roles) + device_module = get_torch_device() + device_module.reset_peak_memory_stats() + start = time.perf_counter() + if not data.batch_size: + raise ValueError("A distillation phase requires a TensorDict with a leading batch dimension.") + total_samples = data.batch_size[0] + if total_samples <= 0: + raise ValueError("A distillation phase cannot execute an empty batch.") + micro_batch_size = self.runtime.micro_batch_size(request.kind) + accumulated_metrics: dict[str, float] = {} + accumulated_losses: dict[str, float] = {} + loss_denominators: dict[str, float] = {} + with ExitStack() as stack: + for engine in self.runtime.engines.values(): + context = engine.train_mode() if engine.optimizers else engine.eval_mode() + stack.enter_context(context) + forward_duration = 0.0 + backward_duration = 0.0 + for micro_batch in data.split(micro_batch_size, dim=0): + weight = micro_batch.batch_size[0] / total_samples + micro_batch = micro_batch.to(get_device_id()) + forward_start = time.perf_counter() + computation = self.dm_computer.compute_phase(request, micro_batch, self.runtime) + forward_duration += time.perf_counter() - forward_start + backward_start = time.perf_counter() + self.runtime.backward_micro_batch(request, computation, weight=weight) + backward_duration += time.perf_counter() - backward_start + normalizer = computation.loss_normalizer + loss_weight = weight if normalizer is None else normalizer + for name, value in computation.metrics.items(): + metric_weight = weight + if normalizer is not None and name.endswith("/loss"): + metric_weight = normalizer + loss_denominators[name] = loss_denominators.get(name, 0.0) + normalizer + accumulated_metrics[name] = accumulated_metrics.get(name, 0.0) + float(value) * metric_weight + for role, loss in computation.losses.items(): + accumulated_losses[role] = ( + accumulated_losses.get(role, 0.0) + float(loss.detach().float()) * loss_weight + ) + if normalizer is not None: + name = f"{role}/loss" + loss_denominators[name] = loss_denominators.get(name, 0.0) + normalizer + optimizer_steps, step_metrics = self.runtime.step_phase(request) + metrics = {**accumulated_metrics, **step_metrics} + metrics.update({f"{role}/loss": loss for role, loss in accumulated_losses.items()}) + metrics[f"perf/{request.kind}_forward_s"] = forward_duration + metrics[f"perf/{request.kind}_backward_s"] = backward_duration + metrics[f"perf/{request.kind}_s"] = time.perf_counter() - start + metrics["memory/max_allocated_gb"] = device_module.max_memory_allocated() / (1024**3) + metrics["memory/max_reserved_gb"] = device_module.max_memory_reserved() / (1024**3) + metrics.update(self.runtime.group_metrics()) + metrics = self.runtime.reduce_metrics(metrics) + for name, denominator in self.runtime.reduce_metrics(loss_denominators).items(): + metrics[name] /= denominator + return tu.get_tensordict( + tensor_dict={}, + non_tensor_dict={"metrics": metrics, "optimizer_steps": optimizer_steps}, + ) + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def save_checkpoint(self, local_path: str, global_step: int) -> None: + """Save every physical role group under one checkpoint root.""" + if self.runtime is None: + raise RuntimeError("init_model() must be called before save_checkpoint().") + os.makedirs(local_path, exist_ok=True) + for group_name, engine in self.runtime.engines.items(): + engine.save_role_group_checkpoint(os.path.join(local_path, "role_groups", group_name), global_step) + computer_state = self.dm_computer.state_dict() if hasattr(self.dm_computer, "state_dict") else {} + torch.save(computer_state, os.path.join(local_path, f"dm_computer_rank_{self.rank}.pt")) + if self.rank == 0: + with open(os.path.join(local_path, "worker_manifest.json"), "w", encoding="utf-8") as file: + json.dump( + { + "plan_name": self.plan.name, + "plan_version": self.plan.version, + "groups": [asdict(group) for group in self.plan.role_layout.groups], + "bindings": [asdict(binding) for binding in self.plan.role_layout.bindings], + }, + file, + indent=2, + sort_keys=True, + ) + torch.distributed.barrier() + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def load_checkpoint(self, local_path: str) -> None: + """Restore every role group after validating the worker manifest.""" + if self.runtime is None: + raise RuntimeError("init_model() must be called before load_checkpoint().") + manifest_path = os.path.join(local_path, "worker_manifest.json") + if not os.path.isfile(manifest_path): + raise FileNotFoundError(f"Missing distillation worker manifest: {manifest_path}") + with open(manifest_path, encoding="utf-8") as file: + manifest = json.load(file) + if manifest.get("plan_name") != self.plan.name or manifest.get("plan_version") != self.plan.version: + raise ValueError("Checkpoint recipe identity does not match the active distillation plan.") + expected_groups = [asdict(group) for group in self.plan.role_layout.groups] + expected_bindings = [asdict(binding) for binding in self.plan.role_layout.bindings] + if manifest.get("groups") != expected_groups or manifest.get("bindings") != expected_bindings: + raise ValueError("Checkpoint role layout does not match the active distillation plan.") + for group_name, engine in self.runtime.engines.items(): + engine.load_role_group_checkpoint(os.path.join(local_path, "role_groups", group_name)) + computer_state_path = os.path.join(local_path, f"dm_computer_rank_{self.rank}.pt") + if not os.path.isfile(computer_state_path): + raise FileNotFoundError(f"Missing phase-runner state: {computer_state_path}") + computer_state = torch.load(computer_state_path, map_location="cpu", weights_only=False) + if computer_state: + if not hasattr(self.dm_computer, "load_state_dict"): + raise ValueError("Checkpoint contains phase-runner state, but the active runner cannot restore it.") + self.dm_computer.load_state_dict(computer_state) + + +class DiffusionDistillationWorkerGroup: + """Driver-side executor facade over a Ray worker group.""" + + def __init__(self, worker_group) -> None: + self.worker_group = worker_group + + def execute_phase(self, request: PhaseRequest, batch: TensorDict) -> PhaseResult: + """Run one distributed phase and convert the collected TensorDict result.""" + batch = batch.copy() + tu.assign_non_tensor(batch, phase_request=request) + output = self.worker_group.execute_phase(batch) + if isinstance(output, DataProtoFuture): + output = output.get() + metrics = dict(tu.get(output, "metrics")) + optimizer_steps = dict(tu.get(output, "optimizer_steps")) + return PhaseResult(metrics=metrics, optimizer_steps=optimizer_steps) + + def save_checkpoint(self, local_path: str, global_step: int) -> None: + """Save all worker-local role groups.""" + self.worker_group.save_checkpoint(local_path, global_step) + + def load_checkpoint(self, local_path: str) -> None: + """Restore all worker-local role groups.""" + self.worker_group.load_checkpoint(local_path) diff --git a/verl_omni/workers/engine/fsdp/diffusers_impl.py b/verl_omni/workers/engine/fsdp/diffusers_impl.py index 928141b8c..ceaf362e4 100644 --- a/verl_omni/workers/engine/fsdp/diffusers_impl.py +++ b/verl_omni/workers/engine/fsdp/diffusers_impl.py @@ -811,15 +811,21 @@ def get_per_tensor_param( peft_model = getattr(self.module, "_fsdp_wrapped_module", self.module) if hasattr(peft_model, "peft_config"): # LoRA if not merge_lora: - peft_config = peft_model.peft_config.get("default", None) - adapter_ctx = self.use_adapter(adapter_name) if adapter_name is not None else nullcontext() + resolved_adapter = adapter_name or "default" + peft_config = peft_model.peft_config.get(resolved_adapter, None) + if peft_config is None: + raise ValueError( + f"Cannot export unknown LoRA adapter {resolved_adapter!r}; " + f"available adapters: {sorted(peft_model.peft_config)}." + ) + adapter_ctx = self.use_adapter(resolved_adapter) with adapter_ctx: params = collect_lora_params( module=self.module, layered_summon=layered_summon, base_sync_done=base_sync_done, is_diffusers=True, - adapter_name=adapter_name or "default", + adapter_name=resolved_adapter, layer_prefixes=self.model_config.fsdp_layer_prefixes, ) else: # merge lora diff --git a/verl_omni/workers/engine/fsdp/distillation_impl.py b/verl_omni/workers/engine/fsdp/distillation_impl.py new file mode 100644 index 000000000..f5b689e17 --- /dev/null +++ b/verl_omni/workers/engine/fsdp/distillation_impl.py @@ -0,0 +1,439 @@ +# Copyright 2026 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Multi-role FSDP runtime for distribution-matching distillation.""" + +from __future__ import annotations + +import json +import math +import os +from collections.abc import Iterator, Mapping +from contextlib import contextmanager, nullcontext +from dataclasses import asdict +from typing import Any, Optional + +import torch +from tensordict import TensorDict +from verl.trainer.config import CheckpointConfig +from verl.workers.config import FSDPEngineConfig, FSDPOptimizerConfig +from verl.workers.config.optimizer import build_optimizer +from verl.workers.engine.base import EngineRegistry + +from verl_omni.trainer.diffusion.distillation.contracts import RoleBinding, RoleGroupSpec +from verl_omni.workers.config import DiffusionModelConfig + +from .diffusers_impl import DiffusersFSDPEngine + +__all__ = ["DistillationRoleGroupEngine"] + + +@EngineRegistry.register( + model_type="diffusion_distillation_model", + backend=["fsdp", "fsdp2"], + device=["cuda", "npu"], +) +class DistillationRoleGroupEngine(DiffusersFSDPEngine): + """One physical FSDP model with logical distillation-role bindings.""" + + def __init__( + self, + model_config: DiffusionModelConfig, + engine_config: FSDPEngineConfig, + optimizer_config: FSDPOptimizerConfig, + checkpoint_config: CheckpointConfig, + *, + role_group: RoleGroupSpec, + role_bindings: tuple[RoleBinding, ...], + optimizer_configs: Mapping[str, FSDPOptimizerConfig], + ) -> None: + self.role_group = role_group + self.role_bindings = {binding.role: binding for binding in role_bindings} + self.optimizer_configs = dict(optimizer_configs) + self.optimizers: dict[str, torch.optim.Optimizer] = {} + self.lr_schedulers: dict[str, Any] = {} + self._role_parameters: dict[str, tuple[torch.nn.Parameter, ...]] = {} + self._active_role: Optional[str] = None + self._primary_role: Optional[str] = None + self.validate_constructor_inputs(engine_config) + super().__init__(model_config, engine_config, optimizer_config, checkpoint_config) + + def validate_constructor_inputs(self, engine_config: FSDPEngineConfig) -> None: + """Validate role ownership and supported FSDP adapter layouts.""" + if not self.role_bindings: + raise ValueError(f"Role group {self.role_group.name!r} must contain at least one binding.") + if any(binding.group != self.role_group.name for binding in self.role_bindings.values()): + raise ValueError(f"Every binding passed to {self.role_group.name!r} must reference that group.") + trainable = {role for role, binding in self.role_bindings.items() if binding.trainable} + if set(self.optimizer_configs) != trainable: + raise ValueError( + f"Optimizer configs for group {self.role_group.name!r} must match trainable roles " + f"{sorted(trainable)}, got {sorted(self.optimizer_configs)}." + ) + if self.role_group.storage == "shared_base_adapters" and engine_config.strategy == "fsdp": + if not engine_config.use_orig_params: + raise ValueError("shared_base_adapters with FSDP1 requires engine.use_orig_params=true.") + + def initialize(self) -> None: + """Build the shared/independent module and all role optimizer states.""" + super().initialize() + if self.has_adapters(): + available_adapters = set(getattr(self.peft_model(), "peft_config", {})) + required_adapters = { + binding.adapter for binding in self.role_bindings.values() if binding.adapter is not None + } + missing_adapters = required_adapters - available_adapters + if missing_adapters: + raise ValueError( + f"Role group {self.role_group.name!r} is missing configured adapters {sorted(missing_adapters)}." + ) + if "default" in available_adapters: + for binding in self.role_bindings.values(): + if binding.trainable and binding.adapter not in {None, "default"}: + self.copy_adapter(source="default", target=binding.adapter) + if "student" in self.role_bindings and "student_ema" in self.role_bindings: + student = self.role_bindings["student"] + ema = self.role_bindings["student_ema"] + if student.adapter and ema.adapter: + self.copy_adapter(source=student.adapter, target=ema.adapter) + self.activate_role(self._primary_role or next(iter(self.role_bindings))) + + def _build_model_optimizer(self) -> None: + super()._build_model_optimizer() + self.optimizers = {} + self.lr_schedulers = {} + self._role_parameters = {} + + if not any(binding.trainable for binding in self.role_bindings.values()): + self.module.requires_grad_(False) + + owned_parameters: dict[int, str] = {} + for role, binding in self.role_bindings.items(): + if not binding.trainable: + continue + with self.use_role(role): + parameters = tuple(parameter for parameter in self.module.parameters() if parameter.requires_grad) + if not parameters: + raise ValueError(f"Trainable role {role!r} resolved no trainable parameters.") + overlap = {owned_parameters[id(parameter)] for parameter in parameters if id(parameter) in owned_parameters} + if overlap: + raise ValueError(f"Trainable role {role!r} shares optimizer parameters with {sorted(overlap)}.") + for parameter in parameters: + owned_parameters[id(parameter)] = role + + optimizer_config = self.optimizer_configs[role] + optimizer = build_optimizer(parameters, optimizer_config) + previous_config = self.optimizer_config + try: + self.optimizer_config = optimizer_config + lr_scheduler = self._build_lr_scheduler(optimizer) + finally: + self.optimizer_config = previous_config + self._role_parameters[role] = parameters + self.optimizers[role] = optimizer + self.lr_schedulers[role] = lr_scheduler + + if self.optimizers: + self._primary_role = next(iter(self.optimizers)) + self.optimizer = self.optimizers[self._primary_role] + self.lr_scheduler = self.lr_schedulers[self._primary_role] + self.optimizer_config = self.optimizer_configs[self._primary_role] + else: + self.optimizer = None + self.lr_scheduler = None + + def peft_model(self): + """Unwrap the FSDP1 root to reach the adapter interface.""" + return getattr(self.module, "_fsdp_wrapped_module", self.module) + + def has_adapters(self) -> bool: + """Whether the physical model exposes named adapters.""" + return hasattr(self.peft_model(), "set_adapter") + + def set_adapters_enabled(self, enabled: bool) -> None: + """Toggle adapters through the Diffusers or PEFT interface.""" + peft_model = self.peft_model() + method_name = "enable_adapters" if enabled else "disable_adapters" + method = getattr(peft_model, method_name, None) + if method is None: + fallback_name = "enable_adapter_layers" if enabled else "disable_adapter_layers" + method = getattr(getattr(peft_model, "base_model", peft_model), fallback_name, None) + if method is None: + raise AttributeError(f"PEFT model does not implement {method_name}().") + method() + + def activate_role(self, role: str) -> None: + """Select one logical role and its optimizer without entering a context.""" + try: + binding = self.role_bindings[role] + except KeyError: + raise KeyError( + f"Role {role!r} is not bound to group {self.role_group.name!r}; " + f"bound roles: {sorted(self.role_bindings)}." + ) from None + + if self.has_adapters(): + peft_model = self.peft_model() + if binding.adapter is None: + self.set_adapters_enabled(False) + else: + self.set_adapters_enabled(True) + peft_model.set_adapter(binding.adapter) + elif binding.adapter is not None and self.role_group.storage == "shared_base_adapters": + raise ValueError( + f"Shared-base role {role!r} requires adapter {binding.adapter!r}, but the model has no PEFT adapters." + ) + + self._active_role = role + if role in self.optimizers: + self.optimizer = self.optimizers[role] + self.lr_scheduler = self.lr_schedulers[role] + self.optimizer_config = self.optimizer_configs[role] + + @contextmanager + def use_role(self, role: str, *, grad_enabled: Optional[bool] = None) -> Iterator[torch.nn.Module]: + """Activate one role with explicit train/eval and autograd state.""" + previous_role = self._active_role + previous_training = self.module.training + binding = self.role_bindings[role] + effective_grad = binding.trainable if grad_enabled is None else binding.trainable and grad_enabled + self.activate_role(role) + self.module.train(effective_grad) + grad_context = nullcontext() if effective_grad else torch.no_grad() + try: + with grad_context: + yield self.module + finally: + self.module.train(previous_training) + if previous_role is not None: + self.activate_role(previous_role) + elif self._primary_role is not None: + self.activate_role(self._primary_role) + + def parameters_for_role(self, role: str) -> tuple[torch.nn.Parameter, ...]: + """Return the exact optimizer-owned parameters for a trainable role.""" + try: + return self._role_parameters[role] + except KeyError: + raise ValueError(f"Role {role!r} has no optimizer-owned parameters.") from None + + def optimizer_zero_grad(self, role: Optional[str] = None) -> None: + """Clear one role optimizer or every optimizer in the group.""" + optimizers = self.optimizers.values() if role is None else (self.optimizers[role],) + for optimizer in optimizers: + optimizer.zero_grad() + + def to(self, device: str, model: bool = True, optimizer: bool = True, grad: bool = True) -> None: + """Move one physical model and every role optimizer across the offload boundary.""" + from verl.utils.fsdp_utils import load_fsdp_optimizer, offload_fsdp_optimizer + + super().to(device=device, model=model, optimizer=False, grad=grad) + if not optimizer: + return + if device == "cpu": + for role_optimizer in self.optimizers.values(): + offload_fsdp_optimizer(role_optimizer) + else: + for role_optimizer in self.optimizers.values(): + load_fsdp_optimizer(role_optimizer, device) + + def backward_role(self, role: str, loss: torch.Tensor, *, retain_graph: bool = False) -> None: + """Backpropagate a role loss with the matching adapter active.""" + if loss.ndim != 0: + raise ValueError(f"Role loss must be scalar, got shape {tuple(loss.shape)} for {role!r}.") + with self.use_role(role): + loss.backward(retain_graph=retain_graph) + + def assert_gradient_isolation(self, active_roles: set[str]) -> None: + """Reject gradients on optimizer-owned parameters outside the active phase.""" + leaked_roles = [] + for role, parameters in self._role_parameters.items(): + if role in active_roles: + continue + if any(parameter.grad is not None for parameter in parameters): + leaked_roles.append(role) + if leaked_roles: + raise RuntimeError(f"Gradient leaked into inactive distillation roles: {sorted(leaked_roles)}.") + + def optimizer_step(self, role: Optional[str] = None) -> tuple[bool, float]: + """Clip and step one role, advancing its scheduler only on finite gradients.""" + role = role or self._active_role + if role is None or role not in self.optimizers: + raise ValueError(f"No trainable active role for optimizer step: {role!r}.") + self.activate_role(role) + grad_norm = super().optimizer_step() + stepped = math.isfinite(grad_norm) + if stepped: + self.lr_schedulers[role].step() + self.optimizers[role].zero_grad() + return stepped, grad_norm + + def lr_scheduler_step(self, role: Optional[str] = None) -> float: + """Advance one role scheduler explicitly.""" + role = role or self._active_role + if role is None or role not in self.lr_schedulers: + raise ValueError(f"No scheduler for role {role!r}.") + self.lr_schedulers[role].step() + return self.lr_schedulers[role].get_last_lr()[0] + + def update_role_ema(self, source_role: str, target_role: str, decay: float) -> None: + """EMA-update two adapter roles in the same physical group.""" + source = self.role_bindings[source_role] + target = self.role_bindings[target_role] + if not source.adapter or not target.adapter: + raise ValueError("In-group EMA requires named source and target adapters.") + self.ema_update_adapter(source=source.adapter, target=target.adapter, decay=decay) + + def update_module_ema_from(self, source: DistillationRoleGroupEngine, decay: float) -> None: + """EMA-update this independent module from an identically sharded source.""" + if not 0.0 <= decay <= 1.0: + raise ValueError(f"EMA decay must be in [0, 1], got {decay}.") + source_adapter = ( + next((binding.adapter for binding in source.role_bindings.values() if binding.adapter is not None), None) + if source.has_adapters() + else None + ) + target_adapter = ( + next((binding.adapter for binding in self.role_bindings.values() if binding.adapter is not None), None) + if self.has_adapters() + else None + ) + if (source_adapter is None) != (target_adapter is None): + raise ValueError("Independent EMA source and target must both use LoRA adapters or both use full modules.") + + if source_adapter is not None: + with source._adapter_state_context(), self._adapter_state_context(), torch.no_grad(): + source_parameters = source._active_adapter_trainable_params(source_adapter) + target_parameters = self._active_adapter_trainable_params(target_adapter) + self.ema_parameter_lists(source_parameters, target_parameters, decay) + return + + source_parameters = tuple(source.module.named_parameters()) + target_parameters = tuple(self.module.named_parameters()) + if len(source_parameters) != len(target_parameters): + raise ValueError("Independent EMA source and target parameter counts do not match.") + with torch.no_grad(): + for (source_name, source_parameter), (target_name, target_parameter) in zip( + source_parameters, target_parameters, strict=True + ): + if source_name != target_name: + raise ValueError( + f"Independent EMA parameter names do not match: {source_name!r} and {target_name!r}." + ) + if source_parameter.shape != target_parameter.shape: + raise ValueError("Independent EMA source and target parameter shapes do not match.") + target_parameter.lerp_(source_parameter, 1.0 - decay) + + @staticmethod + def ema_parameter_lists(source_parameters, target_parameters, decay: float) -> None: + """Blend corresponding independent-module adapter parameters in place.""" + if len(source_parameters) != len(target_parameters) or not source_parameters: + raise ValueError("Independent EMA source and target adapter parameter counts must match and be non-empty.") + for source_parameter, target_parameter in zip(source_parameters, target_parameters, strict=True): + if source_parameter.shape != target_parameter.shape: + raise ValueError("Independent EMA source and target parameter shapes do not match.") + target_parameter.lerp_(source_parameter, 1.0 - decay) + + def iter_export_tensors(self, role: str, base_sync_done: bool): + """Return the selected student role's parameter iterator and matching PEFT config.""" + if role not in {"student", "student_ema"}: + raise ValueError(f"Only student or student_ema can be exported, got {role!r}.") + binding = self.role_bindings[role] + return self.get_per_tensor_param( + base_sync_done=base_sync_done, + adapter_name=binding.adapter, + ) + + def additional_state_path(self, local_path: str) -> str: + """Locate this rank's secondary optimizer and scheduler state.""" + return os.path.join(local_path, f"role_state_rank_{self.rank}.pt") + + def save_role_group_checkpoint(self, local_path: str, global_step: int) -> None: + """Save one physical model once plus every role optimizer/scheduler.""" + immutable_teacher_only = set(self.role_bindings) == {"teacher_score"} + if not immutable_teacher_only: + if self._primary_role is not None: + self.activate_role(self._primary_role) + super().save_checkpoint(local_path=local_path, global_step=global_step) + + os.makedirs(local_path, exist_ok=True) + additional_roles = [role for role in self.optimizers if role != self._primary_role] + torch.save( + { + "optimizers": {role: self.optimizers[role].state_dict() for role in additional_roles}, + "schedulers": {role: self.lr_schedulers[role].state_dict() for role in additional_roles}, + "primary_role": self._primary_role, + }, + self.additional_state_path(local_path), + ) + if self.rank == 0: + with open(os.path.join(local_path, "role_group.json"), "w", encoding="utf-8") as file: + json.dump( + { + "group": asdict(self.role_group), + "bindings": [asdict(binding) for binding in self.role_bindings.values()], + "immutable_teacher_only": immutable_teacher_only, + }, + file, + indent=2, + sort_keys=True, + ) + torch.distributed.barrier() + + def load_role_group_checkpoint(self, local_path: str) -> None: + """Restore the model and every role optimizer/scheduler.""" + manifest_path = os.path.join(local_path, "role_group.json") + if not os.path.isfile(manifest_path): + raise FileNotFoundError(f"Missing role-group manifest: {manifest_path}") + with open(manifest_path, encoding="utf-8") as file: + manifest = json.load(file) + expected_bindings = [asdict(binding) for binding in self.role_bindings.values()] + if manifest.get("group") != asdict(self.role_group) or manifest.get("bindings") != expected_bindings: + raise ValueError(f"Checkpoint role layout does not match runtime group {self.role_group.name!r}.") + + if not manifest.get("immutable_teacher_only", False): + if self._primary_role is not None: + self.activate_role(self._primary_role) + super().load_checkpoint(local_path=local_path, del_local_after_load=False) + + role_state = torch.load(self.additional_state_path(local_path), map_location="cpu", weights_only=False) + if role_state.get("primary_role") != self._primary_role: + raise ValueError( + f"Checkpoint primary role {role_state.get('primary_role')!r} does not match {self._primary_role!r}." + ) + expected_additional_roles = set(self.optimizers) - ({self._primary_role} if self._primary_role else set()) + if set(role_state.get("optimizers", {})) != expected_additional_roles: + raise ValueError("Checkpoint secondary optimizer roles do not match the active role group.") + if set(role_state.get("schedulers", {})) != expected_additional_roles: + raise ValueError("Checkpoint secondary scheduler roles do not match the active role group.") + for role, state in role_state["optimizers"].items(): + self.optimizers[role].load_state_dict(state) + for role, state in role_state["schedulers"].items(): + self.lr_schedulers[role].load_state_dict(state) + torch.distributed.barrier() + + def forward_backward_batch(self, data: TensorDict, loss_function, forward_only: bool = False): + """Reject PPO-shaped execution; distillation phases use the phase runner.""" + raise NotImplementedError("DistillationRoleGroupEngine is driven through DiffusionDistillationWorker phases.") + + def prepare_model_inputs(self, micro_batch: TensorDict, step: int): + """Keep model-specific preparation in the architecture phase runner.""" + raise NotImplementedError("Architecture-owned distillation phase runners prepare model inputs.") + + def prepare_model_outputs(self, output, micro_batch: TensorDict): + """Keep model-specific output conversion in the phase runner.""" + raise NotImplementedError("Architecture-owned distillation phase runners prepare model outputs.") + + def forward_step(self, micro_batch: TensorDict, loss_function, forward_only, step): + """Reject the PPO step interface for multi-role computation.""" + raise NotImplementedError("Architecture-owned distillation phase runners execute forwards.") diff --git a/verl_omni/workers/engine/lora_adapter_mixin.py b/verl_omni/workers/engine/lora_adapter_mixin.py index 2590d9f42..f688b5396 100644 --- a/verl_omni/workers/engine/lora_adapter_mixin.py +++ b/verl_omni/workers/engine/lora_adapter_mixin.py @@ -78,6 +78,25 @@ def _build_lora_module(self, module): return module + def active_adapter_selection(self): + """Read the current named adapter selection for context restoration.""" + module = getattr(self.module, "_fsdp_wrapped_module", self.module) + active = getattr(module, "active_adapters", None) + if callable(active): + active = active() + if active is None: + active = getattr(module, "active_adapter", None) + if isinstance(active, list | tuple): + return active[0] if len(active) == 1 else list(active) + return active + + def restore_adapter_selection(self, selection) -> None: + """Restore the preceding adapter rather than always selecting default.""" + if selection: + self._set_adapter(selection) + else: + self._set_adapter("default") + @contextmanager def _adapter_state_context(self): """Open writable adapter parameter access (FSDP summon when applicable).""" @@ -89,6 +108,7 @@ def _adapter_state_context(self): is_fsdp_module = fsdp_version(self.module) in (1, 2) is_offload_param = getattr(self, "_is_offload_param", False) origin_module_device = next(self.module.parameters()).device.type + previous_adapter = self.active_adapter_selection() if is_fsdp_module and (is_offload_param or origin_module_device == "cpu"): load_fsdp_model_to_gpu(self.module) @@ -98,13 +118,13 @@ def _adapter_state_context(self): try: yield finally: - self._set_adapter("default") + self.restore_adapter_selection(previous_adapter) finally: if is_offload_param: offload_fsdp_model_to_cpu(self.module) aggressive_empty_cache(force_sync=True) - def _set_adapter(self, name: str): + def _set_adapter(self, name): module = getattr(self.module, "_fsdp_wrapped_module", self.module) if not hasattr(module, "set_adapter"): raise AttributeError(f"Module does not support set_adapter({name!r})") @@ -117,6 +137,7 @@ def use_adapter(self, name: str): ``"reference"`` is a logical policy state (see ``policy_state_adapters``) that runs with all LoRA adapters disabled, not a registered PEFT adapter. """ + previous_adapter = self.active_adapter_selection() if name == "reference": with self.disable_adapter(): yield @@ -125,7 +146,7 @@ def use_adapter(self, name: str): try: yield finally: - self._set_adapter("default") + self.restore_adapter_selection(previous_adapter) def _active_adapter_trainable_params(self, adapter_name: str) -> list[torch.nn.Parameter]: peft_model = getattr(self.module, "_fsdp_wrapped_module", self.module)