Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 70 additions & 0 deletions areal/api/cli_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -1200,6 +1200,74 @@ def __post_init__(self):
)


@dataclass
class DTEConfig:
"""Configuration for DTE-backed colocated weight updates."""

enabled: bool | None = field(
default=None,
metadata={
"help": (
"Enable DTE colocated weight update. None keeps backward "
"compatibility with environment variables."
)
},
)
transfer: str | None = field(
default=None,
metadata={
"help": (
"Weight transfer type: 'full' sends full weights every update; "
"'delta' sends a full first update, then DTE deltas."
),
"choices": ["full", "delta", None],
},
)
delta_method: str | None = field(
default=None,
metadata={
"help": "How DTE delta mode finds changed weights.",
"choices": ["snapshot", "adamw", None],
},
)
anchor_interval: int | None = field(
default=None,
metadata={"help": "Force a full sync every N deltas. 0 means never."},
)
bytes_ratio: float | None = field(
default=None,
metadata={"help": "Per-tensor sparse-vs-dense fallback ratio."},
)
release_train_weights_after_update: bool | None = field(
default=None,
metadata={
"help": (
"Release training weights after a colocated update so rollout can "
"reuse the shared GPU memory."
)
},
)
sync_model_params_before_payload: bool | None = field(
default=None,
metadata={
"help": (
"Refresh Megatron model-visible params from optimizer main params "
"before building the DTE payload."
)
},
)
inversion_debug: bool | None = field(
default=None,
metadata={"help": "Enable verbose DTE inversion debug logging."},
)
inversion_bf16_margin_rel: float | None = field(
default=None,
metadata={
"help": "Relative BF16 rounding-boundary margin for inversion masks."
},
)


@dataclass
class TrainEngineConfig:
"""Core configuration for model training, including optimization and backend settings."""
Expand Down Expand Up @@ -1296,6 +1364,7 @@ class TrainEngineConfig:
"choices": ["disk", "xccl", "awex"],
},
)
dte: DTEConfig = field(default_factory=DTEConfig)
fsdp: FSDPEngineConfig = field(default_factory=FSDPEngineConfig)
archon: ArchonEngineConfig = field(default_factory=ArchonEngineConfig)
megatron: MegatronEngineConfig = field(default_factory=MegatronEngineConfig)
Expand Down Expand Up @@ -2060,6 +2129,7 @@ class SGLangConfig:
num_continuous_decode_steps: int = 1
load_format: str = "auto"
enable_memory_saver: bool = False
enable_weights_cpu_backup: bool = False
allow_auto_truncate: bool = False
attention_backend: str | None = "fa3"
enable_multimodal: bool = False
Expand Down
2 changes: 2 additions & 0 deletions areal/engine/megatron_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -1090,6 +1090,8 @@ def optimizer_zero_grad(self):
def optimizer_step(self):
with trace_scope("megatron_engine.step"):
update_successful, grad_norm, _ = self.optimizer.step()
for group in self.optimizer.param_groups:
group["_areal_last_step_lr"] = float(group["lr"])
current_lr = self.optimizer.param_groups[0]["lr"]

return dict(
Expand Down
68 changes: 67 additions & 1 deletion areal/trainer/rl_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,7 @@ def _init_impl(
logging.setup_file_logging(StatsLogger.get_log_path(config.stats_logger))

self.config = config
self._apply_dte_config_envvars()
self.processor, self.tokenizer = load_hf_processor_and_tokenizer(
config.tokenizer_path
)
Expand All @@ -154,6 +155,14 @@ def _init_impl(
self._should_offload_actor = (
self._should_offload_rollout or config.actor.offload
)
delta_env = os.environ.get(
"DTE_DELTA_TRANSFER", os.environ.get("AWEX_DELTA_TRANSFER", "0")
)
self._keep_rollout_weights_resident = (
config.actor._version == "v2"
and delta_env.strip().lower() in {"1", "true", "yes", "on"}
and not config.sglang.enable_weights_cpu_backup
)
self._should_offload_critic = (
config.critic is not None and config.critic.offload
)
Expand Down Expand Up @@ -573,7 +582,10 @@ def _offload_rollout(self, is_eval: bool = False):
category=Category.IO,
),
):
rollout.offload()
if self._keep_rollout_weights_resident and not is_eval:
rollout.offload(tags=["kv_cache"])
else:
rollout.offload()

def _onload_rollout(self, is_eval: bool = False) -> None:
cleanup_error: Exception | None = None
Expand Down Expand Up @@ -1056,6 +1068,60 @@ def _init_scheduler(self) -> Scheduler:
return SlurmScheduler(exp_config=self.config)
raise NotImplementedError(f"Unknown scheduler type: {cfg.type}")

def _apply_dte_config_envvars(self) -> None:
dte_config = self.config.actor.dte
transfer = dte_config.transfer
if transfer not in {None, "full", "delta"}:
raise ValueError(
f"actor.dte.transfer must be 'full' or 'delta', got {transfer!r}"
)
delta_enabled = transfer == "delta" if transfer is not None else None
detector = dte_config.delta_method
if detector not in {None, "snapshot", "adamw"}:
raise ValueError(
"actor.dte.delta_method must be 'snapshot' or 'adamw', "
f"got {detector!r}"
)
if detector == "adamw":
detector = "inversion"

enabled = dte_config.enabled
if enabled is not None and delta_enabled is None:
delta_enabled = False
if enabled is False and delta_enabled:
raise ValueError(
"actor.dte.enabled=false conflicts with actor.dte.transfer='delta'"
)
if enabled is None and transfer is not None:
enabled = True

values = {
"DTE_COLOCATE_WEIGHT_UPDATE": enabled,
"DTE_DELTA_TRANSFER": delta_enabled,
"DTE_DELTA_DETECTOR": detector,
"DTE_DELTA_ANCHOR_INTERVAL": dte_config.anchor_interval,
"DTE_DELTA_BYTES_RATIO": dte_config.bytes_ratio,
"DTE_RELEASE_TRAIN_WEIGHTS_AFTER_UPDATE": (
dte_config.release_train_weights_after_update
),
"DTE_SYNC_MODEL_PARAMS_BEFORE_PAYLOAD": (
dte_config.sync_model_params_before_payload
),
"DTE_DELTA_INVERSION_DEBUG": dte_config.inversion_debug,
"DTE_DELTA_INVERSION_BF16_MARGIN_REL": (
dte_config.inversion_bf16_margin_rel
),
}
exported = {
name: "1" if value is True else "0" if value is False else str(value)
for name, value in values.items()
if value is not None
}
os.environ.update(exported)
for engine_config in (self.config.actor, self.config.rollout):
for spec in engine_config.scheduling_spec:
spec.env_vars.update(exported)

def _create_dataloader(
self,
dataset: Dataset,
Expand Down
2 changes: 1 addition & 1 deletion areal/v2/inference_service/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@ def get_resume_request(self) -> HttpRequest:
"""Return the HTTP request that resumes generation on the backend."""
...

def get_offload_request(self) -> HttpRequest:
def get_offload_request(self, tags: list[str] | None = None) -> HttpRequest:
"""Return the HTTP request that offloads model memory on the backend."""
...

Expand Down
9 changes: 5 additions & 4 deletions areal/v2/inference_service/controller/controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -1380,19 +1380,20 @@ def resume(self) -> None:
assert self._workflow_executor is not None
self._workflow_executor.resume()

def offload(self) -> None:
def offload(self, tags: list[str] | None = None) -> None:
"""Offload model memory on all inference workers."""
from areal.infra.utils.concurrent import run_async_task

self._ensure_initialized()
run_async_task(self._async_offload)
run_async_task(self._async_offload, tags)

async def _async_offload(self) -> None:
async def _async_offload(self, tags: list[str] | None = None) -> None:
if not self._data_proxy_addrs:
return
payload: dict = {"tags": tags} if tags is not None else {}
results = await asyncio.gather(
*(
self._async_data_proxy_post(addr, "/release_memory_occupation", {})
self._async_data_proxy_post(addr, "/release_memory_occupation", payload)
for addr in self._data_proxy_addrs
),
return_exceptions=True,
Expand Down
8 changes: 6 additions & 2 deletions areal/v2/inference_service/data_proxy/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -410,14 +410,18 @@ async def continue_generation():
return PauseGenerationResponse(status="ok", paused=False)

@app.post("/release_memory_occupation")
async def release_memory_occupation():
async def release_memory_occupation(request: Request):
inf_bridge: InfBridge | None = app.state.inf_bridge
if inf_bridge is None:
raise HTTPException(
status_code=503,
detail="No inference backend configured (external model mode).",
)
await inf_bridge.offload()
body = await request.json() if await request.body() else {}
tags = body.get("tags")
if tags is not None and not isinstance(tags, list):
raise HTTPException(status_code=400, detail="'tags' must be a list")
await inf_bridge.offload(tags=tags)
return {"status": "ok"}

@app.post("/resume_memory_occupation")
Expand Down
4 changes: 2 additions & 2 deletions areal/v2/inference_service/inf_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,9 +105,9 @@ async def resume(self) -> None:
await self.pause_state.set_paused(False)
logger.info("Resume request sent to %s", self.backend_addr)

async def offload(self) -> None:
async def offload(self, tags: list[str] | None = None) -> None:
"""Offload model memory on the backend inference server."""
http_req = self.backend.get_offload_request()
http_req = self.backend.get_offload_request(tags=tags)
await self._send_request(http_req, timeout=30.0)
logger.info("Offload request sent to %s", self.backend_addr)

Expand Down
5 changes: 3 additions & 2 deletions areal/v2/inference_service/sglang/bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,8 +140,9 @@ def get_pause_request(self) -> HttpRequest:
def get_resume_request(self) -> HttpRequest:
return HttpRequest(endpoint="/continue_generation", payload={})

def get_offload_request(self) -> HttpRequest:
return HttpRequest(endpoint="/release_memory_occupation", payload={})
def get_offload_request(self, tags: list[str] | None = None) -> HttpRequest:
payload = {"tags": tags} if tags is not None else {}
return HttpRequest(endpoint="/release_memory_occupation", payload=payload)

def get_onload_request(self, tags: list[str] | None = None) -> HttpRequest:
payload = {"tags": tags} if tags is not None else {}
Expand Down
3 changes: 2 additions & 1 deletion areal/v2/inference_service/vllm/bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,7 +99,8 @@ def get_pause_request(self) -> HttpRequest:
def get_resume_request(self) -> HttpRequest:
return HttpRequest(endpoint="/areal_continue_generation", payload={})

def get_offload_request(self) -> HttpRequest:
def get_offload_request(self, tags: list[str] | None = None) -> HttpRequest:
del tags
return HttpRequest(endpoint="/sleep", payload={}, method="POST")

def get_onload_request(self, tags: list[str] | None = None) -> HttpRequest:
Expand Down
37 changes: 35 additions & 2 deletions areal/v2/training_service/controller/controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import asyncio
import concurrent.futures
import os
import sys
import threading
import time
Expand Down Expand Up @@ -56,6 +57,7 @@ def __init__(
self._own_process_group = False
self.rollout: Any | None = None
self._weight_update_ctrl: Any | None = None
self._colocate_weight_update = False

# Version management
self._version_lock = Lock()
Expand Down Expand Up @@ -975,6 +977,13 @@ def connect_engine(self, rollout: Any, meta: Any) -> None:

inference_urls: list[str] = rollout.inference_worker_urls
pair_name = f"{self._role}-rollout"
colocate_env = os.environ.get("DTE_COLOCATE_WEIGHT_UPDATE", "0")
self._colocate_weight_update = meta.type == "awex" and colocate_env.lower() in {
"1",
"true",
"yes",
"on",
}

if meta.type == "awex":
# NCCL rendezvous master must live on the rank-0 process's node.
Expand All @@ -995,6 +1004,7 @@ def connect_engine(self, rollout: Any, meta: Any) -> None:
mode="awex",
nccl_master_addr=port_data["host"],
nccl_master_port=port_data["ports"][0],
colocate=self._colocate_weight_update,
)
else: # disk
ctrl.connect(
Expand Down Expand Up @@ -1024,8 +1034,21 @@ def update_weights(self, meta: Any) -> None:
assert meta.version is not None and meta.version > 0, (
f"meta.version must be a positive integer, got {meta.version}"
)
result = self._weight_update_ctrl.update_weights(version=meta.version)
self.rollout.continue_generation()
try:
result = self._weight_update_ctrl.update_weights(version=meta.version)
finally:
if self._colocate_weight_update:
import requests

for url in self._worker_addrs:
response = requests.post(
f"{url}/awex/resume_memory",
json={"tags": ["weights", "optimizer"]},
timeout=120,
)
response.raise_for_status()
else:
self.rollout.continue_generation()
logger.info(
"Weight update v%d completed (%s, %.0fms)",
meta.version,
Expand Down Expand Up @@ -1157,6 +1180,16 @@ def _cleanup_runtime_state(self) -> None:
except Exception:
logger.error("Failed to unregister model: %s", traceback.format_exc())

if self._weight_update_ctrl is not None:
try:
self._weight_update_ctrl.destroy()
except Exception:
logger.error(
"Failed to destroy weight update controller: %s",
traceback.format_exc(),
)
self._weight_update_ctrl = None

self._graceful_shutdown_workers()

for guard_addr, role, worker_index in reversed(self._forked_services):
Expand Down
Loading
Loading