Skip to content
Closed
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
6 changes: 5 additions & 1 deletion examples/train_integrations/harbor_skycap/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,9 @@ uv run --isolated --extra fsdp --extra harbor --extra skycap \
The rest of the configuration is the sibling's: `harbor_trial_config` holds
Harbor's `TrialConfig`, with defaults from `../harbor/harbor_trial_config/default.yaml`.
`skycap.*` sets the record directory (default `{trainer.export_path}/skycap`),
the idle TTL, the port and the renderer pool size.
the idle TTL, the port and the renderer pool size. With `trainer.logger=wandb`, each step's
documents go to W&B, without sidecars, as a version of `skycap-records-train-<run id>`, aliased
`step-N` and `latest`. `skycap.wandb.phases` adds `eval`; `skycap.wandb.enabled=false` turns it off.

## How it fits

Expand All @@ -37,6 +39,7 @@ the idle TTL, the port and the renderer pool size.
| `servers.py` | The server pool: one Ray actor per server, each running a `skycap.CaptureService` on a port of its own. skycap builds how calls reach the model from the options; the integration supplies only its engine wire. |
| `engine.py` | `SkyRLEngine`: skycap's vLLM wire on `/skyrl/v1/generate`, with packed routed experts and sampler support decoded by SkyRL's own `generate_wire`, and sessions released at `/finish_session`. |
| `harbor_generator.py` | Per trial: create a trajectory, point the agent's `api_base` at it, run Harbor, and `finish` with the reward to get the samples. A retry gets a fresh trajectory. |
| `artifacts.py` | `SkycapUploads`, a trainer callback: after each step it fetches every attempt's document from its server and uploads them on a background thread, with a `step.json` that marks each one trained or dropped. An upload failure is logged and does not fail the step. |
| `compose.py` | Samples to a step-wise `GeneratorOutput`: a trial's paths are contiguous under its `TrajectoryID`, the last one marked `is_last_step` and carrying the reward. |

What's imposed on every call:
Expand All @@ -56,6 +59,7 @@ Masking is the sibling's:
`enable_return_routed_experts=true`, and the trainer change is a follow-up.
- **Sampler support** (`enable_return_sample_support_set`) is passed through,
padded to `top_k`.
- **No W&B uploads under the fully-async trainer.** It fires no callbacks.
- **One skycap server per run.** The generator takes a list of URLs and spreads
trajectories over them, for when servers are launched separately.

Expand Down
162 changes: 162 additions & 0 deletions examples/train_integrations/harbor_skycap/artifacts.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,162 @@
"""Each step's skycap documents, uploaded to W&B as one version of a ``skycap-records-<phase>-<run>`` artifact."""

import json
import re
import tempfile
import time
import urllib.request
from concurrent.futures import Future, ThreadPoolExecutor
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any, Dict, List, Optional

import orjson
import zstandard
from loguru import logger

from skyrl.train.utils.callbacks import CallbackInput, TrainingCallback, TrainingControl

ARTIFACT_TYPE = "skycap-records"
FETCH_TIMEOUT = 60.0
UPLOAD_ATTEMPTS = 3
UPLOAD_TIMEOUT = 600.0
RETRY_DELAY = 10.0


@dataclass
class Created:
"""One skycap trajectory the generator opened: a trial attempt."""

server: str
id: str
instance_id: str
repetition_id: int
attempt: int


class RecordLog:
"""The trajectories created since the last upload, per training phase."""

def __init__(self) -> None:
self._created: Dict[str, List[Created]] = {}

def add(self, phase: str, created: Created) -> None:
self._created.setdefault(phase, []).append(created)

def take(self, phase: str) -> List[Created]:
return self._created.pop(phase, [])


def artifact_name(phase: str, run_id: str) -> str:
return f"{ARTIFACT_TYPE}-{phase}-{re.sub(r'[^a-zA-Z0-9_.-]', '-', run_id)}"


def fetch_document(created: Created) -> Optional[dict]:
"""The stored document, from the server that wrote it; None unless it is on that server's disk."""
url = f"{created.server}/trajectories/{created.id}"
try:
with urllib.request.urlopen(url, timeout=FETCH_TIMEOUT) as response:
document = json.load(response)
except Exception as error: # noqa: BLE001 - counted as missing
logger.warning(f"skycap artifact: fetching {url} failed: {type(error).__name__}: {error}")
return None
# A document still in memory (its write failed) has no format_version and no sidecars manifest.
return document if "format_version" in document else None


def upload(
wandb: Any, run_id: str, phase: str, step: int, created: List[Created], trained: Optional[set]
) -> Optional[str]:
"""Upload the documents and a ``step.json`` index, aliased ``step-N`` and ``latest``. Returns the artifact name."""
last_attempt: Dict[tuple, int] = {}
for entry in created:
key = (entry.instance_id, entry.repetition_id)
last_attempt[key] = max(last_attempt.get(key, -1), entry.attempt)
index = []
with tempfile.TemporaryDirectory() as tmp:
artifact = wandb.Artifact(name=artifact_name(phase, run_id), type=ARTIFACT_TYPE)
for entry in created:
document = fetch_document(entry)
key = (entry.instance_id, entry.repetition_id)
row = asdict(entry)
del row["server"]
row["uploaded"] = document is not None
# A superseded attempt never trains; without dynamic sampling every final attempt does.
row["trained"] = entry.attempt == last_attempt[key] and (trained is None or key in trained)
index.append(row)
if document is None:
continue
path = Path(tmp) / f"{entry.id}.json.zst"
path.write_bytes(zstandard.ZstdCompressor().compress(orjson.dumps(document)))
artifact.add_file(str(path), name=path.name)
uploaded = sum(row["uploaded"] for row in index)
missing = len(index) - uploaded
if missing:
logger.warning(f"skycap artifact: {missing} of {len(index)} {phase} documents of step {step} are missing")
if not uploaded:
return None
(Path(tmp) / "step.json").write_bytes(orjson.dumps(index))
artifact.add_file(str(Path(tmp) / "step.json"), name="step.json")
artifact.metadata.update(
{
"global_step": step,
"training_phase": phase,
"run_id": run_id,
"contents": "documents",
"num_uploaded": uploaded,
"num_missing": missing,
"num_trained": sum(row["trained"] for row in index),
}
)
wandb.log_artifact(artifact, aliases=[f"step-{step}", "latest"]).wait(timeout=UPLOAD_TIMEOUT)
logger.info(f"skycap artifact {artifact.name}:step-{step}: {uploaded} documents")
return artifact.name


def upload_with_retries(*args: Any) -> Optional[str]:
"""``upload``, retried from the fetch on any exception, up to ``UPLOAD_ATTEMPTS`` times."""
for attempt in range(1, UPLOAD_ATTEMPTS + 1):
try:
return upload(*args)
except Exception as error: # noqa: BLE001 - retried, then raised
if attempt == UPLOAD_ATTEMPTS:
raise
logger.warning(
f"skycap artifact: attempt {attempt}/{UPLOAD_ATTEMPTS} failed: {type(error).__name__}: {error}"
)
time.sleep(RETRY_DELAY * attempt)


class SkycapUploads(TrainingCallback):
"""Uploads the train records after each step and the eval records after each eval pass, off the step."""

def __init__(self, records: RecordLog, phases: List[str]) -> None:
self.records = records
self.phases = set(phases)
self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="skycap-artifacts")

def _submit(self, trainer: Any, phase: str, step: int, trained: Optional[set]) -> None:
created = self.records.take(phase)
tracker = trainer.tracker
if phase not in self.phases or not created or tracker is None or tracker.backend != "wandb":
return
wandb = tracker.logger
future = self._executor.submit(upload_with_retries, wandb, wandb.run.id, phase, step, created, trained)
future.add_done_callback(_log_failure)

def on_step_end(self, trainer: Any, callback_input: CallbackInput, control: TrainingControl) -> None:
ids = callback_input.trajectory_ids
trained = None if ids is None else {(str(t.instance_id), t.repetition_id) for t in ids}
self._submit(trainer, "train", callback_input.global_step, trained)

def on_eval_end(self, trainer: Any, callback_input: CallbackInput, control: TrainingControl) -> None:
self._submit(trainer, "eval", callback_input.global_step, set())

def on_train_end(self, trainer: Any, callback_input: CallbackInput, control: TrainingControl) -> None:
# The tracker finishes the W&B run right after this event.
self._executor.shutdown(wait=True)


def _log_failure(future: Future) -> None:
if future.exception() is not None:
logger.opt(exception=future.exception()).error("uploading skycap records failed")
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
import os
import sys
from dataclasses import dataclass, field
from typing import Any, Optional
from typing import Any, List, Optional

import ray
import yaml
Expand All @@ -28,10 +28,20 @@
HarborSkyRLConfig,
_deep_merge,
)
from ..artifacts import RecordLog, SkycapUploads
from ..harbor_generator import HarborSkycapGenerator
from ..servers import SkycapServers, start_servers


@dataclass
class SkycapWandbConfig:
enabled: bool = True
"""Upload each step's documents, without sidecars, as a version of ``skycap-records-<phase>-<run id>``,
when ``trainer.logger`` is wandb. The trainer fetches them from the skycap servers."""
phases: List[str] = field(default_factory=lambda: ["train"])
"""The training phases to upload: ``train``, ``eval``."""


@dataclass
class SkycapConfig:
num_servers: int = 1
Expand All @@ -43,6 +53,7 @@ class SkycapConfig:
record_dir: Optional[str] = None
"""Where ended trajectories are written. Defaults to ``{trainer.export_path}/skycap``; each server writes
on its own node, so point it at a shared filesystem to have one directory for the run."""
wandb: SkycapWandbConfig = field(default_factory=SkycapWandbConfig)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit style: comments are not added

ttl: float = 3600.0
"""Seconds an open trajectory may be idle before skycap writes it as abandoned and releases it."""
renderer_pool_size: int = 8
Expand All @@ -59,6 +70,10 @@ class HarborSkycapConfig(HarborSkyRLConfig):
skycap: SkycapConfig = field(default_factory=SkycapConfig)


def record_dir(cfg: Any) -> str:
return cfg.skycap.record_dir or os.path.join(cfg.trainer.export_path, "skycap")


def start_skycap(cfg: Any, engine_url: str) -> SkycapServers:
"""skycap servers in token mode, in front of SkyRL's router."""
ie = cfg.generator.inference_engine
Expand Down Expand Up @@ -86,24 +101,34 @@ def start_skycap(cfg: Any, engine_url: str) -> SkycapServers:
num_servers=cfg.skycap.num_servers,
num_cpus_per_server=cfg.skycap.num_cpus_per_server,
placement_strategy=cfg.skycap.placement_strategy,
record_dir=cfg.skycap.record_dir or os.path.join(cfg.trainer.export_path, "skycap"),
record_dir=record_dir(cfg),
ttl=cfg.skycap.ttl,
)


class HarborSkycapExp(HarborExp):
skycap: Optional[SkycapServers] = None
records: Optional[RecordLog] = None

def get_generator(self, cfg, tokenizer, inference_engine_client):
if self.skycap is None:
self.skycap = start_skycap(cfg, inference_engine_client.get_endpoint_url())
if self.records is None and cfg.skycap.wandb.enabled:
self.records = RecordLog()
return HarborSkycapGenerator(
generator_cfg=cfg.generator,
harbor_cfg=cfg.harbor_trial_config,
capture_urls=self.skycap.urls,
inference_engine_client=inference_engine_client,
records=self.records,
)

def get_trainer(self, *args, **kwargs):
trainer = super().get_trainer(*args, **kwargs)
if self.records is not None:
trainer.add_callback(SkycapUploads(self.records, self.cfg.skycap.wandb.phases))
return trainer

def run(self):
try:
super().run()
Expand Down
16 changes: 14 additions & 2 deletions examples/train_integrations/harbor_skycap/harbor_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
from skyrl.train.generators.utils import build_vllm_cache_salt
from skyrl.train.utils.rate_limiter import create_rate_limiter

from .artifacts import Created, RecordLog
from .compose import TrialOutcome, compose, split

litellm.suppress_debug_info = True
Expand All @@ -53,13 +54,15 @@ def __init__(
harbor_cfg: Dict[str, Any],
capture_urls: List[str],
inference_engine_client: Any = None,
records: Optional[RecordLog] = None,
) -> None:
"""
Args:
generator_cfg: the run's generator config.
harbor_cfg: Harbor's ``TrialConfig`` template.
capture_urls: the skycap servers to spread trajectories over.
inference_engine_client: read for its ``weight_version``, which keys the prefix-cache salt.
records: the log each created trajectory goes into, for the W&B uploads.
"""
if not getattr(generator_cfg, "step_wise_trajectories", False):
raise ValueError(
Expand All @@ -80,6 +83,7 @@ def __init__(
self.generator_cfg = generator_cfg
self.capture_urls = list(capture_urls)
self.inference_engine_client = inference_engine_client
self.records = records
served = generator_cfg.inference_engine.served_model_name
if served is None or "/" in served:
raise ValueError("generator.inference_engine.served_model_name must be set, without '/'")
Expand Down Expand Up @@ -110,6 +114,7 @@ async def generate(self, input_batch: GeneratorInput, disable_tqdm: bool = False
raise ValueError(f"Prompt count ({len(prompts)}) doesn't match trajectory_ids ({len(trajectory_ids)})")
metadata = input_batch.get("batch_metadata")
step = getattr(metadata, "global_step", None)
phase = getattr(metadata, "training_phase", "train")
cache_salt = self._cache_salt()

outcomes: List[Optional[TrialOutcome]] = [None] * len(prompts)
Expand All @@ -125,7 +130,7 @@ async def generate(self, input_batch: GeneratorInput, disable_tqdm: bool = False
async with CapturePool(self.capture_urls) as pool:

async def worker(index: int, prompt: ConversationType, trajectory_id: TrajectoryID) -> None:
outcomes[index] = await self._trial(pool, prompt, trajectory_id, cache_salt, step)
outcomes[index] = await self._trial(pool, prompt, trajectory_id, cache_salt, step, phase)
progress.update(1)

try:
Expand All @@ -149,13 +154,14 @@ async def _trial(
trajectory_id: TrajectoryID,
cache_salt: Optional[str],
step: Optional[int],
phase: str,
) -> TrialOutcome:
"""One rollout, retried on unknown errors. Never raises: one failure must not cancel the batch."""
started = time.monotonic()
for attempt in range(MAX_NUM_RETRIES_PER_TRIAL):
prefix = f"Trajectory {trajectory_id} attempt {attempt + 1}/{MAX_NUM_RETRIES_PER_TRIAL}"
try:
outcome = await self._attempt(pool, prompt, trajectory_id, cache_salt, step, attempt)
outcome = await self._attempt(pool, prompt, trajectory_id, cache_salt, step, phase, attempt)
except Exception as error: # noqa: BLE001 - retried, then masked
logger.warning(f"{prefix} failed: {type(error).__name__}: {error}")
continue
Expand All @@ -172,6 +178,7 @@ async def _attempt(
trajectory_id: TrajectoryID,
cache_salt: Optional[str],
step: Optional[int],
phase: str,
attempt: int,
) -> TrialOutcome:
"""One attempt on its own trajectory, so a retry never shares a graph with the attempt it replaces."""
Expand All @@ -183,6 +190,11 @@ async def _attempt(
"attempt": attempt,
}
async with pool.trajectory(meta) as trajectory:
if self.records is not None:
self.records.add(
phase,
Created(trajectory.server, trajectory.id, meta["instance_id"], meta["repetition_id"], attempt),
)
config = self._trial_config(prompt, trajectory.base_url, cache_salt)
async with self._rate_limiter:
results = await (await Trial.create(TrialConfig.model_validate(config))).run()
Expand Down
Loading