diff --git a/areal/api/cli_args.py b/areal/api/cli_args.py index ae17f49d5b..d16ab2ad85 100644 --- a/areal/api/cli_args.py +++ b/areal/api/cli_args.py @@ -2,7 +2,9 @@ import argparse import json +import math import os +import re import warnings from dataclasses import MISSING as dataclass_missing from dataclasses import asdict, dataclass, field, fields @@ -1449,6 +1451,7 @@ class TrainEngineConfig: "help": "Timeout (seconds) for initialize() to wait for guards to be ready." }, ) + scheduling_strategy: SchedulingStrategy = field( default_factory=SchedulingStrategy, metadata={ @@ -1480,6 +1483,11 @@ def __post_init__(self): raise ValueError( f"_version must be either 'v1' or 'v2', got '{self._version}'" ) + if self.weight_update_mode == "awex" and not self.megatron.wrap_with_ddp: + raise ValueError( + "weight_update_mode='awex' requires megatron.wrap_with_ddp=true " + "because AWEX offloads MCore DDP flat buffers" + ) # Canonicalize common aliases so getattr(torch, ...) works at runtime. # Storage map omits fp16 since float16 is not a valid optimizer_dtype; @@ -3078,6 +3086,48 @@ class SchedulerConfig: reward_model_service_url: str = field(default="http://localhost:30000/classify") +@dataclass +class DatasetSourceConfig: + """One source in a dataset mixture.""" + + path: str = field( + default=MISSING, + metadata={"help": "Local path or HuggingFace name for this dataset source."}, + ) + type: str = field( + default=MISSING, + metadata={"help": "Training data type, for example 'rl'."}, + ) + teacher_group: str | None = field( + default=None, + metadata={"help": "Optional MOPD teacher group applied to this entire source."}, + ) + split: str | None = field( + default=None, + metadata={"help": "Optional split override for this dataset source."}, + ) + max_length: int | None = field( + default=None, + metadata={"help": "Optional maximum sequence length for this source."}, + ) + dataset_kwargs: dict[str, Any] = field( + default_factory=dict, + metadata={"help": "Extra keyword arguments for this source's loader."}, + ) + + def __post_init__(self) -> None: + for name in ("path", "type"): + value = getattr(self, name) + if not isinstance(value, str) or not value.strip() or value == MISSING: + raise ValueError(f"dataset source {name} must be a non-empty string") + if self.teacher_group is not None and ( + not isinstance(self.teacher_group, str) or not self.teacher_group.strip() + ): + raise ValueError( + "dataset source teacher_group must be a non-empty string or null" + ) + + @dataclass class _DatasetConfig: """Configuration for dataset loading and preprocessing.""" @@ -3086,15 +3136,33 @@ class _DatasetConfig: default="train", metadata={"help": "Dataset split to use, e.g., 'train', 'test'."}, ) - path: str = field( - default=MISSING, + path: str | None = field( + default=None, + metadata={"help": "Path to one dataset. Mutually exclusive with sources."}, + ) + type: str | None = field( + default=None, metadata={ - "help": "Path to the dataset. Can be a local path or a HuggingFace dataset name." + "help": "Training data type for path. Mutually exclusive with sources." }, ) - type: str = field( - default=MISSING, - metadata={"help": "Type of training method, e.g., 'sft', 'rl', etc."}, + sources: list[DatasetSourceConfig] = field( + default_factory=list, + metadata={ + "help": "Dataset mixture sources. MOPD requires every source to declare " + "a teacher_group." + }, + ) + mixture_sampling_policy: str = field( + default="proportional", + metadata={ + "help": ( + "How a routed mixture represents sources in one epoch: " + "'proportional' preserves source-size proportions; 'uniform' " + "balances source counts by deterministically cycling shorter sources." + ), + "choices": ["proportional", "uniform"], + }, ) batch_size: int = field( default=1, metadata={"help": "Batch size for the dataloader"} @@ -3151,6 +3219,15 @@ class _DatasetConfig: }, ) + def __post_init__(self) -> None: + if self.mixture_sampling_policy not in ("proportional", "uniform"): + raise ValueError( + "mixture_sampling_policy must be 'proportional' or 'uniform', " + f"got {self.mixture_sampling_policy!r}" + ) + if self.sources and (self.path is not None or self.type is not None): + raise ValueError("dataset path/type cannot be combined with sources") + @dataclass class TrainDatasetConfig(_DatasetConfig): @@ -3395,6 +3472,184 @@ def __post_init__(self): ) +@dataclass +class MOPDTeacherSpec: + """Checkpoint specification for one MOPD teacher.""" + + path: str = field( + default=MISSING, + metadata={"help": "Local or shared-filesystem teacher checkpoint path."}, + ) + + def __post_init__(self): + if not isinstance(self.path, str) or not self.path.strip(): + raise ValueError("MOPD teacher path must be a non-empty string") + + +@dataclass +class MOPDTeacherManagerConfig: + """Checkpoint provider configuration for phase-scoped MOPD teachers.""" + + type: str = field( + default="disk", + metadata={ + "help": "Teacher checkpoint provider.", + "choices": ["disk", "local_memory"], + }, + ) + staging_root: str = field( + default="/dev/shm/areal-mopd", + metadata={"help": "Node-local staging root for local_memory providers."}, + ) + min_free_bytes: int | None = field( + default=None, + metadata={ + "help": "Optional minimum free space required after staging a checkpoint." + }, + ) + + def __post_init__(self): + if self.type not in ("disk", "local_memory"): + raise ValueError( + "mopd.manager.type must be either 'disk' or 'local_memory', " + f"got {self.type!r}" + ) + if not isinstance(self.staging_root, str) or not self.staging_root.strip(): + raise ValueError("mopd.manager.staging_root must be a non-empty string") + if self.min_free_bytes is not None and self.min_free_bytes < 0: + raise ValueError( + "mopd.manager.min_free_bytes must be non-negative or None, " + f"got {self.min_free_bytes}" + ) + + +@dataclass +class MOPDLossConfig: + """Coefficients for joint RL and multi-teacher distillation training.""" + + rl_coefficient: float = field( + default=0.0, + metadata={"help": "Coefficient applied to the RL objective."}, + ) + distillation_coefficient: float = field( + default=1.0, + metadata={"help": "Coefficient applied to the MOPD objective."}, + ) + importance_ratio_cap: float = field( + default=5.0, + metadata={"help": "Positive cap applied to the behavior-policy ratio."}, + ) + + def __post_init__(self): + for name in ("rl_coefficient", "distillation_coefficient"): + value = getattr(self, name) + if not isinstance(value, (int, float)) or isinstance(value, bool): + raise ValueError(f"mopd.loss.{name} must be a finite number") + if not math.isfinite(value) or value < 0: + raise ValueError( + f"mopd.loss.{name} must be finite and non-negative, got {value}" + ) + if self.rl_coefficient == 0 and self.distillation_coefficient == 0: + raise ValueError("MOPD loss coefficients cannot both be zero") + if ( + not isinstance(self.importance_ratio_cap, (int, float)) + or isinstance(self.importance_ratio_cap, bool) + or not math.isfinite(self.importance_ratio_cap) + or self.importance_ratio_cap <= 0 + ): + raise ValueError( + "mopd.loss.importance_ratio_cap must be finite and positive" + ) + + +@dataclass +class MOPDTeacherEngineConfig(TrainEngineConfig): + """Forward-only scoring engine configuration used by MOPD teachers.""" + + disable_dropout: bool = field( + default=True, + metadata={"help": "Disable dropout for deterministic teacher scoring."}, + ) + optimizer: OptimizerConfig | None = field( + default=None, + metadata={"help": "MOPD scoring teachers do not construct an optimizer."}, + ) + + def __post_init__(self) -> None: + super().__post_init__() + if self.optimizer is not None: + raise ValueError("MOPDTeacherEngineConfig.optimizer must be null") + if not self.disable_dropout: + raise ValueError("MOPDTeacherEngineConfig.disable_dropout must be true") + + +@dataclass +class MOPDConfig: + """Configuration for multi-teacher on-policy distillation.""" + + teachers: dict[str, MOPDTeacherSpec] = field(default_factory=dict) + teacher_groups: dict[str, dict[str, float]] = field(default_factory=dict) + teacher_engine: MOPDTeacherEngineConfig = field( + default_factory=MOPDTeacherEngineConfig + ) + manager: MOPDTeacherManagerConfig = field(default_factory=MOPDTeacherManagerConfig) + loss: MOPDLossConfig = field(default_factory=MOPDLossConfig) + + def __post_init__(self): + if not self.teachers: + raise ValueError("mopd.teachers must not be empty") + if not self.teacher_groups: + raise ValueError("mopd.teacher_groups must not be empty") + + for teacher_id, teacher in self.teachers.items(): + if not isinstance(teacher_id, str) or not teacher_id.strip(): + raise ValueError("mopd teacher ids must be non-empty strings") + if re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9_.-]*", teacher_id) is None: + raise ValueError( + "mopd teacher ids must be filename-safe and match " + "[A-Za-z0-9][A-Za-z0-9_.-]*" + ) + if not isinstance(teacher, MOPDTeacherSpec): + raise ValueError( + f"mopd.teachers[{teacher_id!r}] must be an MOPDTeacherSpec" + ) + + for teacher_group, weights in self.teacher_groups.items(): + if not isinstance(teacher_group, str) or not teacher_group.strip(): + raise ValueError("mopd teacher group ids must be non-empty strings") + if not weights: + raise ValueError( + f"mopd.teacher_groups[{teacher_group!r}] must not be empty" + ) + + has_positive_weight = False + for teacher_id, weight in weights.items(): + if teacher_id not in self.teachers: + raise ValueError( + f"mopd.teacher_groups[{teacher_group!r}] references " + f"unknown teacher {teacher_id!r}" + ) + if not isinstance(weight, (int, float)) or isinstance(weight, bool): + raise ValueError( + f"mopd.teacher_groups[{teacher_group!r}]" + f"[{teacher_id!r}] must be a " + "finite non-negative number" + ) + if not math.isfinite(weight) or weight < 0: + raise ValueError( + f"mopd.teacher_groups[{teacher_group!r}]" + f"[{teacher_id!r}] must be finite " + f"and non-negative, got {weight}" + ) + has_positive_weight = has_positive_weight or weight > 0 + + if not has_positive_weight: + raise ValueError( + f"mopd.teacher_groups[{teacher_group!r}] must contain at least " + "one positive weight" + ) + + @dataclass class PPOConfig(BaseExperimentConfig): """Configuration for Proximal Policy Optimization (PPO) reinforcement learning experiments.""" @@ -3423,6 +3678,10 @@ class PPOConfig(BaseExperimentConfig): ) }, ) + mopd: MOPDConfig | None = field( + default=None, + metadata={"help": "Optional multi-teacher on-policy distillation config."}, + ) dynamic_bs: bool = field( default=False, metadata={ @@ -3434,6 +3693,10 @@ class PPOConfig(BaseExperimentConfig): def __post_init__(self): """Validate the eval generation config.""" + if self.teacher is not None and self.mopd is not None: + raise ValueError("teacher and mopd cannot be configured at the same time") + if self.mopd is not None: + self._validate_mopd_config() if self.eval_gconfig is None: self.eval_gconfig = self.gconfig.new() if self.rollout.deterministic_sampling: @@ -3465,6 +3728,112 @@ def __post_init__(self): self.rollout.lora_name = self.gconfig.lora_name super().__post_init__() + def _validate_mopd_config(self): + """Validate MOPD engine topology before any workers are created.""" + from areal.api.alloc_mode import ModelAllocation, ParallelStrategy + + assert self.mopd is not None + self._validate_mopd_dataset_sources("train_dataset", self.train_dataset) + if self.valid_dataset is not None: + self._validate_mopd_dataset_sources("valid_dataset", self.valid_dataset) + if self.mopd.loss.distillation_coefficient == 0: + # A pure-RL MOPD plan only scales the actor objective. It must not + # require teacher workers, checkpoint compatibility, or colocated + # actor/rollout infrastructure to initialize successfully. + return + teacher_engine = self.mopd.teacher_engine + + if not self.actor.backend.startswith("megatron:"): + raise ValueError("mopd requires a Megatron actor backend") + if not teacher_engine.backend.startswith("megatron:"): + raise ValueError("mopd.teacher_engine backend must be Megatron") + if teacher_engine._version != "v1": + raise ValueError("mopd.teacher_engine currently requires _version='v1'") + if not self.rollout.backend.startswith("sglang:"): + raise ValueError("mopd requires an SGLang rollout backend") + if self.actor.weight_update_mode != "awex": + raise ValueError("mopd requires actor.weight_update_mode='awex'") + if teacher_engine.optimizer is not None: + raise ValueError("mopd.teacher_engine.optimizer must be null") + if not teacher_engine.disable_dropout: + raise ValueError("mopd.teacher_engine.disable_dropout must be true") + + teacher_schedule = teacher_engine.scheduling_strategy + if ( + teacher_schedule.type != SchedulingStrategyType.colocation.value + or teacher_schedule.target != "actor" + or not teacher_schedule.fork + ): + raise ValueError( + "the current MOPD v1 runtime supports teacher colocation " + "target='actor' with fork=true" + ) + + rollout_schedule = self.rollout.scheduling_strategy + if ( + rollout_schedule.type != SchedulingStrategyType.colocation.value + or rollout_schedule.target != "actor" + or not rollout_schedule.fork + ): + raise ValueError( + "the current MOPD v1 runtime supports rollout colocation " + "target='actor' with fork=true" + ) + actor_worker_ports = self.actor.scheduling_spec[0].port_count + if actor_worker_ports < 2: + raise ValueError( + "the current MOPD v1 runtime requires actor.scheduling_spec[0]." + f"port_count >= 2, got {actor_worker_ports}" + ) + + actor_alloc = ModelAllocation.from_str(self.actor.backend, name="actor") + teacher_alloc = ModelAllocation.from_str( + teacher_engine.backend, name="mopd_teacher" + ) + if not ParallelStrategy.parallelism_eq( + actor_alloc.parallel, teacher_alloc.parallel + ): + raise ValueError( + "mopd teacher and actor must use the same parallel strategy" + ) + if self.mopd.manager.type == "local_memory": + if self.scheduler.type != "local": + raise ValueError( + "mopd local_memory provider requires scheduler.type='local' " + "so controller and teacher workers share the same host" + ) + if actor_alloc.parallel.world_size > self.cluster.n_gpus_per_node: + raise ValueError( + "mopd local_memory provider only supports a single node; use " + "disk for multi-node runs" + ) + + def _validate_mopd_dataset_sources( + self, + name: str, + dataset_config: TrainDatasetConfig | ValidDatasetConfig, + ) -> None: + """Require one configured teacher group for every MOPD dataset source.""" + assert self.mopd is not None + if not dataset_config.sources: + raise ValueError(f"{name}.sources must not be empty when mopd is enabled") + if dataset_config.path is not None or dataset_config.type is not None: + raise ValueError( + f"{name}.path/type cannot be used with {name}.sources in MOPD" + ) + for index, source in enumerate(dataset_config.sources): + teacher_group = source.teacher_group + if not isinstance(teacher_group, str) or not teacher_group.strip(): + raise ValueError( + f"{name}.sources[{index}].teacher_group must be configured " + "when mopd is enabled" + ) + if teacher_group not in self.mopd.teacher_groups: + raise ValueError( + f"{name}.sources[{index}].teacher_group references unknown " + f"MOPD teacher group {teacher_group!r}" + ) + @dataclass class GRPOConfig(PPOConfig): diff --git a/areal/api/scheduler_api.py b/areal/api/scheduler_api.py index d83bddb235..d3c4e70463 100644 --- a/areal/api/scheduler_api.py +++ b/areal/api/scheduler_api.py @@ -147,6 +147,7 @@ def fork_workers( role: str, target_role: str, command: str | None = None, + env_vars: list[dict[str, str]] | None = None, ) -> list[str]: """Fork new worker processes from existing workers. @@ -163,6 +164,8 @@ def fork_workers( Custom module path to run instead of the default rpc_server. If specified, the forked process runs this module (e.g., "areal.experimental.openai.proxy.proxy_rollout_server"). + env_vars : list[dict[str, str]], optional + Per-worker environment overrides for the forked role. Returns ------- diff --git a/areal/dataset/__init__.py b/areal/dataset/__init__.py index 39b3112a28..410ee0449a 100644 --- a/areal/dataset/__init__.py +++ b/areal/dataset/__init__.py @@ -6,6 +6,14 @@ from typing import TYPE_CHECKING, Optional from areal.api.cli_args import _DatasetConfig +from areal.dataset.mopd import ( + ROUTE_METADATA_KEY, + DatasetRoute, + MOPDDataset, + RoutedDataset, + get_mopd_dataset, + get_routed_dataset, +) from areal.utils import logging if TYPE_CHECKING: @@ -219,6 +227,20 @@ def get_custom_dataset( ) -> "Dataset | RDataset": from areal.utils.environ import is_single_controller + if dataset_config is not None and dataset_config.sources: + return get_routed_dataset( + dataset_config, + tokenizer=tokenizer, + processor=processor, + ) + if dataset_config is not None and ( + not isinstance(dataset_config.path, str) + or not dataset_config.path.strip() + or not isinstance(dataset_config.type, str) + or not dataset_config.type.strip() + ): + raise ValueError("dataset_config.path and dataset_config.type are required") + if ( is_single_controller() and dataset_config is not None @@ -257,6 +279,12 @@ def get_custom_dataset( __all__ = [ + "ROUTE_METADATA_KEY", + "DatasetRoute", + "MOPDDataset", + "RoutedDataset", "VALID_DATASETS", "get_custom_dataset", + "get_mopd_dataset", + "get_routed_dataset", ] diff --git a/areal/dataset/mopd.py b/areal/dataset/mopd.py new file mode 100644 index 0000000000..3e3d359818 --- /dev/null +++ b/areal/dataset/mopd.py @@ -0,0 +1,244 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Generic routed dataset mixtures used by MOPD and other consumers.""" + +from __future__ import annotations + +from bisect import bisect_right +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from typing import Any + +from areal.api.cli_args import _DatasetConfig + +ROUTE_METADATA_KEY = "__areal_route" +MOPD_ROUTE_METADATA_KEY = ROUTE_METADATA_KEY + + +@dataclass(frozen=True) +class DatasetRoute: + """Typed route provenance stripped before a sample reaches its workflow.""" + + source_index: int + route: str + + def __post_init__(self) -> None: + if self.source_index < 0: + raise ValueError("source_index must be non-negative") + if not isinstance(self.route, str) or not self.route.strip(): + raise ValueError("route must be a non-empty string") + + +class RoutedDataset: + """Combine dataset sources with an explicit deterministic sampling policy.""" + + def __init__( + self, + sources: list[tuple[Any, str | None]], + *, + sampling_policy: str = "proportional", + ) -> None: + if not sources: + raise ValueError("Routed dataset sources must not be empty") + if sampling_policy not in ("proportional", "uniform"): + raise ValueError( + "sampling_policy must be 'proportional' or 'uniform', " + f"got {sampling_policy!r}" + ) + self._datasets = [dataset for dataset, _ in sources] + self._teacher_groups = [teacher_group for _, teacher_group in sources] + self._sampling_policy = sampling_policy + self._offsets: list[int] = [] + self._source_lengths: list[int] = [] + + from areal.infra.data_service.rdataset import RDataset + + remote = [isinstance(dataset, RDataset) for dataset in self._datasets] + if any(remote) and not all(remote): + raise ValueError("MOPD dataset cannot mix local and remote sources") + self._is_remote = all(remote) + if not self._is_remote: + self._refresh_offsets() + + @property + def is_remote(self) -> bool: + return self._is_remote + + def _refresh_offsets(self) -> None: + total = 0 + offsets: list[int] = [] + source_lengths: list[int] = [] + for dataset in self._datasets: + source_length = len(dataset) + source_lengths.append(source_length) + total += source_length + offsets.append(total) + self._offsets = offsets + self._source_lengths = source_lengths + if self._sampling_policy == "uniform" and any( + length == 0 for length in source_lengths + ): + raise ValueError("uniform routed mixtures do not support empty sources") + + def __len__(self) -> int: + if not self._offsets: + self._refresh_offsets() + if self._sampling_policy == "uniform": + return max(self._source_lengths) * len(self._datasets) + return self._offsets[-1] + + def _locate(self, index: int) -> tuple[int, int]: + if self._sampling_policy == "uniform": + source_index = index % len(self._datasets) + local_index = (index // len(self._datasets)) % self._source_lengths[ + source_index + ] + return source_index, local_index + source_index = bisect_right(self._offsets, index) + source_start = 0 if source_index == 0 else self._offsets[source_index - 1] + return source_index, index - source_start + + def __getitem__(self, index: int) -> dict[str, Any]: + size = len(self) + if index < 0: + index += size + if index < 0 or index >= size: + raise IndexError(index) + + source_index, local_index = self._locate(index) + sample = self._datasets[source_index][local_index] + if not isinstance(sample, Mapping): + raise TypeError( + f"Dataset mixture samples must be mappings, got {type(sample).__name__}" + ) + teacher_group = self._teacher_groups[source_index] + if teacher_group is None: + return dict(sample) + if ROUTE_METADATA_KEY in sample or "mopd_route" in sample: + raise ValueError( + "teacher group must be configured only on the dataset source" + ) + + routed_sample = dict(sample) + routed_sample[ROUTE_METADATA_KEY] = DatasetRoute( + source_index=source_index, + route=teacher_group, + ) + return routed_sample + + def connect( + self, + controller: Any, + dataset_id: str, + tokenizer_or_processor_path: str = "", + shuffle: bool = True, + drop_last: bool = True, + ) -> None: + """Connect every remote source to one shared data controller.""" + if not self._is_remote: + raise RuntimeError("Only remote MOPD datasets require connect()") + for index, dataset in enumerate(self._datasets): + dataset.connect( + controller, + dataset_id=f"{dataset_id}_source_{index}", + tokenizer_or_processor_path=tokenizer_or_processor_path, + shuffle=shuffle, + drop_last=drop_last, + ) + self._refresh_offsets() + + def _start_prefetch(self, indices: list[int]) -> None: + """Translate global sampler indices to each remote source.""" + if not self._is_remote: + return + source_indices: list[list[int]] = [[] for _ in self._datasets] + for index in indices: + source_index, local_index = self._locate(index) + source_indices[source_index].append(local_index) + for dataset, local_indices in zip(self._datasets, source_indices, strict=True): + dataset._start_prefetch(local_indices) + + def close(self) -> None: + """Close all remote source proxies.""" + if not self._is_remote: + return + for dataset in self._datasets: + dataset.close() + + +def is_remote_dataset(dataset: Any) -> bool: + """Return whether a dataset needs data-service connection and prefetching.""" + from areal.infra.data_service.rdataset import RDataset + + return isinstance(dataset, RDataset) or ( + isinstance(dataset, RoutedDataset) and dataset.is_remote + ) + + +def get_routed_dataset( + dataset_config: _DatasetConfig, + tokenizer: Any = None, + processor: Any = None, + source_loader: Callable[..., Any] | None = None, +) -> RoutedDataset: + """Load configured sources and attach optional teacher-group metadata.""" + if not dataset_config.sources: + raise ValueError("Dataset mixture config must contain at least one source") + + if source_loader is None: + from areal.dataset import get_custom_dataset + + def source_loader(**kwargs): + source_config = kwargs["source_config"] + split = kwargs["split"] + return get_custom_dataset( + split=split, + dataset_config=source_config, + tokenizer=kwargs["tokenizer"], + processor=kwargs["processor"], + ) + + routed_sources: list[tuple[Any, str | None]] = [] + for source in dataset_config.sources: + split = source.split or dataset_config.split + source_config = _DatasetConfig( + path=source.path, + type=source.type, + split=split, + max_length=( + source.max_length + if source.max_length is not None + else dataset_config.max_length + ), + dataset_kwargs=dataset_config.dataset_kwargs | source.dataset_kwargs, + scheduling_spec=dataset_config.scheduling_spec, + ) + dataset = source_loader( + source=source, + source_config=source_config, + split=split, + tokenizer=tokenizer, + processor=processor, + ) + routed_sources.append((dataset, source.teacher_group)) + return RoutedDataset( + routed_sources, + sampling_policy=dataset_config.mixture_sampling_policy, + ) + + +# Compatibility names for the first MOPD consumer of the routed mixture API. +MOPDDataset = RoutedDataset +get_mopd_dataset = get_routed_dataset + + +__all__ = [ + "MOPDDataset", + "MOPD_ROUTE_METADATA_KEY", + "ROUTE_METADATA_KEY", + "DatasetRoute", + "RoutedDataset", + "get_mopd_dataset", + "get_routed_dataset", + "is_remote_dataset", +] diff --git a/areal/engine/__init__.py b/areal/engine/__init__.py index afcd6512ea..62abc74d93 100644 --- a/areal/engine/__init__.py +++ b/areal/engine/__init__.py @@ -8,6 +8,7 @@ "FSDPRWEngine", "FSDPDPOEngine", "MegatronEngine", + "MegatronScoringEngine", "MegatronPPOActor", "MegatronPPOCritic", "MegatronLMEngine", @@ -25,6 +26,7 @@ "FSDPRWEngine": "areal.engine.fsdp_engine", "FSDPDPOEngine": "areal.engine.fsdp_engine", "MegatronEngine": "areal.engine.megatron_engine", + "MegatronScoringEngine": "areal.engine.megatron_engine", "MegatronPPOActor": "areal.engine.megatron_engine", "MegatronPPOCritic": "areal.engine.megatron_engine", "MegatronLMEngine": "areal.engine.megatron_engine", diff --git a/areal/engine/awex/colocate_reader.py b/areal/engine/awex/colocate_reader.py index 5bb8b42984..4392d53445 100644 --- a/areal/engine/awex/colocate_reader.py +++ b/areal/engine/awex/colocate_reader.py @@ -32,46 +32,12 @@ import torch - -def _patch_tms_hook_mode() -> None: - """Make ``torch_memory_saver.hook_mode`` setter a no-op once initialized. - - ``megatron.core.inference.contexts.dynamic_context`` (pulled in transitively - by ``awex.converter.mcore_converter`` -> ``megatron.core``) runs a - module-level ``torch_memory_saver.hook_mode = "torch"``. In the SGLang - scheduler process the memory_saver singleton is already initialized (sglang - ran ``_ensure_initialized``, which ``del``s ``_impl_ctor_kwargs``), so that - late assignment raises ``AttributeError``. awex's model registry swallows the - import error, the BailingMoe converter never registers, and weight transfer - later dies with ``Unsupported attention parameter name: attention.g_proj``. - The singleton's own assert already declares post-init configuration - unsupported, so dropping the late set is the intended behavior. - """ - try: - import torch_memory_saver as _tms - except Exception: - return - inst = getattr(_tms, "torch_memory_saver", None) - if inst is None: - return - cls = type(inst) - prop = cls.hook_mode - if getattr(prop.fset, "_awex_safe", False): - return - - def _safe_setter(self, value): - if not hasattr(self, "_impl_ctor_kwargs"): - return # singleton already initialized; late set is a design no-op - prop.fset(self, value) - - _safe_setter._awex_safe = True - cls.hook_mode = property(prop.fget, _safe_setter) - +from areal.engine.awex.memory_saver import patch_tms_hook_mode # Must run before any awex import: awex.models.registry auto-imports model # modules at module load, and the BailingMoe module's transitive megatron import # trips the hook_mode race above. -_patch_tms_hook_mode() +patch_tms_hook_mode() from awex.meta.infer_meta_resolver import InferParamMetaResolver # noqa: E402 from awex.meta.meta_resolver import ParamMetaResolver # noqa: E402 @@ -84,6 +50,47 @@ def _safe_setter(self, value): logger = getLogger("AwexColocateReader") +class _PhysicalDeviceMetaServerClient: + """Use physical GPU ids in AWEX colocate metadata and handshake keys.""" + + _DEVICE_KEY_PREFIXES = ( + "training_serialized_weights_", + "weights_update_finished_", + "write_finished_", + ) + + def __init__(self, client: Any, physical_gpu_id: int): + self._client = client + self._physical_gpu_id = physical_gpu_id + + def __getattr__(self, name: str) -> Any: + return getattr(self._client, name) + + def _rewrite_device_key(self, key: str) -> str: + if not key.startswith(self._DEVICE_KEY_PREFIXES): + return key + prefix_and_ip, step = key.rsplit("_", 1) + prefix_and_ip, _logical_gpu_id = prefix_and_ip.rsplit("_", 1) + return f"{prefix_and_ip}_{self._physical_gpu_id}_{step}" + + def add_object_to_set(self, key: str, value: Any) -> Any: + if key == "inference_device_rank_entries": + ip_address, _logical_gpu_id, transfer_rank = value + value = (ip_address, self._physical_gpu_id, transfer_rank) + return self._client.add_object_to_set(key, value) + + def get_object(self, key: str, *args: Any, **kwargs: Any) -> Any: + return self._client.get_object(self._rewrite_device_key(key), *args, **kwargs) + + def put_object(self, key: str, *args: Any, **kwargs: Any) -> Any: + return self._client.put_object(self._rewrite_device_key(key), *args, **kwargs) + + def get_object_then_delete(self, key: str, *args: Any, **kwargs: Any) -> Any: + return self._client.get_object_then_delete( + self._rewrite_device_key(key), *args, **kwargs + ) + + def _get_router_dtype(config): """Read router dtype from a flat or multimodal Hugging Face config.""" router_dtype = getattr(config, "router_dtype", None) @@ -289,7 +296,7 @@ def _compute_local_raw_meta(self) -> dict: """Per-rank raw meta via awex's own staticmethod (HF-converted names).""" server_args = self._scheduler.server_args model_context = self._build_model_context() - return InferParamMetaResolver._get_model_param_info( + raw_meta = InferParamMetaResolver._get_model_param_info( "sglang", server_args, convert_params=True, @@ -297,6 +304,7 @@ def _compute_local_raw_meta(self) -> dict: model=self._get_model(), model_context=model_context, ) + return raw_meta def _build_instance_params_meta(self): """Gather single-instance raw meta via the MetaServer, then aggregate. @@ -489,6 +497,13 @@ def _ensure_reader(self) -> NCCLWorkerWeightsReader: ipc_backend="cuda", enable_debug_mode=False, ) + # SGLang t1 workers are CUDA_VISIBLE_DEVICES-isolated, so AWEX observes + # logical device 0 on every process while the colocated trainer + # publishes IPC handles under node-local physical ids. Keep CUDA calls + # on logical device 0, but translate only MetaServer identities/keys. + reader.meta_server_client = _PhysicalDeviceMetaServerClient( + reader.meta_server_client, self._local_gpu_id + ) reader.initialize() self._reader = reader logger.info( @@ -500,6 +515,7 @@ def _ensure_reader(self) -> NCCLWorkerWeightsReader: ) return reader + @torch.no_grad() def update_weights(self, version: int) -> None: """Run one colocate weight update via the native awex worker reader. diff --git a/areal/engine/awex/colocate_writer.py b/areal/engine/awex/colocate_writer.py index ed226bd638..e47fe0ec0d 100644 --- a/areal/engine/awex/colocate_writer.py +++ b/areal/engine/awex/colocate_writer.py @@ -28,11 +28,13 @@ import torch import torch.distributed as dist +from areal.engine.megatron_utils.weight_residency import MegatronWeightResidency +from areal.utils.environ import get_float_env_var +from areal.utils.logging import getLogger + if TYPE_CHECKING: from areal.engine.megatron_engine import MegatronEngine -from areal.utils.logging import getLogger - logger = getLogger("AwexColocate") @@ -43,46 +45,46 @@ def resolve_physical_gpu_id(relative_gpu_id: int) -> int: transfer have to agree on physical GPU ids. Inside a process that was given a device mask, ``torch.cuda.current_device()`` and SGLang's ``gpu_id`` are indices into that mask rather than physical ids, so the - mask itself is the only ground truth. Falls back to the relative index - when the mask is absent or holds GPU UUIDs. + mask itself is the only ground truth. UUID masks and invalid indices are + rejected because they cannot produce the node-local ordinal AWEX keys use. """ cuda_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "") if not cuda_visible: return relative_gpu_id - try: - gpu_ids = [int(x) for x in cuda_visible.split(",") if x.strip()] - return gpu_ids[relative_gpu_id] - except (ValueError, IndexError): - return relative_gpu_id + visible_devices = [item.strip() for item in cuda_visible.split(",") if item.strip()] + if not all(item.isdigit() for item in visible_devices): + raise ValueError( + "AWEX colocate requires numeric CUDA_VISIBLE_DEVICES entries; " + f"got {visible_devices!r}" + ) + if relative_gpu_id >= len(visible_devices): + raise ValueError( + f"CUDA device {relative_gpu_id} is outside " + f"CUDA_VISIBLE_DEVICES={visible_devices!r}" + ) + return int(visible_devices[relative_gpu_id]) def awex_colocate_timeout_s(default: float = 1800.0) -> float: - value = os.environ.get("AWEX_COLOCATE_TIMEOUT_S", "").strip() - if not value: - return default - try: - return float(value) - except ValueError: - logger.warning( - "Invalid AWEX_COLOCATE_TIMEOUT_S=%r; using default %.1fs", - value, - default, - ) - return default + return get_float_env_var("AWEX_COLOCATE_TIMEOUT_S", default) -class AwexMegatronAdapter: - """Training-side adapter for AWEX colocated weight transfer. +class AwexWeightPublisher: + """Publish Megatron weights to a colocated SGLang engine through AWEX. Uses CUDA IPC (share_memory + ForkingPickler serialization) for zero-copy weight transfer to the colocated SGLang process on the same GPU. The infer - side handles redistribution among infer ranks via its own NCCL group. + side handles redistribution among infer ranks via its own NCCL group. GPU + residency is delegated to one shared :class:`MegatronWeightResidency`. """ - def __init__(self, engine: MegatronEngine): + def __init__( + self, + engine: MegatronEngine, + residency: MegatronWeightResidency | None = None, + ) -> None: self._engine = engine - self._offloaded_weights: dict[str, torch.Tensor] = {} - self._released_tags: set[str] = set() + self._residency = residency or MegatronWeightResidency(engine) self._meta_server_addr: str | None = None self._meta_server_client = None self._transfer_rank: int | None = None @@ -94,6 +96,11 @@ def __init__(self, engine: MegatronEngine): self._num_infer_engines: int | None = None self._logical_train_rank: int | None = None + @property + def residency(self) -> MegatronWeightResidency: + """Return the sole residency manager used during publication.""" + return self._residency + def init_colocate_weight_update( self, meta_server_addr: str | None = None, @@ -132,11 +139,31 @@ def init_colocate_weight_update( ) logger.info( - "AwexMegatronAdapter initialized: meta_server=%s, transfer_rank=%d", + "AwexWeightPublisher initialized: meta_server=%s, transfer_rank=%d", meta_server_addr, transfer_rank, ) + def eager_publish_train_info(self, meta_server_addr: str | None) -> None: + """Publish train world metadata before the colocated reader starts.""" + addr = meta_server_addr or os.environ.get("AWEX_META_SERVER_ADDR", "") + if not addr or (dist.is_initialized() and dist.get_rank() != 0): + return + try: + from awex.meta.meta_server import MetaServerClient + + host, port = addr.rsplit(":", 1) + client = MetaServerClient(host, int(port)) + world = dist.get_world_size() if dist.is_initialized() else 1 + client.put_object("awex_train_info", {"train_world_size": world}) + logger.info( + "Eager-published awex_train_info (train_world_size=%d) to %s", + world, + addr, + ) + except Exception as exc: + logger.warning("Eager publish awex_train_info failed: %s", exc) + def _lazy_initialize(self) -> None: """Perform deferred initialization: metadata exchange and weight converter setup. @@ -241,33 +268,13 @@ def resume_memory_occupation(self, tags=None): training_world_size, ) - def _release_grad_memory(self) -> None: - """Release gradient buffers to free GPU memory before weight conversion. - - Mirrors the AWEX reference release_grad_memory(). - Saves original sizes to buffer.grad_data_size for later restoration. - """ - from megatron.core.distributed import DistributedDataParallel as DDP - - model = self._engine.model - if model is None: - return - if not isinstance(model, (list, tuple)): - model = [model] - count = 0 - for chunk in model: - if isinstance(chunk, DDP): - for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: - for buf in buffers: - if buf.grad_data.storage().size() > 0: - buf.grad_data_size = buf.grad_data.storage().size() - buf.grad_data.storage().resize_(0) - count += 1 - if count > 0: - torch.cuda.synchronize() - gc.collect() - torch.cuda.empty_cache() - logger.info("Released %d grad buffers", count) + def _prepare_residency_for_publish(self) -> None: + """Free optimizer/grad memory before making weights resident.""" + weights_were_offloaded = self._residency.is_released("weights") + self._residency.release_memory(tags=["optimizer"]) + self._residency.release_grad_memory() + if weights_were_offloaded: + self._residency.resume_memory(tags=["weights"]) @torch.no_grad() def execute_colocate_weight_update(self, version: int) -> None: @@ -289,8 +296,6 @@ def execute_colocate_weight_update(self, version: int) -> None: release_tensors, ) - weights_were_offloaded = "weights" in self._released_tags - # Reclaim any IPC-exported blocks from the previous version whose # peer mappings closed after our last collect (belt-and-braces with the # collect in this method's finally block). @@ -306,12 +311,7 @@ def execute_colocate_weight_update(self, version: int) -> None: # _release_memory_for_weights_exchange). # Optimizer/grad offload operate on independent Megatron buffers and do # not require the weights to be resumed, so reordering is safe. - self.release_memory(tags=["optimizer"]) - - self._release_grad_memory() - - if weights_were_offloaded: - self.resume_memory(tags=["weights"]) + self._prepare_residency_for_publish() # _lazy_initialize AFTER the weights resume — its meta resolver # runs convert_param over live params, which dies with CUDA invalid @@ -369,7 +369,7 @@ def execute_colocate_weight_update(self, version: int) -> None: del tensors, owned parameters.clear() - self.release_memory(tags=["weights"]) + self._residency.release_memory(tags=["weights"]) ip_address = self._ip_address device_id = self._physical_gpu_id @@ -408,7 +408,7 @@ def execute_colocate_weight_update(self, version: int) -> None: update_finished_key = f"weights_update_finished{key_suffix}" try: try: - self._meta_server_client.get_object( + completion = self._meta_server_client.get_object( update_finished_key, timeout=self._timeout_s ) except Exception: @@ -421,6 +421,12 @@ def execute_colocate_weight_update(self, version: int) -> None: update_finished_key, ) raise + if isinstance(completion, dict) and not completion.get("ok", True): + error = completion.get("error", "unknown inference-side error") + raise RuntimeError( + "Inference rejected AWEX weights before IPC release: " + f"version={version}, device={device_id}, error={error}" + ) self._meta_server_client.delete_if_exists(update_finished_key) self._meta_server_client.delete_if_exists(serialized_weights_key) logger.info("Got done signal from infer side: %s", update_finished_key) @@ -503,273 +509,25 @@ def _convert_parameters(self) -> dict[str, torch.Tensor]: # ── Memory management (manual offload) ──────────────────────────────── def release_memory(self, tags: list[str] | None = None) -> None: - tags = tags or ["optimizer", "weights"] - tags_to_release = [t for t in tags if t not in self._released_tags] - if not tags_to_release: - return - - if "optimizer" in tags_to_release: - self._offload_optimizer_states() - self._released_tags.add("optimizer") - - if "weights" in tags_to_release: - self._offload_model_weights() - self._released_tags.add("weights") - - torch.cuda.synchronize() - gc.collect() - torch.cuda.empty_cache() - logger.info("release_memory done: tags=%s", tags_to_release) + """Compatibility delegate for callers of the former combined adapter.""" + self._residency.release_memory(tags) def resume_memory(self, tags: list[str] | None = None) -> None: - tags = tags or ["optimizer", "weights"] - tags_to_resume = [t for t in tags if t in self._released_tags] - if not tags_to_resume: - return - - if "weights" in tags_to_resume: - self._reload_model_weights(load_grad=False) - self._released_tags.discard("weights") - - if "optimizer" in tags_to_resume: - self._reload_optimizer_states() - self._released_tags.discard("optimizer") + """Compatibility delegate for callers of the former combined adapter.""" + self._residency.resume_memory(tags) - torch.cuda.synchronize() - logger.info("resume_memory done: tags=%s", tags_to_resume) - - def _offload_model_weights(self) -> None: - from megatron.core.distributed import DistributedDataParallel as DDP + @property + def _released_tags(self) -> set[str]: + """Compatibility view of the composed residency state.""" + return set(self._residency.released_tags) - model = self._engine.model - if model is None: - return - if not isinstance(model, (list, tuple)): - model = [model] - count = 0 - for chunk in model: - if isinstance(chunk, DDP): - for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: - for buf in buffers: - if hasattr(buf, "offload_to_cpu"): - buf.offload_to_cpu() - count += 1 - continue - if buf.param_data.storage().size() > 0: - if not hasattr(buf.param_data, "cpu_data"): - buf.param_data.cpu_data = torch.zeros( - buf.param_data.data.shape, - dtype=buf.param_data.data.dtype, - pin_memory=True, - device="cpu", - ) - buf.param_data.cpu_data.copy_(buf.param_data.data) - buf.param_data_size = buf.param_data.storage().size() - buf.param_data.storage().resize_(0) - count += 1 - if buf.grad_data.storage().size() > 0: - buf.grad_data_size = buf.grad_data.storage().size() - buf.grad_data.storage().resize_(0) - else: - for name, param in chunk.named_parameters(): - if param.data.is_cuda: - self._offloaded_weights[name] = param.data.detach().to( - "cpu", non_blocking=True - ) - param.data = torch.empty(0, device="cpu") - count += 1 - torch.cuda.synchronize() - logger.info("Offloaded %d weight buffers to CPU", count) - - def _reload_model_weights(self, load_grad: bool = False) -> None: - from megatron.core.distributed import DistributedDataParallel as DDP - - model = self._engine.model - if model is None: - return - if not isinstance(model, (list, tuple)): - model = [model] - device = self._engine.device - for chunk in model: - if isinstance(chunk, DDP): - for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: - for buf in buffers: - if hasattr(buf, "reload_from_cpu"): - buf.reload_from_cpu(move_grads=load_grad) - continue - if buf.param_data.storage().size() == 0: - buf.param_data.storage().resize_(buf.param_data_size) - buf.param_data.copy_(buf.param_data.cpu_data, non_blocking=True) - if ( - load_grad - and hasattr(buf, "grad_data_size") - and buf.grad_data.storage().size() == 0 - ): - buf.grad_data.storage().resize_(buf.grad_data_size) - buf.grad_data.zero_() - else: - for name, param in chunk.named_parameters(): - if name in self._offloaded_weights: - param.data = self._offloaded_weights[name].to( - device, non_blocking=True - ) - self._offloaded_weights.clear() - torch.cuda.synchronize() - logger.info("Reloaded model weights to GPU (load_grad=%s)", load_grad) + def _release_grad_memory(self) -> None: + """Compatibility delegate for the former combined adapter.""" + self._residency.release_grad_memory() def ensure_grad_buffers(self) -> None: - """Allocate grad buffers if they were freed during offload. - - Called before forward_backward (training) to ensure grad storage - is available for backward pass. Separate from _reload_model_weights - because compute_logp (inference-only) should not allocate grad buffers. - """ - from megatron.core.distributed import DistributedDataParallel as DDP - - model = self._engine.model - if model is None: - return - if not isinstance(model, (list, tuple)): - model = [model] - count = 0 - for chunk in model: - if isinstance(chunk, DDP): - for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: - for buf in buffers: - if ( - hasattr(buf, "grad_data_size") - and buf.grad_data.storage().size() == 0 - ): - buf.grad_data.storage().resize_(buf.grad_data_size) - buf.grad_data.zero_() - count += 1 - if count > 0: - torch.cuda.synchronize() - logger.info("Allocated %d grad buffers for training", count) - - def _get_inner_optimizers(self): - optimizer = self._engine.optimizer - if optimizer is None: - return [] - if hasattr(optimizer, "chained_optimizers"): - inner_optimizers = optimizer.chained_optimizers - elif hasattr(optimizer, "optimizers"): - inner_optimizers = optimizer.optimizers - else: - inner_optimizers = [optimizer] - return inner_optimizers - - def _offload_optimizer_states(self) -> None: - optimizer = self._engine.optimizer - if optimizer is None: - return - # Default path mirrors the AWEX reference optimizer offload - # (megatron_util.offload_megatron_optimizer): swap .data / state-dict - # references to CPU, never resize_ storages, then purge TE's global - # _dummy_wgrads cache and synchronize. Megatron HybridDeviceOptimizer's - # offload_to_cpu/restore_from_cpu is kept only as an opt-in fallback — - # its internal pointer bookkeeping is hard to validate and the AWEX - # reference integration deliberately avoids it. - if os.environ.get("AWEX_OPT_OFFLOAD_VIA_HDO", "").strip() == "1" and hasattr( - optimizer, "offload_to_cpu" - ): - optimizer.offload_to_cpu() - logger.info("Offloaded optimizer via offload_to_cpu()") - return - - inner_optimizers = self._get_inner_optimizers() - if not inner_optimizers: - return - - count = 0 - for opt in inner_optimizers: - # Offload FP32 main parameter copies (shard_fp32_from_float16_groups) - if hasattr(opt, "shard_fp32_from_float16_groups"): - for group in opt.shard_fp32_from_float16_groups: - if isinstance(group, list): - for t in group: - if t is not None and t.data.is_cuda: - t.data = t.data.to("cpu", non_blocking=True) - count += 1 - elif group is not None and group.data.is_cuda: - group.data = group.data.to("cpu", non_blocking=True) - count += 1 - - # Offload Adam states (exp_avg, exp_avg_sq) - base_opt = getattr(opt, "optimizer", opt) - if not hasattr(base_opt, "state") or base_opt.state is None: - continue - for state in base_opt.state.values(): - for key in ("exp_avg", "exp_avg_sq"): - if ( - key in state - and isinstance(state[key], torch.Tensor) - and state[key].is_cuda - ): - state[key] = state[key].to("cpu", non_blocking=True) - count += 1 - - # Targeted fix from the AWEX reference: transformer_engine caches dummy wgrad - # tensors in a module-global dict; without purging it the GPU memory - # is never actually freed and stale references survive the offload. - try: - from transformer_engine.pytorch.module.base import _dummy_wgrads - - purged = len(_dummy_wgrads) - for k in list(_dummy_wgrads): - del _dummy_wgrads[k] - if purged: - logger.info("Purged %d TE _dummy_wgrads cache entries", purged) - except ImportError: - pass - torch.cuda.synchronize() - logger.info("Offloaded %d optimizer state tensors to CPU", count) - - def _reload_optimizer_states(self) -> None: - optimizer = self._engine.optimizer - if optimizer is None: - return - if os.environ.get("AWEX_OPT_OFFLOAD_VIA_HDO", "").strip() == "1" and hasattr( - optimizer, "restore_from_cpu" - ): - optimizer.restore_from_cpu() - logger.info("Reloaded optimizer via restore_from_cpu()") - return - - inner_optimizers = self._get_inner_optimizers() - if not inner_optimizers: - return - - device = self._engine.device - count = 0 - for opt in inner_optimizers: - # Reload FP32 main parameter copies - if hasattr(opt, "shard_fp32_from_float16_groups"): - for group in opt.shard_fp32_from_float16_groups: - if isinstance(group, list): - for t in group: - if t is not None and not t.data.is_cuda: - t.data = t.data.to(device, non_blocking=True) - count += 1 - elif group is not None and not group.data.is_cuda: - group.data = group.data.to(device, non_blocking=True) - count += 1 - - # Reload Adam states - base_opt = getattr(opt, "optimizer", opt) - if not hasattr(base_opt, "state") or base_opt.state is None: - continue - for state in base_opt.state.values(): - for key in ("exp_avg", "exp_avg_sq"): - if ( - key in state - and isinstance(state[key], torch.Tensor) - and not state[key].is_cuda - ): - state[key] = state[key].to(device, non_blocking=True) - count += 1 - torch.cuda.synchronize() - logger.info("Reloaded %d optimizer state tensors to GPU", count) + """Compatibility delegate for the former combined adapter.""" + self._residency.ensure_grad_buffers() def _get_tf_config(models): @@ -781,3 +539,15 @@ def _get_tf_config(models): if cfg is not None: return cfg return None + + +# Backward-compatible name for callers that imported the former combined +# adapter. New engine code uses AwexWeightPublisher and injects its residency. +AwexMegatronAdapter = AwexWeightPublisher + +__all__ = [ + "AwexMegatronAdapter", + "AwexWeightPublisher", + "awex_colocate_timeout_s", + "resolve_physical_gpu_id", +] diff --git a/areal/engine/awex/memory_saver.py b/areal/engine/awex/memory_saver.py new file mode 100644 index 0000000000..2188eb16d2 --- /dev/null +++ b/areal/engine/awex/memory_saver.py @@ -0,0 +1,49 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Compatibility helpers for SGLang's torch-memory-saver integration.""" + +from __future__ import annotations + +import os + + +def patch_tms_hook_mode() -> None: + """Keep pauseable CUDA graphs on torch-memory-saver's preload hook. + + ``megatron.core.inference.contexts.dynamic_context`` assigns + ``torch_memory_saver.hook_mode = "torch"`` at module import time. SGLang's + pauseable CUDA graphs require the default ``preload`` hook, so drop that + assignment when graph saving is enabled. Also ignore unsafe attempts to + reconfigure the singleton after its implementation has been initialized. + + This must run from the SGLang entry module, before the scheduler imports + Megatron transitively. Calling it later from the AWEX weight reader is too + late because CUDA graphs are captured before that reader is constructed. + """ + try: + import torch_memory_saver as tms + except Exception: + return + + instance = getattr(tms, "torch_memory_saver", None) + if instance is None: + return + cls = type(instance) + prop = cls.hook_mode + if getattr(prop.fset, "_awex_safe", False): + return + + def safe_setter(self, value): + if value == "torch" and os.environ.get( + "SGLANG_MEMORY_SAVER_CUDA_GRAPH", "" + ).lower() in {"1", "true", "yes", "on"}: + return + if not hasattr(self, "_impl_ctor_kwargs"): + return + prop.fset(self, value) + + safe_setter._awex_safe = True + cls.hook_mode = property(prop.fget, safe_setter) + + +__all__ = ["patch_tms_hook_mode"] diff --git a/areal/engine/awex/sglang_plugin.py b/areal/engine/awex/sglang_plugin.py index 4766423f03..4353ee9ccc 100644 --- a/areal/engine/awex/sglang_plugin.py +++ b/areal/engine/awex/sglang_plugin.py @@ -7,7 +7,7 @@ fetches IPC handles from MetaServer (CPU I/O) and queues them for the scheduler's main loop to process (CUDA copy on main thread). -Weight transfer flow (mirrors the AWEX reference colocate mode): +Weight transfer flow (aligned with Asystem colocate mode): 1. Training side: convert params → cuda_ipc_serialize → MetaServer put 2. Background thread: MetaServer get → queue IPC data (CPU only) 3. Scheduler main loop: release_memory → deserialize + copy → resume_memory @@ -24,24 +24,26 @@ from __future__ import annotations +import importlib import os import queue import threading import time from collections.abc import Callable +from copy import copy from dataclasses import dataclass, field from typing import Any +from areal.engine.awex.memory_saver import patch_tms_hook_mode -def assert_alloc_conf_supports_memory_saver(conf: str) -> None: - """Reject allocator configs that silently disable SGLang's memory saver. +# Must run before importing SGLang. Its scheduler may import Megatron while +# initializing the model, and Megatron otherwise switches torch-memory-saver +# away from the preload hook required by pauseable CUDA graphs. +patch_tms_hook_mode() - torch_memory_saver disables itself when it sees expandable_segments, so - release/resume becomes a no-op: the rollout never hands its GPU back and - weight pages stay mapped, which surfaces much later as a colocate OOM or an - invalid CUDA IPC target. Colocated roles must therefore carry their own - scheduling_spec env_vars rather than share the actor's. - """ + +def assert_alloc_conf_supports_memory_saver(conf: str) -> None: + """Reject allocator configs that silently disable SGLang's memory saver.""" if "expandable_segments:true" in conf.lower().replace(" ", ""): raise RuntimeError( "SGLang's memory saver cannot unmap/remap expandable segments, so " @@ -53,44 +55,119 @@ def assert_alloc_conf_supports_memory_saver(conf: str) -> None: assert_alloc_conf_supports_memory_saver(os.environ.get("PYTORCH_CUDA_ALLOC_CONF", "")) - from areal.utils import pkg_version # noqa: E402 +from areal.utils.environ import ( # noqa: E402 + get_bool_env_var, + get_float_env_var, + get_int_env_var, +) from areal.utils.logging import getLogger # noqa: E402 logger = getLogger("AwexSGLangPlugin") - - SUPPORTED_SGLANG_VERSIONS = ("0.5.9", "0.5.10.post1") def assert_supported_sglang_version() -> None: - """Refuse to patch a SGLang build whose internals were not verified. - - This plugin reaches into scheduler internals: it wraps Scheduler.__init__, - replaces both event loops, and backports a model-worker task dispatcher. - Those touch points move between releases, so a silent mismatch surfaces as - a hang or a corrupt transfer rather than an import error. Fail loudly - instead, and extend the tuple only after re-checking the patched surfaces. - """ + """Refuse to patch a SGLang build whose internals were not verified.""" installed = pkg_version.get_version("sglang") if installed not in SUPPORTED_SGLANG_VERSIONS: raise RuntimeError( - f"AWEX colocate patches SGLang internals and was verified against " + "AWEX colocate patches SGLang internals and was verified against " f"{', '.join(SUPPORTED_SGLANG_VERSIONS)}, but found {installed}. " - f"Re-check Scheduler.__init__, the event loops, and " - f"execute_task_in_model_worker before allowing this version." + "Re-check Scheduler.__init__, the event loops, and " + "execute_task_in_model_worker before allowing this version." ) -def _float_env(name: str, default: float) -> float: - value = os.environ.get(name, "") - if not value: - return default +def _load_sglang_plugins_if_available() -> bool: + """Load SGLang runtime plugins when supported by the installed version. + + ``sglang.srt.plugins`` was added after the 0.5.10 runtime currently pinned + by AReaL. AWEX does not depend on that registry because it injects its + scheduler entry point directly through ``launch_server``. Treat the + registry as optional so the same entry module works with both APIs. + """ try: - return float(value) - except ValueError: - logger.warning("Invalid %s=%r; using %.3f", name, value, default) - return default + plugins = importlib.import_module("sglang.srt.plugins") + except ModuleNotFoundError as exc: + if exc.name != "sglang.srt.plugins": + raise + logger.info( + "[AWEX] SGLang plugin registry is unavailable; using the " + "launch_server scheduler hook" + ) + return False + + load_plugins = getattr(plugins, "load_plugins", None) + if not callable(load_plugins): + logger.info( + "[AWEX] SGLang plugin registry has no load_plugins entry point; " + "using the launch_server scheduler hook" + ) + return False + + load_plugins() + return True + + +def _resolve_transfer_rank( + *, + infer_world_size: int, + gpu_id: int, + node_id: int, + nnodes: int, + instance_world_size: int, +) -> int: + """Resolve the inference rank in AWEX's global transfer world. + + A one-GPU SGLang server can use the colocated actor's inherited global + rank when CUDA device isolation remaps its only GPU to device zero. For a + multi-GPU server, every TP/PP scheduler inherits the same environment, so + its scheduler-local physical GPU identity must be used instead. + """ + explicit_rank = os.environ.get("AWEX_TRANSFER_RANK") + if explicit_rank is not None: + transfer_rank = int(explicit_rank) + else: + env_rank = os.environ.get("RANK") + env_world_size = os.environ.get("WORLD_SIZE") + if ( + instance_world_size == 1 + and env_rank is not None + and env_world_size is not None + and int(env_world_size) == infer_world_size + ): + transfer_rank = int(env_rank) + else: + n_gpus_per_node = max(1, infer_world_size // nnodes) + transfer_rank = node_id * n_gpus_per_node + gpu_id + + if not 0 <= transfer_rank < infer_world_size: + raise ValueError( + "AWEX transfer rank must be in " + f"[0, {infer_world_size}), got {transfer_rank}" + ) + return transfer_rank + + +def _writer_version_key(ip_address: str, physical_gpu_id: int) -> str: + return f"awex_writer_version_{ip_address}_{physical_gpu_id}" + + +def _try_get_writer_version( + meta_server_client: Any, + key: str, + timeout_s: float, +) -> int | None: + """Return the writer's current version if published, otherwise None.""" + + try: + wait_key = getattr(meta_server_client, "wait_key", None) + if callable(wait_key): + wait_key(key, timeout=timeout_s) + return int(meta_server_client.get_object(key, timeout=timeout_s)) + except Exception: + return None class AwexSchedulerPlugin: @@ -107,7 +184,72 @@ def __init__(self, scheduler: Any) -> None: self._weight_queue: queue.Queue = queue.Queue() self._version = 0 self._paused_poll_interval_s = max( - 0.0, _float_env("AWEX_PAUSED_POLL_INTERVAL_S", 0.01) + 0.0, get_float_env_var("AWEX_PAUSED_POLL_INTERVAL_S", 0.01) + ) + self._process_queue_when_idle = get_bool_env_var( + "AREAL_AWEX_PROCESS_QUEUE_WHEN_IDLE", "true" + ) + # Idle-poll throttle in *loop iterations*, not wall-clock time: TP + # ranks run the scheduler loop in lockstep, so a loop-count gate is + # deterministic across ranks (a time-based gate deadlocks, see + # _maybe_process_awex_queue_when_idle). + self._idle_poll_loops = max(1, get_int_env_var("AWEX_IDLE_POLL_LOOPS", 64)) + + @staticmethod + def _int_attr(scheduler: Any, name: str, default: int) -> int: + for obj in ( + scheduler, + getattr(scheduler, "ps", None), + getattr(scheduler, "server_args", None), + ): + if obj is None or not hasattr(obj, name): + continue + value = getattr(obj, name) + if value is not None: + return int(value) + return default + + @staticmethod + def _callable(scheduler: Any, name: str) -> Callable: + for obj in (scheduler, getattr(scheduler, "weight_updater", None)): + if obj is None: + continue + method = getattr(obj, name, None) + if callable(method): + return method + raise AttributeError(f"Scheduler has no callable {name!r}") + + def _logical_gpu_id(self) -> int: + return self._int_attr(self._scheduler, "gpu_id", 0) + + def _physical_gpu_id(self) -> int: + """Return the node-local physical GPU id used by AWEX keys.""" + logical_gpu_id = self._logical_gpu_id() + visible_devices = [ + item.strip() + for item in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") + if item.strip() + ] + if visible_devices: + if logical_gpu_id >= len(visible_devices): + raise ValueError( + f"SGLang gpu_id={logical_gpu_id} is outside " + f"CUDA_VISIBLE_DEVICES={visible_devices!r}" + ) + physical_gpu_id = visible_devices[logical_gpu_id] + if not physical_gpu_id.isdigit(): + raise ValueError( + "AWEX colocate requires numeric CUDA_VISIBLE_DEVICES entries " + "so MetaServer keys can use node-local physical GPU ids; " + f"got {physical_gpu_id!r}" + ) + return int(physical_gpu_id) + return logical_gpu_id + + def _instance_world_size(self) -> int: + """Return scheduler GPU processes in one SGLang server.""" + return self._int_attr(self._scheduler, "tp_size", 1) * self._int_attr( + self._scheduler, "pp_size", 1 ) def bind(self) -> None: @@ -122,6 +264,7 @@ def bind(self) -> None: ] for name in methods: setattr(self._scheduler, name, getattr(self, name)) + self._patch_memory_transitions() logger.info( f"[AWEX] AwexSchedulerPlugin bound {len(methods)} methods to scheduler", ) @@ -138,6 +281,55 @@ def _require_receiver(self): self._receiver = AwexColocateReader(self._scheduler) return self._receiver + def _patch_memory_transitions(self) -> None: + """Make AWEX release/resume requests idempotent across retries.""" + scheduler = self._scheduler + if getattr(scheduler, "_areal_awex_memory_transitions_patched", False): + return + original_release = getattr(scheduler, "release_memory_occupation", None) + original_resume = getattr(scheduler, "resume_memory_occupation", None) + if original_release is None or original_resume is None: + return + + def _filtered_request(request: Any, *, release: bool) -> Any | None: + tags = getattr(request, "tags", None) + offload_tags = getattr(scheduler, "offload_tags", None) + if tags is None or offload_tags is None: + return request + effective_tags = [ + tag + for tag in tags + if (tag not in offload_tags if release else tag in offload_tags) + ] + if not effective_tags: + logger.info( + "[AWEX] skipping duplicate %s_memory_occupation(tags=%s)", + "release" if release else "resume", + tags, + ) + return None + if effective_tags == list(tags): + return request + filtered = copy(request) + filtered.tags = effective_tags + return filtered + + def _release(request: Any, *args: Any, **kwargs: Any) -> Any: + filtered = _filtered_request(request, release=True) + if filtered is None: + return None + return original_release(filtered, *args, **kwargs) + + def _resume(request: Any, *args: Any, **kwargs: Any) -> Any: + filtered = _filtered_request(request, release=False) + if filtered is None: + return None + return original_resume(filtered, *args, **kwargs) + + scheduler.release_memory_occupation = _release + scheduler.resume_memory_occupation = _resume + scheduler._areal_awex_memory_transitions_patched = True + def awex_init_receiver(self, **kwargs: Any) -> None: self._require_receiver().initialize(**kwargs) @@ -158,21 +350,24 @@ def awex_get_parallelism(self) -> dict: # ── Main loop hook: process queued weight updates ───────────────── - def process_awex_queue(self) -> None: + def process_awex_queue(self, extra_ready: bool = True) -> None: """Called from scheduler main loop. Processes pending weight updates. This is a TP-collective operation: ALL TP ranks must call it together (since it's called between recv_requests() calls which use broadcast_pyobj). - Uses all_reduce(MIN) to check if all TP ranks have a pending update. - Only proceeds when ALL ranks have queued an update, preventing the deadlock - where one rank blocks in CUDA ops while others wait in broadcast_pyobj. - - We act as the awex *driver* layer (the community SGLang scheduler has no - ``execute_task_in_model_worker`` driver). The collect-IPC + StreamBatch - transport + writer handshake is delegated to the awex-native worker reader - (``AwexColocateReader.update_weights`` -> ``NCCLWorkerWeightsReader``). We - only own the driver-equivalent steps around it: + Uses all_reduce(MIN) to check if all TP ranks are ready to process a + pending update. ``extra_ready`` lets callers fold per-rank conditions + (e.g. idle state) into the collective decision instead of returning + early, which would desynchronize the ranks. Only proceeds when ALL + ranks are ready, preventing the deadlock where one rank blocks in CUDA + ops while others wait in broadcast_pyobj. + + We act as the awex *driver* layer for the queued colocate update. The + collect-IPC + StreamBatch transport + writer handshake is delegated to the + awex-native worker reader (``AwexColocateReader.update_weights`` -> + ``NCCLWorkerWeightsReader``). We only own the driver-equivalent steps + around it: 1. Wait for all_training_offloaded_weights (= driver _pre_update_weights) 2. resume_memory_occupation(weights) — re-allocate infer weight buffers 3. reader.update_weights(version) — awex worker reader does the rest: @@ -184,9 +379,9 @@ def process_awex_queue(self) -> None: import torch.distributed tp_cpu_group = self._scheduler.tp_cpu_group - tp_size = self._scheduler.tp_size + tp_size = self._int_attr(self._scheduler, "tp_size", 1) - has_item = 1 if not self._weight_queue.empty() else 0 + has_item = 1 if (extra_ready and not self._weight_queue.empty()) else 0 if tp_size > 1: has_item_tensor = torch.tensor([has_item], dtype=torch.int32) @@ -227,7 +422,10 @@ def process_awex_queue(self) -> None: # Step 2: Resume weight memory (memory_saver re-allocates buffers). resume_req = ResumeMemoryOccupationReqInput(tags=["weights"]) - self._scheduler.resume_memory_occupation(resume_req) + resume_memory_occupation = self._callable( + self._scheduler, "resume_memory_occupation" + ) + resume_memory_occupation(resume_req) logger.info( f"[AWEX] main loop: resumed weight memory for v{version} (gpu_id={gpu_id})", ) @@ -264,40 +462,56 @@ def _patch_event_loop(self) -> None: scheduler = self._scheduler plugin = self - _decode_hooks_available = hasattr(scheduler, "log_decode_stats") and hasattr( - scheduler, "log_decode_stats_every_iteration" + decode_stats_name = next( + ( + name + for name in ("log_decode_stats", "report_decode_stats") + if callable(getattr(scheduler, name, None)) + ), + None, ) - if _decode_hooks_available: - _orig_log_decode_stats = scheduler.log_decode_stats - _orig_log_decode_stats_every_iteration = ( - scheduler.log_decode_stats_every_iteration - ) + _orig_decode_stats = ( + getattr(scheduler, decode_stats_name) if decode_stats_name else None + ) + every_iteration_name = "log_decode_stats_every_iteration" + _orig_decode_stats_every_iteration = getattr( + scheduler, every_iteration_name, None + ) + has_decode_stats = callable(_orig_decode_stats) + has_decode_iter_stats = callable(_orig_decode_stats_every_iteration) - def _tracked_log_decode_stats(*args, **kwargs): - scheduler._areal_awex_last_decode_stats_ct = getattr( - scheduler, "forward_ct_decode", None - ) - return _orig_log_decode_stats(*args, **kwargs) + def _tracked_decode_stats(*args, **kwargs): + scheduler._areal_awex_last_decode_stats_ct = getattr( + scheduler, "forward_ct_decode", None + ) + return _orig_decode_stats(*args, **kwargs) - def _tracked_log_decode_stats_every_iteration(*args, **kwargs): - scheduler._areal_awex_last_decode_stats_every_iter_ct = getattr( - scheduler, "forward_ct_decode", None - ) - return _orig_log_decode_stats_every_iteration(*args, **kwargs) + def _tracked_decode_stats_every_iteration(*args, **kwargs): + scheduler._areal_awex_last_decode_stats_every_iter_ct = getattr( + scheduler, "forward_ct_decode", None + ) + return _orig_decode_stats_every_iteration(*args, **kwargs) - scheduler.log_decode_stats = _tracked_log_decode_stats - scheduler.log_decode_stats_every_iteration = ( - _tracked_log_decode_stats_every_iteration + if has_decode_stats: + setattr(scheduler, decode_stats_name, _tracked_decode_stats) + else: + logger.info( + "[AWEX] Scheduler has no decode stats method; " + "skipping decode metrics restore hook", + ) + if has_decode_iter_stats: + setattr( + scheduler, + every_iteration_name, + _tracked_decode_stats_every_iteration, ) else: - logger.warning( - "[AWEX] sglang scheduler has no log_decode_stats hooks " - "(removed in sglang>=0.5.10); skipping decode-stats tracking" + logger.info( + "[AWEX] Scheduler has no log_decode_stats_every_iteration; " + "skipping per-iteration decode metrics restore hook", ) def _maybe_restore_decode_metrics(stage, batch, result): - if not _decode_hooks_available: - return if os.environ.get("AREAL_AWEX_FORCE_SGLANG_METRICS", "1") != "1": return if stage != "after_process_batch_result" or batch is None: @@ -318,27 +532,103 @@ def _maybe_restore_decode_metrics(stage, batch, result): should_log_decode = current_ct is not None and current_ct % interval == 0 if ( - should_log_decode + has_decode_stats + and callable(getattr(scheduler, decode_stats_name, None)) + and should_log_decode and getattr(scheduler, "_areal_awex_last_decode_stats_ct", None) != current_ct ): can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) logger.debug( - f"[AWEX-METRICS] restoring native log_decode_stats " + f"[AWEX-METRICS] restoring native {decode_stats_name} " f"gpu_id={getattr(scheduler, 'gpu_id', '?')} " f"forward_ct_decode={current_ct}", ) - scheduler.log_decode_stats(can_run_cuda_graph, running_batch=batch) + decode_stats_kwargs = {"running_batch": batch} + if decode_stats_name == "report_decode_stats": + decode_stats_kwargs["num_accepted_tokens"] = getattr( + result, "num_accepted_tokens", 0 + ) + getattr(scheduler, decode_stats_name)( + can_run_cuda_graph, **decode_stats_kwargs + ) if ( - getattr(scheduler, "_areal_awex_last_decode_stats_every_iter_ct", None) + has_decode_iter_stats + and callable( + getattr(scheduler, "log_decode_stats_every_iteration", None) + ) + and getattr( + scheduler, "_areal_awex_last_decode_stats_every_iter_ct", None + ) != current_ct ): - scheduler.log_decode_stats_every_iteration( + getattr(scheduler, every_iteration_name)( batch, num_accepted_tokens=getattr(result, "num_accepted_tokens", 0), ) + def _recv_requests(): + if hasattr(scheduler, "recv_requests"): + return scheduler.recv_requests() + return scheduler.request_receiver.recv_requests() + + def _on_idle(): + if hasattr(scheduler, "self_check_during_idle"): + scheduler.self_check_during_idle() + else: + scheduler.on_idle() + + def _is_idle_for_awex_update() -> bool: + is_fully_idle = getattr(scheduler, "is_fully_idle", None) + if callable(is_fully_idle): + try: + return bool(is_fully_idle()) + except TypeError: + return bool(is_fully_idle(for_health_check=False)) + + for attr in ("cur_batch", "last_batch"): + if getattr(scheduler, attr, None) is not None: + return False + + result_queue = getattr(scheduler, "result_queue", None) + if result_queue is not None and len(result_queue) > 0: + return False + + running_batch = getattr(scheduler, "running_batch", None) + if running_batch is not None: + is_empty = getattr(running_batch, "is_empty", None) + if callable(is_empty) and not is_empty(): + return False + + return True + + def _maybe_process_awex_queue_when_idle(loop_count: int) -> None: + if not plugin._process_queue_when_idle: + return + # DEADLOCK WARNING: everything gating the all_reduce inside + # process_awex_queue() MUST be deterministic and identical across + # TP ranks. Loop iterations are lockstep (every iteration goes + # through the recv_requests broadcast), so a loop-count throttle + # is safe. A wall-clock throttle (time.monotonic) is NOT: ranks + # hit the window at different times, some skip the all_reduce + # while others enter it, and the next recv_requests broadcast + # cross-deadlocks against the pending all_reduce (observed as + # TP0 stuck in broadcast_pyobj vs TP1-7 stuck in all_reduce). + if loop_count % plugin._idle_poll_loops != 0: + return + + tp_size = self._int_attr(scheduler, "tp_size", 1) + is_idle = _is_idle_for_awex_update() + if tp_size == 1: + if is_idle and not plugin._weight_queue.empty(): + plugin.process_awex_queue() + return + + # Rank-local idle state is folded into the collective vote instead + # of gating it, so all ranks always enter the all_reduce together. + plugin.process_awex_queue(extra_ready=is_idle) + # Patch event_loop_overlap (the one actually used by SGLang) _orig_overlap = scheduler.event_loop_overlap @@ -362,7 +652,7 @@ def pop_and_process(): ) while True: - recv_reqs = scheduler.recv_requests() + recv_reqs = _recv_requests() if recv_reqs: req_types = [type(r).__name__ for r in recv_reqs] has_control = any( @@ -373,7 +663,7 @@ def pop_and_process(): ) for t in req_types ) - if has_control or _loop_count % 500 == 0: + if has_control: logger.info( f"[AWEX] loop gpu_id={getattr(scheduler, 'gpu_id', '?')}: " f"recv {len(recv_reqs)} reqs, types={req_types[:5]}, " @@ -419,7 +709,9 @@ def pop_and_process(): if not disable_overlap_for_batch: pop_and_process() elif batch is None: - scheduler.self_check_during_idle() + _on_idle() + + _maybe_process_awex_queue_when_idle(_loop_count) if scheduler.is_generation: scheduler.launch_batch_sample_if_needed(batch_result) @@ -435,13 +727,15 @@ def _patched_normal(): logger.info( f"[AWEX] _patched_normal STARTING (gpu_id={getattr(scheduler, 'gpu_id', '?')})", ) + _loop_count = 0 while True: - recv_reqs = scheduler.recv_requests() + recv_reqs = _recv_requests() scheduler.process_input_requests(recv_reqs) if scheduler._engine_paused: plugin.process_awex_queue() time.sleep(plugin._paused_poll_interval_s) continue + _loop_count += 1 batch = scheduler.get_next_batch_to_run() scheduler.cur_batch = batch if batch: @@ -451,7 +745,8 @@ def _patched_normal(): "after_process_batch_result", batch, result ) else: - scheduler.self_check_during_idle() + _on_idle() + _maybe_process_awex_queue_when_idle(_loop_count) scheduler.last_batch = batch scheduler.event_loop_normal = _patched_normal @@ -468,10 +763,12 @@ def _start_background_worker(self, meta_server_addr: str) -> None: daemon=True, ) self._bg_thread.start() - gpu_id = int(getattr(self._scheduler, "gpu_id", -1)) + gpu_id = self._logical_gpu_id() + physical_gpu_id = self._physical_gpu_id() logger.info( f"[AWEX] Started background worker thread " - f"(gpu_id={gpu_id}, meta_server={meta_server_addr})", + f"(gpu_id={gpu_id}, physical_gpu_id={physical_gpu_id}, " + f"meta_server={meta_server_addr})", ) def _background_worker(self, meta_server_addr: str) -> None: @@ -486,9 +783,13 @@ def _background_worker(self, meta_server_addr: str) -> None: """ import torch - gpu_id = int(getattr(self._scheduler, "gpu_id", 0)) + gpu_id = self._logical_gpu_id() + physical_gpu_id = self._physical_gpu_id() torch.cuda.set_device(gpu_id) - logger.info(f"[AWEX] background worker: set CUDA device to {gpu_id}") + logger.info( + f"[AWEX] background worker: set CUDA device to {gpu_id} " + f"(physical_gpu_id={physical_gpu_id})", + ) try: self._init_receiver_from_meta_server(meta_server_addr) @@ -506,21 +807,36 @@ def _background_worker(self, meta_server_addr: str) -> None: _host, _port = meta_server_addr.rsplit(":", 1) _ver_client = _MSC(_host, int(_port)) - from areal.engine.awex.colocate_writer import ( - awex_colocate_timeout_s, - resolve_physical_gpu_id, + _ver_key = _writer_version_key(_get_ip(), physical_gpu_id) + writer_version_poll_timeout_s = max( + 0.1, + get_float_env_var("AWEX_WRITER_VERSION_POLL_TIMEOUT_S", 5.0), ) - - # Shares a key namespace with the training writer, so it needs the - # physical GPU id rather than the mask-relative index. - _ver_key = f"awex_writer_version_{_get_ip()}_{resolve_physical_gpu_id(gpu_id)}" - - version = int( - _ver_client.get_object( + writer_version_log_interval_s = max( + writer_version_poll_timeout_s, + get_float_env_var("AWEX_WRITER_VERSION_LOG_INTERVAL_S", 60.0), + ) + last_writer_wait_log_s = 0.0 + version = None + while version is None: + version = _try_get_writer_version( + _ver_client, _ver_key, - timeout=awex_colocate_timeout_s(), + timeout_s=writer_version_poll_timeout_s, ) - ) + if version is not None: + break + + now = time.monotonic() + if now - last_writer_wait_log_s >= writer_version_log_interval_s: + logger.info( + "[AWEX] background worker: waiting for first writer version " + "key %s; initial RL rollout can run before actor.update_weights", + _ver_key, + ) + last_writer_wait_log_s = now + time.sleep(min(1.0, writer_version_poll_timeout_s)) + logger.info( f"[AWEX] background worker: writer stream starts at v{version}", ) @@ -602,22 +918,23 @@ def _init_receiver_from_meta_server(self, meta_server_addr: str): receiver = self._require_receiver() - # `gpu_id` is node-local. Multi-node colocate needs a globally unique - # transfer rank that stays physically paired with the training process. - gpu_id = int(getattr(self._scheduler, "gpu_id", 0)) + # `physical_gpu_id` is node-local. Multi-node colocate needs a globally + # unique transfer rank that stays physically paired with the training + # process. SGLang may run with logical gpu_id=0 under CUDA_VISIBLE_DEVICES + # isolation, so do not use scheduler.gpu_id for AWEX keys. + gpu_id = self._logical_gpu_id() + physical_gpu_id = self._physical_gpu_id() node_id = int(os.environ.get("SLURM_NODEID", "0")) nnodes = int(os.environ.get("SLURM_NNODES", "1")) logger.info( f"[AWEX] background worker: waiting for awex_train_info " - f"(gpu_id={gpu_id}, node_id={node_id}, nnodes={nnodes})", + f"(gpu_id={gpu_id}, physical_gpu_id={physical_gpu_id}, " + f"node_id={node_id}, nnodes={nnodes})", ) # The driver publishes awex_train_info only after rollout init finishes, # so large models need the same timeout budget as the weight path. - from areal.engine.awex.colocate_writer import ( - awex_colocate_timeout_s, - resolve_physical_gpu_id, - ) + from areal.engine.awex.colocate_writer import awex_colocate_timeout_s train_info = client.get_object( "awex_train_info", @@ -632,23 +949,26 @@ def _init_receiver_from_meta_server(self, meta_server_addr: str): infer_world_size = train_world_size n_gpus_per_node = max(1, infer_world_size // nnodes) - transfer_rank = node_id * n_gpus_per_node + gpu_id + transfer_rank = _resolve_transfer_rank( + infer_world_size=infer_world_size, + gpu_id=gpu_id, + node_id=node_id, + nnodes=nnodes, + instance_world_size=self._instance_world_size(), + ) logger.info( f"[AWEX] background worker: got train_world_size={train_world_size}, " f"infer_world_size={infer_world_size}, n_gpus_per_node={n_gpus_per_node}, " - f"transfer_rank={transfer_rank}", + f"transfer_rank={transfer_rank}, physical_gpu_id={physical_gpu_id}", ) - # transfer_rank stays a relative index so it lines up with the infer - # NCCL world, but the CUDA IPC keys are shared with the training writer - # and must therefore use physical GPU ids. receiver.initialize( meta_server_addr=meta_server_addr, transfer_rank=transfer_rank, infer_world_size=infer_world_size, train_world_size=train_world_size, - local_gpu_id=resolve_physical_gpu_id(gpu_id), + local_gpu_id=physical_gpu_id, ) logger.info( f"[AWEX] background worker: receiver initialized " @@ -672,27 +992,61 @@ def register_awex_plugin() -> None: start method, which doesn't inherit parent-process monkey-patches. """ assert_supported_sglang_version() - from sglang.srt.managers.scheduler import Scheduler _orig_init = Scheduler.__init__ def _patched_init(self, *args, **kwargs): - _orig_init(self, *args, **kwargs) - AwexSchedulerPlugin(self).bind() - _patch_execute_task_in_model_worker(self) + logger.info( + "[AWEX] Scheduler.__init__ entering " + "(pid=%s, CUDA_VISIBLE_DEVICES=%s, AWEX_META_SERVER_ADDR=%s)", + os.getpid(), + os.environ.get("CUDA_VISIBLE_DEVICES", ""), + os.environ.get("AWEX_META_SERVER_ADDR", ""), + ) + try: + _orig_init(self, *args, **kwargs) + except BaseException: + logger.exception("[AWEX] Scheduler.__init__ original init failed") + raise + plugin = AwexSchedulerPlugin(self) + logger.info( + "[AWEX] Scheduler.__init__ original init complete " + "(gpu_id=%s, tp_rank=%s, tp_size=%s)", + getattr(self, "gpu_id", "?"), + getattr(self, "tp_rank", "?"), + plugin._int_attr(self, "tp_size", 1), + ) + try: + plugin.bind() + _patch_execute_task_in_model_worker(self, plugin) + except BaseException: + logger.exception("[AWEX] Scheduler.__init__ AWEX bind failed") + raise + logger.info("[AWEX] Scheduler.__init__ AWEX bind complete") Scheduler.__init__ = _patched_init logger.info("[AWEX] Patched Scheduler.__init__ with awex plugin") -def _patch_execute_task_in_model_worker(scheduler) -> None: +def _patch_execute_task_in_model_worker( + scheduler: Any, plugin: AwexSchedulerPlugin +) -> None: """Add execute_task_in_model_worker to Scheduler (backport from PR #13595).""" - def execute_task_in_model_worker(task_spec: ModelWorkerTask): + if callable(getattr(scheduler, "execute_task_in_model_worker", None)): + logger.info( + "[AWEX] Scheduler already has native execute_task_in_model_worker; " + "skipping legacy backport", + ) + return + + task_cls = _get_model_worker_task_cls() + + def execute_task_in_model_worker(task_spec): model_context = dict( - tp_rank=scheduler.tp_rank, - tp_size=scheduler.tp_size, + tp_rank=plugin._int_attr(scheduler, "tp_rank", 0), + tp_size=plugin._int_attr(scheduler, "tp_size", 1), server_args=scheduler.server_args, scheduler=scheduler, ) @@ -705,23 +1059,36 @@ def execute_task_in_model_worker(task_spec: ModelWorkerTask): scheduler.execute_task_in_model_worker = execute_task_in_model_worker if hasattr(scheduler, "_request_dispatcher"): - scheduler._request_dispatcher._mapping[ModelWorkerTask] = ( - execute_task_in_model_worker - ) + scheduler._request_dispatcher._mapping[task_cls] = execute_task_in_model_worker logger.info("[AWEX] Registered execute_task_in_model_worker in dispatcher") +def _get_model_worker_task_cls(): + try: + from sglang.srt.managers.io_struct import ModelWorkerTask as SGLangTask + + return SGLangTask + except (ImportError, AttributeError): + return ModelWorkerTask + + def awex_run_scheduler_process(*args, **kwargs): """Scheduler process entry point that registers awex plugin. Memory management (pause/resume weights, KV cache, CUDA graphs) is handled - at runtime by AWEX's release_memory/resume_memory, matching the AWEX - reference integration. + at runtime by AWEX's release_memory/resume_memory, matching HybridEngine. No init-time memory patching needed. """ import os meta_addr = os.environ.get("AWEX_META_SERVER_ADDR") + logger.info( + "[AWEX] awex_run_scheduler_process starting " + "(pid=%s, meta_server=%s, CUDA_VISIBLE_DEVICES=%s)", + os.getpid(), + meta_addr or "", + os.environ.get("CUDA_VISIBLE_DEVICES", ""), + ) if meta_addr: register_awex_plugin() else: @@ -730,20 +1097,43 @@ def awex_run_scheduler_process(*args, **kwargs): ) from sglang.srt.managers.scheduler import run_scheduler_process - return run_scheduler_process(*args, **kwargs) + try: + return run_scheduler_process(*args, **kwargs) + except BaseException: + logger.exception("[AWEX] run_scheduler_process failed") + raise if __name__ == "__main__": import os import sys - logger.info("[AWEX] awex_sglang_plugin __main__ starting") + logger.info( + "[AWEX] awex_sglang_plugin __main__ starting " + "(pid=%s, CUDA_VISIBLE_DEVICES=%s, AWEX_META_SERVER_ADDR=%s)", + os.getpid(), + os.environ.get("CUDA_VISIBLE_DEVICES", ""), + os.environ.get("AWEX_META_SERVER_ADDR", ""), + ) from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.server_args import prepare_server_args from sglang.srt.utils import kill_process_tree + _load_sglang_plugins_if_available() server_args = prepare_server_args(sys.argv[1:]) + # In AWEX colocated mode the scheduler may not be able to serve requests + # until the first weight sync (actor holds GPUs / writer version not yet + # published). SGLang's warmup request would then hang for 600s + # (_execute_server_warmup default timeout) and kill_process_tree() the + # whole server. Skip the warmup: _wait_and_warmup marks the server Up + # directly when skip_server_warmup is set. + if not server_args.skip_server_warmup: + logger.info( + "[AWEX] forcing skip_server_warmup=True to avoid warmup-timeout " + "suicide before the first weight sync" + ) + server_args.skip_server_warmup = True try: launch_server( server_args, diff --git a/areal/engine/fsdp_engine.py b/areal/engine/fsdp_engine.py index 1767192094..8549aefba3 100644 --- a/areal/engine/fsdp_engine.py +++ b/areal/engine/fsdp_engine.py @@ -2264,6 +2264,9 @@ def __init__(self, config: PPOActorConfig): super().__init__(config) self.actor = PPOActor(config, self) + def configure_mopd_loss(self, config) -> None: + self.actor.configure_mopd_loss(config) + @torch.no_grad() def compute_logp(self, *args, **kwargs) -> list[torch.Tensor] | None: return self.actor.compute_logp(*args, **kwargs) @@ -2272,6 +2275,9 @@ def compute_logp(self, *args, **kwargs) -> list[torch.Tensor] | None: def compute_advantages(self, *args, **kwargs) -> list[dict[str, Any]]: return self.actor.compute_advantages(*args, **kwargs) + def prepare_mopd_batch(self, *args, **kwargs) -> list[dict[str, Any]]: + return self.actor.prepare_mopd_batch(*args, **kwargs) + def ppo_update(self, *args, **kwargs) -> None: self.actor.ppo_update(*args, **kwargs) diff --git a/areal/engine/megatron_engine.py b/areal/engine/megatron_engine.py index 2d7cc67bfc..aa147ca924 100644 --- a/areal/engine/megatron_engine.py +++ b/areal/engine/megatron_engine.py @@ -126,6 +126,7 @@ MicroBatchItem, MicroBatchList, amend_position_ids, + batched_call, broadcast_tensor, concat_batch, pack_tensor_dict, @@ -151,7 +152,14 @@ if TYPE_CHECKING: from areal.api import Scheduler - from areal.api.cli_args import DPOEngineConfig, PPOActorConfig, PPOCriticConfig + from areal.api.cli_args import ( + DPOEngineConfig, + MOPDTeacherEngineConfig, + PPOActorConfig, + PPOCriticConfig, + ) + from areal.engine.awex.colocate_writer import AwexWeightPublisher + from areal.engine.megatron_utils.weight_residency import MegatronWeightResidency # `model.named_modules()` yields LOCAL layer indices on each PP rank, while @@ -367,7 +375,8 @@ def __init__(self, config: TrainEngineConfig): self.own_global_group: bool = False self.is_offload: bool = False self._offload_depth: int = 0 - self._awex_adapter = None # AwexMegatronAdapter for colocate mode + self._weight_residency: MegatronWeightResidency | None = None + self._awex_publisher: AwexWeightPublisher | None = None self._dte_runtime_config = DTERuntimeConfig.from_env() self._warned_unbounded_microbatch = False self.enable_tree_training: bool = self.config.enable_tree_training @@ -969,11 +978,8 @@ def connect_engine(self, engine: InferenceEngine, meta: WeightUpdateMeta): self._init_weight_update_from_distributed(meta) self.weight_update_group_initialized = True elif meta.type == "awex": - from areal.engine.awex.colocate_writer import AwexMegatronAdapter - - if self._awex_adapter is None: - self._awex_adapter = AwexMegatronAdapter(self) - self._awex_adapter.init_colocate_weight_update( + publisher = self._ensure_awex_publisher() + publisher.init_colocate_weight_update( meta_server_addr=meta.nccl_master_address, pair_name=meta.nccl_group_name or "default", transfer_rank=self.rank or 0, @@ -1033,23 +1039,27 @@ def update_weights(self, meta: WeightUpdateMeta): # weights → signal offloaded → IPC serialize → wait reader done → # cleanup shared → signal write_finished # 2. finish: wait all infer engines done → cleanup MetaServer keys - # 3. resume kv_cache + continue generation - self._awex_adapter.execute_colocate_weight_update(meta.version or 0) - # Do NOT flip is_offload here: the AWEX adapter tracks released - # memory via _released_tags, and the trainer onloads explicitly + # Restoring the rollout must happen in the controller after every + # actor worker returns. Calling back into the rollout from this RPC + # creates a nested controller/rollout call while the actor collective + # is still active and deadlocks at the final barrier. + if self._awex_publisher is None: + raise RuntimeError( + "AWEX weight update requested before publisher initialization" + ) + self._awex_publisher.execute_colocate_weight_update(meta.version or 0) + # Do NOT flip is_offload here: residency tracks released memory, + # and the trainer onloads explicitly # at the next train phase. Marking is_offload would make every # _offload_aware_context RPC (e.g. export_stats) reload optimizer # states onto a GPU already fully occupied by the resumed rollout. dist.barrier(group=self.cpu_group) - self._awex_adapter.finish_colocate_weight_update( + self._awex_publisher.finish_colocate_weight_update( training_world_size=dist.get_world_size(self.cpu_group) ) - if dist.get_rank() == 0: - self.rollout_engine.onload(tags=["kv_cache"]) - self.rollout_engine.continue_generation() dist.barrier(group=self.cpu_group) return with self._offload_aware_context(): @@ -1068,12 +1078,12 @@ def get_version(self) -> int: return self._version def save(self, meta: SaveLoadMeta): - if self._awex_adapter is not None: + if self._weight_residency is not None: # Post-ppo_update the fp32 grad buffers (~2x param bytes) are # dead weight until the next train_batch (which rebuilds them via # ensure_grad_buffers); drop them here to fund the HF saver's TP # coalesced all-gather transient. - self._awex_adapter._release_grad_memory() + self._weight_residency.release_grad_memory() gc.collect() torch.cuda.empty_cache() with self._offload_aware_context(): @@ -1222,12 +1232,9 @@ def forward_step(batch_iter, model): tree_attn_keys = list(tree_kwargs.keys()) cp_size = mpu.get_context_parallel_world_size() - # forward_batch (compute_logp / compute_values) passes - # gather_cp_output=True so CP-local outputs are gathered back to the - # full sequence length inside forward. This matches downstream - # labels / output_seqlens and fixes split_with_sizes mismatches when - # compute_logp runs with CP > 1. Train/eval keeps the default False - # value, so the CP-local loss path (_cp_local_labels) is unchanged. + # CP-local forward keeps the vocabulary logits sharded by sequence. + # Consumers reconstruct token scalars only; gathering logits here + # creates the full-vocabulary CP memory spike MOPD must avoid. cp_local = cp_size > 1 and not gather_cp_output model_vp_stage = getattr(model, "vp_stage", 0) @@ -1408,9 +1415,8 @@ def train_batch( loss_weight_fn: Callable[[dict[str, Any]], torch.Tensor], ) -> dict[str, float]: self._ensure_ready() - if self._awex_adapter is not None: - self._awex_adapter.ensure_grad_buffers() - + if self._weight_residency is not None: + self._weight_residency.ensure_grad_buffers() self.optimizer_zero_grad() input_batched, _ = self._normalize_batch_input(input_) @@ -1553,7 +1559,7 @@ def process_output(output: torch.Tensor, inputs: dict[str, Any]) -> None: return None self.forward_backward_batch( - mb_list, process_output, forward_only=True, gather_cp_output=True + mb_list, process_output, forward_only=True, gather_cp_output=False ) # Step 4: Aggregate, reorder, and broadcast outputs @@ -1595,53 +1601,70 @@ def export_stats(self) -> dict[str, float]: return data def init_awex_adapter(self, meta_server_addr: str | None = None) -> None: - """Create awex adapter early for selective memory management. + """Create the AWEX publisher early for colocated weight transfer. Must be called before offload() in colocate mode so that offload uses - the adapter's tag-based mechanism instead of TMS (which is all-or-nothing - and causes OOM on resume when SGLang occupies GPU memory). + flat-buffer residency instead of TMS, which is all-or-nothing and can + OOM when SGLang already occupies the GPU. """ - if self._awex_adapter is None: - from areal.engine.awex.colocate_writer import AwexMegatronAdapter - - self._awex_adapter = AwexMegatronAdapter(self) - self.logger.info("Created AWEX adapter for memory management") + publisher = self._ensure_awex_publisher() + publisher.eager_publish_train_info(meta_server_addr) - self._eager_publish_awex_train_info(meta_server_addr) - - def _eager_publish_awex_train_info(self, meta_server_addr: str | None) -> None: - addr = meta_server_addr or os.environ.get("AWEX_META_SERVER_ADDR", "") - if not addr or (dist.is_initialized() and dist.get_rank() != 0): - return - try: - from awex.meta.meta_server import MetaServerClient - - host, port = addr.rsplit(":", 1) - client = MetaServerClient(host, int(port)) - world = dist.get_world_size() if dist.is_initialized() else 1 - client.put_object("awex_train_info", {"train_world_size": world}) - self.logger.info( - "[AWEX] eager-published awex_train_info (train_world_size=%d) to %s", - world, - addr, + def _ensure_weight_residency(self) -> MegatronWeightResidency: + if self._weight_residency is None: + from areal.engine.megatron_utils.weight_residency import ( + MegatronWeightResidency, ) - except Exception as e: - self.logger.warning("[AWEX] eager publish awex_train_info failed: %s", e) + + self._weight_residency = MegatronWeightResidency(self) + self.logger.info("Created Megatron weight residency manager") + return self._weight_residency + + def _ensure_awex_publisher(self) -> AwexWeightPublisher: + residency = self._ensure_weight_residency() + if self._awex_publisher is None: + from areal.engine.awex.colocate_writer import AwexWeightPublisher + + self._awex_publisher = AwexWeightPublisher(self, residency) + self.logger.info("Created AWEX weight publisher") + elif self._awex_publisher.residency is not residency: + raise RuntimeError("AWEX publisher does not own the engine residency") + return self._awex_publisher + + def init_weight_residency_adapter(self) -> None: + """Enable DDP-flat-buffer residency without AWEX publication state.""" + self._ensure_weight_residency() + + def _log_weight_residency_stats(self, phase: str) -> None: + """Log per-rank CUDA residency for persistent/AWEX flat buffers.""" + stats = self.get_device_stats() + rank = dist.get_rank(self.cpu_group) + self.logger.info( + "[Megatron residency] rank=%d phase=%s allocated_gb=%.3f " + "reserved_gb=%.3f allocator_conf=%r", + rank, + phase, + stats.mem_allocated, + stats.mem_reserved, + os.environ.get("PYTORCH_CUDA_ALLOC_CONF", ""), + ) def offload(self) -> None: """Offload model memory to CPU. - In colocate mode (awex adapter active): manual tag-based offload via - storage resize. Otherwise: torch_memory_saver pause. + With explicit Megatron residency: manual tag-based flat-buffer offload. + Otherwise: torch_memory_saver pause. Ref: https://github.com/THUDM/slime/blob/main/slime/backends/megatron_utils/actor.py """ - if self._awex_adapter is not None: + if self._weight_residency is not None: + self._log_weight_residency_stats("before_offload") self.get_device_stats().log("before offload model") current_platform.clear_memory() - self._awex_adapter.release_memory(tags=["optimizer", "weights"]) + self._weight_residency.release_memory(tags=["optimizer", "weights"]) current_platform.synchronize() dist.barrier(group=self.cpu_group) + self._log_weight_residency_stats("after_offload") self.get_device_stats().log("after offload model") self.is_offload = True return @@ -1661,7 +1684,8 @@ def offload(self) -> None: # memory left for TMS to back up. if self.mcore_config.disable_grad_buffers_cpu_backup: for m in self.model: - m.offload_grad_buffers(synchronize=False, empty_cache=False) + if isinstance(m, DDP): + m.offload_grad_buffers(synchronize=False, empty_cache=False) current_platform.clear_memory() torch_memory_saver.pause() @@ -1676,16 +1700,16 @@ def offload(self) -> None: def onload(self) -> None: """Onload model memory from CPU back to GPU. - Uses awex adapter (selective, tag-based) when available, otherwise - torch_memory_saver. + Uses explicit Megatron residency when available, otherwise TMS. Ref: https://github.com/THUDM/slime/blob/main/slime/backends/megatron_utils/actor.py """ - if self._awex_adapter is not None: - self._awex_adapter.resume_memory(tags=["optimizer", "weights"]) + if self._weight_residency is not None: + self._weight_residency.resume_memory(tags=["optimizer", "weights"]) current_platform.clear_memory() current_platform.synchronize() dist.barrier(group=self.cpu_group) + self._log_weight_residency_stats("after_onload") self.get_device_stats().log("after onload model") self.is_offload = False return @@ -1696,7 +1720,8 @@ def onload(self) -> None: # storage and zeroes it; param.main_grad views become valid again. if self.mcore_config.disable_grad_buffers_cpu_backup: for m in self.model: - m.restore_grad_buffers(synchronize=False) + if isinstance(m, DDP): + m.restore_grad_buffers(synchronize=False) current_platform.clear_memory() @@ -1707,7 +1732,7 @@ def onload(self) -> None: self.is_offload = False - def clear_batches(self, shard_ids: list[str] | None = None) -> None: + def clear_batches(self, shard_ids: list[str] | None = None) -> int: """Drain this worker's client-side RTensor fetch buffer. Called via RPC by ``TrainController.clear_batches`` at step end so @@ -1718,13 +1743,19 @@ def clear_batches(self, shard_ids: list[str] | None = None) -> None: """ from areal.infra.rpc.rtensor import clear_fetch_buffer - if shard_ids: - clear_fetch_buffer(shard_ids) + if not shard_ids: + return 0 + return clear_fetch_buffer(shard_ids) - def fetch_buffer_stats(self) -> dict[str, int]: + def fetch_buffer_stats(self, shard_ids: list[str] | None = None) -> dict[str, int]: """Expose local fetch-buffer stats for post-step drain verification.""" - from areal.infra.rpc.rtensor import fetch_buffer_stats + from areal.infra.rpc.rtensor import ( + fetch_buffer_matching_stats, + fetch_buffer_stats, + ) + if shard_ids is not None: + return fetch_buffer_matching_stats(shard_ids) return fetch_buffer_stats() def _normalize_adam_bf16_config(self) -> None: @@ -3074,7 +3105,9 @@ def _compute_forward_result( chunk_size=self.config.logprobs_chunk_size, ) return logprobs - labels = torch.roll(inputs["input_ids"], shifts=-1, dims=-1) + labels = inputs.get("_cp_local_labels") + if labels is None: + labels = torch.roll(inputs["input_ids"], shifts=-1, dims=-1) logprobs = gather_logprobs( output, labels, @@ -3084,10 +3117,36 @@ def _compute_forward_result( else None, chunk_size=self.config.logprobs_chunk_size, ) - return logprobs + return self._reassemble_cp_forward_scalars(logprobs, inputs) else: values = output.squeeze(-1) - return values + return self._reassemble_cp_forward_scalars(values, inputs) + + @staticmethod + def _reassemble_cp_forward_scalars( + local_values: torch.Tensor, inputs: dict[str, Any] + ) -> torch.Tensor: + """Reassemble CP token scalars without ever gathering vocabulary logits.""" + padded_cu_seqlens = inputs.get("_cp_padded_cu_seqlens") + if padded_cu_seqlens is None: + return local_values + values = reassemble_cp_packed_logprobs(local_values, padded_cu_seqlens) + return unpad_logits( + values, + inputs.get("_cp_padding_length", 0), + padded_cu_seqlens, + inputs.get("_cp_old_cu_seqlens"), + ) + + def assert_mopd_runtime_topology(self) -> None: + """Verify that MOPD scoring uses the configured MCore pipeline size.""" + configured_pp_size = self.parallel_strategy.pipeline_parallel_size + runtime_pp_size = mpu.get_pipeline_model_parallel_world_size() + if runtime_pp_size != configured_pp_size: + raise RuntimeError( + "MOPD compute_logp pipeline topology mismatch: " + f"configured PP={configured_pp_size}, runtime PP={runtime_pp_size}" + ) # ============================================================================= @@ -3095,6 +3154,38 @@ def _compute_forward_result( # ============================================================================= +class MegatronScoringEngine(MegatronEngine): + """Forward-only Megatron engine used by persistent MOPD teachers.""" + + def __init__(self, config: MOPDTeacherEngineConfig): + super().__init__(config) + + @torch.no_grad() + def compute_logp(self, data: list[dict[str, Any]]) -> list[torch.Tensor] | None: + return batched_call(self._compute_logp, data) + + def _compute_logp(self, data: dict[str, Any]) -> torch.Tensor | None: + self.eval() + return self.forward( + input_=data, + aggregate_fn=lambda xs: torch.cat(xs, dim=-1), + ) + + @classmethod + def as_controller( + cls, + config: MOPDTeacherEngineConfig, + scheduler: Scheduler, + ): + from areal.trainer.mopd.scoring import MOPDTeacherController + + return MOPDTeacherController( + train_engine=cls, + config=config, + scheduler=scheduler, + ) + + class MegatronPPOActor(MegatronEngine): """PPO Actor implementation using Megatron backend.""" @@ -3104,6 +3195,18 @@ def __init__(self, config: PPOActorConfig): super().__init__(config) self.actor = PPOActor(config, self) + def initialize( + self, + addr: str | None, + ft_spec: FinetuneSpec, + *args, + **kwargs, + ) -> None: + super().initialize(addr, ft_spec, *args, **kwargs) + + def configure_mopd_loss(self, config) -> None: + self.actor.configure_mopd_loss(config) + @torch.no_grad() def compute_logp(self, *args, **kwargs) -> list[torch.Tensor] | None: return self.actor.compute_logp(*args, **kwargs) @@ -3112,6 +3215,12 @@ def compute_logp(self, *args, **kwargs) -> list[torch.Tensor] | None: def compute_advantages(self, *args, **kwargs) -> list[dict[str, Any]]: return self.actor.compute_advantages(*args, **kwargs) + def prepare_mopd_batch(self, *args, **kwargs) -> list[dict[str, Any]]: + return self.actor.prepare_mopd_batch(*args, **kwargs) + + def aggregate_mopd_targets(self, *args, **kwargs): + return self.actor.aggregate_mopd_targets(*args, **kwargs) + def ppo_update(self, *args, **kwargs) -> None: self.actor.ppo_update(*args, **kwargs) diff --git a/areal/engine/megatron_utils/weight_residency.py b/areal/engine/megatron_utils/weight_residency.py new file mode 100644 index 0000000000..e7f461181e --- /dev/null +++ b/areal/engine/megatron_utils/weight_residency.py @@ -0,0 +1,311 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Megatron flat-buffer GPU residency management. + +This module deliberately has no AWEX transport or publication state. It owns +the single source of truth for model-weight, optimizer-state, and gradient +buffer residency used by both persistent scoring workers and AWEX publishers. +""" + +from __future__ import annotations + +import gc +import os +from typing import TYPE_CHECKING, Any + +import torch + +from areal.utils.logging import getLogger + +if TYPE_CHECKING: + from areal.engine.megatron_engine import MegatronEngine + + +logger = getLogger("MegatronResidency") + + +class MegatronWeightResidency: + """Own CPU/GPU residency for MCore DDP flat buffers and optimizer state.""" + + def __init__(self, engine: MegatronEngine) -> None: + self._engine = engine + self._released_tags: set[str] = set() + + @property + def released_tags(self) -> frozenset[str]: + """Return an immutable snapshot of currently offloaded state tags.""" + return frozenset(self._released_tags) + + def is_released(self, tag: str) -> bool: + """Return whether one residency tag is currently offloaded.""" + return tag in self._released_tags + + def release_memory(self, tags: list[str] | None = None) -> None: + """Offload the requested state classes to CPU exactly once.""" + tags = tags or ["optimizer", "weights"] + tags_to_release = [tag for tag in tags if tag not in self._released_tags] + if not tags_to_release: + return + + if "optimizer" in tags_to_release: + self._offload_optimizer_states() + self._released_tags.add("optimizer") + + if "weights" in tags_to_release: + self._offload_model_weights() + self._released_tags.add("weights") + + torch.cuda.synchronize() + gc.collect() + torch.cuda.empty_cache() + logger.info("release_memory done: tags=%s", tags_to_release) + + def resume_memory(self, tags: list[str] | None = None) -> None: + """Restore the requested state classes to GPU exactly once.""" + tags = tags or ["optimizer", "weights"] + tags_to_resume = [tag for tag in tags if tag in self._released_tags] + if not tags_to_resume: + return + + if "weights" in tags_to_resume: + self._reload_model_weights(load_grad=False) + self._released_tags.discard("weights") + + if "optimizer" in tags_to_resume: + self._reload_optimizer_states() + self._released_tags.discard("optimizer") + + torch.cuda.synchronize() + logger.info("resume_memory done: tags=%s", tags_to_resume) + + def release_grad_memory(self) -> None: + """Release gradient buffers while retaining sizes for training restore.""" + from megatron.core.distributed import DistributedDataParallel as DDP + + model = self._engine.model + if model is None: + return + if not isinstance(model, (list, tuple)): + model = [model] + count = 0 + for chunk in model: + if isinstance(chunk, DDP): + for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: + for buf in buffers: + if buf.grad_data.storage().size() > 0: + buf.grad_data_size = buf.grad_data.storage().size() + buf.grad_data.storage().resize_(0) + count += 1 + if count > 0: + torch.cuda.synchronize() + gc.collect() + torch.cuda.empty_cache() + logger.info("Released %d grad buffers", count) + + def ensure_grad_buffers(self) -> None: + """Reallocate discarded gradient buffers before training.""" + from megatron.core.distributed import DistributedDataParallel as DDP + + model = self._engine.model + if model is None: + return + if not isinstance(model, (list, tuple)): + model = [model] + count = 0 + for chunk in model: + if isinstance(chunk, DDP): + for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: + for buf in buffers: + if ( + hasattr(buf, "grad_data_size") + and buf.grad_data.storage().size() == 0 + ): + buf.grad_data.storage().resize_(buf.grad_data_size) + buf.grad_data.zero_() + count += 1 + if count > 0: + torch.cuda.synchronize() + logger.info("Allocated %d grad buffers for training", count) + + def _offload_model_weights(self) -> None: + from megatron.core.distributed import DistributedDataParallel as DDP + + model = self._engine.model + if model is None: + return + if not isinstance(model, (list, tuple)): + model = [model] + count = 0 + for chunk in model: + if isinstance(chunk, DDP): + for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: + for buf in buffers: + if hasattr(buf, "offload_to_cpu"): + buf.offload_to_cpu() + count += 1 + continue + if buf.param_data.storage().size() > 0: + if not hasattr(buf, "cpu_param_data"): + buf.cpu_param_data = torch.zeros( + buf.param_data.data.shape, + dtype=buf.param_data.data.dtype, + pin_memory=True, + device="cpu", + ) + buf.cpu_param_data.copy_(buf.param_data.data) + buf.param_data_size = buf.param_data.storage().size() + buf.param_data.storage().resize_(0) + count += 1 + if buf.grad_data.storage().size() > 0: + buf.grad_data_size = buf.grad_data.storage().size() + buf.grad_data.storage().resize_(0) + else: + raise RuntimeError( + "Megatron flat-buffer residency requires MCore DDP; " + "per-parameter weight offload is forbidden" + ) + torch.cuda.synchronize() + logger.info("Offloaded %d weight buffers to CPU", count) + + def _reload_model_weights(self, load_grad: bool = False) -> None: + from megatron.core.distributed import DistributedDataParallel as DDP + + model = self._engine.model + if model is None: + return + if not isinstance(model, (list, tuple)): + model = [model] + for chunk in model: + if isinstance(chunk, DDP): + for buffers in [chunk.buffers, chunk.expert_parallel_buffers]: + for buf in buffers: + if hasattr(buf, "reload_from_cpu"): + buf.reload_from_cpu(move_grads=load_grad) + continue + if buf.param_data.storage().size() == 0: + buf.param_data.storage().resize_(buf.param_data_size) + buf.param_data.copy_(buf.cpu_param_data, non_blocking=True) + if ( + load_grad + and hasattr(buf, "grad_data_size") + and buf.grad_data.storage().size() == 0 + ): + buf.grad_data.storage().resize_(buf.grad_data_size) + buf.grad_data.zero_() + else: + raise RuntimeError( + "Cannot reload Megatron weights without MCore DDP flat buffers" + ) + torch.cuda.synchronize() + logger.info("Reloaded model weights to GPU (load_grad=%s)", load_grad) + + def _get_inner_optimizers(self) -> list[Any]: + optimizer = self._engine.optimizer + if optimizer is None: + return [] + if hasattr(optimizer, "chained_optimizers"): + return optimizer.chained_optimizers + if hasattr(optimizer, "optimizers"): + return optimizer.optimizers + return [optimizer] + + def _offload_optimizer_states(self) -> None: + optimizer = self._engine.optimizer + if optimizer is None: + return + if os.environ.get("AWEX_OPT_OFFLOAD_VIA_HDO", "").strip() == "1" and hasattr( + optimizer, "offload_to_cpu" + ): + optimizer.offload_to_cpu() + logger.info("Offloaded optimizer via offload_to_cpu()") + return + + inner_optimizers = self._get_inner_optimizers() + if not inner_optimizers: + return + + count = 0 + for opt in inner_optimizers: + if hasattr(opt, "shard_fp32_from_float16_groups"): + for group in opt.shard_fp32_from_float16_groups: + if isinstance(group, list): + for tensor in group: + if tensor is not None and tensor.data.is_cuda: + tensor.data = tensor.data.to("cpu", non_blocking=True) + count += 1 + elif group is not None and group.data.is_cuda: + group.data = group.data.to("cpu", non_blocking=True) + count += 1 + + base_opt = getattr(opt, "optimizer", opt) + if not hasattr(base_opt, "state") or base_opt.state is None: + continue + for state in base_opt.state.values(): + for key in ("exp_avg", "exp_avg_sq"): + if ( + key in state + and isinstance(state[key], torch.Tensor) + and state[key].is_cuda + ): + state[key] = state[key].to("cpu", non_blocking=True) + count += 1 + + try: + from transformer_engine.pytorch.module.base import _dummy_wgrads + + purged = len(_dummy_wgrads) + for key in list(_dummy_wgrads): + del _dummy_wgrads[key] + if purged: + logger.info("Purged %d TE _dummy_wgrads cache entries", purged) + except ImportError: + pass + torch.cuda.synchronize() + logger.info("Offloaded %d optimizer state tensors to CPU", count) + + def _reload_optimizer_states(self) -> None: + optimizer = self._engine.optimizer + if optimizer is None: + return + if os.environ.get("AWEX_OPT_OFFLOAD_VIA_HDO", "").strip() == "1" and hasattr( + optimizer, "restore_from_cpu" + ): + optimizer.restore_from_cpu() + logger.info("Reloaded optimizer via restore_from_cpu()") + return + + inner_optimizers = self._get_inner_optimizers() + if not inner_optimizers: + return + + device = self._engine.device + count = 0 + for opt in inner_optimizers: + if hasattr(opt, "shard_fp32_from_float16_groups"): + for group in opt.shard_fp32_from_float16_groups: + if isinstance(group, list): + for tensor in group: + if tensor is not None and not tensor.data.is_cuda: + tensor.data = tensor.data.to(device, non_blocking=True) + count += 1 + elif group is not None and not group.data.is_cuda: + group.data = group.data.to(device, non_blocking=True) + count += 1 + + base_opt = getattr(opt, "optimizer", opt) + if not hasattr(base_opt, "state") or base_opt.state is None: + continue + for state in base_opt.state.values(): + for key in ("exp_avg", "exp_avg_sq"): + if ( + key in state + and isinstance(state[key], torch.Tensor) + and not state[key].is_cuda + ): + state[key] = state[key].to(device, non_blocking=True) + count += 1 + torch.cuda.synchronize() + logger.info("Reloaded %d optimizer state tensors to GPU", count) + + +__all__ = ["MegatronWeightResidency"] diff --git a/areal/engine/sglang_remote.py b/areal/engine/sglang_remote.py index 97891f3969..9c2919de8b 100644 --- a/areal/engine/sglang_remote.py +++ b/areal/engine/sglang_remote.py @@ -44,6 +44,9 @@ class SGLangBackend: """SGLang-specific backend implementation for remote inference.""" + def __init__(self) -> None: + self._readiness_endpoint = "/health" + @staticmethod def build_server_env(env: Mapping[str, str]) -> dict[str, str]: _env = dict(env) @@ -372,8 +375,8 @@ def get_abort_all_request(self) -> HttpRequest: return HttpRequest(endpoint="/abort_request", payload={"abort_all": True}) def get_health_check_request(self) -> HttpRequest: - """Get SGLang health check request.""" - return HttpRequest(endpoint="/health", payload={}, method="GET") + """Get SGLang readiness check request.""" + return HttpRequest(endpoint=self._readiness_endpoint, payload={}, method="GET") def get_offload_request(self, tags: list[str] | None = None) -> HttpRequest: """Get SGLang offload request.""" @@ -386,7 +389,7 @@ def get_onload_request(self, tags: list[str] | None = None) -> HttpRequest: Parameters: ---------- tags: list[str], optional - Available tags for multi-stage resume: weights, kv_cache + Available tags for multi-stage resume: weights, kv_cache, cuda_graph """ payload = {"tags": tags} if tags is not None else {} return HttpRequest(endpoint="/resume_memory_occupation", payload=payload) @@ -402,8 +405,13 @@ def launch_server(self, server_args: dict[str, Any]) -> subprocess.Popen: "input IDs together with image data may be processed again by the " "server. Set skip_tokenizer_init=True for VLM rollout recipes." ) - awex_meta_addr = server_args.pop("awex_meta_server_addr", None) + awex_meta_addr = server_args.pop( + "awex_meta_server_addr", None + ) or os.environ.get("AWEX_META_SERVER_ADDR") awex_colocate = server_args.pop("awex_colocate_mode", False) + self._readiness_endpoint = ( + "/model_info" if awex_colocate or awex_meta_addr else "/health" + ) # Colocate placement: derive base_gpu_id from SLURM_LOCALID so two SGLang # servers sharing a node never claim the same GPU range. The controller # cannot do this reliably because its global rank -> node-slot mapping is @@ -428,9 +436,40 @@ def launch_server(self, server_args: dict[str, Any]) -> subprocess.Popen: ) cmd = SGLangConfig.build_cmd_from_args(server_args) _env = self.build_server_env(os.environ) + _env.setdefault("PYTHONFAULTHANDLER", "1") + + awex_graph_memory_saver = bool( + (awex_colocate or awex_meta_addr) + and server_args.get("enable_memory_saver", False) + ) + if awex_graph_memory_saver: + # SGLang 0.5.10 gates CUDA-graph region registration separately + # from --enable-memory-saver. This must be set before graph capture. + _env.setdefault("SGLANG_MEMORY_SAVER_CUDA_GRAPH", "1") + + drop_ld_preload_default = "1" if (awex_colocate or awex_meta_addr) else "0" + if _env.get( + "AREAL_SGLANG_DROP_LD_PRELOAD", drop_ld_preload_default + ).strip().lower() in { + "1", + "true", + "yes", + "y", + "on", + }: + dropped = _env.pop("LD_PRELOAD", "") + dropped_tms = { + key: _env.pop(key) + for key in ("TMS_INIT_ENABLE", "TMS_INIT_ENABLE_CPU_BACKUP") + if key in _env + } + logger.info( + "Dropping training TMS environment for SGLang child: " + "LD_PRELOAD=%s, TMS=%s", + dropped, + dropped_tms, + ) - if not awex_meta_addr: - awex_meta_addr = os.environ.get("AWEX_META_SERVER_ADDR") if awex_colocate or awex_meta_addr: sglang_entrypoints = ( "sglang.launch_server", diff --git a/areal/experimental/engine/archon_engine.py b/areal/experimental/engine/archon_engine.py index a0402d4f48..c80f05912d 100644 --- a/areal/experimental/engine/archon_engine.py +++ b/areal/experimental/engine/archon_engine.py @@ -1422,6 +1422,9 @@ def __init__(self, config): super().__init__(config) self.actor = PPOActor(config, self) + def configure_mopd_loss(self, config) -> None: + self.actor.configure_mopd_loss(config) + @torch.no_grad() def compute_logp(self, *args, **kwargs) -> list[torch.Tensor] | None: return self.actor.compute_logp(*args, **kwargs) @@ -1430,6 +1433,9 @@ def compute_logp(self, *args, **kwargs) -> list[torch.Tensor] | None: def compute_advantages(self, *args, **kwargs) -> list[dict[str, Any]]: return self.actor.compute_advantages(*args, **kwargs) + def prepare_mopd_batch(self, *args, **kwargs) -> list[dict[str, Any]]: + return self.actor.prepare_mopd_batch(*args, **kwargs) + def ppo_update(self, *args, **kwargs) -> None: self.actor.ppo_update(*args, **kwargs) diff --git a/areal/experimental/openai/client.py b/areal/experimental/openai/client.py index 21aeaf77f0..6a5e05c26f 100644 --- a/areal/experimental/openai/client.py +++ b/areal/experimental/openai/client.py @@ -852,6 +852,7 @@ def _build_chat_completion( self, completion_id: str, current_time: int, + model: str, output_text: str, tool_calls: list | None, response: ModelResponse, @@ -885,7 +886,7 @@ def _build_chat_completion( ) ], created=current_time, - model="None", + model=model, object="chat.completion", service_tier=None, system_fingerprint=None, @@ -903,6 +904,7 @@ async def create( *, messages: Iterable[ChatCompletionMessageParam], stream: Literal[True], + model: str | NotGiven = NOT_GIVEN, frequency_penalty: float | None | NotGiven = NOT_GIVEN, max_completion_tokens: int | None | NotGiven = NOT_GIVEN, max_tokens: int | None | NotGiven = NOT_GIVEN, @@ -926,6 +928,7 @@ async def create( self, *, messages: Iterable[ChatCompletionMessageParam], + model: str | NotGiven = NOT_GIVEN, stream: Literal[False] | NotGiven = NOT_GIVEN, frequency_penalty: float | None | NotGiven = NOT_GIVEN, max_completion_tokens: int | None | NotGiven = NOT_GIVEN, @@ -949,6 +952,7 @@ async def create( self, *, messages: Iterable[ChatCompletionMessageParam], + model: str | NotGiven = NOT_GIVEN, stream: bool | NotGiven = NOT_GIVEN, frequency_penalty: float | None | NotGiven = NOT_GIVEN, max_completion_tokens: int | None | NotGiven = NOT_GIVEN, @@ -970,6 +974,7 @@ async def create( """Override create method to use AReaL engine and cache responses.""" is_streaming = not is_omitted(stream) and stream is True + response_model = "default" if is_omitted(model) else str(model) # Extract and validate supported parameters cache, interaction = None, None @@ -1196,6 +1201,7 @@ async def create( chat_completion, output_message = self._build_chat_completion( completion_id=completion_id, current_time=current_time, + model=response_model, output_text=output_text, tool_calls=tool_calls, response=response, @@ -1208,6 +1214,7 @@ async def create( return self._create_stream( completion_id=completion_id, current_time=current_time, + model=response_model, output_text=output_text, tool_calls=tool_calls, response=response, @@ -1217,6 +1224,7 @@ async def create( chat_completion, output_message = self._build_chat_completion( completion_id=completion_id, current_time=current_time, + model=response_model, output_text=output_text, tool_calls=tool_calls, response=response, @@ -1234,6 +1242,7 @@ async def _create_stream( self, completion_id: str, current_time: int, + model: str, output_text: str, tool_calls: list | None, response: ModelResponse, @@ -1259,7 +1268,7 @@ async def _create_stream( ) ], created=current_time, - model="None", + model=model, object="chat.completion.chunk", ) @@ -1276,7 +1285,7 @@ async def _create_stream( ) ], created=current_time, - model="None", + model=model, object="chat.completion.chunk", ) @@ -1314,7 +1323,7 @@ async def _create_stream( ) ], created=current_time, - model="None", + model=model, object="chat.completion.chunk", ) # Chunk 2: arguments only, emitted as input_json_delta by @@ -1338,7 +1347,7 @@ async def _create_stream( ) ], created=current_time, - model="None", + model=model, object="chat.completion.chunk", ) @@ -1353,7 +1362,7 @@ async def _create_stream( ) ], created=current_time, - model="None", + model=model, object="chat.completion.chunk", usage=CompletionUsage( completion_tokens=len(response.output_tokens), @@ -1399,6 +1408,7 @@ def __init__( async def create( self, *, + model: str | NotGiven = NOT_GIVEN, include: list[str] | None | NotGiven = NOT_GIVEN, input: str | ResponseInputParam | NotGiven = NOT_GIVEN, instructions: str | None | NotGiven = NOT_GIVEN, @@ -1415,6 +1425,7 @@ async def create( **kwargs: Any, ) -> Response: """Override create method to use AReaL engine""" + response_model = "default" if is_omitted(model) else str(model) # Initialize IDs and timestamps resp_id = f"resp-{uuid.uuid4().hex[:29]}" msg_id = f"msg-{uuid.uuid4().hex[:29]}" @@ -1641,7 +1652,7 @@ async def create( incomplete_details=None, instructions=None if is_omitted(instructions) else instructions, metadata=None if is_omitted(metadata) else metadata, - model="None", + model=response_model, object="response", output=resp_output, parallel_tool_calls=False, diff --git a/areal/experimental/openai/proxy/proxy_rollout_server.py b/areal/experimental/openai/proxy/proxy_rollout_server.py index 8394bde7d5..6ddeaf727a 100644 --- a/areal/experimental/openai/proxy/proxy_rollout_server.py +++ b/areal/experimental/openai/proxy/proxy_rollout_server.py @@ -126,6 +126,8 @@ def _warn_once(msg: str) -> None: # Server address (set at startup) _server_host: str = "0.0.0.0" _server_port: int = 8000 +_worker_role: str | None = None +_worker_index: int | None = None # Port allocation tracking _allocated_ports: set[int] = set() @@ -141,6 +143,24 @@ def _warn_once(msg: str) -> None: _nfs_record_root: str = "/tmp/areal/name_resolve" _etcd3_addr: str = "localhost:2379" + +def _resolve_worker_index(cli_worker_index: int) -> int: + """Resolve identity without clobbering an explicit scheduler value. + + A local launch can inherit ``SLURM_PROCID`` from its parent login shell. + Treating that stale value as an unconditional override makes every local + proxy identify as rank 0, so exact fork readiness checks reject ranks + 1..N. Slurm-only launchers still use the environment fallback when they + leave ``--worker-index`` at its sentinel value. + """ + worker_index = cli_worker_index + if worker_index == -1 and "SLURM_PROCID" in os.environ: + worker_index = int(os.environ["SLURM_PROCID"]) + if worker_index == -1: + raise ValueError("Invalid worker index. Not found from SLURM environ or args.") + return worker_index + + # ============================================================================= # Request Validation # ============================================================================= @@ -229,7 +249,12 @@ def _remove_api_keys_for_session(session_id: str) -> None: @app.get("/health") def health(): - return {"status": "ok", "initialized": _engine is not None} + return { + "status": "ok", + "initialized": _engine is not None, + "role": _worker_role, + "worker_index": _worker_index, + } @app.post("/alloc_ports") @@ -592,7 +617,11 @@ async def _call_client_create( ) sig = inspect.signature(create_fn) - areal_client_ignored_args = ["model"] + (extra_ignored_args or []) + # Keep the request model when the AReaL client supports it. Anthropic's + # LiteLLM response adapter uses this field to select the model family; + # dropping it leaves the generated ChatCompletion with no usable model + # identity and fails with ``Model type must be specified``. + areal_client_ignored_args = extra_ignored_args or [] areal_client_disallowed_args = ["areal_cache"] areal_client_allowed_args = list( k @@ -994,7 +1023,7 @@ def main(): args, _ = parser.parse_known_args() # Set global server address variables - global _server_host, _server_port + global _server_host, _server_port, _worker_role, _worker_index global \ _experiment_name, \ _trial_name, \ @@ -1014,13 +1043,9 @@ def main(): # Get worker identity worker_role = args.role - worker_index = args.worker_index - - if "SLURM_PROCID" in os.environ: - # Overwriting with slurm task id - worker_index = int(os.environ["SLURM_PROCID"]) - if worker_index == -1: - raise ValueError("Invalid worker index. Not found from SLURM environ or args.") + worker_index = _resolve_worker_index(args.worker_index) + _worker_role = worker_role + _worker_index = worker_index worker_id = f"{worker_role}/{worker_index}" # Determine port diff --git a/areal/experimental/openai/tool_call_parser.py b/areal/experimental/openai/tool_call_parser.py index 4425e241a5..17fd83dc3a 100644 --- a/areal/experimental/openai/tool_call_parser.py +++ b/areal/experimental/openai/tool_call_parser.py @@ -258,7 +258,7 @@ def _process_tool_calls_sglang( text: str, tools: list[Any], tool_call_parser: str, - reasoning_parser: str, + reasoning_parser: str | None, finish_reason: str, use_responses: bool = False, ) -> tuple[ @@ -266,10 +266,18 @@ def _process_tool_calls_sglang( str, str, ]: - from sglang.srt.entrypoints.openai.protocol import Function as SglFunction - from sglang.srt.entrypoints.openai.protocol import Tool as SglTool - from sglang.srt.function_call.function_call_parser import FunctionCallParser - from sglang.srt.parser.reasoning_parser import ReasoningParser + try: + from sglang.srt.entrypoints.openai.protocol import Function as SglFunction + from sglang.srt.entrypoints.openai.protocol import Tool as SglTool + from sglang.srt.function_call.function_call_parser import ( + FunctionCallParser, + ) + from sglang.srt.parser.reasoning_parser import ReasoningParser + except ImportError: + # Let the backend dispatcher try vLLM before giving up. Returning raw + # text here would make an installed vLLM parser unreachable whenever + # SGLang is absent. + raise ModuleNotFoundError("SGLang tool-call parser is unavailable") from None if use_responses: tools = [ @@ -289,14 +297,37 @@ def _process_tool_calls_sglang( for tool in tools ] - parser_p = FunctionCallParser(tools, tool_call_parser) - reasoning_parser_p = ReasoningParser(reasoning_parser) - - reasoning_text, content_text = _detect_think_and_return_ori_think( - text, - reasoning_parser_p.detector.think_start_token, - reasoning_parser_p.detector.think_end_token, - ) + try: + parser_p = FunctionCallParser(tools, tool_call_parser) + except (KeyError, ValueError) as e: + raise ValueError( + "Invalid rollout.openai.tool_call_parser=" + f"{tool_call_parser!r} for the SGLang backend: {e}. Use a parser " + "supported by SGLang (for Qwen3-Coder use 'qwen3_coder')." + ) from e + + # An empty reasoning parser is a supported way to disable reasoning + # extraction. Tool-call parsing is independent of reasoning parsing and + # must still work on the complete model output in that mode. + if reasoning_parser: + try: + reasoning_parser_p = ReasoningParser(reasoning_parser) + except (KeyError, ValueError) as e: + raise ValueError( + "Invalid rollout.openai.reasoning_parser=" + f"{reasoning_parser!r} while using tool_call_parser=" + f"{tool_call_parser!r}: {e}. Use a supported SGLang model type " + "such as 'qwen3', or leave reasoning_parser empty to disable " + "reasoning extraction." + ) from e + + reasoning_text, content_text = _detect_think_and_return_ori_think( + text, + reasoning_parser_p.detector.think_start_token, + reasoning_parser_p.detector.think_end_token, + ) + else: + reasoning_text, content_text = "", text if parser_p.has_tool_call(content_text): if finish_reason == "stop": @@ -341,7 +372,7 @@ def _process_tool_calls_vllm( text: str, tools: list[Any], tool_call_parser: str, - reasoning_parser: str, + reasoning_parser: str | None, finish_reason: str, use_responses: bool = False, tokenizer: Any = None, @@ -453,7 +484,7 @@ def process_tool_calls( text: str, tools: list[Any], tool_call_parser: str, - reasoning_parser: str, + reasoning_parser: str | None, finish_reason: str, use_responses: bool = False, tokenizer: Any = None, diff --git a/areal/infra/controller/rollout_controller.py b/areal/infra/controller/rollout_controller.py index cd0cab62b8..4e57e0e35f 100644 --- a/areal/infra/controller/rollout_controller.py +++ b/areal/infra/controller/rollout_controller.py @@ -36,6 +36,7 @@ SchedulingSpec, SchedulingStrategyType, ) +from areal.dataset.mopd import MOPD_ROUTE_METADATA_KEY, DatasetRoute from areal.infra.rpc.serialization import deserialize_value from areal.infra.utils.concurrent import run_async_task from areal.utils import logging, perf_tracer @@ -59,6 +60,7 @@ class _RemoteRolloutTaskInput: workflow: str | None workflow_kwargs: dict[str, Any] should_accept_fn: str | None + mopd_route: str | None = None is_eval: bool = False group_size: int = 1 proxy_addr: str | None = None @@ -100,6 +102,7 @@ def __init__( self._version = 0 self._task_id_generator = TaskIdGenerator() + self._mopd_routing_enabled = False # Use provided staleness manager or create a default one # The manager will be properly initialized in initialize() @@ -159,6 +162,43 @@ def _engine_name(self, rank: int) -> str: """ return f"{self._worker_role}/{rank}" + def enable_mopd_routing(self) -> None: + """Require dataset-source route metadata for training rollouts.""" + self._mopd_routing_enabled = True + + def _extract_mopd_route( + self, data: dict[str, Any], *, required: bool + ) -> tuple[dict[str, Any], str | None]: + """Remove internal route metadata before data reaches the workflow.""" + if MOPD_ROUTE_METADATA_KEY not in data: + if required: + raise ValueError("MOPD dataset-source route metadata is missing") + return data, None + + provenance = data[MOPD_ROUTE_METADATA_KEY] + if not isinstance(provenance, DatasetRoute): + raise ValueError( + "MOPD dataset-source route must contain DatasetRoute provenance" + ) + prepared = dict(data) + prepared.pop(MOPD_ROUTE_METADATA_KEY) + return prepared, provenance.route + + @staticmethod + def _propagate_mopd_route( + route: str | None, trajectory: dict[str, Any] + ) -> dict[str, Any]: + """Attach one source route to every trajectory derived from that source.""" + if route is None: + return trajectory + existing = trajectory.get("mopd_route") + if existing is not None and str(existing) != route: + raise ValueError( + f"Workflow changed mopd_route from {route!r} to {existing!r}" + ) + trajectory["mopd_route"] = route + return trajectory + def initialize( self, role: str, @@ -321,6 +361,22 @@ async def _async_initialize( "base_gpu_id": (rank % slots_per_node) * self._gpus_per_server, "_awex_gpus_per_server": self._gpus_per_server, } + # A non-forked colocated rollout aliases the actor worker: + # [0] is RPC and [1] is the actor rendezvous TCPStore, so + # SGLang needs a separately reserved third port. A forked + # rollout owns its worker and can use its second port. Always + # override the global server argument because replicas on the + # same node cannot safely share one explicit NCCL port. + port_index = 1 if self.config.scheduling_strategy.fork else 2 + if len(worker.worker_ports) <= port_index: + required = port_index + 1 + raise ValueError( + f"Colocated rollout worker {worker.id!r} needs at least " + f"{required} allocated ports, but has " + f"{len(worker.worker_ports)}. Set SchedulingSpec.port_count=" + f"{required} for the colocated target role." + ) + per_worker_args["nccl_port"] = int(worker.worker_ports[port_index]) launch_tasks.append( self.scheduler.async_call_engine( worker_id=worker.id, @@ -370,33 +426,68 @@ async def _async_initialize( logger.info("All engines are initialized...") def destroy(self): - # Stop background threads and shutdown the async task runner + errors: list[str] = [] + + # Stop externally reachable/background work before deleting any role. + self._stop_proxy_gateway() if self._dispatcher is not None: self._dispatcher.destroy() - self._stop_callback_server() - self._collective_rpc("destroy", http_timeout=60.0) + # Proxy is a logical leaf of rollout. Destroy its engines and workers + # while the rollout alias still resolves to the actual process owner. + if self._proxy_started or self.proxy_workers: + if self.proxy_workers: + try: + + async def _destroy_proxy_engines(): + tasks = [ + self.scheduler.async_call_engine( + worker_id=worker.id, + method="destroy", + engine_name=self._proxy_engine_name(rank), + ) + for rank, worker in enumerate(self.proxy_workers) + ] + return await asyncio.gather(*tasks, return_exceptions=True) + + results = run_async_task(_destroy_proxy_engines) + errors.extend( + f"proxy engine {rank}: {result}" + for rank, result in enumerate(results) + if isinstance(result, BaseException) + ) + except Exception as exc: # noqa: BLE001 + errors.append(f"proxy engine destroy: {exc}") + try: + self.scheduler.delete_workers(role=self._proxy_role) + except Exception as exc: # noqa: BLE001 + errors.append(f"proxy worker delete: {exc}") + else: + self.proxy_workers.clear() + self.proxy_addrs.clear() + self._proxy_started = False + logger.info("Proxy workers deleted") + + try: + self._collective_rpc("destroy", http_timeout=60.0) + except Exception as exc: # noqa: BLE001 + errors.append(f"rollout engine destroy: {exc}") # Delete workers via scheduler if hasattr(self, "_worker_role"): try: self.scheduler.delete_workers(role=self._worker_role) + except Exception as exc: # noqa: BLE001 + errors.append(f"rollout worker delete: {exc}") + else: self.workers.clear() logger.info("Workers deleted") - except Exception: - logger.error(f"Error deleting workers: {traceback.format_exc()}") - # Delete proxy workers if initialized - if self._proxy_started: - try: - self.scheduler.delete_workers(role=self._proxy_role) - self.proxy_workers.clear() - self.proxy_addrs.clear() - self._proxy_started = False - logger.info("Proxy workers deleted") - except Exception: - logger.error(f"Error deleting proxy workers: {traceback.format_exc()}") + with self._futures_lock: + self._pending_futures.clear() + if errors: + raise RuntimeError("RolloutController cleanup failed: " + "; ".join(errors)) # Shutdown proxy gateway if initialized self._stop_proxy_gateway() @@ -420,8 +511,21 @@ def start_proxy(self) -> None: "Call initialize() first." ) - run_async_task(self._async_start_proxy) - self._proxy_started = True + try: + run_async_task(self._async_start_proxy) + except Exception: + try: + self.scheduler.delete_workers(role=self._proxy_role) + except Exception: # noqa: BLE001 + logger.error( + "Failed to rollback partially initialized proxy workers:\n%s", + traceback.format_exc(), + ) + self.proxy_workers.clear() + self.proxy_addrs.clear() + raise + else: + self._proxy_started = True async def _async_start_proxy(self) -> None: """Async implementation of proxy worker initialization.""" @@ -906,6 +1010,7 @@ async def _submit_then_wait() -> _RemoteRolloutResult | None: traj = result if traj is not None: + traj = self._propagate_mopd_route(pending_task.mopd_route, traj) manager.on_rollout_accepted() if self.config.enable_rollout_tracing: logger.info( @@ -951,6 +1056,9 @@ def submit( reward_normalization: bool = False, drop_incomplete_group: bool = False, ) -> int: + data, mopd_route = self._extract_mopd_route( + data, required=self._mopd_routing_enabled and not is_eval + ) workflow_str = self._resolve_workflow_str(workflow) should_accept_fn = self._resolve_should_accept_fn(should_accept_fn) if workflow_kwargs is None: @@ -967,6 +1075,7 @@ def submit( workflow_kwargs=workflow_kwargs, should_accept_fn=should_accept_fn, task_id=task_id, + mopd_route=mopd_route, is_eval=is_eval, group_size=group_size, proxy_addr=proxy_addr, @@ -1046,12 +1155,16 @@ def prepare_batch( def task_input_generator(): for data in cycle_dataloader(dataloader): for item in data: + workflow_data, mopd_route = self._extract_mopd_route( + item, required=self._mopd_routing_enabled + ) yield _RemoteRolloutTaskInput( - data=item, + data=workflow_data, workflow=workflow_str, workflow_kwargs=workflow_kwargs, should_accept_fn=should_accept_fn, task_id=self._task_id_generator.next(), + mopd_route=mopd_route, group_size=group_size, reward_normalization=reward_normalization, drop_incomplete_group=drop_incomplete_group, diff --git a/areal/infra/controller/train_controller.py b/areal/infra/controller/train_controller.py index d68d43841f..8cbd0f7b40 100644 --- a/areal/infra/controller/train_controller.py +++ b/areal/infra/controller/train_controller.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 import asyncio +import math from threading import Lock from typing import Any @@ -21,7 +22,7 @@ ) from areal.api.alloc_mode import ModelAllocation from areal.api.cli_args import PerfTracerConfig, TrainEngineConfig -from areal.infra.rpc.rtensor import RTensor +from areal.infra.rpc.rtensor import RTensor, RTensorDrainReceipt from areal.infra.utils.concurrent import run_async_task from areal.utils import logging, stats_tracker from areal.utils.data import make_dummy_eval_item @@ -126,29 +127,48 @@ def _dispatch_tensors( def _pad_eval_batch( - args: tuple[Any, ...], dp_size: int, group_size: int = 1 + args: tuple[Any, ...], + dp_size: int, + group_size: int = 1, + *, + min_items_per_dp: int = 1, + items_per_dp_divisor: int = 1, + active_dummies: bool = False, ) -> tuple[Any, ...]: - """Pad the first tensor-like arg to a multiple of ``dp_size * group_size``. + """Pad the first tensor-like arg evenly across all DP replicas. Called before dispatch for explicit evaluation controller paths so that ``balanced_greedy_partition`` always receives a divisible input. - Dummy items have zero attention/loss masks and contribute nothing - to metrics or loss. + ``min_items_per_dp`` can reserve enough inputs for pipeline microbatches. + Active dummies contain one attended token, which keeps pipeline forwards + valid; their outputs must be discarded by the caller. """ + if min_items_per_dp < 1: + raise ValueError("min_items_per_dp must be positive") + if items_per_dp_divisor < 1: + raise ValueError("items_per_dp_divisor must be positive") result = list(args) - pad_target = dp_size * group_size + per_dp_quantum = math.lcm(group_size, items_per_dp_divisor) + pad_quantum = dp_size * per_dp_quantum + minimum_size = dp_size * min_items_per_dp for i, arg in enumerate(result): if isinstance(arg, list) and arg and _is_tensor_like(arg): n = len(arg) - pad_count = (-n) % pad_target + padded_size = max(n, minimum_size) + padded_size += (-padded_size) % pad_quantum + pad_count = padded_size - n if pad_count > 0: padded = list(arg) template = arg[0] - padded.extend(make_dummy_eval_item(template) for _ in range(pad_count)) + padded.extend( + make_dummy_eval_item(template, active_attention=active_dummies) + for _ in range(pad_count) + ) result[i] = padded logger.info( f"Eval dispatch: padded {pad_count} dummy items " - f"(total {len(padded)}) for dp_size={dp_size}" + f"(total {len(padded)}) for dp_size={dp_size}, " + f"min_items_per_dp={min_items_per_dp}" ) break # only pad the first tensor-like arg return tuple(result) @@ -294,14 +314,10 @@ def initialize( worker_ids = self.scheduler.create_workers(job=job) logger.info(f"Workers created: {worker_ids}") - # Wait for workers to be ready logger.info("Waiting for workers to be ready...") self.workers = self.scheduler.get_workers(role=job.role) logger.info(f"Workers ready: {[w.id for w in self.workers]}") - # Determine distributed training master address and port from rank 0 worker - # These are used for PyTorch distributed initialization across workers - # Prefer engine_ports[1] if available, fallback to worker_ports[1] rank0_worker = self.workers[0] if rank0_worker.engine_ports: self._master_port = int(rank0_worker.engine_ports[1]) @@ -313,18 +329,15 @@ def initialize( f"Distributed training: MASTER_ADDR={self._master_addr}, MASTER_PORT={self._master_port}" ) - # Construct engine class import path for dynamic loading on workers - # Workers will import and instantiate the engine class using this path engine_class = self.train_engine - - # Create and initialize engines on workers run_async_task( self._async_create_engines, f"{engine_class.__module__}.{engine_class.__name__}", ) - run_async_task(self._async_initialize_engines, ft_spec, **kwargs) + engine_init_kwargs = dict(kwargs) + engine_init_kwargs.setdefault("role", role) + run_async_task(self._async_initialize_engines, ft_spec, **engine_init_kwargs) - # Identify DP head workers self._identify_dp_heads() logger.info("TrainController initialization complete") @@ -466,7 +479,6 @@ async def _destroy_all_engines(): except Exception as e: logger.error(f"Error deleting workers: {e}") - # Clear worker lists self.workers.clear() self.workers_is_dp_head.clear() @@ -488,6 +500,26 @@ def _custom_function_call( ) return self._collect_results(results, group_indices) + def _custom_function_call_all_dp_heads( + self, + method: str, + *args, + rpc_meta: dict[str, Any] | None = None, + **kwargs, + ) -> list[Any]: + """Call every rank and retain one result for every DP head.""" + dp_args, dp_kwargs, group_indices = self._prepare_dispatch(*args, **kwargs) + if group_indices is not None: + raise ValueError("all-DP-head calls only support replicated inputs") + results = run_async_task( + self._call_workers, method, dp_args, dp_kwargs, rpc_meta=rpc_meta + ) + return [ + result + for result, is_head in zip(results, self.workers_is_dp_head, strict=True) + if is_head + ] + async def _async_custom_function_call( self, method: str, @@ -508,11 +540,19 @@ def _pad_eval_dispatch_args( kwargs: dict[str, Any], *, group_size: int, + min_items_per_dp: int = 1, + items_per_dp_divisor: int = 1, + active_dummies: bool = False, ) -> tuple[tuple[Any, ...], dict[str, Any]]: """Pad eval batches for explicit algorithm-level evaluation dispatch.""" kwargs = dict(kwargs) args = _pad_eval_batch( - args, self.parallel_strategy.dp_size, group_size=group_size + args, + self.parallel_strategy.dp_size, + group_size=group_size, + min_items_per_dp=min_items_per_dp, + items_per_dp_divisor=items_per_dp_divisor, + active_dummies=active_dummies, ) return args, kwargs @@ -719,6 +759,10 @@ def init_awex_adapter(self, meta_server_addr: str | None = None): "init_awex_adapter", meta_server_addr=meta_server_addr ) + def init_weight_residency_adapter(self): + """Create the Megatron DDP-flat-buffer residency adapter.""" + self._custom_function_call("init_weight_residency_adapter") + def step_lr_scheduler(self): """Step the learning rate scheduler. @@ -734,11 +778,35 @@ def update_weights(self, meta: WeightUpdateMeta): def offload(self) -> None: """Offload model parameters to CPU across all train workers.""" - self._custom_function_call("offload") + self._collective_lifecycle_call("offload") def onload(self) -> None: """Onload model parameters to GPU across all train workers.""" - self._custom_function_call("onload") + self._collective_lifecycle_call("onload") + + def _collective_lifecycle_call(self, method: str) -> None: + """Wait for every rank once; lifecycle collectives are not retry-safe.""" + + async def _call_all_workers(): + tasks = [ + self.scheduler.async_call_engine( + worker.id, + method, + self._engine_name(rank), + max_retries=1, + ) + for rank, worker in enumerate(self.workers) + ] + return await asyncio.gather(*tasks, return_exceptions=True) + + results = run_async_task(_call_all_workers) + failures = [ + RuntimeError(f"{method} failed on rank {rank}: {result}") + for rank, result in enumerate(results) + if isinstance(result, BaseException) + ] + if failures: + raise ExceptionGroup(f"Train worker collective {method} failed", failures) def get_device_stats(self): return self._custom_function_call("get_device_stats") @@ -1001,3 +1069,59 @@ def _clear_batches_locked(self, *targets: dict[str, RTensor]) -> None: "RTensor storage cleanup failed across two clear_batches calls " f"({summary})" ) + + def strict_clear_batches(self, *targets: Any) -> RTensorDrainReceipt: + """Clear source shards and prove every consumer DP head drained them. + + Unlike :meth:`clear_batches`, any source-delete or consumer RPC failure is + fatal. The returned receipt is suitable for the MOPD teacher lifecycle + gate; callers must not destroy teacher workers before this succeeds. + """ + shards_by_node = { + node_addr: list(dict.fromkeys(shard_ids)) + for node_addr, shard_ids in RTensor.collect_shards(targets).items() + } + shard_ids = [ + shard_id + for node_shards in shards_by_node.values() + for shard_id in node_shards + ] + if not shard_ids: + return RTensorDrainReceipt( + consumer_role=self._worker_role, + shard_count=0, + source_node_count=0, + consumer_dp_head_count=0, + ) + + async def _strict_clear_sources() -> None: + await asyncio.gather( + *[ + RTensor.clear_node(node_addr, node_shards) + for node_addr, node_shards in shards_by_node.items() + ] + ) + + run_async_task(_strict_clear_sources) + self._custom_function_call_all_dp_heads( + "clear_batches", shard_ids, rpc_meta={"broadcast": False} + ) + stats = self._custom_function_call_all_dp_heads( + "fetch_buffer_stats", shard_ids, rpc_meta={"broadcast": False} + ) + leaking_heads = [ + index + for index, stat in enumerate(stats) + if not isinstance(stat, dict) or stat.get("matching_entries") != 0 + ] + if leaking_heads: + raise RuntimeError( + f"RTensor drain incomplete on {self._worker_role} DP heads " + f"{leaking_heads}: stats={stats}" + ) + return RTensorDrainReceipt( + consumer_role=self._worker_role, + shard_count=len(shard_ids), + source_node_count=len(shards_by_node), + consumer_dp_head_count=len(stats), + ) diff --git a/areal/infra/remote_inf_engine.py b/areal/infra/remote_inf_engine.py index 061e586bc4..8bf9243f74 100644 --- a/areal/infra/remote_inf_engine.py +++ b/areal/infra/remote_inf_engine.py @@ -1399,6 +1399,10 @@ def prepare_batch( @trace_perf("remote_inf_engine.pause_generation", category="misc") def pause_generation(self): """Pause request submission for async rollout.""" + # SGLang needs a two-stage pause before colocated memory can be + # released: ``abort`` closes admission and waits for in-flight work to + # drain, then ``retract`` puts the now-idle scheduler into its paused + # state. Other backends keep their existing single-request protocol. get_pause_requests = getattr(self.backend, "get_pause_requests", None) pause_requests = ( get_pause_requests() diff --git a/areal/infra/rpc/rtensor.py b/areal/infra/rpc/rtensor.py index d60ade2d54..0c836665e3 100644 --- a/areal/infra/rpc/rtensor.py +++ b/areal/infra/rpc/rtensor.py @@ -25,6 +25,34 @@ logger = logging.getLogger("HttpRTensor") +@dataclass(frozen=True) +class RTensorDrainReceipt: + """Typed proof that one controller drained its registered consumers. + + This receipt describes one controller fan-out. Cross-role lease ownership + remains with the caller until the runtime has a dynamic consumer registry; + callers must collect a receipt from every role that localized the batch + before releasing its storage owner. + """ + + consumer_role: str + shard_count: int + source_node_count: int + consumer_dp_head_count: int + + def __post_init__(self) -> None: + if not self.consumer_role: + raise ValueError("consumer_role must be a non-empty string") + for name in ( + "shard_count", + "source_node_count", + "consumer_dp_head_count", + ): + value = getattr(self, name) + if value < 0: + raise ValueError(f"{name} must be non-negative, got {value}") + + class RTensorBackend(Protocol): def fetch(self, shards: list[TensorShardInfo]) -> list[torch.Tensor]: """Fetch multiple tensors concurrently. @@ -343,6 +371,16 @@ def fetch_buffer_stats() -> dict[str, int]: return {"num_entries": len(_fetch_buffer)} +def fetch_buffer_matching_stats(shard_ids: Iterable[Any]) -> dict[str, int]: + """Return total and requested-shard counts for strict drain verification.""" + requested = set(shard_ids) + with _fetch_buffer_lock: + return { + "num_entries": len(_fetch_buffer), + "matching_entries": sum(sid in _fetch_buffer for sid in requested), + } + + def flatten_shard_ids(obj: Any) -> list[Any]: """Collect all RTensor shard IDs from a nested structure as a flat list. diff --git a/areal/infra/scheduler/local.py b/areal/infra/scheduler/local.py index 953b77b5fb..0dc0a028cd 100644 --- a/areal/infra/scheduler/local.py +++ b/areal/infra/scheduler/local.py @@ -309,6 +309,7 @@ async def _fork_single_worker( target_wi: WorkerInfo, target_role: str, command: str | None = None, + env: dict[str, str] | None = None, ) -> WorkerInfo: """Fork a single worker asynchronously. @@ -382,6 +383,7 @@ async def _fork_single_worker( "role": role, "worker_index": idx, "raw_cmd": raw_cmd, + "env": env or {}, } async with session.post( f"{guard_url}/fork", @@ -462,7 +464,7 @@ async def _fork_single_worker( gpu_devices=target_wi.gpu_devices, # Inherited from target created_at=time.time(), log_file=str(self.log_dir / f"{role}.log"), - env_vars=target_wi.env_vars.copy(), # Inherited from target + env_vars={**target_wi.env_vars, **(env or {})}, ) async def _kill_forked_worker( @@ -533,6 +535,7 @@ async def _create_forked_workers_async( target_role: str, target_workers: list[WorkerInfo], command: str | None = None, + env_vars: list[dict[str, str]] | None = None, ) -> list[str]: """Create forked workers concurrently using async requests. @@ -550,7 +553,13 @@ async def _create_forked_workers_async( # Launch all fork requests concurrently with exception handling tasks = [ self._fork_single_worker( - session, role, idx, target_wi, target_role, command + session, + role, + idx, + target_wi, + target_role, + command, + None if env_vars is None else env_vars[idx], ) for idx, target_wi in enumerate(target_workers) ] @@ -613,6 +622,7 @@ def fork_workers( role: str, target_role: str, command: str | None = None, + env_vars: list[dict[str, str]] | None = None, ) -> list[str]: """Fork new worker processes from existing workers. @@ -645,6 +655,7 @@ def fork_workers( target_role, target_workers, command, + env_vars, ) except Exception: # Cleanup on failure @@ -729,7 +740,11 @@ def create_workers(self, job: Job, *args, **kwargs) -> list[str]: # Check if fork mode is enabled if strategy.fork: # Fork mode: spawn new processes on same GPUs via /fork endpoint - worker_ids = self.fork_workers(role, colocate_role) + worker_ids = self.fork_workers( + role, + colocate_role, + env_vars=[scheduling.env_vars for scheduling in schedulings], + ) else: # Reuse existing workers - no new processes spawned worker_ids = [w.worker.id for w in target_workers] diff --git a/areal/infra/scheduler/ray.py b/areal/infra/scheduler/ray.py index 9cd96d93d2..0d9d7ebc20 100644 --- a/areal/infra/scheduler/ray.py +++ b/areal/infra/scheduler/ray.py @@ -1003,6 +1003,7 @@ async def _fork_single_worker( target_wi: RayWorkerInfo, target_role: str, command: str | None = None, + env: dict[str, str] | None = None, ) -> RayWorkerInfo: """Fork a single worker asynchronously. @@ -1077,6 +1078,7 @@ async def _fork_single_worker( "role": role, "worker_index": idx, "raw_cmd": raw_cmd, + "env": env or {}, } async with session.post( f"{guard_url}/fork", @@ -1222,6 +1224,7 @@ async def _create_forked_workers_async( target_role: str, target_workers: list[RayWorkerInfo], command: str | None = None, + env_vars: list[dict[str, str]] | None = None, ) -> list[str]: """Create forked workers concurrently using async requests. @@ -1239,7 +1242,13 @@ async def _create_forked_workers_async( # Launch all fork requests concurrently with exception handling tasks = [ self._fork_single_worker( - session, role, idx, target_wi, target_role, command + session, + role, + idx, + target_wi, + target_role, + command, + None if env_vars is None else env_vars[idx], ) for idx, target_wi in enumerate(target_workers) ] @@ -1301,6 +1310,7 @@ def fork_workers( role: str, target_role: str, command: str | None = None, + env_vars: list[dict[str, str]] | None = None, ) -> list[str]: """Fork new worker processes from existing workers. @@ -1333,6 +1343,7 @@ def fork_workers( target_role, target_workers, command, + env_vars, ) except Exception: # Cleanup on failure @@ -1563,7 +1574,11 @@ def create_workers(self, job: Job, *args, **kwargs) -> list[str]: # Check if fork mode is enabled if strategy.fork: # Fork mode: spawn new processes on same nodes via /fork endpoint - return self.fork_workers(role, colocate_role) + return self.fork_workers( + role, + colocate_role, + env_vars=[scheduling.env_vars for scheduling in schedulings], + ) # Reuse existing workers - no new Ray launchers created worker_ids = [w.worker.id for w in target_workers] diff --git a/areal/infra/scheduler/slurm.py b/areal/infra/scheduler/slurm.py index 10a2eaa128..7cf91a55e1 100644 --- a/areal/infra/scheduler/slurm.py +++ b/areal/infra/scheduler/slurm.py @@ -494,6 +494,7 @@ async def _fork_single_worker( target_wi: SlurmWorkerInfo, target_role: str, command: str | None = None, + env: dict[str, str] | None = None, ) -> SlurmWorkerInfo: """Fork a single worker asynchronously. @@ -568,6 +569,7 @@ async def _fork_single_worker( "role": role, "worker_index": idx, "raw_cmd": raw_cmd, + "env": env or {}, } async with session.post( f"{guard_url}/fork", @@ -715,6 +717,7 @@ async def _create_forked_workers_async( target_role: str, target_workers: list[SlurmWorkerInfo], command: str | None = None, + env_vars: list[dict[str, str]] | None = None, ) -> list[str]: """Create forked workers concurrently using async requests. @@ -732,7 +735,13 @@ async def _create_forked_workers_async( # Launch all fork requests concurrently with exception handling tasks = [ self._fork_single_worker( - session, role, idx, target_wi, target_role, command + session, + role, + idx, + target_wi, + target_role, + command, + None if env_vars is None else env_vars[idx], ) for idx, target_wi in enumerate(target_workers) ] @@ -795,6 +804,7 @@ def fork_workers( role: str, target_role: str, command: str | None = None, + env_vars: list[dict[str, str]] | None = None, ) -> list[str]: """Fork new worker processes from existing workers. @@ -827,6 +837,7 @@ def fork_workers( target_role, target_workers, command, + env_vars, ) except Exception: # Cleanup on failure @@ -1042,7 +1053,11 @@ def create_workers(self, job: Job, *args, **kwargs) -> list[str]: # Check if fork mode is enabled if strategy.fork: # Fork mode: spawn new processes on same nodes via /fork endpoint - return self.fork_workers(role, colocate_role) + return self.fork_workers( + role, + colocate_role, + env_vars=[scheduling.env_vars for scheduling in schedulings], + ) # Reuse existing workers - no new Slurm job submitted worker_ids = [w.worker.id for w in target_workers] diff --git a/areal/infra/workflow_executor.py b/areal/infra/workflow_executor.py index 349ca77f34..2993998513 100644 --- a/areal/infra/workflow_executor.py +++ b/areal/infra/workflow_executor.py @@ -795,7 +795,6 @@ def __init__( self.config = config self.inference_engine = inference_engine - # Use provided staleness manager or create a default one # The manager will be properly initialized in initialize() self._staleness_manager = staleness_manager @@ -1207,9 +1206,11 @@ async def _execute_workflow() -> _RolloutResult | None: reason: str | None = None try: - traj = await pending_task.workflow.arun_episode( - self.inference_engine, pending_task.data - ) + workflow_data = pending_task.data + if workflow_data is not None: + traj = await pending_task.workflow.arun_episode( + self.inference_engine, workflow_data + ) # Trajectory format checking if self.config.check_trajectory_format and traj is not None: diff --git a/areal/reward/__init__.py b/areal/reward/__init__.py index d6b8e42359..fc06b83bbe 100644 --- a/areal/reward/__init__.py +++ b/areal/reward/__init__.py @@ -107,6 +107,7 @@ def get_math_verify_worker() -> MathVerifyWorker: "gsm8k_reward_fn", "geometry3k_reward_fn", "clevr_count_70k_reward_fn", + "if_gap_reward_fn", ] @@ -114,6 +115,7 @@ def get_math_verify_worker() -> MathVerifyWorker: "gsm8k_reward_fn": "areal.reward.gsm8k", "geometry3k_reward_fn": "areal.reward.geometry3k", "clevr_count_70k_reward_fn": "areal.reward.clevr_count_70k", + "if_gap_reward_fn": "areal.reward.if_gap", } diff --git a/areal/reward/if_gap.py b/areal/reward/if_gap.py new file mode 100644 index 0000000000..ab8e6f5563 --- /dev/null +++ b/areal/reward/if_gap.py @@ -0,0 +1,181 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Constraint reward used by the mixed SWE/IF MOPD agent. + +The input dataset contains two verifier formats. ``ifeval_g`` rows use the +``verifiable_instructions`` registry, while ``ifrl_recovery`` rows use the +checkers vendored in IFDataSynthesis. Both dependencies are optional and are +loaded lazily from the paths supplied by the launch script. +""" + +from __future__ import annotations + +import os +import sys +from typing import Any + +from areal.utils import logging + +logger = logging.getLogger("IFGapReward") + +STRICT_BONUS = 0.5 + +_ifeval_registry: dict[str, Any] | None = None +_ifrl_engines: tuple[Any, Any, Any] | None = None + + +def _load_ifeval_registry() -> dict[str, Any]: + """Load the generalized IFEval instruction registry lazily.""" + global _ifeval_registry + if _ifeval_registry is None: + from verifiable_instructions import instructions_registry + + _ifeval_registry = instructions_registry.INSTRUCTION_DICT + logger.info( + "Loaded verifiable_instructions registry with %d entries", + len(_ifeval_registry), + ) + return _ifeval_registry + + +def _load_ifrl_engines() -> tuple[Any, Any, Any]: + """Load the IFDataSynthesis verifier and recovery checkers lazily.""" + global _ifrl_engines + if _ifrl_engines is not None: + return _ifrl_engines + + configured_root = os.getenv("IF_SYNTH_ROOT", "").strip() + if not configured_root: + raise FileNotFoundError("IF_SYNTH_ROOT is not configured") + root = os.path.abspath(configured_root) + if not os.path.isdir(root): + raise FileNotFoundError( + f"IF_SYNTH_ROOT does not point to IFDataSynthesis: {root!r}" + ) + if root not in sys.path: + sys.path.insert(0, root) + + import verifier as generalized_verifier # type: ignore[import-not-found] + from ifrl_recovery import ( # type: ignore[import-not-found] + bench_map, + simple_checkers, + ) + + _ifrl_engines = (generalized_verifier, simple_checkers, bench_map) + logger.info("Loaded IF recovery verifier from %s", root) + return _ifrl_engines + + +def extract_visible_answer(completion: str) -> str: + """Remove model thinking; truncated thinking has no scoreable answer.""" + if "" in completion: + return completion.split("", 1)[1] + if "" in completion: + return "" + return completion + + +def _score_ifeval(spec: dict[str, Any], answer: str) -> list[bool]: + registry = _load_ifeval_registry() + instruction_ids = spec.get("instruction_id_list") or [] + kwargs_list = spec.get("kwargs") or [{} for _ in instruction_ids] + results: list[bool] = [] + for instruction_id, instruction_kwargs in zip( + instruction_ids, kwargs_list, strict=True + ): + passed = False + try: + instruction = registry[instruction_id](instruction_id) + params = { + key: value + for key, value in (instruction_kwargs or {}).items() + if value is not None + } + instruction.build_description(**params) + passed = bool(instruction.check_following(answer)) + except KeyError: + logger.warning( + "Unknown ifeval_g instruction id (scored as fail): %s", + instruction_id, + ) + except Exception: # noqa: BLE001 + logger.warning( + "IFEval instruction failed (scored as fail): %s", + instruction_id, + exc_info=True, + ) + results.append(passed) + return results + + +def _score_ifrl_recovery(spec: dict[str, Any], answer: str) -> list[bool]: + verifier, simple_checkers, bench_map = _load_ifrl_engines() + if_type = spec.get("if_type") + is_zh = simple_checkers.is_zh(answer) + results: list[bool] = [] + for constraint in spec.get("constraints") or []: + name = constraint.get("constraint_name") + params = constraint.get("params") + passed = False + mapped = ( + bench_map.map_constraint(name, params) if if_type == "if_bench" else None + ) + if mapped is not None: + instruction_id, instruction_kwargs = mapped + full_kwargs = dict(verifier.blank_kwargs()) + full_kwargs.update(instruction_kwargs) + try: + passed = bool( + verifier.verify([instruction_id], [full_kwargs], answer)[0][1] + ) + except Exception: # noqa: BLE001 + passed = False + elif name in simple_checkers.CHECKERS: + try: + passed = bool(simple_checkers.check(name, params, answer, is_zh)) + except Exception: # noqa: BLE001 + passed = False + else: + logger.warning( + "Unknown ifrl_recovery constraint (scored as fail): %s", name + ) + results.append(passed) + return results + + +def score_if_gap_spec( + verify_engine: str, spec: dict[str, Any], answer: str +) -> list[bool]: + """Return one pass/fail result for every constraint in an IF row.""" + if verify_engine == "ifeval_g": + return _score_ifeval(spec, answer) + if verify_engine == "ifrl_recovery": + return _score_ifrl_recovery(spec, answer) + raise ValueError(f"Unsupported IF verify engine: {verify_engine!r}") + + +def if_gap_reward_fn( + prompt: str, + completions: str, + prompt_ids: list[int] | None = None, + completion_ids: list[int] | None = None, + verify_engine: str = "", + spec: dict[str, Any] | None = None, + **kwargs: Any, +) -> float: + """Score a single IF answer using pass rate plus a strict-pass bonus.""" + del prompt, prompt_ids, completion_ids, kwargs + try: + spec = spec or {} + answer = extract_visible_answer(str(completions)).strip() + if not answer: + return 0.0 + results = score_if_gap_spec(verify_engine, spec, answer) + if not results: + return 0.0 + pass_rate = sum(results) / len(results) + strict = float(all(results)) + return float((pass_rate + STRICT_BONUS * strict) / (1.0 + STRICT_BONUS)) + except Exception: # noqa: BLE001 + logger.warning("Exception in IF gap reward", exc_info=True) + return 0.0 diff --git a/areal/trainer/mopd/__init__.py b/areal/trainer/mopd/__init__.py new file mode 100644 index 0000000000..691a1dd631 --- /dev/null +++ b/areal/trainer/mopd/__init__.py @@ -0,0 +1,24 @@ +# SPDX-License-Identifier: Apache-2.0 + +from areal.infra.rpc.rtensor import RTensorDrainReceipt +from areal.trainer.mopd.execution import MOPDExecutionPlan +from areal.trainer.mopd.loss import compose_mopd_loss, mopd_loss_fn +from areal.trainer.mopd.scoring import MOPDTeacherController +from areal.trainer.mopd.targets import aggregate_mopd_targets +from areal.trainer.mopd.teacher_manager import ( + PersistentTeacherManager, + TeacherManagerState, +) +from areal.trainer.mopd.teacher_phase import MOPDTeacherPhase + +__all__ = [ + "RTensorDrainReceipt", + "PersistentTeacherManager", + "MOPDTeacherPhase", + "MOPDTeacherController", + "MOPDExecutionPlan", + "TeacherManagerState", + "aggregate_mopd_targets", + "compose_mopd_loss", + "mopd_loss_fn", +] diff --git a/areal/trainer/mopd/compatibility.py b/areal/trainer/mopd/compatibility.py new file mode 100644 index 0000000000..fb4ff65618 --- /dev/null +++ b/areal/trainer/mopd/compatibility.py @@ -0,0 +1,180 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import hashlib +import json +from collections.abc import Mapping +from pathlib import Path +from typing import Any + +_NON_MODEL_CONFIG_KEYS = { + "_name_or_path", + "dtype", + "finetuning_task", + "id2label", + "label2id", + "output_attentions", + "output_hidden_states", + "problem_type", + "return_dict", + "return_dict_in_generate", + "task_specific_params", + "tokenizer_class", + "torch_dtype", + "transformers_version", + "use_cache", +} + + +def _read_json(path: Path) -> Any: + if not path.is_file(): + raise FileNotFoundError(f"Missing MOPD model metadata: {path}") + return json.loads(path.read_text(encoding="utf-8")) + + +def _sha256_json(value: Any) -> str: + canonical = json.dumps( + value, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ).encode() + return hashlib.sha256(canonical).hexdigest() + + +def _architecture_fingerprint(model_path: Path) -> tuple[dict[str, Any], str]: + config = _read_json(model_path / "config.json") + if not isinstance(config, dict): + raise ValueError(f"MOPD config.json must contain an object: {model_path}") + architecture = { + key: value for key, value in config.items() if key not in _NON_MODEL_CONFIG_KEYS + } + return architecture, _sha256_json(architecture) + + +def _tokenizer_payload(tokenizer_path: Path) -> dict[str, Any]: + tokenizer_config_path = tokenizer_path / "tokenizer_config.json" + tokenizer_config = ( + _read_json(tokenizer_config_path) if tokenizer_config_path.is_file() else {} + ) + if not isinstance(tokenizer_config, dict): + raise ValueError( + "MOPD tokenizer_config.json must contain an object: " + f"{tokenizer_config_path}" + ) + + vocab_path = tokenizer_path / "vocab.json" + tokenizer_json_path = tokenizer_path / "tokenizer.json" + tokenizer_model_path = tokenizer_path / "tokenizer.model" + if vocab_path.is_file(): + vocabulary = _read_json(vocab_path) + tokenizer_added_tokens = None + elif tokenizer_json_path.is_file(): + tokenizer_json = _read_json(tokenizer_json_path) + try: + vocabulary = tokenizer_json["model"]["vocab"] + except (KeyError, TypeError) as exc: + raise ValueError( + f"Cannot find model.vocab in {tokenizer_json_path}" + ) from exc + tokenizer_added_tokens = tokenizer_json.get("added_tokens") + elif tokenizer_model_path.is_file(): + vocabulary = { + "sentencepiece_sha256": hashlib.sha256( + tokenizer_model_path.read_bytes() + ).hexdigest() + } + tokenizer_added_tokens = None + else: + raise FileNotFoundError( + "MOPD tokenizer must contain vocab.json, tokenizer.json, or " + f"tokenizer.model: {tokenizer_path}" + ) + + return { + "vocabulary": vocabulary, + "added_tokens": tokenizer_added_tokens, + "added_tokens_decoder": tokenizer_config.get("added_tokens_decoder", {}), + } + + +def model_fingerprint( + model_path: str | Path, + *, + tokenizer_path: str | Path | None = None, +) -> dict[str, object]: + """Fingerprint checkpoint structure and token-ID mappings without weights.""" + resolved_model_path = Path(model_path) + resolved_tokenizer_path = ( + Path(tokenizer_path) if tokenizer_path else resolved_model_path + ) + architecture, architecture_sha256 = _architecture_fingerprint(resolved_model_path) + tokenizer = _tokenizer_payload(resolved_tokenizer_path) + return { + "path": str(resolved_model_path.resolve()), + "tokenizer_path": str(resolved_tokenizer_path.resolve()), + "architecture": architecture, + "architecture_sha256": architecture_sha256, + "tokenizer_sha256": _sha256_json(tokenizer), + } + + +def validate_mopd_model_compatibility( + actor_path: str | Path, + teacher_paths: Mapping[str, str | Path], + *, + actor_tokenizer_path: str | Path | None = None, +) -> dict[str, dict[str, object]]: + """Validate the persistent-teacher model and tokenizer invariants. + + Teachers share one resident controller, so every teacher must have the same + architecture. The actor may use a different architecture, but all models + must map tokens to the same IDs because teacher log-probabilities score actor + trajectories directly. + """ + if not teacher_paths: + raise ValueError("MOPD compatibility validation requires at least one teacher") + + if "actor" in teacher_paths: + raise ValueError("MOPD teacher ID 'actor' is reserved for the actor model") + + actor_fingerprint = model_fingerprint( + actor_path, + tokenizer_path=actor_tokenizer_path, + ) + teacher_fingerprints = { + teacher_id: model_fingerprint(path) + for teacher_id, path in teacher_paths.items() + } + fingerprints = {"actor": actor_fingerprint, **teacher_fingerprints} + + actor_tokenizer = actor_fingerprint["tokenizer_sha256"] + tokenizer_mismatches = [ + name + for name, fingerprint in fingerprints.items() + if fingerprint["tokenizer_sha256"] != actor_tokenizer + ] + if tokenizer_mismatches: + raise ValueError( + "MOPD actor and teachers must use compatible token-ID mappings; " + f"mismatched={tokenizer_mismatches}" + ) + + teacher_ids = list(teacher_paths) + reference_id = teacher_ids[0] + reference_architecture = teacher_fingerprints[reference_id]["architecture_sha256"] + architecture_mismatches = [ + teacher_id + for teacher_id in teacher_ids[1:] + if teacher_fingerprints[teacher_id]["architecture_sha256"] + != reference_architecture + ] + if architecture_mismatches: + raise ValueError( + "All MOPD teachers must share one architecture because checkpoints " + "are loaded into a persistent controller; reference=" + f"{reference_id!r}, mismatched={architecture_mismatches}" + ) + + return fingerprints diff --git a/areal/trainer/mopd/execution.py b/areal/trainer/mopd/execution.py new file mode 100644 index 0000000000..470f85dffd --- /dev/null +++ b/areal/trainer/mopd/execution.py @@ -0,0 +1,46 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Static execution planning for optional MOPD objectives.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from areal.api.cli_args import PPOConfig + + +@dataclass(frozen=True) +class MOPDExecutionPlan: + """Infrastructure and objective phases required by configured coefficients.""" + + requires_teacher_scoring: bool + requires_rl: bool + requires_critic: bool + requires_ref: bool + requires_prox_logp: bool + + @classmethod + def from_config(cls, config: PPOConfig) -> MOPDExecutionPlan | None: + """Derive one immutable plan before any engines are initialized.""" + if config.mopd is None: + return None + requires_rl = config.mopd.loss.rl_coefficient > 0 + requires_distillation_filter = ( + config.mopd.loss.distillation_coefficient > 0 + and ( + config.actor.m2_threshold is not None + or config.actor.rejection_sampling is not None + ) + ) + return cls( + requires_teacher_scoring=(config.mopd.loss.distillation_coefficient > 0), + requires_rl=requires_rl, + requires_critic=requires_rl and config.critic is not None, + requires_ref=( + requires_rl and config.actor.kl_ctl > 0 and config.ref is not None + ), + requires_prox_logp=( + (requires_rl or requires_distillation_filter) + and config.actor.should_compute_prox_logp() + ), + ) diff --git a/areal/trainer/mopd/loss.py b/areal/trainer/mopd/loss.py new file mode 100644 index 0000000000..90bdaf3ac7 --- /dev/null +++ b/areal/trainer/mopd/loss.py @@ -0,0 +1,185 @@ +# SPDX-License-Identifier: Apache-2.0 + +import math +from typing import Any + +import torch + +from areal.api.cli_args import MOPDLossConfig + +DEFAULT_MOPD_IMPORTANCE_RATIO_CAP = 5.0 + + +def _validate_mopd_tensor_shapes( + logprobs: torch.Tensor, + old_logprobs: torch.Tensor, + teacher_logp_sum: torch.Tensor, + teacher_weight_sum: torch.Tensor, + loss_mask: torch.Tensor, +) -> None: + expected_shape = logprobs.shape + tensors = { + "old_logprobs": old_logprobs, + "teacher_logp_sum": teacher_logp_sum, + "teacher_weight_sum": teacher_weight_sum, + "loss_mask": loss_mask, + } + for name, tensor in tensors.items(): + if tensor.shape != expected_shape: + raise ValueError( + f"{name} must have token shape {expected_shape}, got {tensor.shape}" + ) + + +def mopd_loss_fn( + logprobs: torch.Tensor, + old_logprobs: torch.Tensor, + teacher_logp_sum: torch.Tensor, + teacher_weight_sum: torch.Tensor, + loss_mask: torch.Tensor, + normalization_mask: torch.Tensor | None = None, + importance_ratio_cap: float = DEFAULT_MOPD_IMPORTANCE_RATIO_CAP, +) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: + """Compute a truncated-IS multi-teacher reverse-KL surrogate. + + ``teacher_logp_sum`` is ``sum_j(w_j * log pi_Tj)`` and + ``teacher_weight_sum`` is ``sum_j(w_j)`` at every response token. Teacher + weights are deliberately not normalized. The behavior-policy importance + ratio is capped with a stop-gradient weight to prevent exponential + overflow while retaining the score-function gradient. + """ + if ( + not isinstance(importance_ratio_cap, (int, float)) + or isinstance(importance_ratio_cap, bool) + or not math.isfinite(importance_ratio_cap) + or importance_ratio_cap <= 0 + ): + raise ValueError("importance_ratio_cap must be a finite positive number") + _validate_mopd_tensor_shapes( + logprobs, + old_logprobs, + teacher_logp_sum, + teacher_weight_sum, + loss_mask, + ) + if normalization_mask is not None and normalization_mask.shape != logprobs.shape: + raise ValueError( + "normalization_mask must have token shape " + f"{logprobs.shape}, got {normalization_mask.shape}" + ) + + mask = loss_mask.bool() + safe_logprobs = torch.where(mask, logprobs, torch.zeros_like(logprobs)) + detached_teacher_logp = torch.where( + mask, teacher_logp_sum.detach(), torch.zeros_like(teacher_logp_sum) + ) + detached_teacher_weight = torch.where( + mask, teacher_weight_sum.detach(), torch.zeros_like(teacher_weight_sum) + ) + detached_old_logprobs = torch.where( + mask, old_logprobs.detach(), torch.zeros_like(old_logprobs) + ) + active_inputs_are_finite = ( + torch.isfinite(safe_logprobs) + & torch.isfinite(detached_teacher_logp) + & torch.isfinite(detached_teacher_weight) + & torch.isfinite(detached_old_logprobs) + ).all() + torch._assert_async( + active_inputs_are_finite, + "MOPD loss inputs must be finite at active tokens", + ) + # Sanitize masked positions before nonlinear operations. Applying the mask + # only after exp() can leave an infinite intermediate whose backward is + # 0 * inf = NaN, even though that token contributes zero to the loss. + detached_log_ratio = safe_logprobs.detach() - detached_old_logprobs + bounded_importance_weight = torch.exp( + detached_log_ratio.clamp(max=math.log(importance_ratio_cap)) + ) + score_reward = torch.where( + mask, + detached_teacher_logp - (detached_teacher_weight * safe_logprobs.detach()), + torch.zeros_like(logprobs), + ) + # At forward time the carrier is exactly one. Its derivative is + # d log pi_theta, preserving the score-function gradient even when + # the detached importance ratio was clipped. + score_function_carrier = torch.exp(safe_logprobs - safe_logprobs.detach()) + importance_weight = bounded_importance_weight * score_function_carrier + per_token_loss = -(importance_weight * score_reward) + masked_per_token_loss = per_token_loss + denominator_mask = mask if normalization_mask is None else normalization_mask.bool() + denominator = denominator_mask.count_nonzero().clamp_min(1) + loss = masked_per_token_loss.sum() / denominator + + reverse_kl = bounded_importance_weight * (-score_reward) + stats = { + "loss": loss.detach(), + "loss_per_token": masked_per_token_loss.detach(), + "score_reward": torch.where(mask, score_reward, torch.zeros_like(score_reward)), + "importance_weight": torch.where( + mask, importance_weight.detach(), torch.zeros_like(importance_weight) + ), + "teacher_weight_sum": torch.where( + mask, + detached_teacher_weight, + torch.zeros_like(detached_teacher_weight), + ), + "reverse_kl": torch.where(mask, reverse_kl, torch.zeros_like(reverse_kl)), + } + return loss, stats + + +def compose_mopd_loss( + rl_loss: torch.Tensor, + *, + config: MOPDLossConfig | None, + logprobs: torch.Tensor | None = None, + old_logprobs: torch.Tensor | None = None, + teacher_logp_sum: torch.Tensor | None = None, + teacher_weight_sum: torch.Tensor | None = None, + loss_mask: torch.Tensor | None = None, + normalization_mask: torch.Tensor | None = None, +) -> tuple[torch.Tensor, dict[str, Any]]: + """Compose RL and MOPD objectives without changing the disabled RL path.""" + if config is None: + return rl_loss, {} + + required_tensors = { + "logprobs": logprobs, + "old_logprobs": old_logprobs, + "teacher_logp_sum": teacher_logp_sum, + "teacher_weight_sum": teacher_weight_sum, + "loss_mask": loss_mask, + } + missing = [name for name, tensor in required_tensors.items() if tensor is None] + if missing: + raise ValueError(f"MOPD loss inputs are missing: {', '.join(missing)}") + + assert logprobs is not None + assert old_logprobs is not None + assert teacher_logp_sum is not None + assert teacher_weight_sum is not None + assert loss_mask is not None + mopd_loss, stats = mopd_loss_fn( + logprobs=logprobs, + old_logprobs=old_logprobs, + teacher_logp_sum=teacher_logp_sum, + teacher_weight_sum=teacher_weight_sum, + loss_mask=loss_mask, + normalization_mask=normalization_mask, + importance_ratio_cap=config.importance_ratio_cap, + ) + + if config.rl_coefficient == 0: + total_loss = config.distillation_coefficient * mopd_loss + elif config.distillation_coefficient == 0: + total_loss = config.rl_coefficient * rl_loss + else: + total_loss = ( + config.rl_coefficient * rl_loss + + config.distillation_coefficient * mopd_loss + ) + + stats["total_loss"] = total_loss.detach() + return total_loss, stats diff --git a/areal/trainer/mopd/scoring.py b/areal/trainer/mopd/scoring.py new file mode 100644 index 0000000000..c29cf474f3 --- /dev/null +++ b/areal/trainer/mopd/scoring.py @@ -0,0 +1,45 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Controller surface for forward-only MOPD teacher scoring.""" + +from __future__ import annotations + +from typing import Any + +from areal.infra import TrainController + + +class MOPDTeacherController(TrainController): + """Dispatch scoring with pipeline-safe padding and active dummy outputs.""" + + def compute_logp_padded(self, data: list[dict[str, Any]]): + original_size = len(data) + pp_size = self.parallel_strategy.pp_size + min_microbatches = max( + 2 * pp_size if pp_size > 1 else 1, + self.config.mb_spec.n_mbs, + ) + min_items_per_dp = ( + ((min_microbatches + pp_size - 1) // pp_size) + * pp_size + * self.config.mb_spec.granularity + ) + args, kwargs = self._pad_eval_dispatch_args( + (data,), + {}, + group_size=1, + min_items_per_dp=min_items_per_dp, + items_per_dp_divisor=pp_size * self.config.mb_spec.granularity, + active_dummies=True, + ) + results = self._custom_function_call( + "compute_logp", *args, rpc_meta={"broadcast": True}, **kwargs + ) + if results is None: + return None, [] + return results[:original_size], results[original_size:] + + def assert_mopd_runtime_topology(self) -> None: + self._custom_function_call( + "assert_mopd_runtime_topology", rpc_meta={"broadcast": False} + ) diff --git a/areal/trainer/mopd/targets.py b/areal/trainer/mopd/targets.py new file mode 100644 index 0000000000..19d8e807a9 --- /dev/null +++ b/areal/trainer/mopd/targets.py @@ -0,0 +1,66 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import math +from typing import Any + +import torch + +MOPD_CONTRIBUTIONS_KEY = "_mopd_teacher_contributions" + + +def aggregate_mopd_targets( + data: list[dict[str, Any]] | None, +) -> list[dict[str, Any]] | None: + """Aggregate raw weighted teacher log-probabilities on an actor DP head.""" + if data is None: + return None + for trajectory in data: + route = trajectory.get("mopd_route") + contributions = trajectory.pop(MOPD_CONTRIBUTIONS_KEY, None) + if not isinstance(contributions, dict) or not contributions: + raise ValueError( + f"MOPD route {route!r} has no positive teacher contributions" + ) + + logp_sum: torch.Tensor | None = None + weight_sum: torch.Tensor | None = None + for teacher_id, contribution in contributions.items(): + if not isinstance(contribution, dict): + raise TypeError( + f"MOPD contribution from {teacher_id!r} must be a mapping" + ) + teacher_logp = contribution.get("logp") + weight = contribution.get("weight") + if not isinstance(teacher_logp, torch.Tensor): + raise TypeError( + f"MOPD contribution from {teacher_id!r} has no tensor logp" + ) + if ( + not isinstance(weight, (int, float)) + or isinstance(weight, bool) + or not math.isfinite(weight) + or weight <= 0 + ): + raise ValueError( + f"MOPD contribution weight from {teacher_id!r} must be " + "finite and positive" + ) + if logp_sum is not None and teacher_logp.shape != logp_sum.shape: + raise ValueError( + f"MOPD teacher logp shape mismatch for route {route!r}: " + f"expected {tuple(logp_sum.shape)}, got {tuple(teacher_logp.shape)}" + ) + weighted_logp = teacher_logp * weight + token_weight = torch.full_like(teacher_logp, weight) + logp_sum = weighted_logp if logp_sum is None else logp_sum + weighted_logp + weight_sum = ( + token_weight if weight_sum is None else weight_sum + token_weight + ) + + assert logp_sum is not None and weight_sum is not None + trajectory["mopd_teacher_logp_sum"] = logp_sum + trajectory["mopd_teacher_weight_sum"] = weight_sum + trajectory.pop("mopd_route", None) + return data diff --git a/areal/trainer/mopd/teacher_manager.py b/areal/trainer/mopd/teacher_manager.py new file mode 100644 index 0000000000..49a41847b4 --- /dev/null +++ b/areal/trainer/mopd/teacher_manager.py @@ -0,0 +1,414 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import json +import os +import shutil +import uuid +from collections.abc import Callable +from concurrent.futures import Future, ThreadPoolExecutor +from enum import Enum, auto +from pathlib import Path +from threading import Lock +from typing import Any, Protocol + +from areal.api import SaveLoadMeta +from areal.api.cli_args import MOPDConfig +from areal.infra.rpc.rtensor import RTensorDrainReceipt + + +class TeacherController(Protocol): + def compute_logp_padded( + self, data: list[dict[str, Any]] + ) -> tuple[list[Any] | None, list[Any]]: ... + + def assert_mopd_runtime_topology(self) -> None: ... + + def load(self, meta: SaveLoadMeta) -> None: ... + + def onload(self) -> None: ... + + def offload(self) -> None: ... + + def strict_clear_batches(self, *targets: Any) -> RTensorDrainReceipt: ... + + def destroy(self) -> None: ... + + +class TeacherManager(Protocol): + def pre_fetch(self, teacher_id: str) -> None: ... + + def load(self, teacher_id: str) -> TeacherController: ... + + def release(self, receipt: RTensorDrainReceipt) -> None: ... + + def close(self) -> None: ... + + +class TeacherManagerState(Enum): + """GPU residency and lifecycle state of a persistent teacher companion.""" + + EMPTY = auto() + RESIDENT = auto() + OFFLOADED = auto() + BROKEN = auto() + CLOSED = auto() + + +class DiskCheckpointProvider: + """Resolve teacher snapshots already available on shared storage.""" + + def __init__(self, config: MOPDConfig): + self._config = config + + def pre_fetch(self, teacher_id: str) -> None: + self._path(teacher_id) + + def resolve(self, teacher_id: str) -> Path: + return self._path(teacher_id) + + def consumed(self, teacher_id: str) -> None: + del teacher_id + + def close(self) -> None: + return + + def _path(self, teacher_id: str) -> Path: + try: + path = Path(self._config.teachers[teacher_id].path) + except KeyError as exc: + raise KeyError(f"Unknown MOPD teacher {teacher_id!r}") from exc + if not path.is_dir(): + raise FileNotFoundError( + f"Teacher checkpoint {teacher_id!r} is not a local directory: {path}" + ) + return path + + +class LocalMemoryCheckpointProvider: + """Stage at most one next teacher snapshot using atomic ready directories.""" + + _RUN_PREFIX = ".run-" + + def __init__(self, config: MOPDConfig): + self._config = config + self._root = Path(config.manager.staging_root) + self._root.mkdir(parents=True, exist_ok=True) + self._sweep_stale_runs() + self._run_dir = ( + self._root / f"{self._RUN_PREFIX}{os.getpid()}-{uuid.uuid4().hex}" + ) + self._run_dir.mkdir() + self._write_manifest() + self._executor = ThreadPoolExecutor( + max_workers=1, thread_name_prefix="mopd-copy" + ) + self._future: Future[Path] | None = None + self._future_teacher: str | None = None + self._ready_paths: dict[str, Path] = {} + self._closed = False + self._lock = Lock() + + def pre_fetch(self, teacher_id: str) -> None: + source = self._source_path(teacher_id) + with self._lock: + self._ensure_open() + if teacher_id in self._ready_paths or self._future_teacher == teacher_id: + return + if self._ready_paths: + ready_teacher = next(iter(self._ready_paths)) + raise RuntimeError( + "Local-memory MOPD provider already holds ready checkpoint " + f"{ready_teacher!r}" + ) + if self._future is not None: + if not self._future.done(): + raise RuntimeError( + "Local-memory MOPD provider already has one checkpoint in flight" + ) + self._finalize_future_locked() + if self._ready_paths: + raise RuntimeError( + "Local-memory MOPD provider already holds one ready checkpoint" + ) + self._check_capacity(source) + self._future_teacher = teacher_id + self._future = self._executor.submit( + self._copy_snapshot, teacher_id, source + ) + + def resolve(self, teacher_id: str) -> Path: + with self._lock: + self._ensure_open() + needs_prefetch = ( + teacher_id not in self._ready_paths + and self._future_teacher != teacher_id + ) + if needs_prefetch: + self.pre_fetch(teacher_id) + with self._lock: + self._finalize_future_locked(expected_teacher=teacher_id) + try: + return self._ready_paths[teacher_id] + except KeyError as exc: + raise RuntimeError( + f"Teacher checkpoint {teacher_id!r} was not staged" + ) from exc + + def consumed(self, teacher_id: str) -> None: + with self._lock: + path = self._ready_paths.pop(teacher_id, None) + if path is not None: + shutil.rmtree(path, ignore_errors=False) + + def close(self) -> None: + with self._lock: + if self._closed: + return + self._closed = True + future = self._future + if future is not None: + future.cancel() + if not future.cancelled(): + try: + future.result() + except Exception: # noqa: BLE001 + pass + self._executor.shutdown(wait=True, cancel_futures=True) + shutil.rmtree(self._run_dir, ignore_errors=True) + + def _ensure_open(self) -> None: + if self._closed: + raise RuntimeError("Local-memory MOPD provider is closed") + + def _source_path(self, teacher_id: str) -> Path: + try: + path = Path(self._config.teachers[teacher_id].path) + except KeyError as exc: + raise KeyError(f"Unknown MOPD teacher {teacher_id!r}") from exc + if not path.is_dir(): + raise FileNotFoundError( + f"Teacher checkpoint {teacher_id!r} is not a local directory: {path}" + ) + return path + + def _check_capacity(self, source: Path) -> None: + required = sum( + path.stat().st_size for path in source.rglob("*") if path.is_file() + ) + reserve = self._config.manager.min_free_bytes or 0 + available = shutil.disk_usage(self._root).free + if available < required + reserve: + raise OSError( + f"Insufficient staging space under {self._root}: need " + f"{required + reserve} bytes, have {available}" + ) + + def _copy_snapshot(self, teacher_id: str, source: Path) -> Path: + tmp = self._run_dir / f"{teacher_id}.tmp.{uuid.uuid4().hex}" + ready = self._run_dir / f"{teacher_id}.ready" + try: + shutil.copytree(source, tmp) + self._fsync_tree(tmp) + os.replace(tmp, ready) + self._fsync_directory(self._run_dir) + return ready + except BaseException: + shutil.rmtree(tmp, ignore_errors=True) + raise + + def _finalize_future_locked(self, expected_teacher: str | None = None) -> None: + if self._future is None: + return + teacher_id = self._future_teacher + if expected_teacher is not None and teacher_id != expected_teacher: + raise RuntimeError( + f"Checkpoint {teacher_id!r} is staged, not {expected_teacher!r}" + ) + ready = self._future.result() + assert teacher_id is not None + self._ready_paths[teacher_id] = ready + self._future = None + self._future_teacher = None + + def _write_manifest(self) -> None: + manifest = self._run_dir / "owner.json" + manifest.write_text(json.dumps({"pid": os.getpid()}), encoding="utf-8") + with manifest.open("rb") as stream: + os.fsync(stream.fileno()) + self._fsync_directory(self._run_dir) + + def _sweep_stale_runs(self) -> None: + for run_dir in self._root.glob(f"{self._RUN_PREFIX}*"): + manifest = run_dir / "owner.json" + try: + owner_pid = int(json.loads(manifest.read_text(encoding="utf-8"))["pid"]) + except ( + FileNotFoundError, + KeyError, + TypeError, + ValueError, + json.JSONDecodeError, + ): + continue + if not _pid_exists(owner_pid): + shutil.rmtree(run_dir, ignore_errors=True) + + @staticmethod + def _fsync_tree(root: Path) -> None: + for path in root.rglob("*"): + if path.is_file(): + with path.open("rb") as stream: + os.fsync(stream.fileno()) + for path in sorted( + (entry for entry in root.rglob("*") if entry.is_dir()), + key=lambda entry: len(entry.parts), + reverse=True, + ): + LocalMemoryCheckpointProvider._fsync_directory(path) + LocalMemoryCheckpointProvider._fsync_directory(root) + + @staticmethod + def _fsync_directory(path: Path) -> None: + fd = os.open(path, os.O_RDONLY) + try: + os.fsync(fd) + finally: + os.close(fd) + + +def _pid_exists(pid: int) -> bool: + try: + os.kill(pid, 0) + except ProcessLookupError: + return False + except PermissionError: + return True + return True + + +class PersistentTeacherManager: + """Keep one isolated teacher controller alive across training phases.""" + + def __init__( + self, + config: MOPDConfig, + controller_factory: Callable[[str], TeacherController], + ): + self._controller_factory = controller_factory + self._provider = ( + DiskCheckpointProvider(config) + if config.manager.type == "disk" + else LocalMemoryCheckpointProvider(config) + ) + self._controller: TeacherController | None = None + self._loaded_teacher: str | None = None + self._state = TeacherManagerState.EMPTY + + @property + def controller(self) -> TeacherController | None: + return self._controller + + @property + def state(self) -> TeacherManagerState: + return self._state + + def pre_fetch(self, teacher_id: str) -> None: + self._ensure_usable() + if teacher_id != self._loaded_teacher: + self._provider.pre_fetch(teacher_id) + + def load(self, teacher_id: str) -> TeacherController: + self._ensure_usable() + needs_checkpoint = ( + self._state is TeacherManagerState.EMPTY + or self._loaded_teacher != teacher_id + ) + path = self._provider.resolve(teacher_id) if needs_checkpoint else None + loaded = False + try: + if self._state is TeacherManagerState.EMPTY: + assert path is not None + self._controller = self._controller_factory(str(path)) + self._state = TeacherManagerState.RESIDENT + else: + assert self._controller is not None + if self._state is TeacherManagerState.OFFLOADED: + self._controller.onload() + self._state = TeacherManagerState.RESIDENT + if self._loaded_teacher != teacher_id: + assert path is not None + self._controller.load( + SaveLoadMeta( + path=str(path), + weight_format="hf", + with_optim=False, + ) + ) + self._loaded_teacher = teacher_id + loaded = True + assert self._controller is not None + return self._controller + except BaseException as exc: + self._break_controller(exc) + raise + finally: + if loaded: + self._provider.consumed(teacher_id) + + def release(self, receipt: RTensorDrainReceipt) -> None: + self._ensure_usable() + if receipt.consumer_role != "actor": + raise RuntimeError( + "Cannot release MOPD teacher without an actor RTensor drain receipt" + ) + if self._state in ( + TeacherManagerState.EMPTY, + TeacherManagerState.OFFLOADED, + ): + return + assert self._controller is not None + try: + self._controller.offload() + self._state = TeacherManagerState.OFFLOADED + except BaseException as exc: + self._break_controller(exc) + raise + + def close(self) -> None: + if self._state is TeacherManagerState.CLOSED: + return + try: + self._destroy_controller() + finally: + try: + self._provider.close() + finally: + self._state = TeacherManagerState.CLOSED + + def _break_controller(self, cause: BaseException) -> None: + self._state = TeacherManagerState.BROKEN + try: + self._destroy_controller() + except BaseException as cleanup_error: + cause.add_note( + "Persistent MOPD teacher cleanup also failed: " + f"{type(cleanup_error).__name__}: {cleanup_error}" + ) + + def _destroy_controller(self) -> None: + controller = self._controller + if controller is None: + return + try: + controller.destroy() + finally: + self._controller = None + self._loaded_teacher = None + + def _ensure_usable(self) -> None: + if self._state is TeacherManagerState.CLOSED: + raise RuntimeError("MOPD TeacherManager is closed") + if self._state is TeacherManagerState.BROKEN: + raise RuntimeError("Persistent MOPD teacher companion is broken") diff --git a/areal/trainer/mopd/teacher_phase.py b/areal/trainer/mopd/teacher_phase.py new file mode 100644 index 0000000000..24d1550598 --- /dev/null +++ b/areal/trainer/mopd/teacher_phase.py @@ -0,0 +1,259 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Narrow lifecycle transaction for one MOPD teacher-scoring phase.""" + +from __future__ import annotations + +from typing import Any, Protocol + +from areal.api.cli_args import MOPDConfig +from areal.infra.rpc.rtensor import RTensorDrainReceipt +from areal.trainer.mopd.targets import MOPD_CONTRIBUTIONS_KEY +from areal.trainer.mopd.teacher_manager import ( + PersistentTeacherManager, + TeacherController, + TeacherManagerState, +) +from areal.utils import logging + +logger = logging.getLogger("MOPDTeacherPhase") + + +class MOPDTargetActor(Protocol): + """Actor operations required by the teacher transaction.""" + + def assert_mopd_runtime_topology(self) -> None: ... + + def aggregate_mopd_targets( + self, batch: list[dict[str, Any]] + ) -> list[dict[str, Any]]: ... + + def strict_clear_batches(self, *targets: Any) -> RTensorDrainReceipt: ... + + +class BatchDrainer(Protocol): + """One role that may have localized an RTensor batch.""" + + def strict_clear_batches(self, *targets: Any) -> RTensorDrainReceipt: ... + + +class MOPDTeacherPhase: + """Select, score, aggregate, drain, and release MOPD teachers atomically.""" + + def __init__( + self, + *, + config: MOPDConfig, + manager: PersistentTeacherManager, + actor: MOPDTargetActor, + critic: BatchDrainer | None = None, + ref: BatchDrainer | None = None, + ) -> None: + self._config = config + self._manager = manager + self._actor = actor + self._critic = critic + self._ref = ref + + def materialize(self, rollout_batch: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Materialize actor-owned MOPD targets and release teacher residency.""" + teacher_group_weights = self._resolve_teacher_group_weights(rollout_batch) + required_teachers = [ + teacher_id + for teacher_id in self._config.teachers + if any( + weights.get(teacher_id, 0.0) > 0 for weights in teacher_group_weights + ) + ] + if not required_teachers: + raise ValueError("MOPD batch does not require any positive-weight teacher") + + teacher_outputs: list[Any] = [] + teacher_controllers: list[TeacherController] = [] + receipts: dict[str, RTensorDrainReceipt] = {} + release_attempted = False + + def drain_critic() -> None: + if self._critic is not None: + receipts["critic"] = self._critic.strict_clear_batches(rollout_batch) + + def drain_ref() -> None: + if self._ref is not None: + receipts["ref"] = self._ref.strict_clear_batches(rollout_batch) + + def drain_teachers() -> None: + for index, controller in enumerate(teacher_controllers): + receipts[f"teacher:{index}"] = controller.strict_clear_batches( + rollout_batch, teacher_outputs + ) + + def drain_actor() -> None: + receipts["actor"] = self._actor.strict_clear_batches( + rollout_batch, teacher_outputs + ) + + try: + self._actor.assert_mopd_runtime_topology() + self._manager.pre_fetch(required_teachers[0]) + for teacher_index, teacher_id in enumerate(required_teachers): + was_offloaded = self._manager.state is TeacherManagerState.OFFLOADED + controller = self._manager.load(teacher_id) + if all(existing is not controller for existing in teacher_controllers): + teacher_controllers.append(controller) + if was_offloaded: + logger.info("[MOPD] teacher onload complete") + controller.assert_mopd_runtime_topology() + if teacher_index + 1 < len(required_teachers): + self._manager.pre_fetch(required_teachers[teacher_index + 1]) + + indices = [ + index + for index, weights in enumerate(teacher_group_weights) + if weights.get(teacher_id, 0.0) > 0 + ] + subset = [ + { + key: value + for key, value in rollout_batch[index].items() + if key != MOPD_CONTRIBUTIONS_KEY + } + for index in indices + ] + logps, dummy_logps = controller.compute_logp_padded(subset) + if logps is None or len(logps) != len(indices): + raise RuntimeError( + f"MOPD teacher {teacher_id!r} returned an invalid logp batch" + ) + teacher_outputs.extend(logps) + teacher_outputs.extend(dummy_logps) + for index, logp in zip(indices, logps, strict=True): + rollout_batch[index].setdefault(MOPD_CONTRIBUTIONS_KEY, {})[ + teacher_id + ] = { + "logp": logp, + "weight": teacher_group_weights[index][teacher_id], + } + + aggregated = self._actor.aggregate_mopd_targets(rollout_batch) + drain_critic() + drain_ref() + drain_teachers() + drain_actor() + self._require_all_receipts(receipts, teacher_controllers) + release_attempted = True + self._manager.release(receipts["actor"]) + logger.info("[MOPD] teacher offload complete") + return aggregated + except BaseException: + if self._critic is not None and "critic" not in receipts: + self._emergency_drain("critic", drain_critic) + if self._ref is not None and "ref" not in receipts: + self._emergency_drain("reference", drain_ref) + if any( + f"teacher:{index}" not in receipts + for index in range(len(teacher_controllers)) + ): + self._emergency_drain("teacher", drain_teachers) + if "actor" not in receipts: + self._emergency_drain("actor", drain_actor) + try: + self._require_all_receipts(receipts, teacher_controllers) + except RuntimeError: + try: + self._manager.close() + except Exception: + logger.error( + "MOPD teacher phase forced close failed", exc_info=True + ) + else: + if release_attempted: + try: + self._manager.close() + except Exception: + logger.error( + "MOPD teacher phase release rollback failed", + exc_info=True, + ) + else: + try: + self._manager.release(receipts["actor"]) + except Exception: + try: + self._manager.close() + except Exception: + logger.error( + "MOPD teacher phase release rollback failed", + exc_info=True, + ) + raise + + def close(self) -> None: + """Close the persistent teacher manager owned by this phase.""" + self._manager.close() + + def _resolve_teacher_group_weights( + self, rollout_batch: list[dict[str, Any]] + ) -> list[dict[str, float]]: + teacher_group_weights: list[dict[str, float]] = [] + for trajectory in rollout_batch: + teacher_group = trajectory.get("mopd_route") + if ( + not isinstance(teacher_group, str) + or teacher_group not in self._config.teacher_groups + ): + raise ValueError( + f"Unknown or missing MOPD teacher group {teacher_group!r}" + ) + teacher_group_weights.append(self._config.teacher_groups[teacher_group]) + return teacher_group_weights + + @staticmethod + def _emergency_drain( + role: str, + drain: Any, + ) -> None: + try: + drain() + except Exception: + logger.error( + "MOPD emergency %s RTensor drain failed; forcing phase teardown", + role, + exc_info=True, + ) + + def _require_all_receipts( + self, + receipts: dict[str, RTensorDrainReceipt], + teacher_controllers: list[TeacherController], + ) -> None: + expected = {"actor"} + if self._critic is not None: + expected.add("critic") + if self._ref is not None: + expected.add("ref") + expected.update(f"teacher:{index}" for index in range(len(teacher_controllers))) + missing = sorted(expected - receipts.keys()) + if missing: + raise RuntimeError( + "MOPD RTensor consumers did not acknowledge drain: " + + ", ".join(missing) + ) + expected_roles = { + "actor": "actor", + "critic": "critic", + "ref": "ref", + **{ + f"teacher:{index}": "mopd-teacher" + for index in range(len(teacher_controllers)) + }, + } + mismatched = [ + f"{key}={receipts[key].consumer_role!r}" + for key, role in expected_roles.items() + if key in expected and receipts[key].consumer_role != role + ] + if mismatched: + raise RuntimeError( + "MOPD RTensor drain receipts have unexpected consumer roles: " + + ", ".join(mismatched) + ) diff --git a/areal/trainer/ppo/actor.py b/areal/trainer/ppo/actor.py index ffa21ce17f..f5f7289f9e 100644 --- a/areal/trainer/ppo/actor.py +++ b/areal/trainer/ppo/actor.py @@ -7,10 +7,17 @@ import torch from areal.api import TrainEngine -from areal.api.cli_args import MicroBatchSpec, PPOActorConfig, RejectionSamplingConfig +from areal.api.cli_args import ( + MicroBatchSpec, + MOPDLossConfig, + PPOActorConfig, + RejectionSamplingConfig, +) from areal.engine.core import stage_batch_for_engine from areal.infra import TrainController from areal.infra.rpc.serialization import serialize_value +from areal.trainer.mopd.loss import compose_mopd_loss +from areal.trainer.mopd.targets import aggregate_mopd_targets from areal.trainer.ppo.gae import ( _build_gae_lambda_context, _compute_token_level_gae, @@ -37,6 +44,7 @@ split_padded_tensor_dict_into_mb_list, ) from areal.utils.functional import ( + apply_rejection_sampling, cispo_loss_fn, ppo_actor_loss_fn, reward_overlong_penalty, @@ -112,10 +120,17 @@ def __init__(self, config: PPOActorConfig, engine: TrainEngine): self.temperature = config.temperature self.m2_threshold = config.m2_threshold + self._mopd_loss_config: MOPDLossConfig | None = None # Log critical GSPO/GRPO configuration for reproducibility self._log_configuration() + def configure_mopd_loss(self, config: MOPDLossConfig) -> None: + """Bind static MOPD loss settings once on each actor worker.""" + if self._mopd_loss_config is not None and self._mopd_loss_config != config: + raise RuntimeError("MOPD loss configuration is already bound") + self._mopd_loss_config = config + def _log_configuration(self): """Log PPO configuration including how proximal policy is computed.""" config = self.config @@ -189,10 +204,40 @@ def _compute_logp(self, data: dict[str, Any]) -> torch.Tensor | None: aggregate_fn=lambda xs: torch.cat(xs, dim=-1), ) + def aggregate_mopd_targets( + self, + data: list[dict[str, Any]] | None = None, + ) -> list[dict[str, Any]] | None: + """Fetch-localized teacher contributions and create actor-owned targets.""" + return aggregate_mopd_targets(data) + + def assert_mopd_runtime_topology(self) -> None: + """Validate the live Megatron process groups used for MOPD scoring.""" + self.engine.assert_mopd_runtime_topology() + @trace_perf("ppo_actor.compute_advantages", category="compute") def compute_advantages(self, data: list[dict[str, Any]]) -> list[dict[str, Any]]: return batched_call(self._compute_advantages, data, pass_meta=True) + @trace_perf("ppo_actor.prepare_mopd_batch", category="compute") + def prepare_mopd_batch(self, data: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Align pure-distillation inputs without computing rewards or GAE.""" + return batched_call(self._prepare_mopd_batch, data) + + def _prepare_mopd_batch(self, data: dict[str, Any]) -> dict[str, Any]: + if self._mopd_loss_config is None: + raise RuntimeError("MOPD loss configuration is not bound") + if self._mopd_loss_config.rl_coefficient != 0: + raise RuntimeError("prepare_mopd_batch is only valid for pure distillation") + if "mopd_teacher_logp_sum" not in data: + raise RuntimeError("Pure MOPD distillation requires teacher targets") + loss_mask = torch.roll(data["loss_mask"].float(), shifts=-1, dims=-1) + behavior_logp = torch.roll(data["logprobs"], shifts=-1, dims=-1) + data["mopd_behavior_logprobs"] = (behavior_logp * loss_mask).detach() + data["logprobs"] = behavior_logp * loss_mask + data["loss_mask"] = loss_mask + return data + def _compute_advantages( self, data: dict[str, Any], meta: TrajBatchMeta | None = None ) -> dict[str, Any]: @@ -232,6 +277,14 @@ def _compute_advantages( loss_mask = data["loss_mask"].float() loss_mask = torch.roll(loss_mask, shifts=-1, dims=-1) + if "mopd_teacher_logp_sum" in data: + # MOPD's correction is always relative to the immutable rollout + # behavior policy, even when standard PPO recomputes its proximal + # policy and overwrites ``logprobs`` below. + data["mopd_behavior_logprobs"] = ( + torch.roll(data["logprobs"], shifts=-1, dims=-1) * loss_mask + ).detach() + # Align structural turn IDs to the same next-token prediction # convention used by loss_mask and log probabilities. turn_ids = data.get("turn_ids") @@ -429,12 +482,17 @@ def _ppo_update(self, data: dict[str, Any]) -> None: incorrect_seq_len=seqlens.float(), denominator="incorrect_n_seqs" ) - stats = dict( - advantages=data["advantages"], - kl_rewards=data["kl_rewards"], - final_reward=data["tot_rewards"], + pure_mopd_distillation = ( + self._mopd_loss_config is not None + and self._mopd_loss_config.rl_coefficient == 0 ) - stats_tracker.stat(**stats, denominator="n_valid_tokens") + if not pure_mopd_distillation: + stats_tracker.stat( + advantages=data["advantages"], + kl_rewards=data["kl_rewards"], + final_reward=data["tot_rewards"], + denominator="n_valid_tokens", + ) prompt_lens = _infer_prompt_lens(data["attention_mask"], data["loss_mask"]) seq_truncated_mask = _get_truncated_mask(data, seqlens) @@ -505,6 +563,7 @@ def _ppo_update(self, data: dict[str, Any]) -> None: sapo_tau_neg=self.config.sapo_tau_neg, use_cispo_loss=self.config.use_cispo_loss, use_decoupled_loss=self.config.use_decoupled_loss, + mopd_loss_config=self._mopd_loss_config, ), loss_weight_fn=lambda x: x["loss_mask"].count_nonzero(), ) @@ -512,16 +571,67 @@ def _ppo_update(self, data: dict[str, Any]) -> None: class PPOActorController(TrainController): + def configure_mopd_loss(self, config: MOPDLossConfig) -> None: + self._custom_function_call( + "configure_mopd_loss", config, rpc_meta={"broadcast": True} + ) + def compute_logp(self, *args, **kwargs): return self._custom_function_call( "compute_logp", *args, rpc_meta={"broadcast": True}, **kwargs ) + def compute_logp_padded(self, data: list[dict[str, Any]]): + """Compute logp with DP/PP padding and retain dummy outputs for drain.""" + original_size = len(data) + pp_size = self.parallel_strategy.pp_size + min_microbatches = max( + 2 * pp_size if pp_size > 1 else 1, + self.config.mb_spec.n_mbs, + ) + min_items_per_dp = ( + ((min_microbatches + pp_size - 1) // pp_size) + * pp_size + * self.config.mb_spec.granularity + ) + args, kwargs = self._pad_eval_dispatch_args( + (data,), + {}, + group_size=1, + min_items_per_dp=min_items_per_dp, + items_per_dp_divisor=pp_size * self.config.mb_spec.granularity, + active_dummies=True, + ) + results = self._custom_function_call( + "compute_logp", *args, rpc_meta={"broadcast": True}, **kwargs + ) + if results is None: + return None, [] + return results[:original_size], results[original_size:] + def compute_advantages(self, *args, **kwargs): return self._custom_function_call( "compute_advantages", *args, rpc_meta={"broadcast": True}, **kwargs ) + def prepare_mopd_batch(self, *args, **kwargs): + return self._custom_function_call( + "prepare_mopd_batch", *args, rpc_meta={"broadcast": True}, **kwargs + ) + + def aggregate_mopd_targets(self, *args, **kwargs): + return self._custom_function_call( + "aggregate_mopd_targets", + *args, + rpc_meta={"broadcast": False}, + **kwargs, + ) + + def assert_mopd_runtime_topology(self) -> None: + self._custom_function_call( + "assert_mopd_runtime_topology", rpc_meta={"broadcast": False} + ) + def ppo_update(self, *args, **kwargs) -> None: self._custom_function_call( "ppo_update", *args, rpc_meta={"broadcast": True}, **kwargs @@ -568,6 +678,7 @@ def grpo_loss_fn( sapo_tau_neg: float = 1.05, use_cispo_loss: bool = False, use_decoupled_loss: bool = False, + mopd_loss_config: MOPDLossConfig | None = None, vocab_min_logits: torch.Tensor | None = None, vocab_max_logits: torch.Tensor | None = None, vocab_mean_logits: torch.Tensor | None = None, @@ -575,9 +686,70 @@ def grpo_loss_fn( ): """Loss function for actor step, all inputs should be splitted into pipeline micro batches, returns loss and logging stats.""" + loss_mask = input_data["loss_mask"].bool() + if mopd_loss_config is not None and mopd_loss_config.rl_coefficient == 0: + teacher_logp_sum = input_data.get("mopd_teacher_logp_sum") + teacher_weight_sum = input_data.get("mopd_teacher_weight_sum") + behavior_logp = input_data.get("mopd_behavior_logprobs") + if teacher_logp_sum is None or teacher_weight_sum is None: + raise RuntimeError("Pure MOPD distillation requires teacher targets") + if behavior_logp is None: + raise RuntimeError( + "MOPD targets require immutable rollout behavior log-probabilities" + ) + normalization_mask = loss_mask + prox_logp_gt = input_data.get("prox_logp") + if m2_threshold is not None or rejection_sampling is not None: + prox_logp = _resolve_proximal_logp( + prox_logp_gt=prox_logp_gt, + prox_logp_method=prox_logp_method, + old_logp=input_data["logprobs"], + logprobs=logprobs.detach(), + versions=input_data.get("versions"), + current_version=current_version, + ) + if m2_threshold is not None: + loss_mask = _apply_m2po_masking( + input_data["logprobs"], prox_logp, loss_mask, m2_threshold + ) + normalization_mask = loss_mask + if rejection_sampling is not None: + loss_mask = apply_rejection_sampling( + proximal_logprobs=prox_logp, + old_logprobs=input_data["logprobs"], + loss_mask=loss_mask, + cu_seqlens=input_data.get("cu_seqlens"), + config=rejection_sampling, + ).loss_mask + loss, mopd_stats = compose_mopd_loss( + logprobs.new_zeros(()), + config=mopd_loss_config, + logprobs=logprobs, + old_logprobs=behavior_logp, + teacher_logp_sum=teacher_logp_sum, + teacher_weight_sum=teacher_weight_sum, + loss_mask=loss_mask, + normalization_mask=normalization_mask, + ) + stats_tracker.denominator( + n_tokens=infer_token_denominator(input_data, loss_mask), + n_valid_tokens=normalization_mask, + n_mopd_tokens=normalization_mask, + ) + stats_tracker.stat( + mopd_loss=mopd_stats["loss_per_token"].float(), + mopd_reward=mopd_stats["score_reward"].float(), + mopd_importance_weight=mopd_stats["importance_weight"].float(), + mopd_teacher_weight_sum=mopd_stats["teacher_weight_sum"].float(), + new_logp=logprobs.detach(), + old_logp=behavior_logp, + entropy=entropy.detach().float(), + denominator="n_mopd_tokens", + ) + return loss + old_logp = input_data["logprobs"] advantages = input_data["advantages"] - loss_mask = input_data["loss_mask"].bool() prox_logp_gt = input_data.get("prox_logp") # Could be None if skipped entropy = entropy.detach() @@ -653,36 +825,67 @@ def grpo_loss_fn( cu_seqlens=input_data.get("cu_seqlens"), ) - # Joint Distillation KL Loss + # M2 is part of the shared training-validity contract. Behavioral + # rejection may narrow the MOPD numerator further, while its denominator + # stays at the pre-rejection count to avoid amplifying accepted tokens. + mopd_normalization_mask = loss_mask + mopd_loss_mask = stat.get("behave_mask", loss_mask).bool() + + # Multi-teacher on-policy distillation. The deprecated single-teacher + # fields retain their original joint-loss semantics for compatibility. teacher_logp = input_data.get("teacher_logp") + mopd_teacher_logp_sum = input_data.get("mopd_teacher_logp_sum") rkl_stat = None - if teacher_logp is not None: - # Coefficients for RL and Knowledge Distillation + mopd_stats = {} + if teacher_logp is not None and mopd_teacher_logp_sum is not None: + raise ValueError( + "teacher_logp and mopd_teacher_logp_sum cannot both be provided" + ) + if mopd_teacher_logp_sum is not None: + teacher_weight_sum = input_data.get("mopd_teacher_weight_sum") + if mopd_loss_config is None: + raise RuntimeError( + "MOPD targets require actor-local MOPDLossConfig initialization" + ) + behavior_logp = input_data.get("mopd_behavior_logprobs") + if behavior_logp is None: + raise RuntimeError( + "MOPD targets require immutable rollout behavior log-probabilities" + ) + loss, mopd_stats = compose_mopd_loss( + loss, + config=mopd_loss_config, + logprobs=logprobs, + old_logprobs=behavior_logp, + teacher_logp_sum=mopd_teacher_logp_sum, + teacher_weight_sum=teacher_weight_sum, + loss_mask=mopd_loss_mask, + normalization_mask=mopd_normalization_mask, + ) + rkl_stat = mopd_stats["reverse_kl"].float() + elif mopd_loss_config is not None: + if mopd_loss_config.distillation_coefficient != 0: + raise RuntimeError("MOPD distillation is enabled but targets are missing") + loss = mopd_loss_config.rl_coefficient * loss + elif teacher_logp is not None: rl_loss_weight = input_data.get("rl_loss_weight", 1.0) distill_loss_weight = input_data.get("distill_loss_weight", 0.005) - - teacher_logp = ( - teacher_logp.detach() - ) # detach to prevent gradient backprop to teacher + teacher_logp = teacher_logp.detach() if rl_loss_weight == 0: - # Pure KD using reverse KL (importance-sampling) rkl_reward = teacher_logp - logprobs.detach() importance_weight = torch.exp(logprobs - old_logp) - rkl_weighted_term = importance_weight * rkl_reward * loss_mask - - kd_coef = -1 * distill_loss_weight - loss = kd_coef * rkl_weighted_term.sum() / loss_mask.sum().clamp(min=1) - - rkl_stat = -1 * rkl_weighted_term + loss = ( + -distill_loss_weight + * rkl_weighted_term.sum() + / loss_mask.sum().clamp(min=1) + ) + rkl_stat = -rkl_weighted_term else: - # KDRL: Knowledge Distillation + Reinforcement Learning (joint loss) rkl_penalty_per_token = (logprobs - teacher_logp) * loss_mask rkl_penalty = rkl_penalty_per_token.sum() / loss_mask.sum().clamp(min=1) - loss = rl_loss_weight * loss + distill_loss_weight * rkl_penalty - rkl_stat = rkl_penalty_per_token # Log training statistics @@ -694,10 +897,20 @@ def grpo_loss_fn( ) if rkl_stat is not None: - stats_tracker.stat( - rkl_loss=rkl_stat, - denominator="n_valid_tokens", - ) + if mopd_stats: + stats_tracker.denominator(n_mopd_tokens=mopd_normalization_mask.bool()) + stats_tracker.stat( + mopd_loss=mopd_stats["loss_per_token"].float(), + mopd_reward=mopd_stats["score_reward"].float(), + mopd_importance_weight=mopd_stats["importance_weight"].float(), + mopd_teacher_weight_sum=mopd_stats["teacher_weight_sum"].float(), + denominator="n_mopd_tokens", + ) + else: + stats_tracker.stat( + rkl_loss=rkl_stat, + denominator="n_valid_tokens", + ) logp_diff = (old_logp - logprobs.detach()) * loss_mask stats_tracker.stat( diff --git a/areal/trainer/rl_trainer.py b/areal/trainer/rl_trainer.py index 0a32623454..793862b54a 100644 --- a/areal/trainer/rl_trainer.py +++ b/areal/trainer/rl_trainer.py @@ -34,6 +34,7 @@ ValidDatasetConfig, vLLMConfig, ) +from areal.dataset.mopd import RoutedDataset, is_remote_dataset from areal.engine import RemoteSGLangEngine, RemotevLLMEngine from areal.infra import ( LocalScheduler, @@ -46,6 +47,12 @@ from areal.infra.data_service.controller.config import DataServiceConfig from areal.infra.data_service.rdataset import RDataset from areal.infra.utils.concurrent import call_maybe_async +from areal.trainer.mopd.compatibility import validate_mopd_model_compatibility +from areal.trainer.mopd.execution import MOPDExecutionPlan +from areal.trainer.mopd.teacher_manager import ( + PersistentTeacherManager, +) +from areal.trainer.mopd.teacher_phase import MOPDTeacherPhase from areal.utils import logging, perf_tracer, seeding, stats_tracker from areal.utils.cleanup import run_batch_cleanups from areal.utils.dataloader import create_dataloader @@ -134,6 +141,7 @@ def _init_impl( logging.setup_file_logging(StatsLogger.get_log_path(config.stats_logger)) self.config = config + self.mopd_execution_plan = MOPDExecutionPlan.from_config(config) self._apply_dte_config_envvars() self.processor, self.tokenizer = load_hf_processor_and_tokenizer( config.tokenizer_path @@ -142,8 +150,8 @@ def _init_impl( if is_single_controller(): self.scheduler = self._init_scheduler() self.data_controller: DataController | None = None - self._train_rdataset: RDataset | None = None - self._valid_rdataset: RDataset | None = None + self._train_rdataset: RDataset | RoutedDataset | None = None + self._valid_rdataset: RDataset | RoutedDataset | None = None # Set seed. seeding.set_random_seed(config.seed, key=f"trainer{rank}") @@ -157,10 +165,22 @@ def _init_impl( self._should_offload_actor = ( self._should_offload_rollout or config.actor.offload ) - self._should_offload_critic = ( - config.critic is not None and config.critic.offload + requires_critic = ( + self.mopd_execution_plan.requires_critic + if self.mopd_execution_plan is not None + else config.critic is not None + ) + requires_ref = ( + self.mopd_execution_plan.requires_ref + if self.mopd_execution_plan is not None + else config.actor.kl_ctl > 0 and config.ref is not None + ) + self._should_offload_critic = bool( + requires_critic and config.critic is not None and config.critic.offload + ) + self._should_offload_ref = bool( + requires_ref and config.ref is not None and config.ref.offload ) - self._should_offload_ref = config.ref is not None and config.ref.offload self._should_offload_teacher = ( config.teacher is not None and config.teacher.offload ) @@ -197,13 +217,13 @@ def _init_impl( # Create models: actor, critic, ref — each with its own allocation. self.actor = self._create_train_engine(config.actor, self.actor_alloc) self.critic = None - if config.critic is not None: + if requires_critic and config.critic is not None: critic_alloc = ModelAllocation.from_str( config.critic.backend, name="critic" ) self.critic = self._create_critic(config.critic, critic_alloc) self.ref = None - if config.actor.kl_ctl > 0 and config.ref is not None: + if requires_ref and config.ref is not None: ref_alloc = ModelAllocation.from_str(config.ref.backend, name="ref") self.ref = self._create_train_engine(config.ref, ref_alloc) @@ -246,7 +266,7 @@ def _init_impl( ) else: assert train_dataset is not None - if is_single_controller() and isinstance(train_dataset, RDataset): + if is_single_controller() and is_remote_dataset(train_dataset): ds_cfg = DataServiceConfig.from_dataset_config( config.train_dataset, seed=config.seed ) @@ -275,7 +295,7 @@ def _init_impl( self.valid_dataloader: StatefulDataLoader | None = None if self.config.valid_dataset is not None and valid_dataset is not None: assert self.config.valid_dataset is not None - if is_single_controller() and isinstance(valid_dataset, RDataset): + if is_single_controller() and is_remote_dataset(valid_dataset): assert self.data_controller is not None valid_dataset.connect( self.data_controller, @@ -308,11 +328,14 @@ def _init_impl( * config.train_dataset.batch_size, train_batch_size=config.train_dataset.batch_size, ) + self._ft_spec = ft_spec # Initialize engines first — the scheduler must know about roles # before the data controller can colocate with them. engine_init_kwargs = {"addr": None, "ft_spec": ft_spec} self.actor.initialize(**engine_init_kwargs, role="actor") + if self.config.mopd is not None: + self.actor.configure_mopd_loss(self.config.mopd.loss) if self.critic is not None: self.critic.initialize(**engine_init_kwargs, role="critic") if self.ref is not None: @@ -371,6 +394,30 @@ def _init_impl( ): self.teacher = self._init_teacher_rollout(self.config.teacher.rollout) + self.mopd_teacher_manager: PersistentTeacherManager | None = None + self.mopd_teacher_phase: MOPDTeacherPhase | None = None + if ( + self.config.mopd is not None + and self.mopd_execution_plan is not None + and self.mopd_execution_plan.requires_teacher_scoring + ): + if self.config.mopd.manager.type == "local_memory" and not isinstance( + self.scheduler, LocalScheduler + ): + raise RuntimeError( + "MOPD local_memory staging requires a same-host LocalScheduler" + ) + self.mopd_teacher_manager = PersistentTeacherManager( + self.config.mopd, self._create_mopd_teacher_controller + ) + self.mopd_teacher_phase = MOPDTeacherPhase( + config=self.config.mopd, + manager=self.mopd_teacher_manager, + actor=self.actor, + critic=self.critic, + ref=self.ref, + ) + # Proxy worker initialization (lazy, for AgentWorkflow support) self._proxy_started = False @@ -546,6 +593,32 @@ def _offload_model(self, engine, role: str) -> None: ): engine.offload() + def _update_weights_and_publish_version( + self, meta: WeightUpdateMeta, new_version: int + ) -> None: + """Update weights, publish their version, then restore AWEX rollout.""" + self.actor.update_weights(meta) + + self.actor.set_version(new_version) + if self.critic is not None: + self.critic.set_version(new_version) + self.rollout.set_version(new_version) + if self.eval_rollout is not None: + self.eval_rollout.set_version(new_version) + + if not self._is_v1_awex_colocate(self.config): + return + + # The AWEX reader flushes all old cache entries while installing the + # new weights. Reallocate an empty KV pool only after every actor worker + # has returned, then let SGLang serve requests again. This must remain a + # controller-side call: invoking rollout RPCs from an actor worker creates + # a nested controller call while its update_weights collective is active. + self.rollout.abort_all_requests() + self.rollout.onload(tags=["cuda_graph"]) + self.rollout.onload(tags=["kv_cache"]) + call_maybe_async(self.rollout.continue_generation) + def _offload_rollout(self, is_eval: bool = False): rollout = self.rollout if not is_eval else self.eval_rollout if rollout is None: @@ -787,12 +860,49 @@ def train( self.rollout.offload(tags=["kv_cache"]) logger.info("[AWEX] colocate: offload weights...") self.rollout.offload(tags=["weights"]) - logger.info("[AWEX] colocate: offload done, onloading actor...") - self.actor.onload() + logger.info("[AWEX] colocate: offload cuda_graph...") + self.rollout.offload(tags=["cuda_graph"]) + try: + if self.mopd_teacher_phase is not None: + rollout_batch = self.mopd_teacher_phase.materialize( + rollout_batch + ) + logger.info("[AWEX] colocate: offload done, onloading actor...") + self.actor.onload() + except BaseException: + logger.error( + "AWEX teacher/train ownership transition failed; " + "restoring rollout owner", + exc_info=True, + ) + try: + self.actor.offload() + except Exception: + logger.error( + "Failed to re-offload actor during rollback", exc_info=True + ) + try: + self.rollout.onload(tags=["cuda_graph"]) + self.rollout.onload(tags=["weights"]) + self.rollout.onload(tags=["kv_cache"]) + call_maybe_async(self.rollout.continue_generation) + self.rollout.resume() + except Exception: + logger.error( + "Failed to restore rollout during AWEX rollback; " + "leaving large owners offloaded", + exc_info=True, + ) + raise if self._should_offload_actor: self._onload_model(self.actor, role="actor") - if config.actor.should_compute_prox_logp(): + should_compute_prox_logp = ( + self.mopd_execution_plan.requires_prox_logp + if self.mopd_execution_plan is not None + else config.actor.should_compute_prox_logp() + ) + if should_compute_prox_logp: with ( stats_tracker.record_timing("recompute_logp"), perf_tracer.trace_scope( @@ -814,8 +924,15 @@ def train( args={"global_step": global_step}, ), ): - adv_batch = self.actor.compute_advantages(rollout_batch) - self.actor.get_device_stats().log("compute advantages") + if ( + self.mopd_execution_plan is not None + and not self.mopd_execution_plan.requires_rl + ): + adv_batch = self.actor.prepare_mopd_batch(rollout_batch) + self.actor.get_device_stats().log("prepare MOPD batch") + else: + adv_batch = self.actor.compute_advantages(rollout_batch) + self.actor.get_device_stats().log("compute advantages") # Wait for async checkpoint staging to complete before modifying parameters self.saver.maybe_wait_for_staging() @@ -898,14 +1015,7 @@ def train( # Use versioned path for weight updates new_version = global_step + 1 versioned_meta = self.weight_update_meta.with_version(new_version) - self.actor.update_weights(versioned_meta) - - self.actor.set_version(new_version) - if self.critic is not None: - self.critic.set_version(new_version) - self.rollout.set_version(new_version) - if self.eval_rollout is not None: - self.eval_rollout.set_version(new_version) + self._update_weights_and_publish_version(versioned_meta, new_version) if not self._is_v1_awex_colocate(config): self._save_training_state( @@ -988,8 +1098,22 @@ def train( epoch=epoch, epoch_step=step, global_step=global_step ) - # Resume rollout - self.rollout.resume() + # Resume rollout only when another train step will consume it. + # + # The dispatcher may have queued overlap rollouts while producing + # the current batch. Resuming it after the final step lets those + # stale tasks hit the inference server while close() is already + # tearing workers down, producing noisy ConnectionRefused errors. + if not self._is_final_train_step( + global_step=global_step, max_steps=max_steps + ): + self.rollout.resume() + else: + logger.info( + "Skipping rollout resume after final training step " + "(global_step=%s)", + global_step, + ) self._save_perf_tracer(step=global_step) @@ -1024,25 +1148,65 @@ def _save_training_state( global_step=global_step, ) + def _is_final_train_step(self, *, global_step: int, max_steps: int) -> bool: + next_step = global_step + 1 + if next_step >= max_steps: + return True + return ( + self.config.total_train_steps is not None + and next_step >= self.config.total_train_steps + ) + def close(self): - self.saver.finalize() - if hasattr(self, "_train_rdataset") and self._train_rdataset is not None: - self._train_rdataset.close() - if hasattr(self, "_valid_rdataset") and self._valid_rdataset is not None: - self._valid_rdataset.close() - if hasattr(self, "data_controller") and self.data_controller is not None: - self.data_controller.destroy() - self.stats_logger.close() - if self.eval_rollout is not None: - self.eval_rollout.destroy() - self.rollout.destroy() - if self.teacher is not None: - self.teacher.destroy() - if self.ref is not None: - self.ref.destroy() - if self.critic is not None: - self.critic.destroy() - self.actor.destroy() + # P87: must tolerate a partially-constructed trainer (called from + # __init__'s failure path), and one engine's destroy() failure must + # not keep the remaining workers alive. + saver = getattr(self, "saver", None) + if saver is not None: + try: + saver.finalize() + except Exception: + logger.warning("saver.finalize() failed during close", exc_info=True) + for attr in ("_train_rdataset", "_valid_rdataset"): + rdataset = getattr(self, attr, None) + if rdataset is not None: + try: + rdataset.close() + except Exception: + logger.warning(f"{attr}.close() failed during close", exc_info=True) + data_controller = getattr(self, "data_controller", None) + if data_controller is not None: + try: + data_controller.destroy() + except Exception: + logger.warning( + "data_controller.destroy() failed during close", exc_info=True + ) + stats_logger = getattr(self, "stats_logger", None) + if stats_logger is not None: + try: + stats_logger.close() + except Exception: + logger.warning( + "stats_logger.close() failed during close", exc_info=True + ) + mopd_phase = getattr(self, "mopd_teacher_phase", None) + if mopd_phase is not None: + try: + mopd_phase.close() + except Exception: + logger.warning( + "mopd_teacher_phase.close() failed during close", exc_info=True + ) + for attr in ("eval_rollout", "rollout", "teacher", "ref", "critic", "actor"): + engine = getattr(self, attr, None) + if engine is not None: + try: + engine.destroy() + except Exception: + logger.warning( + f"{attr}.destroy() failed during close", exc_info=True + ) perf_tracer.save(force=True) def _config_perf_tracer(self): @@ -1175,6 +1339,44 @@ def _create_train_engine( actor.create_process_group(parallel_strategy=alloc.parallel) return actor + def _create_mopd_teacher_controller(self, checkpoint_path: str): + """Create one persistent fork teacher from its first checkpoint.""" + from areal.engine import MegatronScoringEngine + + assert self.config.mopd is not None + teacher_config = deepcopy(self.config.mopd.teacher_engine) + teacher_config.path = checkpoint_path + teacher_config.experiment_name = self.config.experiment_name + teacher_config.trial_name = self.config.trial_name + if teacher_config.optimizer is None and teacher_config.backend.startswith( + "megatron:" + ): + teacher_config.megatron.disable_grad_buffers_cpu_backup = True + teacher_alloc = ModelAllocation.from_str( + teacher_config.backend, name="mopd-teacher" + ) + if is_single_controller(): + controller = MegatronScoringEngine.as_controller( + teacher_config, self.scheduler + ) + else: + controller = MegatronScoringEngine(config=teacher_config) + controller.create_process_group(parallel_strategy=teacher_alloc.parallel) + try: + controller.initialize( + addr=None, + ft_spec=self._ft_spec, + role="mopd-teacher", + ) + # The scoring-only teacher needs DDP-flat-buffer CPU residency, + # without enabling TMS in the AWEX actor processes. + if teacher_config.backend.startswith("megatron:"): + controller.init_weight_residency_adapter() + except BaseException: + controller.destroy() + raise + return controller + def _create_critic( self, critic_config: PPOCriticConfig, alloc: ModelAllocation ) -> FSDPPPOCritic | MegatronPPOCritic | ArchonPPOCritic | PPOCriticController: @@ -1245,6 +1447,12 @@ def _init_rollout( if self._is_v1_awex_colocate(self.config): server_args["awex_colocate_mode"] = True server_args["awex_meta_server_addr"] = self._awex_meta_server_addr + # SGLang's release/resume endpoints are no-ops unless its + # torch-memory-saver regions were enabled at server startup. + # AWEX relies on those endpoints to hand the colocated GPU to + # teachers and the actor, so this is a correctness requirement + # rather than an optional inference tuning flag. + server_args["enable_memory_saver"] = True elif rollout_backend == "vllm": if self.config.rollout.return_routed_experts: raise ValueError( @@ -1281,6 +1489,12 @@ def _init_rollout( ) else: controller = engine_cls.as_controller(config, self.scheduler) + if ( + self.mopd_execution_plan is not None + and self.mopd_execution_plan.requires_teacher_scoring + and not is_eval + ): + controller.enable_mopd_routing() init_kwargs = dict( role="rollout", server_args=server_args, @@ -1593,6 +1807,29 @@ def _validate_cfg(self): f"actor._version ('{actor_version}') and rollout._version " f"('{rollout_version}') must match. Both must be 'v1' or both 'v2'." ) + if self.config.mopd is not None: + requires_teacher_scoring = ( + self.mopd_execution_plan is not None + and self.mopd_execution_plan.requires_teacher_scoring + ) + if requires_teacher_scoring and not is_single_controller(): + raise ValueError("MOPD currently requires single-controller mode") + if actor_version != "v1": + raise ValueError( + "MOPD objective configuration currently requires the v1 actor " + "and rollout controller API" + ) + if requires_teacher_scoring: + validate_mopd_model_compatibility( + self.config.actor.path, + { + teacher_id: teacher.path + for teacher_id, teacher in self.config.mopd.teachers.items() + }, + actor_tokenizer_path=( + self.config.tokenizer_path or self.config.actor.path + ), + ) def _requires_proxy_workflow(self, workflow: WorkflowLike | None) -> bool: """Check if workflow requires proxy workers (i.e., not a RolloutWorkflow). diff --git a/areal/utils/data.py b/areal/utils/data.py index 1b06dfe076..4f2ffeee94 100644 --- a/areal/utils/data.py +++ b/areal/utils/data.py @@ -1792,33 +1792,58 @@ def _compute_approx_kl( return log_ratio -def make_dummy_eval_item(template: dict[str, Any]) -> dict[str, Any]: +def make_dummy_eval_item( + template: dict[str, Any], *, active_attention: bool = False +) -> dict[str, Any]: """Create a zero-contribution dummy item matching *template*'s schema. Every tensor field is replaced with a minimal all-zeros tensor that - preserves dtype and device. ``attention_mask`` and ``loss_mask`` are - set to zero so that downstream loss/metric code treats the item as - contributing nothing. + preserves dtype, device, and all leading dimensions. Keeping the + trajectory group dimension is required when distributed ranks synchronize + their microbatch counts: a padded rank must be able to create as many + microbatches as a rank holding a real multi-sample trajectory. + ``attention_mask`` and ``loss_mask`` are normally zero so downstream + loss/metric code treats the item as contributing nothing. + ``active_attention=True`` creates one attended token per sequence for + pipeline evaluation; callers must discard its output. """ + from areal.infra.rpc.rtensor import RTensor - def _zero_tensor_like(tensor: torch.Tensor) -> torch.Tensor: - return torch.zeros((1, 1), dtype=tensor.dtype, device=tensor.device) + def _minimal_tensor_like( + tensor: torch.Tensor | RTensor, *, fill_value: int = 0 + ) -> torch.Tensor: + if isinstance(tensor, RTensor): + device = torch.device("cpu") + else: + device = tensor.device + shape = (*tensor.shape[:-1], 1) if tensor.ndim > 0 else (1,) + return torch.full(shape, fill_value, dtype=tensor.dtype, device=device) + + group_size = 1 + attention_mask = template.get("attention_mask") + if isinstance(attention_mask, (torch.Tensor, RTensor)) and attention_mask.ndim >= 2: + group_size = attention_mask.shape[0] dummy: dict[str, Any] = {} for key, value in template.items(): if key in {"attention_mask", "loss_mask"}: - if isinstance(value, torch.Tensor): - dummy[key] = _zero_tensor_like(value) + if isinstance(value, (torch.Tensor, RTensor)): + fill_value = int(active_attention and key == "attention_mask") + dummy[key] = _minimal_tensor_like(value, fill_value=fill_value) else: - dummy[key] = torch.zeros((1, 1), dtype=torch.bool) + dummy[key] = torch.full( + (1, 1), + int(active_attention and key == "attention_mask"), + dtype=torch.bool, + ) continue if key.startswith("multi_modal_input"): - dummy[key] = [{}] + dummy[key] = [{} for _ in range(group_size)] continue - if isinstance(value, torch.Tensor): - dummy[key] = _zero_tensor_like(value) + if isinstance(value, (torch.Tensor, RTensor)): + dummy[key] = _minimal_tensor_like(value) else: dummy[key] = copy.deepcopy(value) diff --git a/areal/utils/dataloader.py b/areal/utils/dataloader.py index 3f84389a1d..679bd0505f 100644 --- a/areal/utils/dataloader.py +++ b/areal/utils/dataloader.py @@ -8,6 +8,7 @@ from torchdata.stateful_dataloader import StatefulDataLoader from areal.api.cli_args import ValidDatasetConfig, _DatasetConfig +from areal.dataset.mopd import is_remote_dataset def create_dataloader( @@ -31,15 +32,16 @@ def create_dataloader( f"batch size({dataset_config.batch_size}) must be divisible by world_size({world_size})!" ) - from areal.infra.data_service.rdataset import RDataset, _PrefetchAwareSampler + from areal.infra.data_service.rdataset import _PrefetchAwareSampler drop_sampler_last = True if isinstance(dataset_config, ValidDatasetConfig): drop_sampler_last = False - if isinstance(dataset, RDataset) and isinstance(dataset_config, ValidDatasetConfig): + remote_dataset = is_remote_dataset(dataset) + if remote_dataset and isinstance(dataset_config, ValidDatasetConfig): sampler_cls = _PrefetchAwareEvalSampler - elif isinstance(dataset, RDataset): + elif remote_dataset: sampler_cls = _PrefetchAwareSampler elif isinstance(dataset_config, ValidDatasetConfig): sampler_cls = EvalDistributedSampler diff --git a/areal/utils/environ.py b/areal/utils/environ.py index 6cfbf07058..670d552cf8 100644 --- a/areal/utils/environ.py +++ b/areal/utils/environ.py @@ -7,6 +7,7 @@ logger = logging.getLogger("EnvironUtils") _warned_bool_env_var_keys = set() +_warned_numeric_env_var_values = set() _warned_rank_env_var_values = set() @@ -67,6 +68,34 @@ def get_bool_env_var( return value in truthy_values +def get_float_env_var(name: str, default: float) -> float: + value = os.getenv(name) + if value is None or not value.strip(): + return default + try: + return float(value) + except ValueError: + warn_key = (name, value) + if warn_key not in _warned_numeric_env_var_values: + logger.warning("Invalid %s=%r; using %s", name, value, default) + _warned_numeric_env_var_values.add(warn_key) + return default + + +def get_int_env_var(name: str, default: int) -> int: + value = os.getenv(name) + if value is None or not value.strip(): + return default + try: + return int(value) + except ValueError: + warn_key = (name, value) + if warn_key not in _warned_numeric_env_var_values: + logger.warning("Invalid %s=%r; using %s", name, value, default) + _warned_numeric_env_var_values.add(warn_key) + return default + + def is_in_ci(): return get_bool_env_var("AREAL_IS_IN_CI") diff --git a/areal/utils/logging.py b/areal/utils/logging.py index f54a384e1f..4cbfaa7e16 100644 --- a/areal/utils/logging.py +++ b/areal/utils/logging.py @@ -129,6 +129,8 @@ # AWEX weight exchange - cyan (compute backend) "AwexColocate": "light_cyan", "AwexColocateReader": "light_cyan", + "MegatronResidency": "light_cyan", + "MOPDTeacherPhase": "light_cyan", "AwexSGLangPlugin": "light_cyan", } diff --git a/docs/en/algorithms/mopd.md b/docs/en/algorithms/mopd.md new file mode 100644 index 0000000000..16617e5c4e --- /dev/null +++ b/docs/en/algorithms/mopd.md @@ -0,0 +1,143 @@ +# Multi-Teacher On-Policy Distillation (MOPD) + +MOPD combines on-policy reinforcement learning with token-level targets from one or +more teacher checkpoints. Each dataset source selects one required teacher group, +which assigns non-negative weights to any number of teachers. The weights are applied +directly and are not normalized. + +MOPD currently runs in single-controller mode with Megatron actor and teacher engines, +an SGLang rollout engine, and AWEX colocated weight transfer. Actor, rollout, and a +persistent teacher companion share the same GPUs but never own the large model weights +at the same time. + +## Runtime lifecycle + +Each training step follows three exclusive phases: + +1. **Rollout:** SGLang generates trajectories while AReaL propagates the source's + teacher group as internal task metadata. Dataset samples and workflow inputs need + no routing field. +2. **Teacher:** rollout weights and KV cache are offloaded. A forked Megatron teacher + process onloads, loads each required checkpoint, and scores its routed samples. The + actor materializes and clears all teacher RTensors before the teacher weights are + offloaded again. The companion process stays alive for reuse by the next step. +3. **Train:** the actor computes the configured RL and distillation loss, updates its + weights, and publishes the next version to SGLang through AWEX. SGLang drops stale + KV cache entries and allocates a new empty cache before generation continues. + +The actor and `mopd.teacher_engine` must use identical parallel strategies and world +sizes, including the pipeline parallel size. The current implementation requires +teacher and actor controllers to use v1. PP, TP, CP, DP, and EP values are validated by +the selected Megatron model and allocation. + +## Configuration + +Add `mopd` to a PPO configuration: + +```yaml +actor: + backend: "megatron:(attn:d1p1t4c2|ffn:d1p1e8)" + weight_update_mode: awex + +rollout: + backend: sglang:d8t1 + scheduling_strategy: {type: colocation, target: actor, fork: true} + +train_dataset: + mixture_sampling_policy: proportional + sources: + - {path: /data/code, type: rl, teacher_group: coding} + - {path: /data/mixed, type: rl, teacher_group: mixed} + +mopd: + teachers: + coder: {path: /models/teacher-coder} + reasoning: {path: /models/teacher-reasoning} + teacher_groups: + coding: {coder: 1.0} + mixed: {coder: 0.3, reasoning: 0.7} + teacher_engine: + backend: ${actor.backend} + optimizer: null + disable_dropout: true + scheduling_strategy: {type: colocation, target: actor, fork: true} + scheduling_spec: ${actor.scheduling_spec} + manager: + type: disk + staging_root: /dev/shm/areal-mopd + loss: + rl_coefficient: 0.0 + distillation_coefficient: 1.0 +``` + +When MOPD is enabled, every entry in `train_dataset.sources` must declare a +`teacher_group` that matches a key in `mopd.teacher_groups`; validation datasets follow +the same rule when configured. Outside MOPD, `teacher_group` defaults to `null`. +Samples cannot override their source's teacher group and need no `task_type` field. A +teacher group can reference any number of known teacher IDs and must contain at least +one positive weight. +`mixture_sampling_policy: proportional` preserves source-size proportions. +`uniform` gives every source the same number of samples per epoch by cycling shorter +sources deterministically before the distributed sampler shuffles global indices. + +`manager.type: disk` loads checkpoints from shared storage and supports multi-node +runs. `local_memory` asynchronously stages one upcoming checkpoint below +`staging_root`, atomically publishes it to the persistent teacher, and removes it +after loading. Because this path is visible only on the controller host, +`local_memory` requires `scheduler.type: local` and a single-node actor/teacher +topology. `min_free_bytes` can reserve free space below the staging root. + +For teacher weights $w_j$, define $S_T(a)=\sum_j w_j\log\pi_{T_j}(a)$ and +$W=\sum_j w_j$. MOPD minimizes the raw weighted reverse KL +$\sum_j w_j D_{KL}(\pi_\theta \parallel \pi_{T_j})$ with the on-policy +score-function surrogate: + +```text +rho(a) = min(exp(log pi_theta(a) - log pi_old(a)), importance_ratio_cap) +reward(a) = S_T(a) - W * stop_gradient(log pi_theta(a)) +mopd_loss = -mean(rho(a) * reward(a)) +loss = rl_coefficient * rl_loss + distillation_coefficient * mopd_loss +``` + +`importance_ratio_cap` defaults to `5.0` and bounds the importance-sampling +multiplier to prevent exponential overflow. + +This is a weighted sum of reverse-KL objectives, equivalently a geometric teacher +ensemble up to an additive constant. It is not teacher cross-entropy or an arithmetic +mixture of teacher probabilities. Route weights are applied directly and are not +normalized. + +Set `rl_coefficient: 0.0` for pure distillation. Set both coefficients to positive +values for joint RL and distillation. + +## Examples + +- `examples/mopd/gsm8k_qwen3_14b_to_0_6b.py` provides the local GSM8K entry point and + dry-run validator. +- `examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml` configures a Qwen3-14B teacher + and Qwen3-0.6B actor on one eight-GPU node. + +Validate a local configuration without starting workers: + +```bash +MOPD_STUDENT_MODEL_PATH=/models/Qwen3-0.6B \ +MOPD_TEACHER_MODEL_PATH=/models/Qwen3-14B \ +MOPD_GSM8K_PATH=/data/gsm8k \ +AREAL_ADMIN_API_KEY="$(python -c 'import secrets; print(secrets.token_hex(32))')" \ +python -m examples.mopd.gsm8k_qwen3_14b_to_0_6b \ + --config examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml \ + --dry-run +``` + +## Operational notes + +- Actor and teacher checkpoints must share the same token-ID mapping, and every model + architecture must be supported by its selected Megatron adapter. +- The teacher companion process is persistent, but its model weights must be offloaded + before actor ownership resumes. `DrainReceipt` is the phase reclamation boundary for + teacher RTensors. +- Keep W&B credentials and service endpoints in environment variables; do not add + them to YAML or shell files. +- A cancelled actor or rollout Slurm child job can be normal when the driver has + already completed and is cleaning up its persistent workers. Use the driver exit + code and final training-step log to determine success. diff --git a/docs/en/cli_reference.md b/docs/en/cli_reference.md index ff80f8166e..0d4948b5c7 100644 --- a/docs/en/cli_reference.md +++ b/docs/en/cli_reference.md @@ -76,8 +76,14 @@ For detailed examples, see the experiment configurations in the `examples/` dire - [ArchonFP8 Configuration](section-archon-fp8) - [DPO Configuration](section-dpo) - [DPOEngine Configuration](section-dpo-engine) +- [DatasetSource Configuration](section-dataset-source) - [DistributedDataParallel Configuration](section-distributed-data-parallel) - [FP8Engine Configuration](section-fp8-engine) +- [MOPD Configuration](section-mopd) +- [MOPDLoss Configuration](section-mopd-loss) +- [MOPDTeacherEngine Configuration](section-mopd-teacher-engine) +- [MOPDTeacherManager Configuration](section-mopd-teacher-manager) +- [MOPDTeacher Specification](section-mopd-teacher) - [MegatronEngine Configuration](section-megatron-engine) - [MemoryProfiler Configuration](section-memory-profiler) - [PerfTracer Configuration](section-perf-tracer) @@ -158,6 +164,7 @@ A dummy place holder of GRPO config for backward compatibility. | `ref` | [`PPOActorConfig`](section-ppo-actor) \| None | `None` | - | | `critic` | [`PPOCriticConfig`](section-ppo-critic) \| None | `None` | - | | `teacher` | [`TeacherConfig`](section-teacher) \| None | `None` | Optional teacher model configuration used for on-policy distillation during PPO training. If provided, the actor may be trained to match the teacher in addition to the standard PPO objective. | +| `mopd` | [`MOPDConfig`](section-mopd) \| None | `None` | Optional multi-teacher on-policy distillation config. | | `dynamic_bs` | boolean | `False` | Enable dynamic batch sizing in prepare_batch. When True, batch collection stops when (accepted + rejected) >= batch_size, returning only accepted results. This results in variable-sized batches of valid data. | (section-ppo)= @@ -197,6 +204,7 @@ Configuration for Proximal Policy Optimization (PPO) reinforcement learning expe | `ref` | [`PPOActorConfig`](section-ppo-actor) \| None | `None` | - | | `critic` | [`PPOCriticConfig`](section-ppo-critic) \| None | `None` | - | | `teacher` | [`TeacherConfig`](section-teacher) \| None | `None` | Optional teacher model configuration used for on-policy distillation during PPO training. If provided, the actor may be trained to match the teacher in addition to the standard PPO objective. | +| `mopd` | [`MOPDConfig`](section-mopd) \| None | `None` | Optional multi-teacher on-policy distillation config. | | `dynamic_bs` | boolean | `False` | Enable dynamic batch sizing in prepare_batch. When True, batch collection stops when (accepted + rejected) >= batch_size, returning only accepted results. This results in variable-sized batches of valid data. | (section-rw)= @@ -698,21 +706,23 @@ https://docs.vllm.ai/en/stable/api/index.html for detailed documentation. Configuration for training dataset loading and preprocessing. -| Parameter | Type | Default | Description | -| --------------------- | ---------------------------------------------- | ---------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `split` | string | `"train"` | Dataset split to use, e.g., 'train', 'test'. | -| `path` | string | **Required** | Path to the dataset. Can be a local path or a HuggingFace dataset name. | -| `type` | string | **Required** | Type of training method, e.g., 'sft', 'rl', etc. | -| `batch_size` | integer | `1` | Batch size for the dataloader | -| `shuffle` | boolean | `True` | Whether to shuffle the dataset | -| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | -| `num_workers` | integer | `0` | Number of worker processes for data loading | -| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | -| `drop_last` | boolean | `True` | Drop the last incomplete batch | -| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | -| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | -| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | -| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | +| Parameter | Type | Default | Description | +| ------------------------- | ------------------------------------------------------- | ---------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `split` | string | `"train"` | Dataset split to use, e.g., 'train', 'test'. | +| `path` | string \| None | `None` | Path to one dataset. Mutually exclusive with sources. | +| `type` | string \| None | `None` | Training data type for path. Mutually exclusive with sources. | +| `sources` | list of [`DatasetSourceConfig`](section-dataset-source) | `[]` | Dataset mixture sources. MOPD requires every source to declare a teacher_group. | +| `mixture_sampling_policy` | string | `"proportional"` | How a routed mixture represents sources in one epoch: 'proportional' preserves source-size proportions; 'uniform' balances source counts by deterministically cycling shorter sources. **Choices:** `proportional`, `uniform` | +| `batch_size` | integer | `1` | Batch size for the dataloader | +| `shuffle` | boolean | `True` | Whether to shuffle the dataset | +| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | +| `num_workers` | integer | `0` | Number of worker processes for data loading | +| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | +| `drop_last` | boolean | `True` | Drop the last incomplete batch | +| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | +| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | +| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | +| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | (section-valid-dataset)= @@ -723,21 +733,23 @@ Configuration for validation dataset loading and preprocessing. It has different default values with `TrainDatasetConfig`. `shuffle` and `drop_last` default to False. -| Parameter | Type | Default | Description | -| --------------------- | ---------------------------------------------- | ---------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `split` | string | `"test"` | Dataset split to use, e.g., 'train', 'test'. | -| `path` | string | **Required** | Path to the dataset. Can be a local path or a HuggingFace dataset name. | -| `type` | string | **Required** | Type of training method, e.g., 'sft', 'rl', etc. | -| `batch_size` | integer | `1` | Batch size for the dataloader | -| `shuffle` | boolean | `False` | Whether to shuffle the dataset | -| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | -| `num_workers` | integer | `0` | Number of worker processes for data loading | -| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | -| `drop_last` | boolean | `False` | Drop the last incomplete batch | -| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | -| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | -| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | -| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | +| Parameter | Type | Default | Description | +| ------------------------- | ------------------------------------------------------- | ---------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `split` | string | `"test"` | Dataset split to use, e.g., 'train', 'test'. | +| `path` | string \| None | `None` | Path to one dataset. Mutually exclusive with sources. | +| `type` | string \| None | `None` | Training data type for path. Mutually exclusive with sources. | +| `sources` | list of [`DatasetSourceConfig`](section-dataset-source) | `[]` | Dataset mixture sources. MOPD requires every source to declare a teacher_group. | +| `mixture_sampling_policy` | string | `"proportional"` | How a routed mixture represents sources in one epoch: 'proportional' preserves source-size proportions; 'uniform' balances source counts by deterministically cycling shorter sources. **Choices:** `proportional`, `uniform` | +| `batch_size` | integer | `1` | Batch size for the dataloader | +| `shuffle` | boolean | `False` | Whether to shuffle the dataset | +| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | +| `num_workers` | integer | `0` | Number of worker processes for data loading | +| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | +| `drop_last` | boolean | `False` | Drop the last incomplete batch | +| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | +| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | +| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | +| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | (section-cluster)= @@ -1051,6 +1063,21 @@ fields. | `beta` | float | `0.1` | KL penalty coefficient for DPO loss. | | `loss_type` | string | `"sigmoid"` | DPO loss variant. 'sigmoid': original DPO loss (Rafailov et al. 2023). 'ipo': Identity Preference Optimization with per-token length normalization (Azar et al. 2023). **Choices:** `sigmoid`, `ipo` | +(section-dataset-source)= + +## DatasetSource Configuration + +One source in a dataset mixture. + +| Parameter | Type | Default | Description | +| ---------------- | --------------- | ------------ | ---------------------------------------------------------- | +| `path` | string | **Required** | Local path or HuggingFace name for this dataset source. | +| `type` | string | **Required** | Training data type, for example 'rl'. | +| `teacher_group` | string \| None | `None` | Optional MOPD teacher group applied to this entire source. | +| `split` | string \| None | `None` | Optional split override for this dataset source. | +| `max_length` | integer \| None | `None` | Optional maximum sequence length for this source. | +| `dataset_kwargs` | `dict` | `{}` | Extra keyword arguments for this source's loader. | + (section-distributed-data-parallel)= ## DistributedDataParallel Configuration @@ -1098,6 +1125,103 @@ is disabled. | `num_layers_at_end_in_bf16` | integer | `1` | Number of layers at end to keep in BF16 when first_last_layers_bf16 is True. | | `direct_convert` | boolean | `True` | Whether to use direct FP8 conversion during weight updates and save/load. When True, FP8 parameters are directly converted between TE FP8 and PyTorch FP8 without intermediate dequantization/quantization. | +(section-mopd)= + +## MOPD Configuration + +Configuration for multi-teacher on-policy distillation. + +| Parameter | Type | Default | Description | +| ---------------- | ---------------------------------------------------------- | -------------------------- | ----------- | +| `teachers` | `dict` | `{}` | - | +| `teacher_groups` | `dict` | `{}` | - | +| `teacher_engine` | [`MOPDTeacherEngineConfig`](section-mopd-teacher-engine) | *MOPDTeacherEngineConfig* | - | +| `manager` | [`MOPDTeacherManagerConfig`](section-mopd-teacher-manager) | *MOPDTeacherManagerConfig* | - | +| `loss` | [`MOPDLossConfig`](section-mopd-loss) | *MOPDLossConfig* | - | + +(section-mopd-loss)= + +## MOPDLoss Configuration + +Coefficients for joint RL and multi-teacher distillation training. + +| Parameter | Type | Default | Description | +| -------------------------- | ----- | ------- | -------------------------------------------------- | +| `rl_coefficient` | float | `0.0` | Coefficient applied to the RL objective. | +| `distillation_coefficient` | float | `1.0` | Coefficient applied to the MOPD objective. | +| `importance_ratio_cap` | float | `5.0` | Positive cap applied to the behavior-policy ratio. | + +(section-mopd-teacher-engine)= + +## MOPDTeacherEngine Configuration + +Forward-only scoring engine configuration used by MOPD teachers. + +| Parameter | Type | Default | Description | +| ------------------------------- | --------------------------------------------------- | ---------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `experiment_name` | string | **Required** | - | +| `trial_name` | string | **Required** | - | +| `path` | string | `""` | Path to HuggingFace checkpoint | +| `attn_impl` | string | `"flash_attention_2"` | Attention implementation for huggingface transformers model. Accepts builtin transformers backends or a Hugging Face kernels repo ID formatted as org/repo\[@revision\]\[:entrypoint\]. **Choices:** `eager`, `sdpa`, `flash_attention_2`, `flash_attention_3`, `flex_attention` | +| `use_kernels` | boolean | `False` | Enable Hugging Face kernels model kernelization after model creation. | +| `init_from_scratch` | boolean | `False` | Initialize model weights randomly | +| `is_critic` | boolean | `False` | Whether to use a critic/reward model | +| `temperature` | float | `1.0` | Temperature during generation. | +| `logprobs_chunk_size` | integer | `1024` | Maximum sequence chunk size used to compute log probabilities and entropy. Must be positive. | +| `mb_spec` | [`MicroBatchSpec`](section-micro-batch) | *MicroBatchSpec* | - | +| `pad_to_maximum` | boolean | `False` | Whether to pad each microbatch to the length upper bound specified by mb_spec. Can reduce memory fragmentation but slows down training. | +| `disable_dropout` | boolean | `True` | Disable dropout for deterministic teacher scoring. | +| `gradient_checkpointing` | boolean | `False` | Enable gradient checkpointing | +| `dtype` | string | `"bfloat16"` | Forward/backward compute dtype. | +| `grad_reduce_dtype` | string | `"float32"` | Gradient reduction data type. | +| `optimizer_dtype` | string | `"float32"` | Underlying parameter storage dtype, also the dtype of optimizer states (exp_avg, exp_avg_sq) since torch.optim.AdamW inherits dtype from model.parameters(). Default 'float32' maintains fp32 master weights matching DeepSpeed ZeRO-3 and Megatron precision-aware optimizer behavior. FSDP2's MixedPrecisionPolicy(param_dtype=`dtype`) will still cast forward/backward computation to `dtype` (e.g. bfloat16). Set to 'bfloat16' together with optimizer.type='adam_bf16' to reduce memory at the cost of needing Kahan summation for stability. Currently FSDP-only; Megatron uses use_precision_aware_optimizer instead and ignores this field. | +| `optimizer` | [`OptimizerConfig`](section-optimizer) \| None | `None` | MOPD scoring teachers do not construct an optimizer. | +| `weight_update_mode` | string | `"xccl"` | Weight update backend type. 'awex' requires a Megatron actor and an SGLang rollout. **Choices:** `disk`, `xccl`, `awex` | +| `enable_delta_weight_update` | boolean | `False` | Enable sparse delta weight updates for separation AWEX. | +| `weight_update_delta_method` | string | `"adamw"` | Change detection method used for delta weight transfer. **Choices:** `adamw` | +| `weight_update_anchor_interval` | integer | `0` | Force a full sync every N committed deltas. 0 disables periodic anchors. | +| `fsdp` | [`FSDPEngineConfig`](section-fsdp-engine) | *FSDPEngineConfig* | - | +| `archon` | [`ArchonEngineConfig`](section-archon-engine) | *ArchonEngineConfig* | - | +| `megatron` | [`MegatronEngineConfig`](section-megatron-engine) | *MegatronEngineConfig* | - | +| `offload` | boolean | `False` | Whether to offload model parameters and optimizer states to CPU. | +| `use_lora` | boolean | `False` | Whether to use LoRA. Only support FSDP. Note that should be enabled together with vLLM/SGLang. | +| `lora_rank` | integer | `32` | lora rank | +| `lora_alpha` | integer | `16` | lora alpha | +| `target_modules` | list of string | `[]` | lora target_modules. | +| `peft_type` | string | `"lora"` | peft method type. Only LoRA is supported for now. | +| `enable_tree_training` | boolean | `False` | Enable tree training with flex attention module. | +| `scheduling_spec` | `tuple` | *tuple* | Train engine schedule specs. Can accept 1 or 2 SchedulingSpec: if 1 spec provided, it's used for both worker and engine, engine is embedded in the worker; if 2 specs provided, first one is for worker, second one is for engine. Currently only used by the TrainController. | +| `backend` | string | **Required** | Backend and parallelism strategy. Must include an explicit backend prefix, e.g. 'fsdp:d4', 'megatron:d4t2p2', 'archon:d2'. Required. | +| `_version` | string | `"v1"` | Train controller implementation version. Use 'v1' for legacy TrainController, 'v2' for GatewayTrainController. **Choices:** `v1`, `v2` | +| `admin_api_key` | string | `"areal-admin-key"` | Admin API key used by gateway/router/data-proxy in controller v2. | +| `log_level` | string | `"warning"` | Gateway stack log level for controller v2. | +| `request_timeout` | float | `3600.0` | Gateway request timeout in seconds for controller v2. | +| `setup_timeout` | float | `3600.0` | Gateway setup timeout in seconds for controller v2. | +| `workers_ready_timeout` | float | `30.0` | Timeout (seconds) for initialize() to wait for guards to be ready. | +| `scheduling_strategy` | [`SchedulingStrategy`](section-scheduling-strategy) | *SchedulingStrategy* | The scheduling strategy of this TrainEngine, either separation or colocation. Currently only used by the TrainController. | + +(section-mopd-teacher-manager)= + +## MOPDTeacherManager Configuration + +Checkpoint provider configuration for phase-scoped MOPD teachers. + +| Parameter | Type | Default | Description | +| ---------------- | --------------- | ----------------------- | ---------------------------------------------------------------- | +| `type` | string | `"disk"` | Teacher checkpoint provider. **Choices:** `disk`, `local_memory` | +| `staging_root` | string | `"/dev/shm/areal-mopd"` | Node-local staging root for local_memory providers. | +| `min_free_bytes` | integer \| None | `None` | Optional minimum free space required after staging a checkpoint. | + +(section-mopd-teacher)= + +## MOPDTeacher Specification + +Checkpoint specification for one MOPD teacher. + +| Parameter | Type | Default | Description | +| --------- | ------ | ------------ | --------------------------------------------------- | +| `path` | string | **Required** | Local or shared-filesystem teacher checkpoint path. | + (section-megatron-engine)= ## MegatronEngine Configuration diff --git a/docs/zh/algorithms/mopd.md b/docs/zh/algorithms/mopd.md new file mode 100644 index 0000000000..799c122ab0 --- /dev/null +++ b/docs/zh/algorithms/mopd.md @@ -0,0 +1,130 @@ +# 多 Teacher On-Policy Distillation(MOPD) + +MOPD 将 on-policy 强化学习与一个或多个 teacher checkpoint 的 token 级目标结合。每个 +dataset source 必须选择一个 teacher group,group 可以为任意数量的 teacher 指定非负权重。 +权重按配置值直接使用,不会自动归一化。 + +MOPD 当前运行于 single-controller 模式,使用 Megatron actor 和 teacher engine、SGLang +rollout engine 及 AWEX 共卡权重传输。actor、rollout 和常驻的 teacher companion 共用同一组 +GPU,但大模型权重不会同时驻留。 + +## 运行时生命周期 + +每个训练 step 包含三个互斥阶段: + +1. **Rollout:** SGLang 生成 trajectory,AReaL 将数据源的 teacher group 作为内部 task + metadata 透传;dataset sample 和 workflow 输入都不需要路由字段。 +2. **Teacher:** offload rollout 权重和 KV cache。fork 出的 Megatron teacher 进程 onload, + 依次加载当前 batch 所需的 checkpoint,并为路由到它的 sample 计算分数;actor 完成 + teacher RTensor 的物化与清理后,再次 offload teacher 权重。companion 进程保持常驻,供 + 下一个 step 复用。 +3. **Train:** actor 计算配置的 RL 与 distillation loss,更新权重,并通过 AWEX 向 SGLang + 发布新版本。SGLang 丢弃旧 KV cache,并在继续生成前分配新的空 cache。 + +actor 与 `mopd.teacher_engine` 必须使用相同的并行策略和 world size,包括相同的 +pipeline parallel size。当前实现要求 actor 与 teacher controller 使用 v1;PP、TP、CP、 +DP 和 EP 由所选 Megatron 模型及 allocation 校验。 + +## 配置 + +在 PPO 配置中增加 `mopd`: + +```yaml +actor: + backend: "megatron:(attn:d1p1t4c2|ffn:d1p1e8)" + weight_update_mode: awex + +rollout: + backend: sglang:d8t1 + scheduling_strategy: {type: colocation, target: actor, fork: true} + +train_dataset: + mixture_sampling_policy: proportional + sources: + - {path: /data/code, type: rl, teacher_group: coding} + - {path: /data/mixed, type: rl, teacher_group: mixed} + +mopd: + teachers: + coder: {path: /models/teacher-coder} + reasoning: {path: /models/teacher-reasoning} + teacher_groups: + coding: {coder: 1.0} + mixed: {coder: 0.3, reasoning: 0.7} + teacher_engine: + backend: ${actor.backend} + optimizer: null + disable_dropout: true + scheduling_strategy: {type: colocation, target: actor, fork: true} + scheduling_spec: ${actor.scheduling_spec} + manager: + type: disk + staging_root: /dev/shm/areal-mopd + loss: + rl_coefficient: 0.0 + distillation_coefficient: 1.0 +``` + +启用 MOPD 时,`train_dataset.sources` 中的每个数据源都必须显式设置 +`teacher_group`,且必须匹配 `mopd.teacher_groups` 中的 key;配置 valid dataset 时也遵循 +相同规则。非 MOPD 场景下,`teacher_group` 默认为 `null`。sample 不能覆盖数据源的 +teacher group,也不需要 `task_type` 字段。每个 teacher group 可以引用任意数量的已知 +teacher ID,并且必须至少包含一个正权重。 +`mixture_sampling_policy: proportional` 保持各数据源按长度占比采样;`uniform` +则通过确定性循环较短数据源,使每个数据源在一个 epoch 中贡献相同数量的样本,再由分布式 +sampler 对全局索引进行 shuffle。 + +`manager.type: disk` 从共享存储加载 checkpoint,支持多节点运行。`local_memory` +会在 `staging_root` 下异步暂存下一个 checkpoint,使用原子发布后交给常驻 teacher +加载,并在加载完成后删除。由于该路径只在 controller 所在节点可见, +`local_memory` 要求 `scheduler.type: local` 且 actor/teacher 为单机拓扑;可通过 +`min_free_bytes` 为暂存目录预留可用空间。 + +对 teacher 权重 $w_j$,定义 $S_T(a)=\sum_j w_j\log\pi_{T_j}(a)$ 和 +$W=\sum_j w_j$。MOPD 使用 on-policy score-function surrogate 最小化未归一化的加权 +reverse KL:$\sum_j w_j D_{KL}(\pi_\theta \parallel \pi_{T_j})$。 + +```text +rho(a) = min(exp(log pi_theta(a) - log pi_old(a)), importance_ratio_cap) +reward(a) = S_T(a) - W * stop_gradient(log pi_theta(a)) +mopd_loss = -mean(rho(a) * reward(a)) +loss = rl_coefficient * rl_loss + distillation_coefficient * mopd_loss +``` + +`importance_ratio_cap` 默认值为 `5.0`,用于限制重要性采样乘数并避免指数溢出。 + +该目标是多个 reverse-KL 的加权和;忽略与 student 无关的常数后,也等价于几何 teacher +ensemble。它不是 teacher cross-entropy,也不是 teacher 概率的算术混合。teacher group +权重直接生效, +不会自动归一化。 + +设置 `rl_coefficient: 0.0` 可执行纯 distillation;两个 coefficient 都设为正数时执行 RL 与 +distillation 联合训练。 + +## 示例 + +- `examples/mopd/gsm8k_qwen3_14b_to_0_6b.py` 提供本地 GSM8K 入口和 dry-run 校验。 +- `examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml` 在单机八卡上配置 + Qwen3-14B teacher 和 Qwen3-0.6B actor。 + +无需启动 worker 即可校验本地配置: + +```bash +MOPD_STUDENT_MODEL_PATH=/models/Qwen3-0.6B \ +MOPD_TEACHER_MODEL_PATH=/models/Qwen3-14B \ +MOPD_GSM8K_PATH=/data/gsm8k \ +AREAL_ADMIN_API_KEY="$(python -c 'import secrets; print(secrets.token_hex(32))')" \ +python -m examples.mopd.gsm8k_qwen3_14b_to_0_6b \ + --config examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml \ + --dry-run +``` + +## 运维注意事项 + +- actor 与 teacher checkpoint 必须共享相同的 token-ID 映射,且每个模型架构都必须由所选 + Megatron adapter 支持。 +- teacher companion 进程会跨 phase 常驻,但 actor 恢复显存所有权前必须 offload teacher + 权重。`DrainReceipt` 是 teacher RTensor 的 phase 资源回收边界。 +- W&B 凭据及服务 endpoint 应通过环境变量传入,不要写入 YAML 或 shell 文件。 +- driver 已完成并清理持久 worker 时,actor 或 rollout 的 Slurm 子作业显示 cancelled 可能是 + 正常现象。应以 driver exit code 和最后一个训练 step 日志判断是否成功。 diff --git a/docs/zh/cli_reference.md b/docs/zh/cli_reference.md index bd4661ad8a..61f931434e 100644 --- a/docs/zh/cli_reference.md +++ b/docs/zh/cli_reference.md @@ -74,8 +74,14 @@ python3 train.py --config path/to/config.yaml actor.lr=1e-4 seed=42 - [ArchonFP8 Configuration](section-archon-fp8) - [DPO Configuration](section-dpo) - [DPOEngine Configuration](section-dpo-engine) +- [DatasetSource Configuration](section-dataset-source) - [DistributedDataParallel Configuration](section-distributed-data-parallel) - [FP8Engine Configuration](section-fp8-engine) +- [MOPD Configuration](section-mopd) +- [MOPDLoss Configuration](section-mopd-loss) +- [MOPDTeacherEngine Configuration](section-mopd-teacher-engine) +- [MOPDTeacherManager Configuration](section-mopd-teacher-manager) +- [MOPDTeacher Specification](section-mopd-teacher) - [MegatronEngine Configuration](section-megatron-engine) - [MemoryProfiler Configuration](section-memory-profiler) - [PerfTracer Configuration](section-perf-tracer) @@ -156,6 +162,7 @@ A dummy place holder of GRPO config for backward compatibility. | `ref` | [`PPOActorConfig`](section-ppo-actor) \| None | `None` | - | | `critic` | [`PPOCriticConfig`](section-ppo-critic) \| None | `None` | - | | `teacher` | [`TeacherConfig`](section-teacher) \| None | `None` | Optional teacher model configuration used for on-policy distillation during PPO training. If provided, the actor may be trained to match the teacher in addition to the standard PPO objective. | +| `mopd` | [`MOPDConfig`](section-mopd) \| None | `None` | Optional multi-teacher on-policy distillation config. | | `dynamic_bs` | boolean | `False` | Enable dynamic batch sizing in prepare_batch. When True, batch collection stops when (accepted + rejected) >= batch_size, returning only accepted results. This results in variable-sized batches of valid data. | (section-ppo)= @@ -195,6 +202,7 @@ Configuration for Proximal Policy Optimization (PPO) reinforcement learning expe | `ref` | [`PPOActorConfig`](section-ppo-actor) \| None | `None` | - | | `critic` | [`PPOCriticConfig`](section-ppo-critic) \| None | `None` | - | | `teacher` | [`TeacherConfig`](section-teacher) \| None | `None` | Optional teacher model configuration used for on-policy distillation during PPO training. If provided, the actor may be trained to match the teacher in addition to the standard PPO objective. | +| `mopd` | [`MOPDConfig`](section-mopd) \| None | `None` | Optional multi-teacher on-policy distillation config. | | `dynamic_bs` | boolean | `False` | Enable dynamic batch sizing in prepare_batch. When True, batch collection stops when (accepted + rejected) >= batch_size, returning only accepted results. This results in variable-sized batches of valid data. | (section-rw)= @@ -696,21 +704,23 @@ https://docs.vllm.ai/en/stable/api/index.html for detailed documentation. Configuration for training dataset loading and preprocessing. -| Parameter | Type | Default | Description | -| --------------------- | ---------------------------------------------- | ---------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `split` | string | `"train"` | Dataset split to use, e.g., 'train', 'test'. | -| `path` | string | **Required** | Path to the dataset. Can be a local path or a HuggingFace dataset name. | -| `type` | string | **Required** | Type of training method, e.g., 'sft', 'rl', etc. | -| `batch_size` | integer | `1` | Batch size for the dataloader | -| `shuffle` | boolean | `True` | Whether to shuffle the dataset | -| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | -| `num_workers` | integer | `0` | Number of worker processes for data loading | -| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | -| `drop_last` | boolean | `True` | Drop the last incomplete batch | -| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | -| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | -| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | -| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | +| Parameter | Type | Default | Description | +| ------------------------- | ------------------------------------------------------- | ---------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `split` | string | `"train"` | Dataset split to use, e.g., 'train', 'test'. | +| `path` | string \| None | `None` | Path to one dataset. Mutually exclusive with sources. | +| `type` | string \| None | `None` | Training data type for path. Mutually exclusive with sources. | +| `sources` | list of [`DatasetSourceConfig`](section-dataset-source) | `[]` | Dataset mixture sources. MOPD requires every source to declare a teacher_group. | +| `mixture_sampling_policy` | string | `"proportional"` | How a routed mixture represents sources in one epoch: 'proportional' preserves source-size proportions; 'uniform' balances source counts by deterministically cycling shorter sources. **Choices:** `proportional`, `uniform` | +| `batch_size` | integer | `1` | Batch size for the dataloader | +| `shuffle` | boolean | `True` | Whether to shuffle the dataset | +| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | +| `num_workers` | integer | `0` | Number of worker processes for data loading | +| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | +| `drop_last` | boolean | `True` | Drop the last incomplete batch | +| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | +| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | +| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | +| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | (section-valid-dataset)= @@ -721,21 +731,23 @@ Configuration for validation dataset loading and preprocessing. It has different default values with `TrainDatasetConfig`. `shuffle` and `drop_last` default to False. -| Parameter | Type | Default | Description | -| --------------------- | ---------------------------------------------- | ---------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `split` | string | `"test"` | Dataset split to use, e.g., 'train', 'test'. | -| `path` | string | **Required** | Path to the dataset. Can be a local path or a HuggingFace dataset name. | -| `type` | string | **Required** | Type of training method, e.g., 'sft', 'rl', etc. | -| `batch_size` | integer | `1` | Batch size for the dataloader | -| `shuffle` | boolean | `False` | Whether to shuffle the dataset | -| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | -| `num_workers` | integer | `0` | Number of worker processes for data loading | -| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | -| `drop_last` | boolean | `False` | Drop the last incomplete batch | -| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | -| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | -| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | -| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | +| Parameter | Type | Default | Description | +| ------------------------- | ------------------------------------------------------- | ---------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `split` | string | `"test"` | Dataset split to use, e.g., 'train', 'test'. | +| `path` | string \| None | `None` | Path to one dataset. Mutually exclusive with sources. | +| `type` | string \| None | `None` | Training data type for path. Mutually exclusive with sources. | +| `sources` | list of [`DatasetSourceConfig`](section-dataset-source) | `[]` | Dataset mixture sources. MOPD requires every source to declare a teacher_group. | +| `mixture_sampling_policy` | string | `"proportional"` | How a routed mixture represents sources in one epoch: 'proportional' preserves source-size proportions; 'uniform' balances source counts by deterministically cycling shorter sources. **Choices:** `proportional`, `uniform` | +| `batch_size` | integer | `1` | Batch size for the dataloader | +| `shuffle` | boolean | `False` | Whether to shuffle the dataset | +| `pin_memory` | boolean | `False` | Pin memory for faster data loading (set True for GPU training) | +| `num_workers` | integer | `0` | Number of worker processes for data loading | +| `num_dataset_workers` | integer | `1` | Number of remote data-service worker processes to launch when using scheduling_spec. | +| `drop_last` | boolean | `False` | Drop the last incomplete batch | +| `max_length` | integer \| None | `None` | Maximum token length of sequences in dataset. Longer sequences are filtered out. | +| `dataset_kwargs` | `dict` | `{}` | Additional keyword arguments for dataset loading. These are passed to the dataset loading function `get_custom_dataset`. | +| `scheduling_spec` | [`SchedulingSpec`](section-scheduling) \| None | *SchedulingSpec* | Scheduling spec for remote data loading workers. If set, dataset loading will be offloaded to a data service with remote workers. | +| `setup_timeout` | float | `120.0` | Timeout in seconds for the data service to load and register a dataset. Increase this value when loading large datasets for the first time (e.g. HuggingFace datasets that require downloading and preprocessing). | (section-cluster)= @@ -1049,6 +1061,21 @@ fields. | `beta` | float | `0.1` | KL penalty coefficient for DPO loss. | | `loss_type` | string | `"sigmoid"` | DPO loss variant. 'sigmoid': original DPO loss (Rafailov et al. 2023). 'ipo': Identity Preference Optimization with per-token length normalization (Azar et al. 2023). **Choices:** `sigmoid`, `ipo` | +(section-dataset-source)= + +## DatasetSource Configuration + +One source in a dataset mixture. + +| Parameter | Type | Default | Description | +| ---------------- | --------------- | ------------ | ---------------------------------------------------------- | +| `path` | string | **Required** | Local path or HuggingFace name for this dataset source. | +| `type` | string | **Required** | Training data type, for example 'rl'. | +| `teacher_group` | string \| None | `None` | Optional MOPD teacher group applied to this entire source. | +| `split` | string \| None | `None` | Optional split override for this dataset source. | +| `max_length` | integer \| None | `None` | Optional maximum sequence length for this source. | +| `dataset_kwargs` | `dict` | `{}` | Extra keyword arguments for this source's loader. | + (section-distributed-data-parallel)= ## DistributedDataParallel Configuration @@ -1096,6 +1123,103 @@ is disabled. | `num_layers_at_end_in_bf16` | integer | `1` | Number of layers at end to keep in BF16 when first_last_layers_bf16 is True. | | `direct_convert` | boolean | `True` | Whether to use direct FP8 conversion during weight updates and save/load. When True, FP8 parameters are directly converted between TE FP8 and PyTorch FP8 without intermediate dequantization/quantization. | +(section-mopd)= + +## MOPD Configuration + +Configuration for multi-teacher on-policy distillation. + +| Parameter | Type | Default | Description | +| ---------------- | ---------------------------------------------------------- | -------------------------- | ----------- | +| `teachers` | `dict` | `{}` | - | +| `teacher_groups` | `dict` | `{}` | - | +| `teacher_engine` | [`MOPDTeacherEngineConfig`](section-mopd-teacher-engine) | *MOPDTeacherEngineConfig* | - | +| `manager` | [`MOPDTeacherManagerConfig`](section-mopd-teacher-manager) | *MOPDTeacherManagerConfig* | - | +| `loss` | [`MOPDLossConfig`](section-mopd-loss) | *MOPDLossConfig* | - | + +(section-mopd-loss)= + +## MOPDLoss Configuration + +Coefficients for joint RL and multi-teacher distillation training. + +| Parameter | Type | Default | Description | +| -------------------------- | ----- | ------- | -------------------------------------------------- | +| `rl_coefficient` | float | `0.0` | Coefficient applied to the RL objective. | +| `distillation_coefficient` | float | `1.0` | Coefficient applied to the MOPD objective. | +| `importance_ratio_cap` | float | `5.0` | Positive cap applied to the behavior-policy ratio. | + +(section-mopd-teacher-engine)= + +## MOPDTeacherEngine Configuration + +Forward-only scoring engine configuration used by MOPD teachers. + +| Parameter | Type | Default | Description | +| ------------------------------- | --------------------------------------------------- | ---------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `experiment_name` | string | **Required** | - | +| `trial_name` | string | **Required** | - | +| `path` | string | `""` | Path to HuggingFace checkpoint | +| `attn_impl` | string | `"flash_attention_2"` | Attention implementation for huggingface transformers model. Accepts builtin transformers backends or a Hugging Face kernels repo ID formatted as org/repo\[@revision\]\[:entrypoint\]. **Choices:** `eager`, `sdpa`, `flash_attention_2`, `flash_attention_3`, `flex_attention` | +| `use_kernels` | boolean | `False` | Enable Hugging Face kernels model kernelization after model creation. | +| `init_from_scratch` | boolean | `False` | Initialize model weights randomly | +| `is_critic` | boolean | `False` | Whether to use a critic/reward model | +| `temperature` | float | `1.0` | Temperature during generation. | +| `logprobs_chunk_size` | integer | `1024` | Maximum sequence chunk size used to compute log probabilities and entropy. Must be positive. | +| `mb_spec` | [`MicroBatchSpec`](section-micro-batch) | *MicroBatchSpec* | - | +| `pad_to_maximum` | boolean | `False` | Whether to pad each microbatch to the length upper bound specified by mb_spec. Can reduce memory fragmentation but slows down training. | +| `disable_dropout` | boolean | `True` | Disable dropout for deterministic teacher scoring. | +| `gradient_checkpointing` | boolean | `False` | Enable gradient checkpointing | +| `dtype` | string | `"bfloat16"` | Forward/backward compute dtype. | +| `grad_reduce_dtype` | string | `"float32"` | Gradient reduction data type. | +| `optimizer_dtype` | string | `"float32"` | Underlying parameter storage dtype, also the dtype of optimizer states (exp_avg, exp_avg_sq) since torch.optim.AdamW inherits dtype from model.parameters(). Default 'float32' maintains fp32 master weights matching DeepSpeed ZeRO-3 and Megatron precision-aware optimizer behavior. FSDP2's MixedPrecisionPolicy(param_dtype=`dtype`) will still cast forward/backward computation to `dtype` (e.g. bfloat16). Set to 'bfloat16' together with optimizer.type='adam_bf16' to reduce memory at the cost of needing Kahan summation for stability. Currently FSDP-only; Megatron uses use_precision_aware_optimizer instead and ignores this field. | +| `optimizer` | [`OptimizerConfig`](section-optimizer) \| None | `None` | MOPD scoring teachers do not construct an optimizer. | +| `weight_update_mode` | string | `"xccl"` | Weight update backend type. 'awex' requires a Megatron actor and an SGLang rollout. **Choices:** `disk`, `xccl`, `awex` | +| `enable_delta_weight_update` | boolean | `False` | Enable sparse delta weight updates for separation AWEX. | +| `weight_update_delta_method` | string | `"adamw"` | Change detection method used for delta weight transfer. **Choices:** `adamw` | +| `weight_update_anchor_interval` | integer | `0` | Force a full sync every N committed deltas. 0 disables periodic anchors. | +| `fsdp` | [`FSDPEngineConfig`](section-fsdp-engine) | *FSDPEngineConfig* | - | +| `archon` | [`ArchonEngineConfig`](section-archon-engine) | *ArchonEngineConfig* | - | +| `megatron` | [`MegatronEngineConfig`](section-megatron-engine) | *MegatronEngineConfig* | - | +| `offload` | boolean | `False` | Whether to offload model parameters and optimizer states to CPU. | +| `use_lora` | boolean | `False` | Whether to use LoRA. Only support FSDP. Note that should be enabled together with vLLM/SGLang. | +| `lora_rank` | integer | `32` | lora rank | +| `lora_alpha` | integer | `16` | lora alpha | +| `target_modules` | list of string | `[]` | lora target_modules. | +| `peft_type` | string | `"lora"` | peft method type. Only LoRA is supported for now. | +| `enable_tree_training` | boolean | `False` | Enable tree training with flex attention module. | +| `scheduling_spec` | `tuple` | *tuple* | Train engine schedule specs. Can accept 1 or 2 SchedulingSpec: if 1 spec provided, it's used for both worker and engine, engine is embedded in the worker; if 2 specs provided, first one is for worker, second one is for engine. Currently only used by the TrainController. | +| `backend` | string | **Required** | Backend and parallelism strategy. Must include an explicit backend prefix, e.g. 'fsdp:d4', 'megatron:d4t2p2', 'archon:d2'. Required. | +| `_version` | string | `"v1"` | Train controller implementation version. Use 'v1' for legacy TrainController, 'v2' for GatewayTrainController. **Choices:** `v1`, `v2` | +| `admin_api_key` | string | `"areal-admin-key"` | Admin API key used by gateway/router/data-proxy in controller v2. | +| `log_level` | string | `"warning"` | Gateway stack log level for controller v2. | +| `request_timeout` | float | `3600.0` | Gateway request timeout in seconds for controller v2. | +| `setup_timeout` | float | `3600.0` | Gateway setup timeout in seconds for controller v2. | +| `workers_ready_timeout` | float | `30.0` | Timeout (seconds) for initialize() to wait for guards to be ready. | +| `scheduling_strategy` | [`SchedulingStrategy`](section-scheduling-strategy) | *SchedulingStrategy* | The scheduling strategy of this TrainEngine, either separation or colocation. Currently only used by the TrainController. | + +(section-mopd-teacher-manager)= + +## MOPDTeacherManager Configuration + +Checkpoint provider configuration for phase-scoped MOPD teachers. + +| Parameter | Type | Default | Description | +| ---------------- | --------------- | ----------------------- | ---------------------------------------------------------------- | +| `type` | string | `"disk"` | Teacher checkpoint provider. **Choices:** `disk`, `local_memory` | +| `staging_root` | string | `"/dev/shm/areal-mopd"` | Node-local staging root for local_memory providers. | +| `min_free_bytes` | integer \| None | `None` | Optional minimum free space required after staging a checkpoint. | + +(section-mopd-teacher)= + +## MOPDTeacher Specification + +Checkpoint specification for one MOPD teacher. + +| Parameter | Type | Default | Description | +| --------- | ------ | ------------ | --------------------------------------------------- | +| `path` | string | **Required** | Local or shared-filesystem teacher checkpoint path. | + (section-megatron-engine)= ## MegatronEngine Configuration diff --git a/examples/mopd/gsm8k_qwen3_14b_to_0_6b.py b/examples/mopd/gsm8k_qwen3_14b_to_0_6b.py new file mode 100644 index 0000000000..ef706e9368 --- /dev/null +++ b/examples/mopd/gsm8k_qwen3_14b_to_0_6b.py @@ -0,0 +1,200 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Single-node Qwen3-14B -> Qwen3-0.6B MOPD example on GSM8K.""" + +from __future__ import annotations + +import os +import sys +from pathlib import Path +from typing import Any + +from areal import PPOTrainer +from areal.api import AsyncRewardWrapper +from areal.api.cli_args import GRPOConfig, load_expr_config +from areal.dataset import get_custom_dataset, get_mopd_dataset +from areal.reward import gsm8k_reward_fn +from areal.utils.hf_utils import load_hf_tokenizer + +MOPD_TEACHER_GROUP = "gsm8k" +NO_THINK_SUFFIX = " /no_think" + + +def dynamic_filter(data: dict[str, Any]) -> bool: + """Reject nearly all-correct rollout groups, matching the reference run.""" + return data["rewards"].mean() <= 0.95 + + +class GSM8KRewardDistillationAgent: + """Generate an on-policy response and report its GSM8K verifier reward.""" + + def __init__( + self, + *, + reward_timeout: float = 15.0, + **generation_kwargs: Any, + ): + self.generation_kwargs = generation_kwargs + self._reward = AsyncRewardWrapper( + gsm8k_reward_fn, + timeout_seconds=reward_timeout, + max_workers=1, + max_retries=1, + ) + + async def run(self, data: dict[str, Any], **extra_kwargs: Any) -> dict[str, float]: + from openai import AsyncOpenAI + + client = AsyncOpenAI( + base_url=extra_kwargs.get("base_url") or os.getenv("OPENAI_BASE_URL"), + api_key=extra_kwargs.get("api_key") or os.getenv("OPENAI_API_KEY"), + http_client=extra_kwargs.get("http_client"), + max_retries=0, + ) + response = await client.chat.completions.create( + messages=data["messages"], + model="default", + **self.generation_kwargs, + ) + completion = response.choices[0].message.content or "" + reward = await self._reward( + prompt=str(data["messages"]), + completions=completion, + prompt_ids=[], + completion_ids=[], + answer=data["answer"], + ) + return {response.id: float(reward)} + + +def add_no_think_suffix(sample: dict[str, Any]) -> dict[str, Any]: + """Attach the reference no-think prompt suffix without dataset routing fields.""" + update: dict[str, Any] = {} + messages = sample.get("messages") + if not isinstance(messages, list): + return update + + messages = [dict(message) for message in messages] + for message in reversed(messages): + content = message.get("content") + if message.get("role") == "user" and isinstance(content, str): + if not content.rstrip().endswith("/no_think"): + message["content"] = content.rstrip() + NO_THINK_SUFFIX + break + update["messages"] = messages + return update + + +def _load_gsm8k_source( + source_config: Any, + *, + split: str, + tokenizer: Any, +): + """Load either a local parquet mirror or a standard AReaL dataset snapshot.""" + path = Path(source_config.path) + parquet_files = sorted((path / "main").glob(f"{split}-*.parquet")) + if not parquet_files: + return get_custom_dataset( + split=split, + dataset_config=source_config, + tokenizer=tokenizer, + ).map(add_no_think_suffix, desc="Attach Qwen3-14B no-think suffix") + + from datasets import load_dataset + + dataset = load_dataset( + "parquet", + data_files=[str(parquet_file) for parquet_file in parquet_files], + split="train", + ) + + def process(sample: dict[str, Any]) -> dict[str, Any]: + formatted = { + "messages": [ + { + "role": "user", + "content": sample["question"] + + "\nPlease put your final answer within \\boxed{}.", + } + ], + } + return formatted | add_no_think_suffix(formatted) + + dataset = dataset.map( + process, + remove_columns=["question"], + desc=f"Format routed GSM8K {split} split", + ) + if source_config.max_length is not None: + dataset = dataset.filter( + lambda sample: len(tokenizer.encode(sample["messages"][0]["content"])) + <= source_config.max_length, + desc=f"Filter GSM8K {split} prompts by length", + ) + return dataset + + +def load_routed_gsm8k_dataset( + dataset_config: Any, + *, + tokenizer: Any, +): + """Load GSM8K sources with optional source-level MOPD teacher groups.""" + + def load_source(**kwargs): + return _load_gsm8k_source( + kwargs["source_config"], + split=kwargs["split"], + tokenizer=kwargs["tokenizer"], + ) + + return get_mopd_dataset( + dataset_config, + tokenizer=tokenizer, + source_loader=load_source, + ) + + +def train(argv: list[str]) -> None: + """Run pure MOPD and report GSM8K rewards as rollout metrics.""" + config, _ = load_expr_config(argv, GRPOConfig) + tokenizer = load_hf_tokenizer(config.tokenizer_path) + + train_dataset = load_routed_gsm8k_dataset( + config.train_dataset, + tokenizer=tokenizer, + ) + assert config.valid_dataset is not None + valid_dataset = load_routed_gsm8k_dataset( + config.valid_dataset, + tokenizer=tokenizer, + ) + + workflow_kwargs = { + "temperature": config.gconfig.temperature, + "top_p": config.gconfig.top_p, + "max_completion_tokens": config.gconfig.max_new_tokens, + } + eval_workflow_kwargs = workflow_kwargs | {"temperature": 0.6} + + with PPOTrainer( + config, + train_dataset=train_dataset, + valid_dataset=valid_dataset, + ) as trainer: + trainer.train( + workflow=( + "examples.mopd.gsm8k_qwen3_14b_to_0_6b.GSM8KRewardDistillationAgent" + ), + workflow_kwargs=workflow_kwargs, + eval_workflow=( + "examples.mopd.gsm8k_qwen3_14b_to_0_6b.GSM8KRewardDistillationAgent" + ), + eval_workflow_kwargs=eval_workflow_kwargs, + dynamic_filter_fn=("examples.mopd.gsm8k_qwen3_14b_to_0_6b.dynamic_filter"), + ) + + +if __name__ == "__main__": + train(sys.argv[1:]) diff --git a/examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml b/examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml new file mode 100644 index 0000000000..e574977498 --- /dev/null +++ b/examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml @@ -0,0 +1,153 @@ +# Single-node, eight-GPU heterogeneous MOPD example. +experiment_name: mopd_gsm8k +trial_name: qwen3-14b-to-0p6b-local +seed: 1 +enable_offload: false +total_train_epochs: 10 +tokenizer_path: ${actor.path} + +cluster: + n_nodes: 1 + n_gpus_per_node: 8 + fileroot: ${oc.env:AREAL_FILEROOT,/tmp/areal/experiments} + name_resolve: + type: nfs + nfs_record_root: ${cluster.fileroot}/name_resolve + +scheduler: {type: local} + +gconfig: + n_samples: 4 + min_new_tokens: 0 + max_new_tokens: 2048 + max_tokens: 32768 + greedy: false + top_p: 1.0 + top_k: 100000000 + temperature: 1.0 + +actor: + _version: v1 + experiment_name: ${experiment_name} + trial_name: ${trial_name} + path: ${oc.env:MOPD_STUDENT_MODEL_PATH} + backend: megatron:d1p1t8 + weight_update_mode: awex + disable_dropout: true + gradient_checkpointing: true + dtype: bfloat16 + mb_spec: {n_mbs: 1, max_tokens_per_mb: 10240} + optimizer: + type: adam + lr: 5.0e-6 + weight_decay: 0.017 + beta1: 0.9 + beta2: 0.999 + eps: 1.0e-8 + lr_scheduler_type: constant + warmup_steps_proportion: 0.001 + gradient_clipping: 1.0 + megatron: + wrap_with_ddp: true + ddp: {use_distributed_optimizer: true, grad_reduce_in_fp32: true} + cross_entropy_loss_fusion: false + kl_ctl: 0.0 + ppo_n_minibatches: 1 + # Reuse the training forward as pi_prox while preserving rollout behavior + # logprobs for pure-MOPD pi_theta / pi_behavior importance sampling. + recompute_logprob: false + use_decoupled_loss: true + prox_logp_method: reuse_train_logp + max_new_tokens: ${gconfig.max_new_tokens} + scheduling_spec: + - task_type: worker + port_count: 3 + gpu: 1 + cpu: 8 + mem: 64 + image: ${oc.env:AREAL_IMAGE} + cmd: python3 -m areal.infra.rpc.rpc_server + env_vars: {} + +rollout: + _version: v1 + experiment_name: ${experiment_name} + trial_name: ${trial_name} + backend: sglang:d8t1p1 + scheduling_strategy: {type: colocation, target: actor, fork: true} + scheduling_spec: ${actor.scheduling_spec} + max_concurrent_rollouts: 256 + consumer_batch_size: ${train_dataset.batch_size} + max_head_offpolicyness: 2 + fileroot: ${cluster.fileroot} + tokenizer_path: ${tokenizer_path} + dump_to_file: true + setup_timeout: 7200.0 + agent: + admin_api_key: ${oc.env:AREAL_ADMIN_API_KEY} + +mopd: + teachers: + qwen3_14b: {path: "${oc.env:MOPD_TEACHER_MODEL_PATH}"} + teacher_groups: + gsm8k: {qwen3_14b: 1.0} + teacher_engine: + _version: v1 + experiment_name: ${experiment_name} + trial_name: ${trial_name} + backend: ${actor.backend} + optimizer: null + disable_dropout: true + dtype: ${actor.dtype} + mb_spec: {n_mbs: 1, max_tokens_per_mb: 10240} + megatron: + wrap_with_ddp: true + ddp: {use_distributed_optimizer: true, grad_reduce_in_fp32: true} + cross_entropy_loss_fusion: false + scheduling_strategy: {type: colocation, target: actor, fork: true} + scheduling_spec: ${actor.scheduling_spec} + manager: {type: disk} + loss: {rl_coefficient: 0.0, distillation_coefficient: 1.0} + +sglang: + model_path: ${actor.path} + random_seed: ${seed} + skip_tokenizer_init: true + disable_radix_cache: true + dtype: ${actor.dtype} + context_length: 32768 + max_running_requests: null + mem_fraction_static: 0.8 + +train_dataset: + mixture_sampling_policy: proportional + sources: + - path: ${oc.env:MOPD_GSM8K_PATH} + type: rl + teacher_group: gsm8k + batch_size: 256 + shuffle: true + pin_memory: true + num_workers: 4 + +valid_dataset: + mixture_sampling_policy: proportional + sources: + - path: ${oc.env:MOPD_GSM8K_PATH} + type: rl + teacher_group: gsm8k + batch_size: 256 + shuffle: false + num_workers: 4 + +saver: {experiment_name: "${experiment_name}", trial_name: "${trial_name}", fileroot: "${cluster.fileroot}", freq_epochs: 1} +recover: {mode: disabled, experiment_name: "${experiment_name}", trial_name: "${trial_name}", fileroot: "${cluster.fileroot}"} +evaluator: {experiment_name: "${experiment_name}", trial_name: "${trial_name}", fileroot: "${cluster.fileroot}", freq_epochs: 1} +stats_logger: + experiment_name: ${experiment_name} + trial_name: ${trial_name} + fileroot: ${cluster.fileroot} + wandb: + mode: disabled + wandb_api_key: ${oc.env:WANDB_API_KEY,""} + wandb_base_url: ${oc.env:WANDB_BASE_URL,""} diff --git a/tests/experimental/openai/test_proxy_rollout_server.py b/tests/experimental/openai/test_proxy_rollout_server.py index 9e1295f676..df473b3c34 100644 --- a/tests/experimental/openai/test_proxy_rollout_server.py +++ b/tests/experimental/openai/test_proxy_rollout_server.py @@ -28,6 +28,8 @@ def _reset_server_globals(monkeypatch): monkeypatch.setattr(srv, "_admin_api_key", _ADMIN_KEY) monkeypatch.setattr(srv, "_lock", threading.Lock()) monkeypatch.setattr(srv, "_last_cleanup_time", 0.0) + monkeypatch.setattr(srv, "_worker_role", None) + monkeypatch.setattr(srv, "_worker_index", None) monkeypatch.setattr(srv, "_engine", None) monkeypatch.setattr(srv, "_openai_client", None) @@ -45,6 +47,43 @@ def _admin_headers(): return {"Authorization": f"Bearer {_ADMIN_KEY}"} +# --------------------------------------------------------------------------- +# Tests: health reports forked worker identity +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_health_with_worker_identity_returns_exact_role_and_index(monkeypatch): + """Health identifies the forked worker that owns the listening port.""" + monkeypatch.setattr(srv, "_worker_role", "proxy-rollout") + monkeypatch.setattr(srv, "_worker_index", 7) + + async with _client() as client: + response = await client.get("/health") + + assert response.status_code == 200 + assert response.json() == { + "status": "ok", + "initialized": False, + "role": "proxy-rollout", + "worker_index": 7, + } + + +def test_explicit_proxy_worker_index_wins_over_stale_slurm_env(monkeypatch): + """A local proxy keeps the identity supplied by its scheduler.""" + monkeypatch.setenv("SLURM_PROCID", "0") + + assert srv._resolve_worker_index(7) == 7 + + +def test_proxy_worker_index_falls_back_to_slurm_env(monkeypatch): + """A Slurm proxy can still obtain its identity from the task environment.""" + monkeypatch.setenv("SLURM_PROCID", "5") + + assert srv._resolve_worker_index(-1) == 5 + + # --------------------------------------------------------------------------- # Tests: start_session with provided api_key # --------------------------------------------------------------------------- diff --git a/tests/experimental/openai/test_streaming_chat_completions.py b/tests/experimental/openai/test_streaming_chat_completions.py index c747fb12a7..411554ab51 100644 --- a/tests/experimental/openai/test_streaming_chat_completions.py +++ b/tests/experimental/openai/test_streaming_chat_completions.py @@ -19,6 +19,8 @@ from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice from openai.types.chat.chat_completion_chunk import ChoiceDelta +from areal.api import ModelResponse +from areal.experimental.openai.client import AsyncCompletionsWithReward from areal.experimental.openai.proxy import proxy_rollout_server as srv # --------------------------------------------------------------------------- @@ -67,6 +69,7 @@ def _reset_server_globals(monkeypatch): async def _fake_create( *, messages=None, + model=None, stream=None, temperature=None, top_p=None, @@ -95,7 +98,7 @@ async def _gen(): ) ], created=0, - model="test", + model=model, object="chat.completion.chunk", ) @@ -111,7 +114,7 @@ async def _gen(): ) ], created=0, - model="test", + model=model, object="chat.completion", ) @@ -173,6 +176,7 @@ async def test_streaming_returns_sse_response( # Verify the data payload is a valid ChatCompletionChunk chunk = json.loads(events[0].removeprefix("data: ")) assert chunk["object"] == "chat.completion.chunk" + assert chunk["model"] == "test" assert chunk["choices"][0]["delta"]["content"] == "hello" @pytest.mark.asyncio @@ -201,4 +205,38 @@ async def test_non_streaming_returns_json(self, monkeypatch, _mock_openai_client assert resp.status_code == 200 data = resp.json() assert data["object"] == "chat.completion" + assert data["model"] == "test" assert data["choices"][0]["message"]["content"] == "hello" + + +@pytest.mark.asyncio +async def test_areal_completion_response_uses_request_model(): + """Both response modes retain the model required by Anthropic adapters.""" + client = object.__new__(AsyncCompletionsWithReward) + model_response = ModelResponse( + input_tokens=[1, 2], + output_tokens=[3], + stop_reason="stop", + ) + + completion, _ = client._build_chat_completion( + completion_id="chatcmpl-test", + current_time=0, + model="claude-test-model", + output_text="hello", + tool_calls=None, + response=model_response, + ) + stream = client._create_stream( + completion_id="chatcmpl-test", + current_time=0, + model="claude-test-model", + output_text="hello", + tool_calls=None, + response=model_response, + ) + chunks = [chunk async for chunk in stream] + + assert completion.model == "claude-test-model" + assert chunks + assert all(chunk.model == "claude-test-model" for chunk in chunks) diff --git a/tests/experimental/openai/test_tool_call_parser.py b/tests/experimental/openai/test_tool_call_parser.py index c08c4aebe3..8c954e7efb 100644 --- a/tests/experimental/openai/test_tool_call_parser.py +++ b/tests/experimental/openai/test_tool_call_parser.py @@ -209,6 +209,91 @@ def test_qwen3_coder_xml_literal_closing_tag_is_not_silently_truncated(): assert finish_reason == "stop" +@pytest.mark.sglang +@pytest.mark.parametrize("reasoning_parser", ["", None], ids=["empty", "none"]) +def test_process_tool_calls_without_reasoning_parser_returns_plain_text( + reasoning_parser: str | None, +): + """An empty reasoning parser disables reasoning extraction on CPU.""" + pytest.importorskip( + "sglang.srt.function_call.function_call_parser", + reason="sglang is required for sglang parser tests", + ) + pytest.importorskip( + "sglang.srt.parser.reasoning_parser", + reason="sglang is required for sglang parser tests", + ) + text = "The task is complete." + + tool_calls, new_text, finish_reason = parser_module.process_tool_calls( + text=text, + tools=QWEN3_CODER_TOOLS, + tool_call_parser="qwen3_coder", + reasoning_parser=reasoning_parser, + finish_reason="stop", + use_responses=False, + tokenizer=object(), + ) + + assert tool_calls is None + assert new_text == text + assert finish_reason == "stop" + + +@pytest.mark.sglang +def test_process_tool_calls_invalid_reasoning_parser_reports_config_key(): + """Invalid parser names identify the exact rollout configuration field.""" + pytest.importorskip( + "sglang.srt.function_call.function_call_parser", + reason="sglang is required for sglang parser tests", + ) + pytest.importorskip( + "sglang.srt.parser.reasoning_parser", + reason="sglang is required for sglang parser tests", + ) + + with pytest.raises( + ValueError, + match=r"rollout\.openai\.reasoning_parser='not-a-parser'.*qwen3_coder", + ): + parser_module.process_tool_calls( + text="The task is complete.", + tools=QWEN3_CODER_TOOLS, + tool_call_parser="qwen3_coder", + reasoning_parser="not-a-parser", + finish_reason="stop", + use_responses=False, + tokenizer=object(), + ) + + +@pytest.mark.sglang +def test_process_tool_calls_invalid_tool_parser_reports_config_key(): + """Invalid tool parser names identify the exact rollout config field.""" + pytest.importorskip( + "sglang.srt.function_call.function_call_parser", + reason="sglang is required for sglang parser tests", + ) + pytest.importorskip( + "sglang.srt.parser.reasoning_parser", + reason="sglang is required for sglang parser tests", + ) + + with pytest.raises( + ValueError, + match=r"rollout\.openai\.tool_call_parser='not-a-parser'.*qwen3_coder", + ): + parser_module.process_tool_calls( + text="The task is complete.", + tools=QWEN3_CODER_TOOLS, + tool_call_parser="not-a-parser", + reasoning_parser=None, + finish_reason="stop", + use_responses=False, + tokenizer=object(), + ) + + @pytest.mark.sglang def test_process_tool_calls_qwen25_chat_completions_sglang(): pytest.importorskip( diff --git a/tests/test_awex_colocate_device_id.py b/tests/test_awex_colocate_device_id.py index b8c83f2982..133af63629 100644 --- a/tests/test_awex_colocate_device_id.py +++ b/tests/test_awex_colocate_device_id.py @@ -1,3 +1,5 @@ +import pytest + from areal.engine.awex.colocate_writer import resolve_physical_gpu_id @@ -14,13 +16,15 @@ def test_physical_gpu_id_is_identity_without_visible_devices(monkeypatch): assert resolve_physical_gpu_id(2) == 2 -def test_physical_gpu_id_falls_back_on_non_integer_entries(monkeypatch): +def test_physical_gpu_id_rejects_non_integer_entries(monkeypatch): monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "GPU-abc,GPU-def") - assert resolve_physical_gpu_id(1) == 1 + with pytest.raises(ValueError, match="numeric CUDA_VISIBLE_DEVICES"): + resolve_physical_gpu_id(1) -def test_physical_gpu_id_falls_back_when_index_out_of_range(monkeypatch): +def test_physical_gpu_id_rejects_index_out_of_range(monkeypatch): monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "4") - assert resolve_physical_gpu_id(2) == 2 + with pytest.raises(ValueError, match="outside CUDA_VISIBLE_DEVICES"): + resolve_physical_gpu_id(2) diff --git a/tests/test_awex_sglang_plugin.py b/tests/test_awex_sglang_plugin.py new file mode 100644 index 0000000000..4f575159ea --- /dev/null +++ b/tests/test_awex_sglang_plugin.py @@ -0,0 +1,245 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import pytest +import torch + +from areal.api.cli_args import MegatronEngineConfig, PPOActorConfig +from areal.engine.awex.colocate_reader import ( + _PhysicalDeviceMetaServerClient, +) +from areal.engine.awex.memory_saver import patch_tms_hook_mode +from areal.engine.awex.sglang_plugin import ( + AwexSchedulerPlugin, + _load_sglang_plugins_if_available, + _resolve_transfer_rank, + _writer_version_key, +) + + +def test_load_sglang_plugins_accepts_runtime_without_registry(monkeypatch): + import areal.engine.awex.sglang_plugin as plugin_module + + def _missing_registry(name): + assert name == "sglang.srt.plugins" + raise ModuleNotFoundError(name=name) + + monkeypatch.setattr(plugin_module.importlib, "import_module", _missing_registry) + + assert _load_sglang_plugins_if_available() is False + + +def test_event_loop_patch_supports_current_metrics_api(): + class Scheduler: + def __init__(self): + self.forward_ct_decode = 7 + self.event_loop_overlap = lambda: None + self.event_loop_normal = lambda: None + self.calls = [] + + def report_decode_stats( + self, can_run_cuda_graph, running_batch=None, num_accepted_tokens=0 + ): + self.calls.append((can_run_cuda_graph, running_batch, num_accepted_tokens)) + + scheduler = Scheduler() + AwexSchedulerPlugin(scheduler)._patch_event_loop() + scheduler.report_decode_stats(True, running_batch=object(), num_accepted_tokens=3) + + assert scheduler._areal_awex_last_decode_stats_ct == 7 + assert scheduler.calls[0][0] is True + assert scheduler.calls[0][2] == 3 + + +def test_memory_transitions_are_idempotent(): + class Scheduler: + def __init__(self): + self.offload_tags = set() + self.calls = [] + + def release_memory_occupation(self, request): + self.calls.append(("release", list(request.tags))) + self.offload_tags.update(request.tags) + + def resume_memory_occupation(self, request): + self.calls.append(("resume", list(request.tags))) + self.offload_tags.difference_update(request.tags) + + scheduler = Scheduler() + AwexSchedulerPlugin(scheduler)._patch_memory_transitions() + request = SimpleNamespace(tags=["kv_cache"]) + + scheduler.release_memory_occupation(request) + scheduler.release_memory_occupation(request) + scheduler.resume_memory_occupation(request) + scheduler.resume_memory_occupation(request) + + assert scheduler.calls == [ + ("release", ["kv_cache"]), + ("resume", ["kv_cache"]), + ] + + +def test_tms_hook_mode_stays_preload_after_initialization(monkeypatch): + import sys + + class Saver: + def __init__(self): + self._impl_ctor_kwargs = {} + + @property + def hook_mode(self): + raise AttributeError + + @hook_mode.setter + def hook_mode(self, value): + self._impl_ctor_kwargs["hook_mode"] = value + + saver = Saver() + monkeypatch.setitem( + sys.modules, "torch_memory_saver", SimpleNamespace(torch_memory_saver=saver) + ) + monkeypatch.setenv("SGLANG_MEMORY_SAVER_CUDA_GRAPH", "1") + + patch_tms_hook_mode() + saver.hook_mode = "torch" + + assert saver._impl_ctor_kwargs == {} + + +def test_awex_meta_client_uses_physical_device_for_colocate_identity(): + class Client: + def __init__(self): + self.calls = [] + + def add_object_to_set(self, key, value): + self.calls.append(("add", key, value)) + + def get_object(self, key, *args, **kwargs): + self.calls.append(("get", key, args, kwargs)) + + def put_object(self, key, *args, **kwargs): + self.calls.append(("put", key, args, kwargs)) + + def get_object_then_delete(self, key, *args, **kwargs): + self.calls.append(("delete", key, args, kwargs)) + + client = Client() + physical_client = _PhysicalDeviceMetaServerClient(client, physical_gpu_id=6) + + physical_client.add_object_to_set( + "inference_device_rank_entries", ("10.0.0.1", 0, 6) + ) + physical_client.get_object("training_serialized_weights_10.0.0.1_0_3") + physical_client.put_object("weights_update_finished_10.0.0.1_0_3", True) + physical_client.get_object_then_delete("write_finished_10.0.0.1_0_3") + + assert client.calls == [ + ("add", "inference_device_rank_entries", ("10.0.0.1", 6, 6)), + ("get", "training_serialized_weights_10.0.0.1_6_3", (), {}), + ("put", "weights_update_finished_10.0.0.1_6_3", (True,), {}), + ("delete", "write_finished_10.0.0.1_6_3", (), {}), + ] + + +def test_awex_weight_update_runs_without_grad_tracking(): + from areal.engine.awex.colocate_reader import AwexColocateReader + + grad_modes = [] + reader = SimpleNamespace( + update_weights=lambda step_id: grad_modes.append(torch.is_grad_enabled()) + ) + instance = object.__new__(AwexColocateReader) + instance._initialized = True + instance._ensure_reader = lambda: reader + instance._rebuild_derived_weights = lambda: None + + AwexColocateReader.update_weights(instance, 1) + + assert grad_modes == [False] + + +def test_transfer_rank_uses_global_rank_for_isolated_gpu(monkeypatch): + monkeypatch.setenv("RANK", "7") + monkeypatch.setenv("WORLD_SIZE", "8") + + assert ( + _resolve_transfer_rank( + infer_world_size=8, + gpu_id=0, + node_id=0, + nnodes=1, + instance_world_size=1, + ) + == 7 + ) + + +@pytest.mark.parametrize(("tp_size", "pp_size"), [(4, 1), (1, 4)]) +def test_scheduler_instance_world_size_includes_tp_and_pp(tp_size, pp_size): + scheduler = SimpleNamespace( + server_args=SimpleNamespace(tp_size=tp_size, pp_size=pp_size) + ) + + assert AwexSchedulerPlugin(scheduler)._instance_world_size() == 4 + + +def test_transfer_rank_uses_scheduler_gpu_for_multi_gpu_server(monkeypatch): + monkeypatch.setenv("RANK", "5") + monkeypatch.setenv("WORLD_SIZE", "32") + + ranks = [ + _resolve_transfer_rank( + infer_world_size=32, + gpu_id=gpu_id, + node_id=2, + nnodes=4, + instance_world_size=4, + ) + for gpu_id in range(4) + ] + + assert ranks == [16, 17, 18, 19] + + +def test_transfer_rank_falls_back_to_node_local_identity(monkeypatch): + monkeypatch.delenv("AWEX_TRANSFER_RANK", raising=False) + monkeypatch.delenv("RANK", raising=False) + monkeypatch.delenv("WORLD_SIZE", raising=False) + + assert ( + _resolve_transfer_rank( + infer_world_size=16, + gpu_id=3, + node_id=1, + nnodes=2, + instance_world_size=1, + ) + == 11 + ) + + +def test_physical_gpu_id_uses_noncontiguous_visible_device(monkeypatch): + monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "2,5,6,7") + scheduler = SimpleNamespace(gpu_id=1) + gpu_id = AwexSchedulerPlugin(scheduler)._physical_gpu_id() + + assert gpu_id == 5 + assert _writer_version_key("10.0.0.1", gpu_id) == "awex_writer_version_10.0.0.1_5" + + +def test_physical_gpu_id_rejects_uuid_visible_device(monkeypatch): + monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "GPU-deadbeef") + + with pytest.raises(ValueError, match="numeric CUDA_VISIBLE_DEVICES"): + AwexSchedulerPlugin(SimpleNamespace(gpu_id=0))._physical_gpu_id() + + +def test_awex_rejects_megatron_without_ddp_flat_buffers(): + with pytest.raises(ValueError, match="requires megatron.wrap_with_ddp=true"): + PPOActorConfig( + backend="megatron:d1", + weight_update_mode="awex", + megatron=MegatronEngineConfig(wrap_with_ddp=False), + ) diff --git a/tests/test_environ.py b/tests/test_environ.py index edf016c008..a480d8d34e 100644 --- a/tests/test_environ.py +++ b/tests/test_environ.py @@ -99,3 +99,21 @@ def test_get_bool_env_var_can_strip_legacy_dte_values(monkeypatch): ) is True ) + + +def test_numeric_env_helpers_parse_supported_values(monkeypatch): + environ = _load_environ(monkeypatch) + monkeypatch.setenv("AREAL_TEST_FLOAT", "1.25") + monkeypatch.setenv("AREAL_TEST_INT", "7") + + assert environ.get_float_env_var("AREAL_TEST_FLOAT", 0.0) == 1.25 + assert environ.get_int_env_var("AREAL_TEST_INT", 0) == 7 + + +def test_numeric_env_helpers_use_defaults_for_invalid_values(monkeypatch): + environ = _load_environ(monkeypatch) + monkeypatch.setenv("AREAL_TEST_FLOAT", "many") + monkeypatch.setenv("AREAL_TEST_INT", "several") + + assert environ.get_float_env_var("AREAL_TEST_FLOAT", 2.5) == 2.5 + assert environ.get_int_env_var("AREAL_TEST_INT", 3) == 3 diff --git a/tests/test_eval_dispatch.py b/tests/test_eval_dispatch.py index 5f32edf5fc..3088e41642 100644 --- a/tests/test_eval_dispatch.py +++ b/tests/test_eval_dispatch.py @@ -13,6 +13,7 @@ _dispatch_tensors, _pad_eval_batch, ) +from areal.infra.rpc.rtensor import RTensor, TensorShardInfo from areal.trainer.rw.rw_engine import ( RWController, RWEngine, @@ -99,6 +100,46 @@ def test_pad_eval_batch_pads_when_n_less_than_dp(self): assert len(padded) == 4 assert _count_dummies(padded) == 2 + def test_pad_eval_batch_reserves_active_pp_microbatches_per_dp(self): + items = [_make_item(0)] + + (padded,) = _pad_eval_batch( + (items,), + dp_size=2, + min_items_per_dp=16, + items_per_dp_divisor=8, + active_dummies=True, + ) + splits, _ = _dispatch_tensors(padded, dp_size=2) + + assert len(padded) == 32 + assert [len(split) for split in splits] == [16, 16] + assert all( + cast(torch.Tensor, item["attention_mask"]).count_nonzero() == 1 + for item in padded[1:] + ) + + @pytest.mark.parametrize( + ("input_size", "expected_size"), + [(17, 32), (33, 48), (34, 48)], + ) + def test_pad_eval_batch_aligns_each_dp_shard_for_pp8( + self, input_size, expected_size + ): + items = [_make_item(index) for index in range(input_size)] + + (padded,) = _pad_eval_batch( + (items,), + dp_size=2, + min_items_per_dp=16, + items_per_dp_divisor=8, + active_dummies=True, + ) + splits, _ = _dispatch_tensors(padded, dp_size=2) + + assert len(padded) == expected_size + assert all(len(split) % 8 == 0 for split in splits) + def test_dispatch_tensors_raises_when_not_divisible(self): items = [_make_item(i) for i in range(7)] with pytest.raises(ValueError, match="divisible"): @@ -144,6 +185,47 @@ def test_make_dummy_eval_item_schema(self): cast(dict[str, list[str]], template["meta"])["tag"].append("y") assert dummy["meta"] == {"tag": ["x"]} + def test_make_dummy_eval_item_materializes_rtensor_metadata_locally(self): + remote = RTensor( + shard=TensorShardInfo(shard_id="input-ids", node_addr="node:1"), + data=torch.empty((1, 128), dtype=torch.long, device="meta"), + ) + template = { + "input_ids": remote, + "attention_mask": RTensor( + shard=TensorShardInfo(shard_id="attention-mask", node_addr="node:1"), + data=torch.empty((1, 128), dtype=torch.bool, device="meta"), + ), + } + + dummy = make_dummy_eval_item(template, active_attention=True) + + assert isinstance(dummy["input_ids"], torch.Tensor) + assert dummy["input_ids"].device.type == "cpu" + assert dummy["input_ids"].shape == (1, 1) + assert dummy["attention_mask"].tolist() == [[True]] + + def test_make_dummy_eval_item_preserves_multi_sample_group_for_microbatches(self): + """A padded DP rank can match real multi-sample microbatch counts.""" + template = { + "input_ids": torch.ones((8, 16), dtype=torch.long), + "attention_mask": torch.ones((8, 16), dtype=torch.bool), + "loss_mask": torch.ones((8, 16), dtype=torch.bool), + "multi_modal_input": [{} for _ in range(8)], + } + + dummy = make_dummy_eval_item(template, active_attention=True) + mb_list = split_padded_tensor_dict_into_mb_list( + cast(dict[str, Any], dummy), + MicroBatchSpec(n_mbs=4, max_tokens_per_mb=8), + ) + + assert cast(torch.Tensor, dummy["input_ids"]).shape == (8, 1) + assert cast(torch.Tensor, dummy["attention_mask"]).shape == (8, 1) + assert cast(torch.Tensor, dummy["attention_mask"]).count_nonzero() == 8 + assert len(cast(list[dict[str, Any]], dummy["multi_modal_input"])) == 8 + assert len(mb_list.mbs) == 4 + def test_pad_eval_batch_keeps_multimodal_payload_aligned(self): items: list[dict[str, object]] = [] for _ in range(3): diff --git a/tests/test_if_gap_reward.py b/tests/test_if_gap_reward.py new file mode 100644 index 0000000000..0934270210 --- /dev/null +++ b/tests/test_if_gap_reward.py @@ -0,0 +1,44 @@ +# SPDX-License-Identifier: Apache-2.0 + +import pytest + +from areal.reward import if_gap + + +@pytest.mark.parametrize( + ("completion", "expected"), + [ + ("workanswer", "answer"), + ("unfinished", ""), + ("plain answer", "plain answer"), + ], +) +def test_extract_visible_answer(completion, expected): + assert if_gap.extract_visible_answer(completion) == expected + + +@pytest.mark.parametrize( + ("results", "expected"), + [([True, False], 1.0 / 3.0), ([True, True], 1.0), ([], 0.0)], +) +def test_if_gap_reward_combines_pass_rate_and_strict_bonus( + monkeypatch, results, expected +): + monkeypatch.setattr(if_gap, "score_if_gap_spec", lambda *_: results) + + reward = if_gap.if_gap_reward_fn( + prompt="", + completions="answer", + verify_engine="test", + spec={}, + ) + + assert reward == pytest.approx(expected) + + +def test_ifrl_loader_requires_explicit_root(monkeypatch): + monkeypatch.delenv("IF_SYNTH_ROOT", raising=False) + monkeypatch.setattr(if_gap, "_ifrl_engines", None) + + with pytest.raises(FileNotFoundError, match="not configured"): + if_gap._load_ifrl_engines() diff --git a/tests/test_inference_engines.py b/tests/test_inference_engines.py index ad00296b95..329bcfd566 100644 --- a/tests/test_inference_engines.py +++ b/tests/test_inference_engines.py @@ -1,6 +1,7 @@ """Test suite for remote inference engines (vLLM and SGLang).""" import os +from types import SimpleNamespace import pytest import torch.distributed as dist @@ -27,6 +28,110 @@ IS_SGLANG_INSTALLED = is_available("sglang") +def test_sglang_uses_functional_health_check_by_default() -> None: + from areal.engine.sglang_remote import SGLangBackend + + assert SGLangBackend().get_health_check_request().endpoint == "/health" + + +def test_sglang_awex_uses_metadata_readiness_check(monkeypatch) -> None: + from areal.engine.sglang_remote import SGLangBackend + + backend = SGLangBackend() + monkeypatch.setattr( + "areal.engine.sglang_remote.SGLangConfig.build_cmd_from_args", + lambda args: ["sglang.launch_server"], + ) + monkeypatch.setattr( + "areal.engine.sglang_remote.subprocess.Popen", + lambda *args, **kwargs: SimpleNamespace(pid=1, poll=lambda: None), + ) + + backend.launch_server( + {"model_path": "model", "awex_meta_server_addr": "127.0.0.1:1234"} + ) + + assert backend.get_health_check_request().endpoint == "/model_info" + + +def test_wait_for_server_raises_when_subprocess_exits() -> None: + from areal.infra.remote_inf_engine import RemoteInfEngine + + engine = RemoteInfEngine( + InferenceEngineConfig(setup_timeout=60), + backend=SimpleNamespace(), + ) + engine.check_health = lambda _: False + process = SimpleNamespace( + pid=123, + args=["python3", "-m", "sglang.launch_server"], + returncode=42, + poll=lambda: 42, + ) + + with pytest.raises(RuntimeError, match="pid=123.*code 42"): + engine._wait_for_server("127.0.0.1:1", process=process) + + +def test_awex_sglang_child_drops_training_ld_preload(monkeypatch) -> None: + from areal.engine.sglang_remote import SGLangBackend + + backend = SGLangBackend() + monkeypatch.setenv("LD_PRELOAD", "/tmp/torch_memory_saver.so") + monkeypatch.setenv("TMS_INIT_ENABLE", "1") + monkeypatch.setenv("TMS_INIT_ENABLE_CPU_BACKUP", "1") + monkeypatch.delenv("AREAL_SGLANG_DROP_LD_PRELOAD", raising=False) + monkeypatch.setenv("AWEX_META_SERVER_ADDR", "127.0.0.1:1234") + + captured = {} + + class Process: + pid = 1 + + def poll(self): + return None + + def fake_popen(cmd, *, env, stdout, stderr): + captured["cmd"] = cmd + captured["env"] = env + return Process() + + monkeypatch.setattr("areal.engine.sglang_remote.subprocess.Popen", fake_popen) + monkeypatch.setattr( + "areal.engine.sglang_remote.SGLangConfig.build_cmd_from_args", + lambda args: ["sglang.launch_server"], + ) + backend.launch_server({"model_path": "model"}) + + assert "LD_PRELOAD" not in captured["env"] + assert "TMS_INIT_ENABLE" not in captured["env"] + assert "TMS_INIT_ENABLE_CPU_BACKUP" not in captured["env"] + + +def test_awex_sglang_child_enables_cuda_graph_memory_saver(monkeypatch) -> None: + """AWEX enables graph region capture before the SGLang child starts.""" + from areal.engine.sglang_remote import SGLangBackend + + backend = SGLangBackend() + monkeypatch.setenv("AWEX_META_SERVER_ADDR", "127.0.0.1:1234") + monkeypatch.delenv("SGLANG_MEMORY_SAVER_CUDA_GRAPH", raising=False) + captured = {} + + def fake_popen(cmd, *, env, stdout, stderr): + captured["env"] = env + return SimpleNamespace(pid=1, poll=lambda: None) + + monkeypatch.setattr("areal.engine.sglang_remote.subprocess.Popen", fake_popen) + monkeypatch.setattr( + "areal.engine.sglang_remote.SGLangConfig.build_cmd_from_args", + lambda args: ["sglang.launch_server"], + ) + + backend.launch_server({"model_path": "model", "enable_memory_saver": True}) + + assert captured["env"]["SGLANG_MEMORY_SAVER_CUDA_GRAPH"] == "1" + + def _dummy_reward_fn(*args, **kwargs): """Dummy reward function for testing.""" return 1.0 diff --git a/tests/test_megatron_mopd_teacher.py b/tests/test_megatron_mopd_teacher.py new file mode 100644 index 0000000000..8b315ce3cb --- /dev/null +++ b/tests/test_megatron_mopd_teacher.py @@ -0,0 +1,513 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import os +import subprocess +import sys +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch + +from areal.api import WeightUpdateMeta, Worker +from areal.engine import MegatronEngine, MegatronScoringEngine +from areal.engine.awex.colocate_writer import AwexWeightPublisher +from areal.infra.controller.train_controller import TrainController +from areal.trainer.mopd.scoring import MOPDTeacherController +from areal.trainer.rl_trainer import PPOTrainer + + +def _device_worker_env() -> dict[str, str]: + """Restore controller-hidden devices for a fresh GPU worker subprocess.""" + env = os.environ.copy() + hidden_env = env.get("AREAL_CONTROLLER_HIDDEN_DEVICE_ENV") + if hidden_env: + if env.get("AREAL_CONTROLLER_ORIG_DEVICES_SET") == "1": + env[hidden_env] = env.get("AREAL_CONTROLLER_ORIG_DEVICES", "") + else: + env.pop(hidden_env, None) + env["AREAL_ROLE_WORKER"] = "1" + repo_root = str(Path(__file__).resolve().parents[1]) + pythonpath = env.get("PYTHONPATH", "") + env["PYTHONPATH"] = repo_root if not pythonpath else f"{repo_root}:{pythonpath}" + return env + + +def _assigned_cuda_devices() -> list[str]: + original = os.environ.get("AREAL_CONTROLLER_ORIG_DEVICES", "") + if os.environ.get("AREAL_CONTROLLER_HIDDEN_DEVICE_ENV") == "CUDA_VISIBLE_DEVICES": + return [device for device in original.split(",") if device] + return [str(device) for device in range(torch.cuda.device_count())] + + +def test_mopd_runtime_topology_accepts_configured_pp8(monkeypatch): + engine = object.__new__(MegatronEngine) + engine.parallel_strategy = SimpleNamespace(pipeline_parallel_size=8) + monkeypatch.setattr( + "areal.engine.megatron_engine.mpu.get_pipeline_model_parallel_world_size", + lambda: 8, + ) + + engine.assert_mopd_runtime_topology() + + +def test_mopd_runtime_topology_rejects_config_runtime_mismatch(monkeypatch): + engine = object.__new__(MegatronEngine) + engine.parallel_strategy = SimpleNamespace(pipeline_parallel_size=8) + monkeypatch.setattr( + "areal.engine.megatron_engine.mpu.get_pipeline_model_parallel_world_size", + lambda: 4, + ) + + with pytest.raises(RuntimeError, match="configured PP=8, runtime PP=4"): + engine.assert_mopd_runtime_topology() + + +def test_megatron_scoring_engine_computes_logp_without_ppo_actor(): + events = [] + engine = object.__new__(MegatronScoringEngine) + engine.eval = lambda: events.append("eval") + engine.forward = lambda *, input_, aggregate_fn: aggregate_fn( + [input_["part0"], input_["part1"]] + ) + + result = engine._compute_logp( + {"part0": torch.tensor([1.0]), "part1": torch.tensor([2.0])} + ) + + assert events == ["eval"] + torch.testing.assert_close(result, torch.tensor([1.0, 2.0]), rtol=0.0, atol=0.0) + + +def test_mopd_teacher_controller_preserves_pipeline_active_dummies(monkeypatch): + controller = object.__new__(MOPDTeacherController) + controller.train_alloc = SimpleNamespace( + parallel=SimpleNamespace(pp_size=2, dp_size=1) + ) + controller.config = SimpleNamespace(mb_spec=SimpleNamespace(n_mbs=1, granularity=1)) + calls = [] + + def pad(args, kwargs, **options): + calls.append(options) + return ([{"id": index} for index in range(4)],), kwargs + + monkeypatch.setattr(controller, "_pad_eval_dispatch_args", pad) + monkeypatch.setattr( + controller, + "_custom_function_call", + lambda *_args, **_kwargs: ["real", "dummy1", "dummy2", "dummy3"], + ) + + real, dummies = controller.compute_logp_padded([{"id": 0}]) + + assert real == ["real"] + assert dummies == ["dummy1", "dummy2", "dummy3"] + assert calls == [ + { + "group_size": 1, + "min_items_per_dp": 4, + "items_per_dp_divisor": 2, + "active_dummies": True, + } + ] + + +def test_teacher_weight_residency_adapter_has_no_awex_publication_state(): + engine = object.__new__(MegatronEngine) + engine._weight_residency = None + engine._awex_publisher = None + engine.logger = SimpleNamespace(info=lambda *_args, **_kwargs: None) + + engine.init_weight_residency_adapter() + + assert engine._weight_residency is not None + assert engine._awex_publisher is None + + +def test_awex_publisher_composes_engine_weight_residency(monkeypatch): + engine = object.__new__(MegatronEngine) + engine._weight_residency = None + engine._awex_publisher = None + engine.logger = SimpleNamespace(info=lambda *_args, **_kwargs: None) + monkeypatch.setattr( + "areal.engine.awex.colocate_writer.AwexWeightPublisher.eager_publish_train_info", + lambda *_args, **_kwargs: None, + ) + + engine.init_awex_adapter() + first_publisher = engine._awex_publisher + engine.init_awex_adapter() + + assert engine._weight_residency is not None + assert engine._awex_publisher is first_publisher + assert engine._awex_publisher.residency is engine._weight_residency + + +@pytest.mark.parametrize("weights_released", [False, True]) +def test_awex_publisher_prepares_residency_in_oom_safe_order( + weights_released: bool, +): + events = [] + + class _Residency: + def is_released(self, tag): + events.append(("is_released", tag)) + return weights_released + + def release_memory(self, tags): + events.append(("release", tags)) + + def release_grad_memory(self): + events.append(("release_grad", None)) + + def resume_memory(self, tags): + events.append(("resume", tags)) + + publisher = AwexWeightPublisher(SimpleNamespace(), _Residency()) + + publisher._prepare_residency_for_publish() + + assert events[:3] == [ + ("is_released", "weights"), + ("release", ["optimizer"]), + ("release_grad", None), + ] + if weights_released: + assert events[3:] == [("resume", ["weights"])] + else: + assert len(events) == 3 + + +def test_awex_actor_worker_does_not_reenter_rollout(monkeypatch): + events = [] + + class _Adapter: + def execute_colocate_weight_update(self, version): + events.append(("execute", version)) + + def finish_colocate_weight_update(self, training_world_size): + events.append(("finish", training_world_size)) + + class _Rollout: + def onload(self, tags=None): + raise AssertionError(f"actor worker re-entered rollout onload: {tags}") + + def continue_generation(self): + raise AssertionError("actor worker re-entered rollout generation") + + engine = object.__new__(MegatronEngine) + engine._awex_publisher = _Adapter() + engine.rollout_engine = _Rollout() + engine.rollout_coordinator = object() + engine.process_group_initialized = True + engine._cpu_group = object() + + monkeypatch.setattr( + "areal.engine.megatron_engine.dist.barrier", + lambda group: events.append(("barrier", group)), + ) + monkeypatch.setattr( + "areal.engine.megatron_engine.dist.get_world_size", lambda group: 8 + ) + + engine.update_weights(WeightUpdateMeta(type="awex", version=4)) + + assert events == [ + ("execute", 4), + ("barrier", engine.cpu_group), + ("finish", 8), + ("barrier", engine.cpu_group), + ] + + +def test_awex_weight_update_requires_initialized_publisher(): + engine = object.__new__(MegatronEngine) + engine._awex_publisher = None + engine.rollout_engine = object() + engine.rollout_coordinator = object() + + with pytest.raises(RuntimeError, match="before publisher initialization"): + engine.update_weights(WeightUpdateMeta(type="awex", version=1)) + + +def test_megatron_engine_uses_residency_without_awex_publisher(monkeypatch): + events = [] + + class _Residency: + def release_memory(self, tags): + events.append(("release", tags)) + + def resume_memory(self, tags): + events.append(("resume", tags)) + + class _Stats: + def log(self, message): + events.append(("stats", message)) + + engine = object.__new__(MegatronEngine) + engine._weight_residency = _Residency() + engine._awex_publisher = None + engine.process_group_initialized = True + engine._cpu_group = object() + engine.is_offload = False + engine.get_device_stats = lambda: _Stats() + engine._log_weight_residency_stats = lambda phase: events.append(("log", phase)) + monkeypatch.setattr( + "areal.engine.megatron_engine.current_platform.clear_memory", + lambda: events.append(("clear", None)), + ) + monkeypatch.setattr( + "areal.engine.megatron_engine.current_platform.synchronize", + lambda: events.append(("synchronize", None)), + ) + monkeypatch.setattr( + "areal.engine.megatron_engine.dist.barrier", + lambda group: events.append(("barrier", group)), + ) + + engine.offload() + engine.onload() + + assert ("release", ["optimizer", "weights"]) in events + assert ("resume", ["optimizer", "weights"]) in events + assert engine.is_offload is False + + +def test_awex_controller_discards_unfinished_requests_before_restoring_kv(): + events = [] + + class _Actor: + def update_weights(self, meta): + events.append(("update", meta.version)) + + def set_version(self, version): + events.append(("actor_version", version)) + + class _Rollout: + def set_version(self, version): + events.append(("rollout_version", version)) + + def abort_all_requests(self): + events.append(("abort_all", None)) + + def onload(self, tags=None): + events.append(("onload", tags)) + + async def continue_generation(self): + events.append(("continue", None)) + + trainer = object.__new__(PPOTrainer) + trainer.config = SimpleNamespace( + actor=SimpleNamespace(_version="v1", weight_update_mode="awex") + ) + trainer.actor = _Actor() + trainer.rollout = _Rollout() + trainer.critic = None + trainer.eval_rollout = None + + trainer._update_weights_and_publish_version( + WeightUpdateMeta(type="awex", version=1), 1 + ) + trainer._update_weights_and_publish_version( + WeightUpdateMeta(type="awex", version=2), 2 + ) + + assert events == [ + ("update", 1), + ("actor_version", 1), + ("rollout_version", 1), + ("abort_all", None), + ("onload", ["cuda_graph"]), + ("onload", ["kv_cache"]), + ("continue", None), + ("update", 2), + ("actor_version", 2), + ("rollout_version", 2), + ("abort_all", None), + ("onload", ["cuda_graph"]), + ("onload", ["kv_cache"]), + ("continue", None), + ] + + +def test_awex_controller_does_not_resume_when_discard_fails(): + events = [] + + class _Actor: + def update_weights(self, meta): + events.append(("update", meta.version)) + + def set_version(self, version): + events.append(("actor_version", version)) + + class _Rollout: + def set_version(self, version): + events.append(("rollout_version", version)) + + def abort_all_requests(self): + events.append(("abort_all", None)) + raise RuntimeError("discard failed") + + def onload(self, tags=None): + raise AssertionError(f"restored KV after discard failure: {tags}") + + def continue_generation(self): + raise AssertionError("resumed generation after discard failure") + + trainer = object.__new__(PPOTrainer) + trainer.config = SimpleNamespace( + actor=SimpleNamespace(_version="v1", weight_update_mode="awex") + ) + trainer.actor = _Actor() + trainer.rollout = _Rollout() + trainer.critic = None + trainer.eval_rollout = None + + with pytest.raises(RuntimeError, match="discard failed"): + trainer._update_weights_and_publish_version( + WeightUpdateMeta(type="awex", version=1), 1 + ) + + assert events == [ + ("update", 1), + ("actor_version", 1), + ("rollout_version", 1), + ("abort_all", None), + ] + + +def test_non_awex_weight_update_does_not_restore_rollout(): + events = [] + + class _Actor: + def update_weights(self, meta): + events.append(("update", meta.version)) + + def set_version(self, version): + events.append(("actor_version", version)) + + class _Rollout: + def set_version(self, version): + events.append(("rollout_version", version)) + + def onload(self, tags=None): + raise AssertionError(f"unexpected rollout onload: {tags}") + + def continue_generation(self): + raise AssertionError("unexpected rollout generation resume") + + trainer = object.__new__(PPOTrainer) + trainer.config = SimpleNamespace( + actor=SimpleNamespace(_version="v1", weight_update_mode="disk") + ) + trainer.actor = _Actor() + trainer.rollout = _Rollout() + trainer.critic = None + trainer.eval_rollout = None + + trainer._update_weights_and_publish_version( + WeightUpdateMeta(type="disk", version=3), 3 + ) + + assert events == [ + ("update", 3), + ("actor_version", 3), + ("rollout_version", 3), + ] + + +@pytest.mark.parametrize("method", ["offload", "onload"]) +def test_teacher_lifecycle_collective_waits_all_ranks_without_retry(method): + """A failed teacher rank cannot trigger a partial collective retry.""" + events = [] + + class _Scheduler: + async def async_call_engine(self, worker_id, method, engine_name, **kwargs): + events.append(("start", worker_id, method, engine_name, kwargs)) + if worker_id.endswith("/0"): + raise RuntimeError("rank zero failed") + await asyncio.sleep(0.01) + events.append(("complete", worker_id)) + + controller = object.__new__(TrainController) + controller.scheduler = _Scheduler() + controller._worker_role = "mopd-teacher" + controller.workers = [ + Worker(id="mopd-teacher/0", ip="127.0.0.1"), + Worker(id="mopd-teacher/1", ip="127.0.0.1"), + ] + + with pytest.raises(ExceptionGroup, match=f"collective {method} failed"): + getattr(controller, method)() + + starts = [event for event in events if event[0] == "start"] + assert len(starts) == 2 + assert all(event[2] == method for event in starts) + assert all(event[4]["max_retries"] == 1 for event in starts) + assert ("complete", "mopd-teacher/1") in events + + +def test_megatron_teacher_offload_skips_non_ddp_model_chunks(monkeypatch): + """Scoring teachers can discard DDP grads without assuming every chunk is DDP.""" + events = [] + + class _Stats: + def log(self, message): + events.append(("stats", message)) + + engine = object.__new__(MegatronEngine) + engine._weight_residency = None + engine.mcore_config = SimpleNamespace(disable_grad_buffers_cpu_backup=True) + engine.model = [object()] + engine.process_group_initialized = True + engine._cpu_group = object() + engine.is_offload = False + engine.get_device_stats = lambda: _Stats() + monkeypatch.setattr("areal.engine.megatron_engine.is_tms_enabled", lambda: True) + monkeypatch.setattr( + "areal.engine.megatron_engine.current_platform.clear_memory", + lambda: events.append(("clear", None)), + ) + monkeypatch.setattr( + "areal.engine.megatron_engine.current_platform.synchronize", + lambda: events.append(("synchronize", None)), + ) + monkeypatch.setattr( + "areal.engine.megatron_engine.torch_memory_saver.pause", + lambda: events.append(("pause", None)), + ) + monkeypatch.setattr( + "areal.engine.megatron_engine.dist.barrier", + lambda group: events.append(("barrier", group)), + ) + + engine.offload() + + assert engine.is_offload is True + assert ("pause", None) in events + assert events.index(("synchronize", None)) < events.index( + ("barrier", engine.cpu_group) + ) + + +@pytest.mark.gpu +@pytest.mark.skipif(not _assigned_cuda_devices(), reason="requires one CUDA GPU") +@pytest.mark.parametrize("mode", ["fallback", "native"]) +def test_teacher_residency_adapter_releases_and_restores_cuda_flat_buffer(mode): + """Both MCore paths release CUDA storage and preserve teacher weights.""" + runner = Path(__file__).parent / "torchrun" / "run_mopd_teacher_residency.py" + result = subprocess.run( + [sys.executable, str(runner), "--mode", mode], + capture_output=True, + text=True, + env=_device_worker_env(), + cwd=Path(__file__).resolve().parents[1], + timeout=120, + ) + assert result.returncode == 0, ( + f"CUDA residency worker failed\nSTDOUT:\n{result.stdout}\nSTDERR:\n{result.stderr}" + ) + assert f"Passed mode={mode}" in result.stdout diff --git a/tests/test_mopd_compatibility.py b/tests/test_mopd_compatibility.py new file mode 100644 index 0000000000..059a9a0b1f --- /dev/null +++ b/tests/test_mopd_compatibility.py @@ -0,0 +1,95 @@ +# SPDX-License-Identifier: Apache-2.0 + +import json + +import pytest + +from areal.trainer.mopd.compatibility import validate_mopd_model_compatibility + + +def _write_checkpoint(path, *, hidden_size: int, vocab: dict[str, int]) -> None: + path.mkdir() + (path / "config.json").write_text( + json.dumps( + { + "architectures": ["QwenForCausalLM"], + "model_type": "qwen", + "hidden_size": hidden_size, + "num_hidden_layers": 2, + "num_attention_heads": 2, + "vocab_size": len(vocab), + "torch_dtype": "bfloat16", + "transformers_version": "test-only", + } + ), + encoding="utf-8", + ) + (path / "vocab.json").write_text(json.dumps(vocab), encoding="utf-8") + (path / "tokenizer_config.json").write_text( + json.dumps({"chat_template": f"template for {path.name}"}), + encoding="utf-8", + ) + + +def test_compatibility_allows_heterogeneous_actor_and_matching_teachers(tmp_path): + """Actor structure may differ while resident teacher checkpoints stay aligned.""" + vocab = {"a": 0, "b": 1} + actor = tmp_path / "actor" + teacher_a = tmp_path / "teacher-a" + teacher_b = tmp_path / "teacher-b" + _write_checkpoint(actor, hidden_size=8, vocab=vocab) + _write_checkpoint(teacher_a, hidden_size=32, vocab=vocab) + _write_checkpoint(teacher_b, hidden_size=32, vocab=vocab) + + fingerprints = validate_mopd_model_compatibility( + actor, + {"teacher_a": teacher_a, "teacher_b": teacher_b}, + ) + + assert ( + fingerprints["actor"]["architecture_sha256"] + != fingerprints["teacher_a"]["architecture_sha256"] + ) + assert ( + fingerprints["teacher_a"]["architecture_sha256"] + == fingerprints["teacher_b"]["architecture_sha256"] + ) + + +def test_compatibility_rejects_teacher_architecture_mismatch(tmp_path): + """One persistent controller cannot load structurally different teachers.""" + vocab = {"a": 0, "b": 1} + actor = tmp_path / "actor" + teacher_a = tmp_path / "teacher-a" + teacher_b = tmp_path / "teacher-b" + _write_checkpoint(actor, hidden_size=8, vocab=vocab) + _write_checkpoint(teacher_a, hidden_size=32, vocab=vocab) + _write_checkpoint(teacher_b, hidden_size=64, vocab=vocab) + + with pytest.raises(ValueError, match="teachers must share one architecture"): + validate_mopd_model_compatibility( + actor, + {"teacher_a": teacher_a, "teacher_b": teacher_b}, + ) + + +def test_compatibility_rejects_token_id_mismatch(tmp_path): + """Teacher log-probabilities require the actor's exact token-ID mapping.""" + actor = tmp_path / "actor" + teacher = tmp_path / "teacher" + _write_checkpoint(actor, hidden_size=8, vocab={"a": 0, "b": 1}) + _write_checkpoint(teacher, hidden_size=32, vocab={"a": 1, "b": 0}) + + with pytest.raises(ValueError, match="compatible token-ID mappings"): + validate_mopd_model_compatibility(actor, {"teacher": teacher}) + + +def test_compatibility_rejects_reserved_actor_teacher_id(tmp_path): + """A teacher cannot overwrite the actor fingerprint in the result mapping.""" + actor = tmp_path / "actor-model" + teacher = tmp_path / "teacher-model" + _write_checkpoint(actor, hidden_size=8, vocab={"a": 0, "b": 1}) + _write_checkpoint(teacher, hidden_size=32, vocab={"a": 1, "b": 0}) + + with pytest.raises(ValueError, match="teacher ID 'actor' is reserved"): + validate_mopd_model_compatibility(actor, {"actor": teacher}) diff --git a/tests/test_mopd_config.py b/tests/test_mopd_config.py new file mode 100644 index 0000000000..3cd5d3d701 --- /dev/null +++ b/tests/test_mopd_config.py @@ -0,0 +1,464 @@ +# SPDX-License-Identifier: Apache-2.0 + +import math + +import pytest +from omegaconf import OmegaConf + +from areal.api.cli_args import ( + DatasetSourceConfig, + InferenceEngineConfig, + MOPDConfig, + MOPDLossConfig, + MOPDTeacherEngineConfig, + MOPDTeacherManagerConfig, + MOPDTeacherSpec, + PPOActorConfig, + PPOConfig, + SchedulerConfig, + SchedulingSpec, + SchedulingStrategy, + TeacherConfig, + TrainDatasetConfig, + to_structured_cfg, +) +from areal.trainer.mopd.execution import MOPDExecutionPlan + +MEGATRON_BACKEND = "megatron:(attn:d1p1t2c2|ffn:d1p1e4)" + + +def _colocation(*, fork: bool) -> SchedulingStrategy: + return SchedulingStrategy(type="colocation", target="actor", fork=fork) + + +def _mopd_config(**overrides) -> MOPDConfig: + kwargs = { + "teachers": { + "agriculture": MOPDTeacherSpec(path="checkpoints/agriculture"), + "swe_agent": MOPDTeacherSpec(path="checkpoints/swe-agent"), + }, + "teacher_groups": { + "single": {"agriculture": 1.0}, + "ensemble": {"agriculture": 0.3, "swe_agent": 0.7}, + }, + "teacher_engine": MOPDTeacherEngineConfig( + backend=MEGATRON_BACKEND, + optimizer=None, + disable_dropout=True, + scheduling_strategy=_colocation(fork=True), + ), + } + kwargs.update(overrides) + return MOPDConfig(**kwargs) + + +def _ppo_config(mopd: MOPDConfig | None = None, **overrides) -> PPOConfig: + kwargs = { + "experiment_name": "mopd-test", + "trial_name": "trial", + "actor": PPOActorConfig( + backend=MEGATRON_BACKEND, + weight_update_mode="awex", + scheduling_spec=(SchedulingSpec(port_count=3),), + ), + "rollout": InferenceEngineConfig( + backend="sglang:d4", + scheduling_strategy=_colocation(fork=True), + ), + "mopd": mopd, + } + if mopd is not None: + kwargs["train_dataset"] = TrainDatasetConfig( + sources=[ + DatasetSourceConfig( + path="datasets/train", + type="rl", + teacher_group=next(iter(mopd.teacher_groups)), + ) + ] + ) + kwargs.update(overrides) + return PPOConfig(**kwargs) + + +def test_ppo_config_without_mopd_preserves_disabled_default(): + """MOPD stays opt-in for existing PPO and GRPO configurations.""" + config = _ppo_config() + + assert config.mopd is None + + +def test_mopd_loss_config_defaults_to_unscaled_distillation(): + """The standalone MOPD objective uses a neutral coefficient by default.""" + config = MOPDLossConfig() + + assert config.rl_coefficient == 0.0 + assert config.distillation_coefficient == 1.0 + + +def test_mopd_teacher_config_exposes_only_scoring_engine_fields(): + config = MOPDTeacherEngineConfig(backend=MEGATRON_BACKEND) + + assert config.optimizer is None + assert config.disable_dropout is True + assert not hasattr(config, "ppo_n_minibatches") + assert not hasattr(config, "eps_clip") + + +@pytest.mark.parametrize( + ("rl_coefficient", "distillation_coefficient", "expected"), + [ + (0.0, 1.0, (True, False)), + (1.0, 0.0, (False, True)), + (0.4, 0.6, (True, True)), + ], +) +def test_mopd_execution_plan_short_circuits_disabled_objectives( + rl_coefficient, distillation_coefficient, expected +): + config = _ppo_config( + _mopd_config( + loss=MOPDLossConfig( + rl_coefficient=rl_coefficient, + distillation_coefficient=distillation_coefficient, + ) + ) + ) + + plan = MOPDExecutionPlan.from_config(config) + + assert plan is not None + assert (plan.requires_teacher_scoring, plan.requires_rl) == expected + + +def test_mopd_forked_rollout_requires_worker_and_nccl_ports(): + actor = PPOActorConfig( + backend=MEGATRON_BACKEND, + weight_update_mode="awex", + scheduling_spec=(SchedulingSpec(port_count=1),), + ) + + with pytest.raises(ValueError, match="port_count >= 2"): + _ppo_config(_mopd_config(), actor=actor) + + +def test_mopd_config_valid_preserves_teacher_group_order_and_weights(): + """Teacher order follows insertion order and group weights stay unnormalized.""" + config = _ppo_config(_mopd_config()) + + assert config.enable_offload is False + assert list(config.mopd.teachers) == ["agriculture", "swe_agent"] + assert list(config.mopd.teacher_groups) == ["single", "ensemble"] + assert config.mopd.teacher_groups["ensemble"] == { + "agriculture": 0.3, + "swe_agent": 0.7, + } + assert config.mopd.loss.rl_coefficient == 0.0 + assert config.mopd.loss.distillation_coefficient == 1.0 + assert config.train_dataset.sources[0].teacher_group == "single" + + +def test_mopd_config_requires_routed_dataset_sources(): + """MOPD rejects the legacy single-path dataset shape.""" + with pytest.raises(ValueError, match="train_dataset.sources must not be empty"): + _ppo_config( + _mopd_config(), + train_dataset=TrainDatasetConfig(path="dataset", type="rl"), + ) + + +def test_dataset_source_teacher_group_defaults_to_none_without_mopd(): + """Dataset mixtures outside MOPD do not require teacher selection.""" + source = DatasetSourceConfig(path="dataset", type="rl") + + assert source.teacher_group is None + + +@pytest.mark.parametrize("teacher_group", ["", " ", 7, True]) +def test_dataset_source_rejects_invalid_optional_teacher_group(teacher_group): + """An explicitly configured teacher group must be a non-empty string.""" + with pytest.raises(ValueError, match="non-empty string or null"): + DatasetSourceConfig( + path="dataset", + type="rl", + teacher_group=teacher_group, + ) + + +def test_mopd_dataset_source_requires_teacher_group(): + """Every MOPD source must select a teacher group explicitly.""" + train_dataset = TrainDatasetConfig( + sources=[DatasetSourceConfig(path="dataset", type="rl")] + ) + + with pytest.raises(ValueError, match="teacher_group must be configured"): + _ppo_config(_mopd_config(), train_dataset=train_dataset) + + +def test_mopd_dataset_source_rejects_unknown_teacher_group(): + """A source cannot select a group absent from mopd.teacher_groups.""" + train_dataset = TrainDatasetConfig( + sources=[ + DatasetSourceConfig(path="dataset", type="rl", teacher_group="missing") + ] + ) + + with pytest.raises(ValueError, match="unknown MOPD teacher group 'missing'"): + _ppo_config(_mopd_config(), train_dataset=train_dataset) + + +def test_mopd_dataset_sources_reject_legacy_path_and_type(): + """Mixture sources and the legacy single-source fields are mutually exclusive.""" + with pytest.raises(ValueError, match="path/type cannot be combined"): + TrainDatasetConfig( + path="legacy", + type="rl", + sources=[ + DatasetSourceConfig(path="dataset", type="rl", teacher_group="single") + ], + ) + + +def test_mopd_config_yaml_shape_converts_to_structured_config(): + """A YAML-shaped mapping survives the production OmegaConf conversion path.""" + raw_config = OmegaConf.create( + { + "experiment_name": "mopd-test", + "trial_name": "trial", + "train_dataset": { + "sources": [{"path": "dataset", "type": "rl", "teacher_group": "math"}] + }, + "saver": { + "experiment_name": "mopd-test", + "trial_name": "trial", + "fileroot": "outputs", + }, + "evaluator": { + "experiment_name": "mopd-test", + "trial_name": "trial", + "fileroot": "outputs", + }, + "stats_logger": { + "experiment_name": "mopd-test", + "trial_name": "trial", + "fileroot": "outputs", + }, + "recover": { + "experiment_name": "mopd-test", + "trial_name": "trial", + "fileroot": "outputs", + }, + "actor": { + "experiment_name": "mopd-test", + "trial_name": "trial", + "backend": MEGATRON_BACKEND, + "weight_update_mode": "awex", + "scheduling_spec": [{"port_count": 3}], + }, + "rollout": { + "backend": "sglang:d4", + "scheduling_strategy": { + "type": "colocation", + "target": "actor", + "fork": True, + }, + }, + "mopd": { + "teachers": {"agriculture": {"path": "checkpoints/agriculture"}}, + "teacher_groups": {"math": {"agriculture": 2.0}}, + "teacher_engine": { + "experiment_name": "mopd-test", + "trial_name": "trial", + "backend": MEGATRON_BACKEND, + "optimizer": None, + "disable_dropout": True, + "scheduling_strategy": { + "type": "colocation", + "target": "actor", + "fork": True, + }, + }, + }, + } + ) + + config = OmegaConf.to_object(to_structured_cfg(raw_config, PPOConfig)) + + assert isinstance(config, PPOConfig) + assert isinstance(config.mopd, MOPDConfig) + assert isinstance(config.mopd.teachers["agriculture"], MOPDTeacherSpec) + assert config.mopd.teacher_groups["math"]["agriculture"] == 2.0 + assert isinstance(config.train_dataset.sources[0], DatasetSourceConfig) + assert config.train_dataset.sources[0].teacher_group == "math" + + +@pytest.mark.parametrize("weight", [-1.0, math.inf, math.nan, True, "1.0"]) +def test_mopd_config_invalid_teacher_group_weight_raises(weight): + """Group weights reject negative, non-finite, boolean, and non-numeric values.""" + with pytest.raises(ValueError, match="finite non-negative|finite and non-negative"): + _mopd_config(teacher_groups={"bad": {"agriculture": weight}}) + + +def test_mopd_config_unknown_teacher_raises(): + """Every teacher group entry must reference a configured teacher.""" + with pytest.raises(ValueError, match="unknown teacher"): + _mopd_config(teacher_groups={"bad": {"missing": 1.0}}) + + +def test_mopd_config_all_zero_teacher_group_raises(): + """A teacher group must include at least one positive teacher weight.""" + with pytest.raises(ValueError, match="at least one positive weight"): + _mopd_config(teacher_groups={"bad": {"agriculture": 0.0}}) + + +@pytest.mark.parametrize("coefficient", [-1.0, math.inf, math.nan, True]) +def test_mopd_loss_config_invalid_coefficient_raises(coefficient): + """Loss coefficients must be finite non-negative numbers.""" + with pytest.raises(ValueError, match="finite"): + MOPDLossConfig(rl_coefficient=coefficient) + + +def test_mopd_loss_config_rejects_disabled_objective(): + with pytest.raises(ValueError, match="cannot both be zero"): + MOPDLossConfig(rl_coefficient=0.0, distillation_coefficient=0.0) + + +def test_mopd_config_rollout_colocation_accepts_forked_workers(): + """The current v1 runtime accepts process-isolated rollout workers.""" + rollout = InferenceEngineConfig( + backend="sglang:d4", + scheduling_strategy=_colocation(fork=True), + ) + + config = _ppo_config(_mopd_config(), rollout=rollout) + + assert config.rollout.scheduling_strategy.fork is True + + +def test_mopd_config_rollout_reused_workers_are_rejected(): + """Same-process rollout is outside the current v1 runtime capability.""" + with pytest.raises(ValueError, match="current MOPD v1 runtime"): + _ppo_config( + _mopd_config(), + rollout=InferenceEngineConfig( + backend="sglang:d4", + scheduling_strategy=_colocation(fork=False), + ), + ) + + +def test_mopd_config_teacher_parallelism_mismatch_raises(): + """Actor and teacher must use exactly the same parallel dimensions.""" + teacher_engine = MOPDTeacherEngineConfig( + backend="megatron:d4", + optimizer=None, + disable_dropout=True, + scheduling_strategy=_colocation(fork=True), + ) + + with pytest.raises(ValueError, match="same parallel strategy"): + _ppo_config(_mopd_config(teacher_engine=teacher_engine)) + + +def test_mopd_config_pipeline_parallelism_accepts_matching_topology(): + """MOPD accepts pipeline parallelism shared by the actor and teachers.""" + backend = "megatron:d1p2t2" + teacher_engine = MOPDTeacherEngineConfig( + backend=backend, + optimizer=None, + disable_dropout=True, + scheduling_strategy=_colocation(fork=True), + ) + actor = PPOActorConfig( + backend=backend, + weight_update_mode="awex", + scheduling_spec=(SchedulingSpec(port_count=3),), + ) + + config = _ppo_config( + _mopd_config(teacher_engine=teacher_engine), + actor=actor, + ) + + assert config.actor.backend == backend + assert config.mopd.teacher_engine.backend == backend + + +def test_mopd_config_local_memory_multi_node_raises(): + """Local-memory checkpoint staging cannot span multiple compute nodes.""" + actor = PPOActorConfig( + backend="megatron:d16", + weight_update_mode="awex", + scheduling_spec=(SchedulingSpec(port_count=3),), + ) + teacher_engine = MOPDTeacherEngineConfig( + backend="megatron:d16", + optimizer=None, + disable_dropout=True, + scheduling_strategy=_colocation(fork=True), + ) + manager = MOPDTeacherManagerConfig(type="local_memory") + + with pytest.raises(ValueError, match="single node"): + _ppo_config( + _mopd_config(teacher_engine=teacher_engine, manager=manager), + actor=actor, + scheduler=SchedulerConfig(type="local"), + ) + + +@pytest.mark.parametrize("scheduler_type", ["ray", "slurm", None]) +def test_mopd_config_local_memory_requires_same_host_local_scheduler(scheduler_type): + """Controller-local staging is rejected when workers may run remotely.""" + manager = MOPDTeacherManagerConfig(type="local_memory") + + with pytest.raises(ValueError, match="scheduler.type='local'"): + _ppo_config( + _mopd_config(manager=manager), + scheduler=SchedulerConfig(type=scheduler_type), + ) + + +def test_mopd_config_local_memory_accepts_local_scheduler(): + """LocalScheduler guarantees controller and fork workers share one host.""" + config = _ppo_config( + _mopd_config(manager=MOPDTeacherManagerConfig(type="local_memory")), + scheduler=SchedulerConfig(type="local"), + ) + + assert config.mopd.manager.type == "local_memory" + + +def test_mopd_config_teacher_v2_raises(): + """MOPD fails fast before selecting a controller without v1 lifecycle APIs.""" + teacher_engine = MOPDTeacherEngineConfig( + backend=MEGATRON_BACKEND, + optimizer=None, + disable_dropout=True, + scheduling_strategy=_colocation(fork=True), + _version="v2", + ) + + with pytest.raises(ValueError, match="requires _version='v1'"): + _ppo_config(_mopd_config(teacher_engine=teacher_engine)) + + +@pytest.mark.parametrize("teacher_id", ["../teacher", "/teacher", "org/teacher"]) +def test_mopd_config_rejects_path_like_teacher_ids(teacher_id): + """Teacher IDs cannot escape or create nested staging directories.""" + with pytest.raises(ValueError, match="filename-safe"): + MOPDConfig( + teachers={teacher_id: MOPDTeacherSpec(path="checkpoints/teacher")}, + teacher_groups={"group": {teacher_id: 1.0}}, + ) + + +def test_mopd_and_legacy_teacher_are_mutually_exclusive(): + """Ambiguous legacy and MOPD teacher configuration fails fast.""" + legacy_teacher = TeacherConfig( + engine_type="train", + train=MOPDTeacherEngineConfig(backend=MEGATRON_BACKEND), + ) + + with pytest.raises(ValueError, match="cannot be configured at the same time"): + _ppo_config(_mopd_config(), teacher=legacy_teacher) diff --git a/tests/test_mopd_dataset.py b/tests/test_mopd_dataset.py new file mode 100644 index 0000000000..d7321c8724 --- /dev/null +++ b/tests/test_mopd_dataset.py @@ -0,0 +1,148 @@ +# SPDX-License-Identifier: Apache-2.0 + +import pytest + +from areal.api.cli_args import DatasetSourceConfig, TrainDatasetConfig +from areal.dataset.mopd import ( + MOPD_ROUTE_METADATA_KEY, + MOPDDataset, + get_mopd_dataset, + is_remote_dataset, +) +from areal.infra.data_service.rdataset import RDataset + + +def test_mopd_dataset_routes_each_source_without_mutating_samples(): + """Every item inherits only its source route and stored samples stay unchanged.""" + samples = { + "math": [{"messages": ["m0"]}, {"messages": ["m1"]}], + "code": [{"messages": ["c0"]}], + } + config = TrainDatasetConfig( + sources=[ + DatasetSourceConfig(path="math", type="rl", teacher_group="math_group"), + DatasetSourceConfig(path="code", type="rl", teacher_group="code_group"), + ] + ) + + dataset = get_mopd_dataset( + config, + source_loader=lambda **kwargs: samples[kwargs["source"].path], + ) + + assert len(dataset) == 3 + assert [dataset[index][MOPD_ROUTE_METADATA_KEY].route for index in range(3)] == [ + "math_group", + "math_group", + "code_group", + ] + assert all("task_type" not in dataset[index] for index in range(3)) + assert all( + MOPD_ROUTE_METADATA_KEY not in sample + for data in samples.values() + for sample in data + ) + + +def test_routed_dataset_uniform_policy_balances_unequal_sources(): + """Uniform policy deterministically cycles shorter sources per epoch.""" + config = TrainDatasetConfig( + mixture_sampling_policy="uniform", + sources=[ + DatasetSourceConfig(path="short", type="rl", teacher_group="short-group"), + DatasetSourceConfig(path="long", type="rl", teacher_group="long-group"), + ], + ) + samples = { + "short": [{"id": "s0"}], + "long": [{"id": "l0"}, {"id": "l1"}, {"id": "l2"}], + } + + dataset = get_mopd_dataset( + config, + source_loader=lambda **kwargs: samples[kwargs["source"].path], + ) + + assert len(dataset) == 6 + assert [dataset[index]["id"] for index in range(6)] == [ + "s0", + "l0", + "s0", + "l1", + "s0", + "l2", + ] + + +@pytest.mark.parametrize("field", [MOPD_ROUTE_METADATA_KEY, "mopd_route"]) +def test_mopd_dataset_rejects_sample_level_route(field): + """Samples cannot override the route declared by their source.""" + dataset = MOPDDataset([([{field: "sample-route"}], "source-route")]) + + with pytest.raises(ValueError, match="configured only on the dataset source"): + dataset[0] + + +def test_dataset_mixture_without_teacher_group_has_no_route_metadata(): + """Non-MOPD mixtures leave samples free of teacher-routing metadata.""" + config = TrainDatasetConfig( + sources=[DatasetSourceConfig(path="general", type="rl")] + ) + samples = [{"id": "sample"}] + + dataset = get_mopd_dataset( + config, + source_loader=lambda **kwargs: samples, + ) + + assert dataset[0] == {"id": "sample"} + assert MOPD_ROUTE_METADATA_KEY not in samples[0] + + +class _RemoteSource(RDataset): + def __init__(self, samples): + self.samples = samples + self.connect_calls = [] + self.prefetch_indices = [] + self.closed = False + + def connect(self, controller, dataset_id, **kwargs): + self.connect_calls.append((controller, dataset_id, kwargs)) + + def __len__(self): + return len(self.samples) + + def __getitem__(self, index): + return self.samples[index] + + def _start_prefetch(self, indices): + self.prefetch_indices.append(indices) + + def close(self): + self.closed = True + + +def test_remote_mopd_dataset_connects_and_prefetches_each_source(): + """Global mixture indices are translated to each remote source correctly.""" + first = _RemoteSource([{"id": 0}, {"id": 1}]) + second = _RemoteSource([{"id": 2}, {"id": 3}, {"id": 4}]) + dataset = MOPDDataset([(first, "r0"), (second, "r1")]) + + assert is_remote_dataset(dataset) + dataset.connect( + "controller", + dataset_id="mixture", + tokenizer_or_processor_path="tokenizer", + shuffle=True, + drop_last=True, + ) + dataset._start_prefetch([4, 0, 2, 1, 3]) + + assert first.connect_calls[0][1] == "mixture_source_0" + assert second.connect_calls[0][1] == "mixture_source_1" + assert first.prefetch_indices == [[0, 1]] + assert second.prefetch_indices == [[2, 0, 1]] + assert dataset[2][MOPD_ROUTE_METADATA_KEY].route == "r1" + + dataset.close() + assert first.closed and second.closed diff --git a/tests/test_mopd_example.py b/tests/test_mopd_example.py new file mode 100644 index 0000000000..7ac4c386fa --- /dev/null +++ b/tests/test_mopd_example.py @@ -0,0 +1,184 @@ +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import torch +from omegaconf import OmegaConf + +from examples.mopd.gsm8k_qwen3_14b_to_0_6b import ( + MOPD_TEACHER_GROUP, + GSM8KRewardDistillationAgent, + add_no_think_suffix, + dynamic_filter, + load_routed_gsm8k_dataset, +) + +from areal.api.cli_args import ( + DatasetSourceConfig, + GRPOConfig, + TrainDatasetConfig, + to_structured_cfg, +) +from areal.dataset.mopd import MOPD_ROUTE_METADATA_KEY +from areal.reward import gsm8k_reward_fn + + +class _Tokenizer: + def encode(self, text): + return text.split() + + +def test_qwen3_heterogeneous_local_example_has_expected_topology(monkeypatch): + """The checked-in example resolves to one local node and eight GPUs.""" + monkeypatch.setenv("MOPD_STUDENT_MODEL_PATH", "/models/Qwen3-0.6B") + monkeypatch.setenv("MOPD_TEACHER_MODEL_PATH", "/models/Qwen3-14B") + monkeypatch.setenv("MOPD_GSM8K_PATH", "/datasets/gsm8k") + monkeypatch.setenv("AREAL_ADMIN_API_KEY", "test-only-non-default-key") + monkeypatch.setenv("AREAL_IMAGE", "areal:test") + raw = OmegaConf.load("examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml") + + config = OmegaConf.to_object(to_structured_cfg(raw, GRPOConfig)) + + assert isinstance(config, GRPOConfig) and config.mopd is not None + assert config.enable_offload is False + assert config.scheduler.type == "local" + assert (config.cluster.n_nodes, config.cluster.n_gpus_per_node) == (1, 8) + assert config.actor.backend == "megatron:d1p1t8" + assert config.mopd.teacher_engine.backend == config.actor.backend + assert config.rollout.backend == "sglang:d8t1p1" + assert config.mopd.teacher_groups == {MOPD_TEACHER_GROUP: {"qwen3_14b": 1.0}} + assert config.mopd.loss.rl_coefficient == 0.0 + assert config.mopd.loss.distillation_coefficient == 1.0 + assert config.total_train_epochs == 10 + assert config.total_train_steps is None + assert config.train_dataset.batch_size == 256 + assert config.train_dataset.path is None + assert config.train_dataset.sources[0].teacher_group == MOPD_TEACHER_GROUP + assert config.train_dataset.max_length is None + assert config.valid_dataset is not None + assert config.valid_dataset.batch_size == 256 + assert config.valid_dataset.sources[0].teacher_group == MOPD_TEACHER_GROUP + assert config.valid_dataset.max_length is None + assert config.rollout.max_concurrent_rollouts == 256 + assert config.sglang.max_running_requests is None + assert config.actor.mb_spec.max_tokens_per_mb == 10240 + assert config.mopd.teacher_engine.mb_spec.max_tokens_per_mb == 10240 + assert config.actor.recompute_logprob is False + assert config.actor.use_decoupled_loss is True + assert config.actor.prox_logp_method == "reuse_train_logp" + assert config.actor.should_compute_prox_logp() is False + assert config.actor.reward_norm is None + assert config.actor.adv_norm is None + assert config.actor.rejection_sampling is None + assert add_no_think_suffix({"question": "1+1?"}) == {} + + +def test_add_no_think_suffix_appends_once(): + """Both dataset loaders use the same reference no-think prompt format.""" + sample = {"messages": [{"role": "user", "content": "What is 1 + 1?"}]} + + routed = add_no_think_suffix(sample) + routed_twice = add_no_think_suffix(routed) + + assert sample["messages"][0]["content"] == "What is 1 + 1?" + assert routed["messages"][0]["content"] == "What is 1 + 1? /no_think" + assert routed_twice["messages"] == routed["messages"] + + +def test_dynamic_filter_matches_reference_all_correct_threshold(): + """Keep mixed groups and reject groups whose mean reward exceeds 0.95.""" + assert dynamic_filter({"rewards": torch.tensor([1.0, 1.0, 1.0, 0.0])}) + assert not dynamic_filter({"rewards": torch.ones(4)}) + + +def test_load_routed_gsm8k_dataset_accepts_local_parquet_mirror(tmp_path): + """The example directly consumes the parquet layout from the reference run.""" + from datasets import Dataset + + main_path = tmp_path / "main" + main_path.mkdir() + Dataset.from_dict( + {"question": ["What is 1 + 1?"], "answer": ["#### 2"]} + ).to_parquet(main_path / "train-00000-of-00001.parquet") + config = TrainDatasetConfig( + sources=[ + DatasetSourceConfig( + path=str(tmp_path), + type="rl", + teacher_group=MOPD_TEACHER_GROUP, + max_length=32, + ) + ], + scheduling_spec=None, + ) + + dataset = load_routed_gsm8k_dataset( + config, + tokenizer=_Tokenizer(), + ) + + assert len(dataset) == 1 + assert dataset[0]["answer"] == "#### 2" + assert dataset[0][MOPD_ROUTE_METADATA_KEY].route == MOPD_TEACHER_GROUP + assert "task_type" not in dataset[0] + assert dataset[0]["messages"][0]["content"].startswith("What is 1 + 1?") + assert dataset[0]["messages"][0]["content"].endswith(" /no_think") + + +@pytest.mark.asyncio +async def test_gsm8k_distillation_agent_returns_verifier_reward(monkeypatch): + """Pure distillation reports task quality without using it in the loss.""" + calls = [] + reward_calls = [] + + class _Completions: + async def create(self, **kwargs): + calls.append(kwargs) + message = type("Message", (), {"content": "The answer is \\boxed{2}."})() + choice = type("Choice", (), {"message": message})() + return type( + "Response", + (), + {"id": "completion-1", "choices": [choice]}, + )() + + class _Client: + def __init__(self, **kwargs): + del kwargs + self.chat = type("Chat", (), {"completions": _Completions()})() + + monkeypatch.setattr("openai.AsyncOpenAI", _Client) + agent = GSM8KRewardDistillationAgent(temperature=1.0) + assert agent._reward.reward_fn is gsm8k_reward_fn + + async def _reward(**kwargs): + reward_calls.append(kwargs) + return 1.0 + + agent._reward = _reward + + reward = await agent.run( + { + "messages": [{"role": "user", "content": "What is 1 + 1?"}], + "answer": "#### 2", + }, + base_url="http://localhost:30000/v1", + api_key="test-key", + ) + + assert reward == {"completion-1": 1.0} + assert reward_calls == [ + { + "prompt": "[{'role': 'user', 'content': 'What is 1 + 1?'}]", + "completions": "The answer is \\boxed{2}.", + "prompt_ids": [], + "completion_ids": [], + "answer": "#### 2", + } + ] + assert calls == [ + { + "messages": [{"role": "user", "content": "What is 1 + 1?"}], + "model": "default", + "temperature": 1.0, + } + ] diff --git a/tests/test_mopd_gpu.py b/tests/test_mopd_gpu.py new file mode 100644 index 0000000000..716ce0d887 --- /dev/null +++ b/tests/test_mopd_gpu.py @@ -0,0 +1,282 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Opt-in hardware regression tests for persistent MOPD teacher residency.""" + +from __future__ import annotations + +import math +import os +import re +import secrets +import signal +import socket +import subprocess +import sys +import time +import uuid +from collections import Counter +from pathlib import Path + +import pytest +import torch + +from areal.infra.utils.proc import kill_process_tree + +_RUN_8GPU_SMOKE = os.environ.get("AREAL_RUN_MOPD_8GPU_TEST", "").strip() == "1" +_MODEL_PATH_VARS = ( + "MOPD_STUDENT_MODEL_PATH", + "MOPD_TEACHER_MODEL_PATH", + "MOPD_GSM8K_PATH", +) + + +def _pid_exists(pid: int) -> bool: + try: + os.kill(pid, 0) + except ProcessLookupError: + return False + except PermissionError: + return True + return True + + +def _assigned_cuda_device_count() -> int: + if os.environ.get("AREAL_CONTROLLER_HIDDEN_DEVICE_ENV") == "CUDA_VISIBLE_DEVICES": + original = os.environ.get("AREAL_CONTROLLER_ORIG_DEVICES", "") + return len([device for device in original.split(",") if device]) + return torch.cuda.device_count() + + +def _restore_controller_hidden_devices(env: dict[str, str]) -> None: + """Let a fresh test controller see, then independently hide, its GPUs.""" + hidden_env = env.pop("AREAL_CONTROLLER_HIDDEN_DEVICE_ENV", None) + if hidden_env: + if env.pop("AREAL_CONTROLLER_ORIG_DEVICES_SET", "0") == "1": + env[hidden_env] = env.pop("AREAL_CONTROLLER_ORIG_DEVICES", "") + else: + env.pop(hidden_env, None) + env.pop("AREAL_CONTROLLER_ORIG_DEVICES", None) + + +def _session_members(session_id: int) -> list[int]: + members = [] + for entry in Path("/proc").iterdir(): + if not entry.name.isdigit(): + continue + pid = int(entry.name) + try: + if os.getsid(pid) == session_id: + members.append(pid) + except (PermissionError, ProcessLookupError): + continue + return members + + +def _cleanup_process_session(process: subprocess.Popen) -> None: + if process.poll() is None: + kill_process_tree(process.pid, timeout=10, graceful=False) + for pid in _session_members(process.pid): + kill_process_tree(pid, timeout=5, graceful=False) + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + + +def _tail(path: Path, line_count: int = 200) -> str: + return "\n".join( + path.read_text(encoding="utf-8", errors="replace").splitlines()[-line_count:] + ) + + +def _assert_ports_released(ports: set[int]) -> None: + for port in ports: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + sock.bind(("", port)) + + +@pytest.mark.integration +@pytest.mark.multi_gpu +@pytest.mark.slow +@pytest.mark.skipif( + not _RUN_8GPU_SMOKE, + reason="set AREAL_RUN_MOPD_8GPU_TEST=1 to run the persistent MOPD smoke", +) +def test_persistent_teacher_8gpu_three_step_smoke_releases_memory(tmp_path, request): + """Real Qwen3 phases reuse eight teachers and release their CUDA weights.""" + assert _assigned_cuda_device_count() >= 8, ( + "AREAL_RUN_MOPD_8GPU_TEST=1 requires at least eight visible CUDA GPUs" + ) + paths = {name: Path(os.environ.get(name, "")) for name in _MODEL_PATH_VARS} + missing = [name for name, path in paths.items() if not path.is_dir()] + assert not missing, f"Missing required MOPD directories: {', '.join(missing)}" + + repo_root = Path(__file__).resolve().parents[1] + run_root = tmp_path / "run" + trial_name = f"mopd-gpu-{uuid.uuid4().hex[:10]}" + output_path = tmp_path / "train.log" + env = os.environ.copy() + _restore_controller_hidden_devices(env) + env.setdefault("AREAL_ADMIN_API_KEY", secrets.token_hex(32)) + command = [ + sys.executable, + "-m", + "examples.mopd.gsm8k_qwen3_14b_to_0_6b", + "--config", + "examples/mopd/gsm8k_qwen3_14b_to_0_6b_local.yaml", + "total_train_steps=3", + "total_train_epochs=1", + "gconfig.n_samples=1", + "gconfig.max_new_tokens=64", + "train_dataset.batch_size=1", + "valid_dataset.batch_size=1", + "rollout.max_concurrent_rollouts=2", + "sglang.max_running_requests=2", + f"trial_name={trial_name}", + f"cluster.fileroot={run_root}", + f"cluster.name_resolve.nfs_record_root={run_root / 'name-resolve'}", + ] + + with output_path.open("w", encoding="utf-8") as output: + process = subprocess.Popen( + command, + cwd=repo_root, + env=env, + stdout=output, + stderr=subprocess.STDOUT, + text=True, + start_new_session=True, + ) + request.addfinalizer(lambda: _cleanup_process_session(process)) + try: + return_code = process.wait(timeout=3600) + except subprocess.TimeoutExpired: + _cleanup_process_session(process) + pytest.fail(f"MOPD GPU smoke timed out\n{_tail(output_path)}") + + assert return_code == 0, ( + f"MOPD GPU smoke exited with code {return_code}\n{_tail(output_path)}" + ) + console = output_path.read_text(encoding="utf-8", errors="replace") + ownership_events = re.findall( + r"\[MOPD\] teacher (onload|offload) complete|offload done, (onloading actor)", + console, + ) + assert ownership_events == [ + ("offload", ""), + ("", "onloading actor"), + ("onload", ""), + ("offload", ""), + ("", "onloading actor"), + ("onload", ""), + ("offload", ""), + ("", "onloading actor"), + ] + + teacher_logs = list(run_root.rglob("mopd-teacher.log")) + actor_logs = list(run_root.rglob("actor.log")) + assert len(teacher_logs) == 1, f"Expected one teacher log, found {teacher_logs}" + assert len(actor_logs) == 1, f"Expected one actor log, found {actor_logs}" + teacher_log = teacher_logs[0].read_text(encoding="utf-8", errors="replace") + actor_log = actor_logs[0].read_text(encoding="utf-8", errors="replace") + + spawn_events = re.findall( + r"Forked worker for role 'mopd-teacher' index (\d+) spawned \(pid=(\d+)\)", + actor_log, + ) + assert len(spawn_events) == 8, spawn_events + spawned = {int(rank): int(pid) for rank, pid in spawn_events} + assert set(spawned) == set(range(8)), f"Unexpected teacher workers: {spawned}" + assert teacher_log.count("Created Megatron weight residency adapter") == 8 + + residency_events = re.findall( + r"\[Megatron residency\] rank=(\d+) " + r"phase=(before_offload|after_offload|after_onload) " + r"allocated_gb=([0-9.]+) reserved_gb=([0-9.]+)", + teacher_log, + ) + stats: dict[int, dict[str, list[tuple[float, float]]]] = { + rank: {"before_offload": [], "after_offload": [], "after_onload": []} + for rank in range(8) + } + for rank_text, phase, allocated, reserved in residency_events: + stats[int(rank_text)][phase].append((float(allocated), float(reserved))) + + for rank, rank_stats in stats.items(): + before_offload = rank_stats["before_offload"] + after_offload = rank_stats["after_offload"] + after_onload = rank_stats["after_onload"] + assert len(before_offload) == len(after_offload) == 3, (rank, rank_stats) + assert len(after_onload) == 2, (rank, rank_stats) + assert all( + resident_allocated - offloaded_allocated >= 2.0 + and resident_reserved - offloaded_reserved >= 2.0 + for (resident_allocated, resident_reserved), ( + offloaded_allocated, + offloaded_reserved, + ) in zip(before_offload, after_offload, strict=True) + ), (rank, rank_stats) + assert max(allocated for allocated, _ in after_offload) < 0.5, ( + rank, + rank_stats, + ) + assert max(reserved for _, reserved in after_offload) < 0.5, ( + rank, + rank_stats, + ) + + published_versions = [ + int(version) + for version in re.findall(r"Put writer version .* version=(\d+)", actor_log) + ] + assert Counter(published_versions) == Counter({1: 8, 2: 8, 3: 8}) + + merged_logs = list(run_root.rglob("merged.log")) + assert len(merged_logs) == 1, merged_logs + merged_log = merged_logs[0].read_text(encoding="utf-8", errors="replace") + finite_metrics = { + "mopd_loss": re.findall( + r"ppo_actor/update/mopd_loss/(?:avg|max|min)\s+│\s+(\S+)", + merged_log, + ), + "new_logp": re.findall( + r"ppo_actor/update/new_logp/(?:avg|max|min)\s+│\s+(\S+)", + merged_log, + ), + "grad_norm": re.findall(r"ppo_actor/update/grad_norm\s+│\s+(\S+)", merged_log), + } + assert {name: len(values) for name, values in finite_metrics.items()} == { + "mopd_loss": 9, + "new_logp": 9, + "grad_norm": 3, + } + assert all( + math.isfinite(float(value)) + for values in finite_metrics.values() + for value in values + ), finite_metrics + + deadline = time.monotonic() + 30 + while ( + any(_pid_exists(pid) for pid in spawned.values()) + and time.monotonic() < deadline + ): + time.sleep(0.1) + live_pids = [pid for pid in spawned.values() if _pid_exists(pid)] + assert not live_pids, f"Teacher worker PIDs survived shutdown: {live_pids}" + + deadline = time.monotonic() + 30 + while _session_members(process.pid) and time.monotonic() < deadline: + time.sleep(0.1) + session_pids = _session_members(process.pid) + assert not session_pids, f"MOPD subprocess descendants survived: {session_pids}" + + guard_events = re.findall( + r"Starting Guard on [^: ]+:(\d+) for worker mopd-teacher/(\d+)", + teacher_log, + ) + assert len(guard_events) == 8, guard_events + ports_by_rank = {int(rank): int(port) for port, rank in guard_events} + assert set(ports_by_rank) == set(range(8)), ports_by_rank + _assert_ports_released(set(ports_by_rank.values())) diff --git a/tests/test_mopd_loss.py b/tests/test_mopd_loss.py new file mode 100644 index 0000000000..e69774e1af --- /dev/null +++ b/tests/test_mopd_loss.py @@ -0,0 +1,487 @@ +# SPDX-License-Identifier: Apache-2.0 + +from unittest.mock import MagicMock, patch + +import pytest +import torch + +from areal.api.cli_args import MOPDLossConfig, RejectionSamplingConfig +from areal.trainer.mopd.loss import compose_mopd_loss, mopd_loss_fn +from areal.trainer.ppo.actor import PPOActor, grpo_loss_fn + + +def test_actor_binds_one_mopd_loss_config(): + actor = object.__new__(PPOActor) + actor._mopd_loss_config = None + config = MOPDLossConfig(importance_ratio_cap=1.5) + + actor.configure_mopd_loss(config) + actor.configure_mopd_loss(config) + + with pytest.raises(RuntimeError, match="already bound"): + actor.configure_mopd_loss(MOPDLossConfig(importance_ratio_cap=2.0)) + + +def _loss_inputs(): + logprobs = torch.tensor( + [[-0.7, -1.1, -0.4], [-0.2, -0.9, -1.3]], + dtype=torch.float64, + requires_grad=True, + ) + old_logprobs = torch.tensor( + [[-0.8, -1.0, -0.5], [-0.3, -0.8, -1.4]], + dtype=torch.float64, + ) + teacher_logp_sum = torch.tensor( + [[-0.6, -0.9, -0.8], [-0.4, -1.0, -1.1]], + dtype=torch.float64, + ) + teacher_weight_sum = torch.tensor( + [[1.0, 1.0, 2.0], [1.0, 0.5, 1.5]], + dtype=torch.float64, + ) + loss_mask = torch.tensor( + [[True, True, False], [True, False, True]], + ) + return ( + logprobs, + old_logprobs, + teacher_logp_sum, + teacher_weight_sum, + loss_mask, + ) + + +def test_mopd_loss_matches_exact_weighted_reverse_kl_oracle(): + """The score-function surrogate equals an enumerated categorical RKL.""" + student_logits = torch.tensor( + [0.3, -0.7, 1.1], dtype=torch.float64, requires_grad=True + ) + teacher_logits = torch.tensor( + [[-0.2, 0.6, 0.1], [0.9, -0.4, 0.2]], dtype=torch.float64 + ) + teacher_weights = torch.tensor([0.25, 1.75], dtype=torch.float64) + student_logp = student_logits.log_softmax(dim=0) + teacher_logp = teacher_logits.log_softmax(dim=-1) + old_logp = torch.full_like( + student_logp, -torch.log(torch.tensor(3.0, dtype=torch.float64)) + ) + teacher_logp_sum = (teacher_weights[:, None] * teacher_logp).sum(dim=0) + teacher_weight_sum = torch.ones_like(student_logp) * teacher_weights.sum() + loss_mask = torch.ones_like(student_logp, dtype=torch.bool) + + surrogate, _ = mopd_loss_fn( + student_logp, + old_logp, + teacher_logp_sum, + teacher_weight_sum, + loss_mask, + ) + surrogate_grad = torch.autograd.grad(surrogate, student_logits, retain_graph=True)[ + 0 + ] + exact_reverse_kl = ( + teacher_weights[:, None] + * student_logp.exp()[None, :] + * (student_logp[None, :] - teacher_logp) + ).sum() + exact_grad = torch.autograd.grad(exact_reverse_kl, student_logits)[0] + + torch.testing.assert_close( + surrogate.detach(), exact_reverse_kl.detach(), rtol=1e-12, atol=1e-12 + ) + torch.testing.assert_close(surrogate_grad, exact_grad, rtol=1e-12, atol=1e-12) + + +def test_mopd_loss_masks_extreme_log_ratio_before_exponential(): + """A masked overflowing ratio cannot poison loss, stats, or gradients.""" + logprobs = torch.tensor([-1.0, 1000.0], dtype=torch.float64, requires_grad=True) + old_logprobs = torch.tensor([-1.0, -1000.0], dtype=torch.float64) + teacher_logp_sum = torch.tensor([-2.0, -2.0], dtype=torch.float64) + teacher_weight_sum = torch.ones(2, dtype=torch.float64) + loss_mask = torch.tensor([True, False]) + + loss, stats = mopd_loss_fn( + logprobs, + old_logprobs, + teacher_logp_sum, + teacher_weight_sum, + loss_mask, + ) + loss.backward() + + torch.testing.assert_close( + loss.detach(), torch.tensor(1.0, dtype=torch.float64), rtol=0.0, atol=0.0 + ) + torch.testing.assert_close( + logprobs.grad, + torch.tensor([1.0, 0.0], dtype=torch.float64), + rtol=0.0, + atol=0.0, + ) + assert all(torch.isfinite(value).all() for value in stats.values()) + for value in stats.values(): + if value.shape == logprobs.shape: + torch.testing.assert_close( + value[~loss_mask], + torch.zeros_like(value[~loss_mask]), + rtol=0.0, + atol=0.0, + ) + + +def test_mopd_loss_caps_active_extreme_ratio_with_finite_score_gradient(): + """Truncated IS bounds stale tokens without zeroing their policy gradient.""" + logprobs = torch.tensor([1000.0], dtype=torch.float64, requires_grad=True) + old_logprobs = torch.tensor([-1000.0], dtype=torch.float64) + teacher_logp_sum = torch.tensor([-2.0], dtype=torch.float64) + teacher_weight_sum = torch.ones(1, dtype=torch.float64) + loss_mask = torch.tensor([True]) + + loss, stats = mopd_loss_fn( + logprobs, + old_logprobs, + teacher_logp_sum, + teacher_weight_sum, + loss_mask, + importance_ratio_cap=5.0, + ) + loss.backward() + + assert torch.isfinite(loss) + assert torch.isfinite(logprobs.grad).all() + torch.testing.assert_close( + stats["importance_weight"], + torch.tensor([5.0], dtype=torch.float64), + rtol=1e-12, + atol=1e-12, + ) + torch.testing.assert_close( + logprobs.grad, + torch.tensor([5010.0], dtype=torch.float64), + rtol=1e-12, + atol=1e-12, + ) + + +@pytest.mark.parametrize("invalid", [float("nan"), float("inf"), float("-inf")]) +def test_mopd_loss_rejects_nonfinite_active_logprobs(invalid): + """Invalid active policy values fail explicitly instead of poisoning updates.""" + with pytest.raises(RuntimeError, match="must be finite"): + mopd_loss_fn( + torch.tensor([invalid], requires_grad=True), + torch.tensor([-1.0]), + torch.tensor([-2.0]), + torch.tensor([1.0]), + torch.tensor([True]), + ) + + +def test_mopd_loss_masks_nonfinite_inactive_inputs(): + """Non-finite padding remains harmless after active-token finite checks.""" + logprobs = torch.tensor([-1.0, float("nan")], requires_grad=True) + loss, stats = mopd_loss_fn( + logprobs, + torch.tensor([-1.0, float("inf")]), + torch.tensor([-2.0, float("nan")]), + torch.tensor([1.0, float("inf")]), + torch.tensor([True, False]), + ) + loss.backward() + + assert torch.isfinite(loss) + assert torch.isfinite(logprobs.grad).all() + assert all(torch.isfinite(value).all() for value in stats.values()) + + +def test_mopd_loss_detaches_old_policy_and_teacher_targets(): + """Gradients only flow through current-policy importance weights.""" + inputs = list(_loss_inputs()) + for index in (1, 2, 3): + inputs[index] = inputs[index].requires_grad_() + + loss, _ = mopd_loss_fn(*inputs) + loss.backward() + + assert inputs[0].grad is not None + assert inputs[1].grad is None + assert inputs[2].grad is None + assert inputs[3].grad is None + + +def test_mopd_loss_empty_mask_returns_differentiable_zero(): + """An empty response mask is finite and keeps a zero current-policy graph.""" + inputs = list(_loss_inputs()) + inputs[4] = torch.zeros_like(inputs[4]) + + loss, stats = mopd_loss_fn(*inputs) + loss.backward() + + torch.testing.assert_close( + loss.detach(), torch.zeros_like(loss), rtol=0.0, atol=0.0 + ) + torch.testing.assert_close( + inputs[0].grad, torch.zeros_like(inputs[0]), rtol=0.0, atol=0.0 + ) + torch.testing.assert_close( + stats["reverse_kl"], + torch.zeros_like(inputs[0]), + rtol=0.0, + atol=0.0, + ) + + +def test_compose_mopd_loss_disabled_returns_original_rl_loss(): + """Disabled MOPD returns the identical RL scalar and no MOPD statistics.""" + rl_loss = torch.tensor(1.25, dtype=torch.float64, requires_grad=True) + + total_loss, stats = compose_mopd_loss(rl_loss, config=None) + + assert total_loss is rl_loss + assert stats == {} + + +def test_compose_mopd_loss_joint_combines_objectives_and_gradients(): + """Joint mode applies independent RL and MOPD coefficients.""" + inputs = _loss_inputs() + logprobs = inputs[0] + rl_loss = logprobs.square().mean() + config = MOPDLossConfig( + rl_coefficient=0.4, + distillation_coefficient=0.6, + ) + expected_mopd_loss, _ = mopd_loss_fn(*inputs) + expected_total = 0.4 * rl_loss + 0.6 * expected_mopd_loss + expected_grad = torch.autograd.grad(expected_total, logprobs, retain_graph=True)[0] + + actual_total, _ = compose_mopd_loss( + rl_loss, + config=config, + logprobs=logprobs, + old_logprobs=inputs[1], + teacher_logp_sum=inputs[2], + teacher_weight_sum=inputs[3], + loss_mask=inputs[4], + ) + actual_grad = torch.autograd.grad(actual_total, logprobs)[0] + + torch.testing.assert_close( + actual_total.detach(), expected_total.detach(), rtol=1e-12, atol=1e-12 + ) + torch.testing.assert_close(actual_grad, expected_grad, rtol=1e-12, atol=1e-12) + + +def test_compose_mopd_loss_pure_mode_ignores_nan_rl_loss(): + """Pure MOPD does not multiply a potentially invalid RL loss by zero.""" + inputs = _loss_inputs() + rl_loss = torch.tensor(float("nan"), dtype=torch.float64) + + total_loss, _ = compose_mopd_loss( + rl_loss, + config=MOPDLossConfig(), + logprobs=inputs[0], + old_logprobs=inputs[1], + teacher_logp_sum=inputs[2], + teacher_weight_sum=inputs[3], + loss_mask=inputs[4], + ) + + assert torch.isfinite(total_loss) + + +def test_mopd_loss_rejects_broadcastable_non_token_shape(): + """Teacher targets must be materialized token tensors, not broadcast views.""" + inputs = list(_loss_inputs()) + inputs[3] = inputs[3][:, :1] + + with pytest.raises(ValueError, match="token shape"): + mopd_loss_fn(*inputs) + + +@pytest.mark.parametrize("cap", [0.0, -1.0, float("nan"), float("inf"), True]) +def test_mopd_loss_rejects_invalid_importance_ratio_cap(cap): + """The truncation bound must be finite, positive, and numeric.""" + with pytest.raises(ValueError, match="finite positive"): + mopd_loss_fn(*_loss_inputs(), importance_ratio_cap=cap) + + +def test_grpo_loss_fn_composes_materialized_mopd_targets(): + """Pure distillation consumes MOPD targets without an RL advantage tensor.""" + inputs = _loss_inputs() + logprobs, old_logprobs, teacher_sum, weight_sum, loss_mask = inputs + proximal_logprobs = torch.zeros_like(old_logprobs) + input_data = { + "logprobs": proximal_logprobs, + "prox_logp": proximal_logprobs, + "loss_mask": loss_mask, + "mopd_teacher_logp_sum": teacher_sum, + "mopd_teacher_weight_sum": weight_sum, + "mopd_behavior_logprobs": old_logprobs, + } + expected_loss, _ = mopd_loss_fn(*inputs, importance_ratio_cap=1.05) + + with patch("areal.trainer.ppo.actor.stats_tracker", MagicMock()) as tracker: + actual_loss = grpo_loss_fn( + logprobs=logprobs, + entropy=torch.zeros_like(logprobs), + input_data=input_data, + eps_clip=0.2, + eps_clip_higher=None, + c_clip=None, + mopd_loss_config=MOPDLossConfig(importance_ratio_cap=1.05), + ) + + torch.testing.assert_close( + actual_loss.detach(), expected_loss.detach(), rtol=1e-12, atol=1e-12 + ) + mopd_stat_call = next( + call + for call in tracker.stat.call_args_list + if "mopd_teacher_weight_sum" in call.kwargs + ) + assert "mopd_loss" in mopd_stat_call.kwargs + + +def test_grpo_loss_fn_scales_pure_rl_without_teacher_targets(): + """A disabled distillation objective neither requires targets nor drops RL scale.""" + logprobs = torch.tensor([[-0.3, -0.4]], dtype=torch.float64, requires_grad=True) + input_data = { + "logprobs": torch.tensor([[-0.5, -0.5]], dtype=torch.float64), + "prox_logp": torch.tensor([[-0.5, -0.5]], dtype=torch.float64), + "advantages": torch.ones_like(logprobs), + "loss_mask": torch.ones_like(logprobs, dtype=torch.bool), + } + kwargs = dict( + logprobs=logprobs, + entropy=torch.zeros_like(logprobs), + input_data=input_data, + eps_clip=0.2, + eps_clip_higher=None, + c_clip=None, + ) + + with patch("areal.trainer.ppo.actor.stats_tracker", MagicMock()): + base_loss = grpo_loss_fn(**kwargs) + scaled_loss = grpo_loss_fn( + **kwargs, + mopd_loss_config=MOPDLossConfig( + rl_coefficient=0.25, + distillation_coefficient=0.0, + ), + ) + + torch.testing.assert_close(scaled_loss, 0.25 * base_loss, rtol=1e-12, atol=1e-12) + + +def test_mopd_loss_respects_m2po_filtered_mask(): + """M2PO removes high-variance tokens from both RL and MOPD objectives.""" + logprobs = torch.tensor([[-0.2, -0.4]], dtype=torch.float64, requires_grad=True) + old_logprobs = torch.zeros_like(logprobs) + response_mask = torch.ones_like(logprobs, dtype=torch.bool) + input_data = { + "logprobs": old_logprobs, + "prox_logp": torch.tensor([[2.0, 0.0]], dtype=torch.float64), + "advantages": torch.zeros_like(logprobs), + "loss_mask": response_mask, + "mopd_teacher_logp_sum": torch.tensor([[-1.0, -1.0]], dtype=torch.float64), + "mopd_teacher_weight_sum": torch.ones_like(logprobs), + "mopd_behavior_logprobs": old_logprobs, + } + + with patch("areal.trainer.ppo.actor.stats_tracker", MagicMock()) as tracker: + loss = grpo_loss_fn( + logprobs=logprobs, + entropy=torch.zeros_like(logprobs), + input_data=input_data, + eps_clip=0.2, + eps_clip_higher=None, + c_clip=None, + m2_threshold=1.0, + mopd_loss_config=MOPDLossConfig(), + ) + loss.backward() + + assert logprobs.grad[0, 0] == 0 + assert logprobs.grad[0, 1] != 0 + denominator_call = next( + call + for call in tracker.denominator.call_args_list + if "n_mopd_tokens" in call.kwargs + ) + assert torch.equal( + denominator_call.kwargs["n_mopd_tokens"], + torch.tensor([[False, True]]), + ) + mopd_stat_call = next( + call + for call in tracker.stat.call_args_list + if "mopd_teacher_weight_sum" in call.kwargs + ) + assert mopd_stat_call.kwargs["denominator"] == "n_mopd_tokens" + + +@pytest.mark.parametrize( + ("level", "shape", "prox_logp", "expected_mask"), + [ + ( + "token", + (1, 2), + [[2.0, 0.0]], + [[False, True]], + ), + ( + "sequence", + (2, 2), + [[2.0, 2.0], [0.0, 0.0]], + [[False, False], [True, True]], + ), + ], +) +def test_mopd_loss_respects_behavioral_rejection_without_renormalizing( + level, shape, prox_logp, expected_mask +): + """Rejected stale tokens have zero KD gradient and do not amplify survivors.""" + logprobs = torch.full(shape, -0.2, dtype=torch.float64, requires_grad=True) + old_logprobs = torch.zeros_like(logprobs) + response_mask = torch.ones_like(logprobs, dtype=torch.bool) + expected_mask = torch.tensor(expected_mask, dtype=torch.bool) + input_data = { + "logprobs": old_logprobs, + "prox_logp": torch.tensor(prox_logp, dtype=torch.float64), + "advantages": torch.zeros_like(logprobs), + "loss_mask": response_mask, + "mopd_teacher_logp_sum": torch.full_like(logprobs, -1.0), + "mopd_teacher_weight_sum": torch.ones_like(logprobs), + "mopd_behavior_logprobs": old_logprobs, + } + + with patch("areal.trainer.ppo.actor.stats_tracker", MagicMock()) as tracker: + loss = grpo_loss_fn( + logprobs=logprobs, + entropy=torch.zeros_like(logprobs), + input_data=input_data, + eps_clip=0.2, + eps_clip_higher=None, + c_clip=None, + rejection_sampling=RejectionSamplingConfig( + level=level, + action="mask", + metric="ratio", + upper=5.0, + ), + mopd_loss_config=MOPDLossConfig(), + ) + loss.backward() + + assert torch.all(logprobs.grad[~expected_mask] == 0) + assert torch.all(logprobs.grad[expected_mask] != 0) + expected_loss = ( + torch.exp(logprobs.detach()) * (logprobs.detach() + 1.0) * expected_mask + ).sum() / response_mask.count_nonzero() + torch.testing.assert_close(loss.detach(), expected_loss, rtol=1e-12, atol=1e-12) + denominator_call = next( + call + for call in tracker.denominator.call_args_list + if "n_mopd_tokens" in call.kwargs + ) + assert torch.equal(denominator_call.kwargs["n_mopd_tokens"], response_mask) diff --git a/tests/test_mopd_routing.py b/tests/test_mopd_routing.py new file mode 100644 index 0000000000..dea3f78bcd --- /dev/null +++ b/tests/test_mopd_routing.py @@ -0,0 +1,135 @@ +# SPDX-License-Identifier: Apache-2.0 + +from unittest.mock import MagicMock + +import pytest + +from areal.api.cli_args import InferenceEngineConfig +from areal.dataset.mopd import MOPD_ROUTE_METADATA_KEY, DatasetRoute +from areal.infra.controller.rollout_controller import RolloutController + + +class _InferenceEngine: + pass + + +class _Scheduler: + pass + + +def _controller() -> RolloutController: + controller = RolloutController( + inf_engine=_InferenceEngine, + config=InferenceEngineConfig(backend="sglang:d1"), + scheduler=_Scheduler(), + ) + controller.enable_mopd_routing() + return controller + + +def test_source_route_metadata_is_removed_from_workflow_data(): + """The configured source route travels outside the workflow sample.""" + controller = _controller() + + prepared, route = controller._extract_mopd_route( + { + MOPD_ROUTE_METADATA_KEY: DatasetRoute(0, "gsm8k_single_a"), + "instance_id": "unique-sample", + }, + required=True, + ) + + assert route == "gsm8k_single_a" + assert MOPD_ROUTE_METADATA_KEY not in prepared + assert prepared["instance_id"] == "unique-sample" + + +def test_training_route_metadata_is_required(): + """No sample field substitutes for the configured source route.""" + controller = _controller() + + with pytest.raises(ValueError, match="route metadata is missing"): + controller._extract_mopd_route({"task_type": "gsm8k_single_a"}, required=True) + + +@pytest.mark.parametrize("route", [None, "", 7, 1.5, True, [], {}]) +def test_invalid_source_route_type_raises(route): + """Only non-empty configured route strings are accepted.""" + controller = _controller() + + with pytest.raises(ValueError, match="DatasetRoute provenance"): + controller._extract_mopd_route({MOPD_ROUTE_METADATA_KEY: route}, required=True) + + +def test_concat_trajectory_inherits_source_route(): + """A trajectory produced after OpenAI concat retains its source route.""" + controller = _controller() + concat_trajectory = {"input_ids": "concat-output", "attention_mask": "mask"} + + result = controller._propagate_mopd_route("gsm8k_ensemble", concat_trajectory) + + assert result["mopd_route"] == "gsm8k_ensemble" + + +def test_multiple_derived_trajectories_inherit_same_route(): + """Every trajectory generated from one source sample receives one route.""" + controller = _controller() + trajectories = [ + controller._propagate_mopd_route("gsm8k_single_b", {"trajectory_id": index}) + for index in range(3) + ] + + assert [trajectory["mopd_route"] for trajectory in trajectories] == [ + "gsm8k_single_b", + "gsm8k_single_b", + "gsm8k_single_b", + ] + + +def test_workflow_cannot_change_source_route(): + """A conflicting workflow route is rejected before teacher dispatch.""" + controller = _controller() + + with pytest.raises(ValueError, match="changed mopd_route"): + controller._propagate_mopd_route( + "gsm8k_single_a", {"mopd_route": "gsm8k_single_b"} + ) + + +def test_training_submit_keeps_route_outside_workflow_data(): + """Task metadata retains the source route while workflow data stays clean.""" + controller = _controller() + controller._resolve_workflow_str = MagicMock(return_value="workflow") + controller._resolve_should_accept_fn = MagicMock(return_value=None) + controller._dispatcher = MagicMock() + + controller.submit( + { + "messages": [{"role": "user", "content": "train me"}], + MOPD_ROUTE_METADATA_KEY: DatasetRoute(0, "gsm8k_single_a"), + }, + object(), + ) + + task_input = controller._dispatcher.submit_task_input.call_args.args[0] + assert task_input.mopd_route == "gsm8k_single_a" + assert MOPD_ROUTE_METADATA_KEY not in task_input.data + + +def test_eval_submit_does_not_require_training_route(): + """Validation samples remain usable without the training-only route field.""" + controller = _controller() + controller._resolve_workflow_str = MagicMock(return_value="workflow") + controller._resolve_should_accept_fn = MagicMock(return_value=None) + controller._dispatcher = MagicMock() + + controller.submit( + {"messages": [{"role": "user", "content": "evaluate me"}]}, + object(), + is_eval=True, + ) + + task_input = controller._dispatcher.submit_task_input.call_args.args[0] + assert task_input.is_eval is True + assert task_input.mopd_route is None + assert "mopd_route" not in task_input.data diff --git a/tests/test_mopd_rtensor.py b/tests/test_mopd_rtensor.py new file mode 100644 index 0000000000..113c160c31 --- /dev/null +++ b/tests/test_mopd_rtensor.py @@ -0,0 +1,211 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import FrozenInstanceError +from typing import Any + +import pytest +import torch + +from areal.infra.controller.train_controller import TrainController +from areal.infra.rpc.rtensor import RTensor, RTensorDrainReceipt, TensorShardInfo +from areal.trainer.mopd.targets import ( + MOPD_CONTRIBUTIONS_KEY, + aggregate_mopd_targets, +) + + +def _rtensor(shard_id: str, node_addr: str = "teacher:8000") -> RTensor: + return RTensor( + shard=TensorShardInfo(shard_id=shard_id, node_addr=node_addr), + data=torch.empty(3, device="meta"), + ) + + +def test_megatron_backend_exposes_mopd_rpc_methods(): + """The supported MOPD backend exposes the controller's RPC surface.""" + import importlib + + actor_cls = getattr( + importlib.import_module("areal.engine.megatron_engine"), "MegatronPPOActor" + ) + + assert callable(getattr(actor_cls, "aggregate_mopd_targets", None)) + + +def test_aggregate_mopd_targets_uses_raw_weights_and_removes_teacher_metadata(): + """Actor aggregation preserves raw scale and drops route/contribution keys.""" + batch = [ + { + "mopd_route": "ensemble", + MOPD_CONTRIBUTIONS_KEY: { + "teacher-a": {"logp": torch.tensor([1.0, 2.0]), "weight": 0.5}, + "teacher-b": {"logp": torch.tensor([3.0, 4.0]), "weight": 1.5}, + }, + }, + { + "mopd_route": "single", + MOPD_CONTRIBUTIONS_KEY: { + "teacher-b": {"logp": torch.tensor([5.0, 6.0]), "weight": 2.0} + }, + }, + ] + + result = aggregate_mopd_targets(batch) + + assert result is batch + torch.testing.assert_close( + batch[0]["mopd_teacher_logp_sum"], + torch.tensor([5.0, 7.0]), + rtol=0.0, + atol=0.0, + ) + torch.testing.assert_close( + batch[0]["mopd_teacher_weight_sum"], + torch.tensor([2.0, 2.0]), + rtol=0.0, + atol=0.0, + ) + torch.testing.assert_close( + batch[1]["mopd_teacher_weight_sum"], + torch.tensor([2.0, 2.0]), + rtol=0.0, + atol=0.0, + ) + for trajectory in batch: + assert "mopd_route" not in trajectory + assert MOPD_CONTRIBUTIONS_KEY not in trajectory + assert set(trajectory) >= { + "mopd_teacher_logp_sum", + "mopd_teacher_weight_sum", + } + assert not any(key.endswith("coefficient") for key in trajectory) + assert "mopd_importance_ratio_cap" not in trajectory + + +def test_aggregate_mopd_targets_rejects_mismatched_teacher_shapes(): + """All teachers contributing to one trajectory must use one token shape.""" + batch = [ + { + "mopd_route": "bad", + MOPD_CONTRIBUTIONS_KEY: { + "teacher-a": {"logp": torch.ones(2), "weight": 1.0}, + "teacher-b": {"logp": torch.ones(3), "weight": 1.0}, + }, + } + ] + + with pytest.raises(ValueError, match="shape mismatch"): + aggregate_mopd_targets(batch) + + +def _strict_controller( + monkeypatch: pytest.MonkeyPatch, + *, + stats: list[dict[str, int]], +) -> tuple[TrainController, list[tuple[str, tuple[Any, ...]]]]: + controller = object.__new__(TrainController) + controller.workers_is_dp_head = [True, False, True] + controller._worker_role = "actor" + calls: list[tuple[str, tuple[Any, ...]]] = [] + + def call_all(method: str, *args: Any, **_: Any) -> list[Any]: + calls.append((method, args)) + if method == "clear_batches": + return [1, 1] + if method == "fetch_buffer_stats": + return stats + raise AssertionError(method) + + monkeypatch.setattr(controller, "_custom_function_call_all_dp_heads", call_all) + return controller, calls + + +def test_strict_clear_batches_covers_sources_and_every_actor_dp_head(monkeypatch): + """A receipt is complete only after source and all actor heads are clean.""" + controller, calls = _strict_controller( + monkeypatch, + stats=[ + {"num_entries": 4, "matching_entries": 0}, + {"num_entries": 1, "matching_entries": 0}, + ], + ) + source_calls: list[tuple[str, list[str]]] = [] + + async def clear_node(node_addr: str, shard_ids: list[str]) -> int: + source_calls.append((node_addr, shard_ids)) + return len(shard_ids) + + monkeypatch.setattr(RTensor, "clear_node", clear_node) + targets = [ + {"logp": _rtensor("a")}, + {"logp": _rtensor("a")}, + {"logp": _rtensor("b", "teacher:8001")}, + ] + + receipt = controller.strict_clear_batches(targets) + + assert receipt == RTensorDrainReceipt( + consumer_role="actor", + shard_count=2, + source_node_count=2, + consumer_dp_head_count=2, + ) + assert source_calls == [ + ("teacher:8000", ["a"]), + ("teacher:8001", ["b"]), + ] + assert [method for method, _ in calls] == [ + "clear_batches", + "fetch_buffer_stats", + ] + assert calls[0][1] == (["a", "b"],) + + +def test_rtensor_drain_receipt_is_frozen_and_role_typed(): + receipt = RTensorDrainReceipt( + consumer_role="actor", + shard_count=1, + source_node_count=1, + consumer_dp_head_count=2, + ) + + with pytest.raises(FrozenInstanceError): + receipt.consumer_role = "teacher" + + +def test_strict_clear_batches_source_failure_prevents_receipt(monkeypatch): + """A failed teacher source DELETE is fatal and actor clearing does not hide it.""" + controller, calls = _strict_controller( + monkeypatch, stats=[{"matching_entries": 0}, {"matching_entries": 0}] + ) + + async def clear_node(_: str, __: list[str]) -> int: + raise RuntimeError("source delete failed") + + monkeypatch.setattr(RTensor, "clear_node", clear_node) + + with pytest.raises(RuntimeError, match="source delete failed"): + controller.strict_clear_batches({"logp": _rtensor("a")}) + + assert calls == [] + + +def test_strict_clear_batches_rejects_one_leaking_actor_head(monkeypatch): + """Checking head zero is insufficient when another actor DP head still leaks.""" + controller, _ = _strict_controller( + monkeypatch, + stats=[ + {"num_entries": 0, "matching_entries": 0}, + {"num_entries": 1, "matching_entries": 1}, + ], + ) + + async def clear_node(_: str, shard_ids: list[str]) -> int: + return len(shard_ids) + + monkeypatch.setattr(RTensor, "clear_node", clear_node) + + with pytest.raises(RuntimeError, match=r"DP heads \[1\]"): + controller.strict_clear_batches({"logp": _rtensor("a")}) diff --git a/tests/test_mopd_teacher_manager.py b/tests/test_mopd_teacher_manager.py new file mode 100644 index 0000000000..fe7f564aef --- /dev/null +++ b/tests/test_mopd_teacher_manager.py @@ -0,0 +1,546 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import json +from dataclasses import dataclass, field +from pathlib import Path + +import pytest +import torch + +from areal.api import SaveLoadMeta +from areal.api.cli_args import ( + MOPDConfig, + MOPDTeacherManagerConfig, + MOPDTeacherSpec, +) +from areal.infra.rpc.rtensor import RTensorDrainReceipt +from areal.trainer.mopd.targets import MOPD_CONTRIBUTIONS_KEY, aggregate_mopd_targets +from areal.trainer.mopd.teacher_manager import ( + DiskCheckpointProvider, + LocalMemoryCheckpointProvider, + PersistentTeacherManager, + TeacherManagerState, +) +from areal.trainer.mopd.teacher_phase import MOPDTeacherPhase + + +def _receipt(role: str = "actor") -> RTensorDrainReceipt: + return RTensorDrainReceipt( + consumer_role=role, + shard_count=1, + source_node_count=1, + consumer_dp_head_count=1, + ) + + +def _write_checkpoint(root: Path, teacher_id: str, payload: bytes) -> Path: + checkpoint = root / teacher_id + checkpoint.mkdir() + (checkpoint / "config.json").write_text( + json.dumps({"teacher": teacher_id}), encoding="utf-8" + ) + (checkpoint / "model.safetensors").write_bytes(payload) + return checkpoint + + +def _config( + checkpoints: dict[str, Path], + *, + manager_type: str = "disk", + staging_root: Path | None = None, +) -> MOPDConfig: + return MOPDConfig( + teachers={ + teacher_id: MOPDTeacherSpec(path=str(path)) + for teacher_id, path in checkpoints.items() + }, + teacher_groups={"group": {teacher_id: 1.0 for teacher_id in checkpoints}}, + manager=MOPDTeacherManagerConfig( + type=manager_type, + staging_root=str(staging_root or "/unused"), + ), + ) + + +@dataclass +class _PersistentController: + events: list[str] = field(default_factory=list) + fail_on: str | None = None + destroy_calls: int = 0 + destroy_failures: int = 0 + + def _event(self, name: str) -> None: + self.events.append(name) + if self.fail_on == name: + raise RuntimeError(f"{name} failed") + + def onload(self) -> None: + self._event("onload") + + def offload(self) -> None: + self._event("offload") + + def load(self, meta: SaveLoadMeta) -> None: + self._event(f"load:{Path(meta.path).name}") + assert meta.with_optim is False + + def destroy(self) -> None: + self.destroy_calls += 1 + self.events.append("destroy") + if self.destroy_failures: + self.destroy_failures -= 1 + raise RuntimeError("destroy failed") + + +def test_persistent_manager_reuses_controller_across_phases_and_checkpoints(tmp_path): + """Phase boundaries offload/onload one companion instead of respawning it.""" + t0 = _write_checkpoint(tmp_path, "t0", b"first") + t1 = _write_checkpoint(tmp_path, "t1", b"second") + controller = _PersistentController() + factory_paths: list[str] = [] + + def factory(path: str) -> _PersistentController: + factory_paths.append(path) + return controller + + manager = PersistentTeacherManager(_config({"t0": t0, "t1": t1}), factory) + + first = manager.load("t0") + manager.release(_receipt()) + second = manager.load("t0") + third = manager.load("t1") + + assert first is second is third is controller + assert factory_paths == [str(t0)] + assert controller.events == ["offload", "onload", "load:t1"] + assert controller.destroy_calls == 0 + assert manager.state is TeacherManagerState.RESIDENT + manager.close() + assert controller.destroy_calls == 1 + + +def test_persistent_manager_onloads_before_cross_phase_checkpoint_switch(tmp_path): + """An offloaded companion becomes resident before loading another teacher.""" + t0 = _write_checkpoint(tmp_path, "t0", b"first") + t1 = _write_checkpoint(tmp_path, "t1", b"second") + controller = _PersistentController() + manager = PersistentTeacherManager( + _config({"t0": t0, "t1": t1}), lambda _: controller + ) + manager.load("t0") + manager.release(_receipt()) + controller.events.clear() + + manager.load("t1") + + assert controller.events == ["onload", "load:t1"] + assert manager.state is TeacherManagerState.RESIDENT + manager.close() + + +def test_persistent_manager_repeated_release_does_not_offload_twice(tmp_path): + """An already-offloaded companion treats a duplicate complete receipt as a noop.""" + t0 = _write_checkpoint(tmp_path, "t0", b"first") + controller = _PersistentController() + manager = PersistentTeacherManager(_config({"t0": t0}), lambda _: controller) + manager.load("t0") + + manager.release(_receipt()) + manager.release(_receipt()) + + assert controller.events == ["offload"] + assert manager.state is TeacherManagerState.OFFLOADED + manager.close() + + +def test_persistent_manager_does_not_restage_unchanged_local_checkpoint(tmp_path): + """An offloaded unchanged teacher resumes without copying its snapshot again.""" + source_root = tmp_path / "source" + source_root.mkdir() + t0 = _write_checkpoint(source_root, "t0", b"first") + controller = _PersistentController() + manager = PersistentTeacherManager( + _config( + {"t0": t0}, + manager_type="local_memory", + staging_root=tmp_path / "staging", + ), + lambda _: controller, + ) + manager.pre_fetch("t0") + manager.load("t0") + manager.release(_receipt()) + + manager.pre_fetch("t0") + manager.load("t0") + + assert controller.events == ["offload", "onload"] + manager.close() + + +def test_persistent_manager_rejects_non_actor_receipt(tmp_path): + """A non-actor receipt leaves the resident teacher untouched.""" + t0 = _write_checkpoint(tmp_path, "t0", b"first") + controller = _PersistentController() + manager = PersistentTeacherManager(_config({"t0": t0}), lambda _: controller) + manager.load("t0") + + with pytest.raises(RuntimeError, match="without an actor RTensor drain receipt"): + manager.release(_receipt("teacher")) + + assert controller.events == [] + assert manager.state is TeacherManagerState.RESIDENT + manager.close() + + +@pytest.mark.parametrize( + ("failure", "prepare"), + [ + ("onload", "offload"), + ("load:t1", "resident"), + ("offload", "resident"), + ], +) +def test_persistent_manager_failure_destroys_companion(tmp_path, failure, prepare): + """Load, onload, and offload failures poison the whole group.""" + t0 = _write_checkpoint(tmp_path, "t0", b"first") + t1 = _write_checkpoint(tmp_path, "t1", b"second") + controller = _PersistentController() + manager = PersistentTeacherManager( + _config({"t0": t0, "t1": t1}), lambda _: controller + ) + manager.load("t0") + if prepare == "offload": + manager.release(_receipt()) + controller.events.clear() + controller.fail_on = failure + + with pytest.raises(RuntimeError, match=f"{failure} failed"): + if failure == "offload": + manager.release(_receipt()) + elif failure == "load:t1": + manager.load("t1") + else: + manager.load("t0") + + assert controller.destroy_calls == 1 + assert controller.events[-1] == "destroy" + assert manager.controller is None + assert manager.state is TeacherManagerState.BROKEN + manager.close() + assert controller.destroy_calls == 1 + + +@pytest.mark.parametrize("close_state", ["resident", "offloaded", "broken"]) +def test_persistent_manager_close_is_idempotent_in_every_live_state( + tmp_path, close_state +): + """Closing any persistent state tears down the companion at most once.""" + t0 = _write_checkpoint(tmp_path, "t0", b"first") + controller = _PersistentController() + manager = PersistentTeacherManager(_config({"t0": t0}), lambda _: controller) + manager.load("t0") + if close_state == "offloaded": + manager.release(_receipt()) + elif close_state == "broken": + controller.fail_on = "offload" + with pytest.raises(RuntimeError, match="offload failed"): + manager.release(_receipt()) + + manager.close() + manager.close() + + assert controller.destroy_calls == 1 + assert manager.state is TeacherManagerState.CLOSED + + +def test_disk_provider_requires_existing_local_snapshot(tmp_path): + """Disk mode rejects missing paths instead of attempting a network fetch.""" + provider = DiskCheckpointProvider(_config({"missing": tmp_path / "missing"})) + + with pytest.raises(FileNotFoundError, match="not a local directory"): + provider.resolve("missing") + + +def test_local_memory_provider_uses_atomic_single_ready_checkpoint(tmp_path): + """Staging publishes one ready snapshot and removes it after consumption.""" + source_root = tmp_path / "source" + source_root.mkdir() + t0 = _write_checkpoint(source_root, "t0", b"first") + t1 = _write_checkpoint(source_root, "t1", b"second") + staging_root = tmp_path / "staging" + provider = LocalMemoryCheckpointProvider( + _config( + {"t0": t0, "t1": t1}, + manager_type="local_memory", + staging_root=staging_root, + ) + ) + + provider.pre_fetch("t0") + ready = provider.resolve("t0") + + assert ready.name == "t0.ready" + assert (ready / "model.safetensors").read_bytes() == b"first" + assert not list(staging_root.rglob("*.tmp.*")) + with pytest.raises(RuntimeError, match="already holds ready checkpoint"): + provider.pre_fetch("t1") + + provider.consumed("t0") + assert not ready.exists() + provider.pre_fetch("t1") + second = provider.resolve("t1") + assert (second / "model.safetensors").read_bytes() == b"second" + provider.close() + provider.close() + assert not list(staging_root.iterdir()) + + +def test_local_memory_provider_rejects_insufficient_capacity(tmp_path, monkeypatch): + """Capacity is checked before checkpoint bytes enter the staging root.""" + source_root = tmp_path / "source" + source_root.mkdir() + t0 = _write_checkpoint(source_root, "t0", b"payload") + staging_root = tmp_path / "staging" + provider = LocalMemoryCheckpointProvider( + _config( + {"t0": t0}, + manager_type="local_memory", + staging_root=staging_root, + ) + ) + monkeypatch.setattr( + "areal.trainer.mopd.teacher_manager.shutil.disk_usage", + lambda _: type("Usage", (), {"free": 0})(), + ) + + with pytest.raises(OSError, match="Insufficient staging space"): + provider.pre_fetch("t0") + + provider.close() + assert not list(staging_root.iterdir()) + + +def test_local_memory_provider_sweeps_dead_run(tmp_path): + """Construction removes run directories whose owning process is gone.""" + source_root = tmp_path / "source" + source_root.mkdir() + t0 = _write_checkpoint(source_root, "t0", b"payload") + staging_root = tmp_path / "staging" + stale = staging_root / ".run-stale" + stale.mkdir(parents=True) + (stale / "owner.json").write_text('{"pid": 999999999}', encoding="utf-8") + (stale / "orphan.tmp.data").write_bytes(b"orphan") + + provider = LocalMemoryCheckpointProvider( + _config( + {"t0": t0}, + manager_type="local_memory", + staging_root=staging_root, + ) + ) + + assert not stale.exists() + provider.close() + + +class _PhaseController: + def __init__(self, teacher_id: str, events: list[str]): + self.teacher_id = teacher_id + self.events = events + + def compute_logp_padded(self, subset): + self.events.append(f"compute:{self.teacher_id}:{len(subset)}") + assert all(MOPD_CONTRIBUTIONS_KEY not in trajectory for trajectory in subset) + value = 1.0 if self.teacher_id == "t0" else 3.0 + real = [torch.full((2,), value) for _ in subset] + dummy = [torch.full((1,), -1.0)] if self.teacher_id == "t1" else [] + return real, dummy + + def assert_mopd_runtime_topology(self) -> None: + self.events.append(f"topology:{self.teacher_id}") + + def strict_clear_batches(self, *targets): + target_sizes = ",".join(str(len(target)) for target in targets) + self.events.append(f"clear:{self.teacher_id}:{target_sizes}") + return _receipt("mopd-teacher") + + +class _PhaseManager: + def __init__(self, events: list[str]): + self.events = events + self.closed = False + self.state = TeacherManagerState.RESIDENT + + def pre_fetch(self, teacher_id: str) -> None: + self.events.append(f"prefetch:{teacher_id}") + + def load(self, teacher_id: str) -> _PhaseController: + self.events.append(f"load:{teacher_id}") + return _PhaseController(teacher_id, self.events) + + def release(self, receipt: RTensorDrainReceipt) -> None: + assert receipt.consumer_role == "actor" + self.events.append("release") + + def close(self) -> None: + self.closed = True + self.events.append("close") + + +class _PhaseActor: + def __init__(self, events: list[str]): + self.events = events + + def aggregate_mopd_targets(self, batch): + self.events.append("aggregate") + return aggregate_mopd_targets(batch) + + def assert_mopd_runtime_topology(self) -> None: + self.events.append("topology:actor") + + def strict_clear_batches(self, *targets): + target_sizes = ",".join(str(len(target)) for target in targets) + self.events.append(f"clear:actor:{target_sizes}") + return _receipt("actor") + + +class _PhaseCritic: + def __init__(self, events: list[str]): + self.events = events + + def strict_clear_batches(self, *targets): + target_sizes = ",".join(str(len(target)) for target in targets) + self.events.append(f"clear:critic:{target_sizes}") + return _receipt("critic") + + +class _PhaseRef(_PhaseCritic): + def strict_clear_batches(self, *targets): + target_sizes = ",".join(str(len(target)) for target in targets) + self.events.append(f"clear:ref:{target_sizes}") + return _receipt("ref") + + +class _FailingPhaseActor(_PhaseActor): + def aggregate_mopd_targets(self, batch): + del batch + self.events.append("aggregate") + raise RuntimeError("aggregation failed") + + +class _FailingPhaseRef(_PhaseRef): + def strict_clear_batches(self, *targets): + target_sizes = ",".join(str(len(target)) for target in targets) + self.events.append(f"clear:ref:{target_sizes}") + raise RuntimeError("ref drain failed") + + +def test_trainer_mopd_phase_routes_reuses_drains_then_releases(): + """Teacher scoring isolates prior contributions before aggregation and release.""" + events: list[str] = [] + mopd = MOPDConfig( + teachers={ + "t0": MOPDTeacherSpec(path="/unused/t0"), + "t1": MOPDTeacherSpec(path="/unused/t1"), + }, + teacher_groups={"r0": {"t0": 2.0}, "r1": {"t0": 0.5, "t1": 1.5}}, + ) + phase = MOPDTeacherPhase( + config=mopd, + manager=_PhaseManager(events), + actor=_PhaseActor(events), + critic=_PhaseCritic(events), + ref=_PhaseRef(events), + ) + batch = [{"mopd_route": "r0"}, {"mopd_route": "r1"}] + + result = phase.materialize(batch) + + assert events == [ + "topology:actor", + "prefetch:t0", + "load:t0", + "topology:t0", + "prefetch:t1", + "compute:t0:2", + "load:t1", + "topology:t1", + "compute:t1:1", + "aggregate", + "clear:critic:2", + "clear:ref:2", + "clear:t0:2,4", + "clear:t1:2,4", + "clear:actor:2,4", + "release", + ] + torch.testing.assert_close( + result[0]["mopd_teacher_logp_sum"], + torch.full((2,), 2.0), + rtol=0.0, + atol=0.0, + ) + torch.testing.assert_close( + result[1]["mopd_teacher_logp_sum"], + torch.full((2,), 5.0), + rtol=0.0, + atol=0.0, + ) + assert all("mopd_route" not in trajectory for trajectory in result) + + +def test_trainer_mopd_phase_failure_drains_rollout_before_release(): + """Emergency cleanup drains original shards from teacher and actor workers.""" + events: list[str] = [] + mopd = MOPDConfig( + teachers={"t0": MOPDTeacherSpec(path="/unused/t0")}, + teacher_groups={"r0": {"t0": 1.0}}, + ) + phase = MOPDTeacherPhase( + config=mopd, + manager=_PhaseManager(events), + actor=_FailingPhaseActor(events), + critic=_PhaseCritic(events), + ref=_PhaseRef(events), + ) + + with pytest.raises(RuntimeError, match="aggregation failed"): + phase.materialize([{"mopd_route": "r0"}]) + + assert events == [ + "topology:actor", + "prefetch:t0", + "load:t0", + "topology:t0", + "compute:t0:1", + "aggregate", + "clear:critic:1", + "clear:ref:1", + "clear:t0:1,1", + "clear:actor:1,1", + "release", + ] + + +def test_mopd_phase_closes_teacher_when_any_consumer_drain_has_no_ack(): + """A missing consumer ACK forces teardown instead of teacher offload.""" + events: list[str] = [] + mopd = MOPDConfig( + teachers={"t0": MOPDTeacherSpec(path="/unused/t0")}, + teacher_groups={"r0": {"t0": 1.0}}, + ) + phase = MOPDTeacherPhase( + config=mopd, + manager=_PhaseManager(events), + actor=_PhaseActor(events), + ref=_FailingPhaseRef(events), + ) + + with pytest.raises(RuntimeError, match="ref drain failed"): + phase.materialize([{"mopd_route": "r0"}]) + + assert "close" in events + assert "release" not in events diff --git a/tests/test_ppo_gae.py b/tests/test_ppo_gae.py index b3523651b8..c060f27318 100644 --- a/tests/test_ppo_gae.py +++ b/tests/test_ppo_gae.py @@ -25,12 +25,16 @@ def _make_actor( gae_lambda: float | str = 1.0, gae_lambda_kwargs: dict | None = None, kl_ctl: float = 0.0, + recompute_logprob: bool = False, + use_decoupled_loss: bool = False, ) -> PPOActor: config = PPOActorConfig( gae_timestep_unit=gae_timestep_unit, gae_lambda=gae_lambda, gae_lambda_kwargs=gae_lambda_kwargs or {}, kl_ctl=kl_ctl, + recompute_logprob=recompute_logprob, + use_decoupled_loss=use_decoupled_loss, ) actor = PPOActor.__new__(PPOActor) actor.config = config @@ -53,6 +57,34 @@ def _make_actor( return actor +def test_mopd_preserves_rollout_behavior_logprobs_when_recomputing_proximal(): + actor = _make_actor(recompute_logprob=True, use_decoupled_loss=False) + rollout_logprobs = torch.tensor([[0.0, -1.0, -2.0, -3.0]]) + prox_logprobs = torch.tensor([[-0.2, -0.3, -0.4, -0.5]]) + batch = { + "input_ids": torch.zeros(1, 4, dtype=torch.long), + "loss_mask": torch.tensor([[0, 1, 1, 1]], dtype=torch.float32), + "logprobs": rollout_logprobs.clone(), + "prox_logp": prox_logprobs.clone(), + "attention_mask": torch.ones(1, 4, dtype=torch.bool), + "rewards": torch.tensor([1.0]), + "mopd_teacher_logp_sum": torch.zeros(1, 4), + } + + result = actor._compute_advantages(batch) + + expected_behavior = torch.tensor([[-1.0, -2.0, -3.0, 0.0]]) + torch.testing.assert_close( + result["mopd_behavior_logprobs"], expected_behavior, rtol=0.0, atol=0.0 + ) + torch.testing.assert_close( + result["logprobs"], + prox_logprobs * result["loss_mask"], + rtol=0.0, + atol=0.0, + ) + + def _make_interaction( interaction_id: str, input_tokens: list[int], diff --git a/tests/test_ray_scheduler.py b/tests/test_ray_scheduler.py index 46ea86f8d7..38990c522b 100644 --- a/tests/test_ray_scheduler.py +++ b/tests/test_ray_scheduler.py @@ -479,15 +479,27 @@ def test_create_workers_with_fork_colocation_delegates_to_fork_workers( ] called = [] - def fake_fork_workers(role: str, target_role: str): - called.append((role, target_role)) + def fake_fork_workers( + role: str, + target_role: str, + command: str | None = None, + env_vars: list[dict[str, str]] | None = None, + ): + called.append((role, target_role, command, env_vars)) return ["ref/0", "ref/1"] monkeypatch.setattr(scheduler, "fork_workers", fake_fork_workers) job = Job( role="ref", replicas=2, - tasks=[SchedulingSpec(cpu=1, gpu=1, mem=1)], + tasks=[ + SchedulingSpec( + cpu=1, + gpu=1, + mem=1, + env_vars={"ROLE_ENV": "ref"}, + ) + ], scheduling_strategy=SchedulingStrategy( type=SchedulingStrategyType.colocation, target="actor", fork=True ), @@ -496,7 +508,11 @@ def fake_fork_workers(role: str, target_role: str): worker_ids = scheduler.create_workers(job) assert worker_ids == ["ref/0", "ref/1"] - assert called == [("ref", "actor")] + assert len(called) == 1 + role, target_role, command, env_vars = called[0] + assert (role, target_role, command) == ("ref", "actor", None) + assert env_vars is not None + assert [env["ROLE_ENV"] for env in env_vars] == ["ref", "ref"] def test_colocation_replica_mismatch_raises_error(tmp_path): diff --git a/tests/test_rollout_controller.py b/tests/test_rollout_controller.py index 749a4d193d..44396e6ff9 100644 --- a/tests/test_rollout_controller.py +++ b/tests/test_rollout_controller.py @@ -17,6 +17,7 @@ GenerationHyperparameters, InferenceEngineConfig, SchedulingSpec, + SchedulingStrategy, SGLangConfig, ) from areal.infra import RolloutController @@ -43,6 +44,7 @@ def create_test_config(backend="sglang:d2", **kwargs): class MockScheduler: def __init__(self): self.workers = [] + self.jobs = [] self.call_count = 0 self.engine_calls = [] self._pending_results = {} # worker_id -> dict[task_id -> result] @@ -50,6 +52,7 @@ def __init__(self): def create_workers(self, job, *args, **kwargs): """Create workers based on Job specification.""" + self.jobs.append(job) role = job.role replicas = job.replicas worker_ids = [f"{role}/{i}" for i in range(replicas)] @@ -57,7 +60,11 @@ def create_workers(self, job, *args, **kwargs): Worker( id=wid, ip="127.0.0.1", - worker_ports=["8000", "8001"], + worker_ports=( + ["8000", "8001"] + if job.scheduling_strategy.fork + else ["8000", "8001", "8002"] + ), engine_ports=["9000", "9001"], ) for wid in worker_ids @@ -76,7 +83,9 @@ async def create_engine(self, worker_id, engine, engine_name, config): async def async_call_engine(self, worker_id, method, *args, **kwargs): self.engine_calls.append((worker_id, method, args, kwargs)) self.call_count += 1 - if method == "agenerate": + if method == "launch_server": + return Mock(host="127.0.0.1", port=8000) + elif method == "agenerate": return Mock() # Handle submit method - return a task_id and store the result elif method == "submit": @@ -216,6 +225,93 @@ def test_initialize_creates_workers(self): controller.destroy() + def test_initialize_nonfork_colocation_uses_port_after_actor_rendezvous(self): + """A reused actor worker reserves port 2 for SGLang NCCL.""" + config = create_test_config( + backend="sglang:d2", + scheduling_strategy=SchedulingStrategy( + type="colocation", + target="actor", + fork=False, + ), + ) + scheduler = MockScheduler() + controller = RolloutController( + inf_engine=MockInferenceEngine, + config=config, + scheduler=scheduler, + ) + + controller.initialize(role="rollout", server_args={"dist_init_addr": None}) + + launch_calls = [ + call for call in scheduler.engine_calls if call[1] == "launch_server" + ] + assert len(launch_calls) == 2 + for _, _, _, kwargs in launch_calls: + server_args = kwargs["server_args"] + assert server_args["nccl_port"] == 8002 + assert server_args["dist_init_addr"] is None + + controller.destroy() + + def test_initialize_forked_colocation_uses_owned_rendezvous_port(self): + """A forked rollout worker can use its own port 1 for SGLang NCCL.""" + config = create_test_config( + backend="sglang:d2", + scheduling_strategy=SchedulingStrategy( + type="colocation", + target="actor", + fork=True, + ), + ) + scheduler = MockScheduler() + controller = RolloutController( + inf_engine=MockInferenceEngine, + config=config, + scheduler=scheduler, + ) + + controller.initialize(role="rollout", server_args={}) + + launch_calls = [ + call for call in scheduler.engine_calls if call[1] == "launch_server" + ] + assert len(launch_calls) == 2 + for _, _, _, kwargs in launch_calls: + assert kwargs["server_args"]["nccl_port"] == 8001 + + controller.destroy() + + def test_initialize_nonfork_colocation_without_third_port_fails(self): + """A reused actor worker must not silently reuse its train TCPStore.""" + config = create_test_config( + backend="sglang:d2", + scheduling_strategy=SchedulingStrategy( + type="colocation", + target="actor", + fork=False, + ), + ) + scheduler = MockScheduler() + original_create_workers = scheduler.create_workers + + def create_workers_with_two_ports(job, *args, **kwargs): + worker_ids = original_create_workers(job, *args, **kwargs) + for worker in scheduler.workers: + worker.worker_ports = ["8000", "8001"] + return worker_ids + + scheduler.create_workers = create_workers_with_two_ports + controller = RolloutController( + inf_engine=MockInferenceEngine, + config=config, + scheduler=scheduler, + ) + + with pytest.raises(ValueError, match="needs at least 3 allocated ports"): + controller.initialize(role="rollout", server_args={}) + def test_initialize_creates_staleness_manager(self): config = create_test_config( consumer_batch_size=32, @@ -321,7 +417,8 @@ def test_destroy_handles_scheduler_error(self): controller.initialize(role="rollout", server_args={}) - controller.destroy() + with pytest.raises(RuntimeError, match="rollout worker delete"): + controller.destroy() class TestRolloutControllerCapacity: diff --git a/tests/test_sglang_pp_unit.py b/tests/test_sglang_pp_unit.py index 203023c124..f2058f08d4 100644 --- a/tests/test_sglang_pp_unit.py +++ b/tests/test_sglang_pp_unit.py @@ -819,3 +819,13 @@ def __exit__(self, *a): # Should not raise. aws._init_per_pp_weight_update_groups(state, meta, engine, gen_pp_size) assert state.group_names == ["update_weight_group_0"] + + +def test_pause_requests_abort_before_in_place_pause(): + requests = SGLangBackend().get_pause_requests() + + assert [request.endpoint for request in requests] == [ + "/pause_generation", + "/pause_generation", + ] + assert [request.payload for request in requests] == [{}, {"mode": "in_place"}] diff --git a/tests/torchrun/run_mopd_teacher_residency.py b/tests/torchrun/run_mopd_teacher_residency.py new file mode 100644 index 0000000000..4da12c34d3 --- /dev/null +++ b/tests/torchrun/run_mopd_teacher_residency.py @@ -0,0 +1,136 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Fresh-process CUDA regression for MOPD teacher flat-buffer residency.""" + +from __future__ import annotations + +import argparse +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from areal.engine.megatron_utils.weight_residency import MegatronWeightResidency + + +class _FallbackFlatBuffer: + def __init__(self): + # Large enough to distinguish a real storage release from allocator + # noise while remaining cheap on single-GPU CI runners. + self.param_data = _parameter_pattern() + self.grad_data = torch.ones( + (4 * 1024 * 1024,), dtype=torch.float32, device="cuda" + ) + + +class _NativeFlatBuffer(_FallbackFlatBuffer): + """Exercise the API provided by MCore ParamAndGradBuffer.""" + + def __init__(self): + super().__init__() + self.offload_calls = 0 + self.reload_calls: list[bool] = [] + + def offload_to_cpu(self, move_params: bool = True, move_grads: bool = True) -> None: + self.offload_calls += 1 + if move_params: + self._cpu_param_data = self.param_data.cpu() + self._param_data_size = self.param_data.storage().size() + self.param_data.storage().resize_(0) + if move_grads: + self._grad_data_size = self.grad_data.storage().size() + self.grad_data.storage().resize_(0) + + def reload_from_cpu( + self, move_params: bool = True, move_grads: bool = True + ) -> None: + self.reload_calls.append(move_grads) + if move_params: + self.param_data.storage().resize_(self._param_data_size) + self.param_data.copy_(self._cpu_param_data, non_blocking=True) + if move_grads: + self.grad_data.storage().resize_(self._grad_data_size) + self.grad_data.zero_() + + +class _FakeMCoreDDP: + def __init__(self, flat_buffer): + self.buffers = [flat_buffer] + self.expert_parallel_buffers = [] + + +def _parameter_pattern() -> torch.Tensor: + values = torch.arange(16 * 1024 * 1024, dtype=torch.float32, device="cuda") + return values.remainder_(251).div_(251) + + +def main(mode: str) -> None: + assert torch.cuda.is_available(), "CUDA worker cannot see a GPU" + torch.cuda.empty_cache() + buffer_type = _NativeFlatBuffer if mode == "native" else _FallbackFlatBuffer + flat_buffer = buffer_type() + ddp = _FakeMCoreDDP(flat_buffer) + residency = MegatronWeightResidency( + SimpleNamespace(model=[ddp], optimizer=None, device=torch.device("cuda")) + ) + expected = flat_buffer.param_data.cpu() + param_bytes = flat_buffer.param_data.numel() * flat_buffer.param_data.element_size() + + with patch("megatron.core.distributed.DistributedDataParallel", _FakeMCoreDDP): + try: + for _ in range(2): + torch.cuda.synchronize() + resident_bytes = torch.cuda.memory_allocated() + resident_reserved_bytes = torch.cuda.memory_reserved() + + residency.release_memory(tags=["optimizer", "weights"]) + + torch.cuda.synchronize() + offloaded_bytes = torch.cuda.memory_allocated() + offloaded_reserved_bytes = torch.cuda.memory_reserved() + assert flat_buffer.param_data.untyped_storage().nbytes() == 0 + assert flat_buffer.grad_data.untyped_storage().nbytes() == 0 + if mode == "fallback": + assert hasattr(flat_buffer, "cpu_param_data") + assert not hasattr(flat_buffer.param_data, "cpu_data") + assert resident_bytes - offloaded_bytes >= int(param_bytes * 0.75) + assert resident_reserved_bytes - offloaded_reserved_bytes >= int( + param_bytes * 0.75 + ) + assert residency.released_tags == frozenset({"optimizer", "weights"}) + + residency.resume_memory(tags=["optimizer", "weights"]) + + torch.cuda.synchronize() + restored_bytes = torch.cuda.memory_allocated() + restored_reserved_bytes = torch.cuda.memory_reserved() + assert flat_buffer.param_data.untyped_storage().nbytes() == param_bytes + # Teacher scoring does not restore gradients; training allocates + # them separately through ensure_grad_buffers(). + assert flat_buffer.grad_data.untyped_storage().nbytes() == 0 + assert restored_bytes - offloaded_bytes >= int(param_bytes * 0.75) + assert restored_reserved_bytes - offloaded_reserved_bytes >= int( + param_bytes * 0.75 + ) + torch.testing.assert_close( + flat_buffer.param_data.cpu(), + expected, + rtol=0.0, + atol=0.0, + ) + assert residency.released_tags == frozenset() + if isinstance(flat_buffer, _NativeFlatBuffer): + assert flat_buffer.offload_calls == 2 + assert flat_buffer.reload_calls == [False, False] + finally: + residency.release_memory(tags=["optimizer", "weights"]) + + del ddp, flat_buffer, residency, expected + torch.cuda.empty_cache() + print(f"Passed mode={mode}") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--mode", choices=("fallback", "native"), required=True) + main(parser.parse_args().mode)