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
1,069 changes: 1,069 additions & 0 deletions cookbook/rl/reward_loop/bench_streaming_local.py

Large diffs are not rendered by default.

11 changes: 11 additions & 0 deletions src/twinkle/reward_loop/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,11 @@
from .data import RewardItem, RewardResult, assemble_scores, reorder_by_id, split_items
from .config import RewardLoopArgs
from .metrics import RewardLoopMetrics
from .pipeline import AsyncRewardPipeline, BatchHandle
from .worker import RewardLoopWorker
from .reward_manager import (RewardLoopManager, RewardManagerBase, get_reward_manager_cls, register,
registered_managers)

__all__ = ["RewardItem", "RewardResult", "split_items", "reorder_by_id", "assemble_scores", "RewardLoopArgs",
"RewardLoopMetrics", "AsyncRewardPipeline", "BatchHandle", "RewardLoopWorker", "RewardLoopManager",
"RewardManagerBase", "register", "get_reward_manager_cls", "registered_managers"]
24 changes: 24 additions & 0 deletions src/twinkle/reward_loop/config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
from dataclasses import dataclass, field
from typing import Optional


@dataclass
class RewardLoopArgs:
num_workers: int = 8
custom_reward_function_path: Optional[str] = None
custom_reward_function_name: str = "compute_score"
manager_name: str = "naive"
manager_source: str = "register"
manager_module_path: Optional[str] = None
manager_module_name: str = "RewardLoopManager"
unknown_rewards: str = "warn"
mode: str = "async"
backlog: int = 2
on_backlog_full: str = "block"
on_error: str = "raise"
max_rpm: Optional[int] = None
max_tpm: Optional[int] = None
max_concurrent: int = 1
timeout: float = 300.0
reward_worker_executors: Optional[int] = None
reward_kwargs: dict = field(default_factory=dict)
50 changes: 50 additions & 0 deletions src/twinkle/reward_loop/data.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
"""Framework-independent reward loop data contracts."""
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional


@dataclass
class RewardItem:
item_id: str
data_source: str = ""
solution_str: str = ""
ground_truth: str = ""
extra_info: Dict[str, Any] = field(default_factory=dict)
response_ids: Any = None
attention_mask: Any = None
raw_prompt: Optional[str] = None


@dataclass
class RewardResult:
item_id: str
reward_score: float
reward_extra_info: Dict[str, Any] = field(default_factory=dict)


def split_items(items: List[RewardItem], num_workers: int) -> List[List[RewardItem]]:
if num_workers < 1:
raise ValueError("num_workers must be positive")
chunks = [[] for _ in range(min(num_workers, len(items)))]
for index, item in enumerate(items):
chunks[index % len(chunks)].append(item)
return chunks


def reorder_by_id(results: List[RewardResult]) -> Dict[str, RewardResult]:
ordered: Dict[str, RewardResult] = {}
for result in results:
if result.item_id in ordered:
raise ValueError(f"duplicate reward item_id: {result.item_id}")
ordered[result.item_id] = result
return ordered


def assemble_scores(results: List[RewardResult], items: List[RewardItem], mode: str = "scalar") -> List[float]:
if mode != "scalar":
raise NotImplementedError("token reward assembly is reserved for a future release")
by_id = reorder_by_id(results)
missing = [item.item_id for item in items if item.item_id not in by_id]
if missing:
raise KeyError(f"missing reward results: {missing}")
return [float(by_id[item.item_id].reward_score) for item in items]
26 changes: 26 additions & 0 deletions src/twinkle/reward_loop/default_score.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
"""Default rule-based score dispatch."""
import warnings
from typing import Any, Callable, Dict


_SCORERS: Dict[str, Callable] = {}


def register_score(data_source: str):
def decorator(func):
_SCORERS[data_source.lower()] = func
return func
return decorator


def compute_score(data_source: str, solution_str: str, ground_truth: str, extra_info: dict | None = None,
unknown_rewards: str = "warn"):
scorer = _SCORERS.get((data_source or "").lower())
if scorer is None:
message = f"no default reward scorer registered for data_source={data_source!r}"
if unknown_rewards == "raise":
raise ValueError(message)
if unknown_rewards == "warn":
warnings.warn(message, RuntimeWarning, stacklevel=2)
return 0.0, {"warning": message}
return scorer(solution_str, ground_truth, extra_info or {})
23 changes: 23 additions & 0 deletions src/twinkle/reward_loop/metrics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
from dataclasses import dataclass


@dataclass
class RewardLoopMetrics:
submitted_batches: int = 0
collected_batches: int = 0
submit_time: float = 0.0
collect_wait_time: float = 0.0
reward_time: float = 0.0
max_backlog: int = 0

@property
def overlap_ratio(self):
denominator = self.submit_time + self.reward_time
return 1.0 - self.collect_wait_time / denominator if denominator else 0.0

def snapshot(self):
return dict(self.__dict__, overlap_ratio=self.overlap_ratio)

def reset(self):
for field in self.__dataclass_fields__:
setattr(self, field, 0)
140 changes: 140 additions & 0 deletions src/twinkle/reward_loop/pipeline.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
"""Submission/collection adapter for asynchronous reward computation."""
import concurrent.futures
import threading
import time
from dataclasses import dataclass
from typing import Any, List

from .data import RewardItem, RewardResult, reorder_by_id, split_items
from .metrics import RewardLoopMetrics
from .worker import RewardLoopWorker


@dataclass
class BatchHandle:
items: List[RewardItem]
futures: List[Any]
chunks: List[List[RewardItem]]
collected: bool = False


class AsyncRewardPipeline:
def __init__(self, workers=None, num_workers=1, mode="async", backlog=2, on_backlog_full="block",
on_error="raise", worker_kwargs=None):
if mode not in ("async", "sync"):
raise ValueError("mode must be 'async' or 'sync'")
if on_backlog_full not in ("block", "drop_oldest"):
raise ValueError("on_backlog_full must be 'block' or 'drop_oldest'")
if on_error not in ("raise", "zero"):
raise ValueError("on_error must be 'raise' or 'zero'")
if num_workers < 1 or backlog < 1:
raise ValueError("num_workers and backlog must be positive")
self.mode, self.backlog, self.on_backlog_full, self.on_error = mode, backlog, on_backlog_full, on_error
self.workers = workers or [RewardLoopWorker(**(worker_kwargs or {})) for _ in range(num_workers)]
if not self.workers:
raise ValueError("at least one worker is required")
self.pending: List[BatchHandle] = []
self.metrics = RewardLoopMetrics()
self._executor = concurrent.futures.ThreadPoolExecutor(max_workers=max(1, num_workers))
self._lock = threading.RLock()

@classmethod
def from_args(cls, args, **kwargs):
worker_kwargs = vars(args).copy() if hasattr(args, "__dataclass_fields__") else dict(args)
return cls(num_workers=worker_kwargs.pop("num_workers", 1), worker_kwargs=worker_kwargs, **kwargs)

@staticmethod
def _drain_future(future):
if isinstance(future, concurrent.futures.Future):
if not future.done():
future.cancel()
try:
future.result()
except BaseException:
pass
else:
try:
import ray
ray.cancel(future, force=True)
except Exception:
pass

def _discard(self, handle: BatchHandle):
for future in handle.futures:
self._drain_future(future)
handle.collected = True
if handle in self.pending:
self.pending.remove(handle)

def submit(self, batch: List[RewardItem]) -> BatchHandle:
if not batch:
return BatchHandle([], [], [])
with self._lock:
while len(self.pending) >= self.backlog:
if self.on_backlog_full == "drop_oldest":
self._discard(self.pending[0])
else:
self.collect(self.pending[0])
chunks = split_items(batch, len(self.workers))
futures = []
started = time.monotonic()
for worker, chunk in zip(self.workers, chunks):
if hasattr(worker, "compute_score_batch") and hasattr(worker.compute_score_batch, "remote"):
futures.append(worker.compute_score_batch.remote(chunk))
else:
futures.append(self._executor.submit(worker.compute_score_batch, chunk))
handle = BatchHandle(batch, futures, chunks)
self.pending.append(handle)
self.metrics.submitted_batches += 1
self.metrics.submit_time += time.monotonic() - started
self.metrics.max_backlog = max(self.metrics.max_backlog, len(self.pending))
if self.mode == "sync":
self.collect(handle)
return handle

def collect(self, handle: BatchHandle):
if handle is None:
return []
with self._lock:
if handle.collected:
raise RuntimeError("reward batch was already collected or discarded")
started = time.monotonic()
results = []
errors = []
for index, future in enumerate(handle.futures):
try:
if isinstance(future, concurrent.futures.Future):
values = future.result()
else:
import ray
values = ray.get(future)
results.extend(values)
except BaseException as exc:
errors.append((index, exc))
if self.on_error == "raise":
for remaining in handle.futures[index + 1:]:
self._drain_future(remaining)
with self._lock:
handle.collected = True
if handle in self.pending:
self.pending.remove(handle)
raise
results.extend(RewardResult(item.item_id, 0.0, {"error": str(exc)})
for item in handle.chunks[index])
if errors and self.on_error == "zero":
# Successful chunks remain intact; failed chunks have already received zeros.
pass
with self._lock:
handle.collected = True
if handle in self.pending:
self.pending.remove(handle)
self.metrics.collect_wait_time += time.monotonic() - started
self.metrics.collected_batches += 1
by_id = reorder_by_id(results)
return [by_id[item.item_id] for item in handle.items]

def close(self):
with self._lock:
for handle in list(self.pending):
self._discard(handle)
self._executor.shutdown(wait=True, cancel_futures=True)
13 changes: 13 additions & 0 deletions src/twinkle/reward_loop/reward_manager/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
from .base import RewardManagerBase
from .registry import get_reward_manager_cls, register, registered_managers
from .naive import NaiveRewardManager
from .dapo import DAPORewardManager
from .gdpo import GDPORewardManager
from .limited import AsyncTokenBucket, RateLimitedRewardManager
from .remote import RemoteRewardManager

RewardLoopManager = RewardManagerBase

__all__ = ["RewardManagerBase", "RewardLoopManager", "register", "get_reward_manager_cls", "registered_managers",
"NaiveRewardManager", "DAPORewardManager", "GDPORewardManager", "AsyncTokenBucket",
"RateLimitedRewardManager", "RemoteRewardManager"]
76 changes: 76 additions & 0 deletions src/twinkle/reward_loop/reward_manager/base.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
"""Base reward manager implementation."""
import asyncio
import inspect
import math
from abc import ABC
from collections.abc import Mapping
from typing import Any, Callable, Optional

from ..data import RewardItem, RewardResult


class RewardManagerBase(ABC):
def __init__(self, compute_score: Optional[Callable] = None, max_concurrent: Optional[int] = None, **kwargs):
if max_concurrent is not None and max_concurrent < 1:
raise ValueError("max_concurrent must be positive")
self.compute_score = compute_score
self.kwargs = kwargs
self._semaphore = asyncio.Semaphore(max_concurrent) if max_concurrent else None

async def call_score(self, item: RewardItem) -> Any:
if self.compute_score is None:
raise ValueError(f"compute_score is required for item {item.item_id!r}")
args = (item.data_source, item.solution_str, item.ground_truth, item.extra_info)
if inspect.iscoroutinefunction(self.compute_score):
return await self.compute_score(*args)
loop = asyncio.get_running_loop()
value = await loop.run_in_executor(None, lambda: self.compute_score(*args))
if inspect.isawaitable(value):
return await value
return value

@staticmethod
def normalize_score(value: Any):
if isinstance(value, tuple):
if len(value) != 2 or not isinstance(value[1], Mapping):
raise TypeError("tuple score must be (number, mapping)")
score, extra = value
return RewardManagerBase._finite_score(score), dict(extra)
if isinstance(value, Mapping):
if "score" not in value and "reward" not in value:
raise TypeError("mapping score must contain 'score' or 'reward'")
score = value.get("score", value.get("reward"))
extra = value.get("extra_info", {})
if not isinstance(extra, Mapping):
raise TypeError("score extra_info must be a mapping")
return RewardManagerBase._finite_score(score), dict(extra)
return RewardManagerBase._finite_score(value), {}

@staticmethod
def _finite_score(value: Any) -> float:
score = float(value)
if not math.isfinite(score):
raise ValueError("reward score must be finite")
return score

async def run_single(self, item: RewardItem) -> RewardResult:
async def run():
score, extra = self.normalize_score(await self.call_score(item))
return RewardResult(item.item_id, score, extra)
if self._semaphore is None:
return await run()
async with self._semaphore:
return await run()

async def run_batch(self, items):
tasks = [asyncio.create_task(self.run_single(item)) for item in items]
if not tasks:
return []
try:
return await asyncio.gather(*tasks)
except BaseException:
for task in tasks:
if not task.done():
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
raise
19 changes: 19 additions & 0 deletions src/twinkle/reward_loop/reward_manager/dapo.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
from .base import RewardManagerBase
from .registry import register


@register("dapo")
class DAPORewardManager(RewardManagerBase):
def __init__(self, *args, reward_kwargs=None, **kwargs):
super().__init__(*args, **kwargs)
self.reward_kwargs = reward_kwargs or {}

async def run_single(self, item):
result = await super().run_single(item)
limit = self.reward_kwargs.get("max_response_length")
penalty = self.reward_kwargs.get("overlong_penalty", 0.0)
length = len(item.response_ids) if item.response_ids is not None else len(item.solution_str)
if limit is not None and length > limit:
result.reward_score -= penalty
result.reward_extra_info.update({"overlong": True, "response_length": length})
return result
Loading
Loading