feat: add online rollout framework for RL strategies

- RolloutRunner: generate + score responses with cached re-rollout trigger
- BaseStrategy.__call__ switches online/offline via runner injection
- GRPO/DPO implement prepare_from_rollout; aliases online_grpo/online_dpo
- TrainConfig + train.py add rollout params and CLI flags
- Tests cover generate_responses, RolloutRunner cache, shared __call__
This commit is contained in:
2026-07-20 03:49:56 +08:00
parent 0b6a17330f
commit 754624acf0
7 changed files with 1111 additions and 9 deletions
+272
View File
@@ -0,0 +1,272 @@
"""Online rollout runner for RL training.
Provides:
- :class:`RolloutResult` — universal data container for online sampling
- :class:`BaseRewardModel` — pluggable reward interface
- :class:`RolloutRunner` — generates + scores batches for any RL strategy
"""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Dict, List, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from astrai.inference.sample import SamplingPipeline
@dataclass
class RolloutResult:
"""Universal container produced by :class:`RolloutRunner`.
Fields are designed to cover all common RL algorithms:
GRPO, PPO, Online DPO, Rejection Sampling, etc.
"""
prompts: Tensor
"""Tokenized prompts, shape ``[B, P_len]``."""
responses: Tensor
"""Generated response token IDs, shape ``[B, G, R_max]``."""
response_mask: Tensor
"""Boolean mask for real (non-pad) response tokens, shape ``[B, G, R_max]``."""
rewards: Tensor
"""Reward per response, shape ``[B, G]``."""
logprobs_old: Tensor
"""Per-token log-probs under the behaviour policy, shape ``[B, G, R_max]``."""
prompt_texts: List[str] = field(default_factory=list)
"""Decoded prompt strings (for reward models that need text)."""
response_texts: List[List[str]] = field(default_factory=list)
"""Decoded response strings, shape ``[B, G]`` (for reward models)."""
class BaseRewardModel(ABC):
"""Pluggable reward model interface.
Subclasses should implement ``score()`` to return a ``[B, G]`` float
tensor of rewards. Implementations can be:
* A loaded reward model (e.g. ArmoRM, Skywork-Reward)
* An external API call
* A rule-based function (format, length, keyword matching)
"""
@abstractmethod
def score(self, prompts: List[str], responses: List[List[str]]) -> Tensor:
"""Score each generated response.
Args:
prompts: Raw prompt strings, length ``B``.
responses: Generated response strings, shape ``[B, G]``.
Returns:
Float tensor of shape ``[B, G]``.
"""
...
def generate_responses(
model: nn.Module,
input_ids: Tensor,
attention_mask: Tensor,
max_new_tokens: int,
sampling_pipeline: SamplingPipeline,
stop_ids: List[int],
) -> Dict[str, Tensor]:
"""Autoregressive generation with log-prob tracking.
Args:
model: Policy model (``forward`` returns ``{"logits": ...}``).
input_ids: ``[B, P_len]`` prompt token IDs.
attention_mask: ``[B, P_len]`` boolean mask.
max_new_tokens: Maximum tokens to generate.
sampling_pipeline: Composed sampling strategies.
stop_ids: Token IDs that stop generation (eos, etc.).
Returns:
``dict`` with keys:
- ``generated_ids``: ``[B, max_new_tokens]`` (padded to same length)
- ``generated_mask``: ``[B, max_new_tokens]``
- ``logprobs``: ``[B, max_new_tokens]`` per-token log-probs
"""
_PAD = 0
B, P_len = input_ids.shape
device = input_ids.device
stop_ids_set = set(stop_ids)
done = torch.zeros(B, dtype=torch.bool, device=device)
all_ids = input_ids.clone()
all_mask = attention_mask.clone()
logprob_list: List[Tensor] = []
for _ in range(max_new_tokens):
outputs = model(input_ids=all_ids, input_mask=all_mask)
logits = outputs["logits"][:, -1, :].float()
log_probs = F.log_softmax(logits, dim=-1)
logits = sampling_pipeline.apply(logits, input_ids=all_ids, input_mask=all_mask)
probs = torch.softmax(logits, dim=-1)
next_tokens = torch.multinomial(probs, num_samples=1).squeeze(-1)
next_tokens[done] = _PAD
chosen_logprobs = torch.gather(log_probs, -1, next_tokens.unsqueeze(-1))
logprob_list.append(chosen_logprobs)
all_ids = torch.cat([all_ids, next_tokens.unsqueeze(1)], dim=-1)
all_mask = torch.cat([all_mask, (~done).unsqueeze(1)], dim=-1)
done = done | torch.tensor(
[t.item() in stop_ids_set for t in next_tokens],
device=device,
)
if done.all():
break
logprobs = torch.cat(logprob_list, dim=-1)
if logprobs.size(1) < max_new_tokens:
pad_len = max_new_tokens - logprobs.size(1)
logprobs = F.pad(logprobs, (0, pad_len), value=0.0)
generated_ids = all_ids[:, P_len:]
if generated_ids.size(1) < max_new_tokens:
pad_len = max_new_tokens - generated_ids.size(1)
generated_ids = F.pad(generated_ids, (0, pad_len), value=_PAD)
generated_mask = generated_ids != _PAD
return {
"generated_ids": generated_ids,
"generated_mask": generated_mask,
"logprobs": logprobs,
}
class RolloutRunner:
"""Produces :class:`RolloutResult` from a prompt batch.
Maintains an internal cache so the same batch prompt can be replayed
for multiple gradient steps. A new rollout is triggered every
``rollout_interval`` calls to :meth:`step`.
Usage::
runner = RolloutRunner(policy, old_policy, tokenizer,
reward_model, sampling_pipeline, config)
result = runner(prompt_batch)
"""
def __init__(
self,
policy_model: nn.Module,
old_model: Optional[nn.Module],
tokenizer,
reward_model: BaseRewardModel,
sampling_pipeline: SamplingPipeline,
max_tokens: int = 1024,
group_size: int = 8,
rollout_interval: int = 512,
):
self.policy_model = policy_model
self.old_model = old_model
self.tokenizer = tokenizer
self.reward_model = reward_model
self.sampling_pipeline = sampling_pipeline
self.max_tokens = max_tokens
self.group_size = group_size
self.rollout_interval = rollout_interval
self.stop_ids = getattr(tokenizer, "stop_ids", []) or []
self._cache: Optional[RolloutResult] = None
self._steps_since_rollout: int = 0
def step(self):
"""Advance the internal counter (call once per optimizer step)."""
self._steps_since_rollout += 1
def clear_cache(self):
"""Force next call to re-run rollout."""
self._cache = None
def _tokenize_prompts(self, raw_texts: List[str]) -> Dict[str, Tensor]:
ids_list = self.tokenizer.encode(raw_texts, out_ids=True)
B = len(ids_list)
P_max = max(len(ids) for ids in ids_list) if ids_list else 0
input_ids = torch.zeros(B, P_max, dtype=torch.long)
for i, ids in enumerate(ids_list):
input_ids[i, : len(ids)] = torch.tensor(ids[:P_max], dtype=torch.long)
attention_mask = input_ids != 0
return {"input_ids": input_ids, "attention_mask": attention_mask}
def _decode(self, token_ids: Tensor, mask: Tensor) -> List[List[str]]:
B, G, _ = token_ids.shape
texts = []
for i in range(B):
group_texts = []
for g in range(G):
ids = token_ids[i, g, mask[i, g]].tolist()
group_texts.append(self.tokenizer.decode(ids, skip_special_tokens=True))
texts.append(group_texts)
return texts
@torch.no_grad()
def _run(self, batch: Dict[str, Tensor]) -> RolloutResult:
"""Execute the actual generation + reward scoring."""
prompt_ids = batch["input_ids"] if "input_ids" in batch else batch["prompts"]
prompt_mask = (
batch["attention_mask"] if "attention_mask" in batch else (prompt_ids != 0)
)
B, P_len = prompt_ids.shape
G = self.group_size
device = prompt_ids.device
prompt_texts: List[str] = []
for i in range(B):
ids = prompt_ids[i, prompt_mask[i]].tolist()
prompt_texts.append(self.tokenizer.decode(ids, skip_special_tokens=True))
expanded_ids = prompt_ids.unsqueeze(1).expand(-1, G, -1).reshape(B * G, P_len)
expanded_mask = prompt_mask.unsqueeze(1).expand(-1, G, -1).reshape(B * G, P_len)
gen_out = generate_responses(
model=self.policy_model,
input_ids=expanded_ids,
attention_mask=expanded_mask,
max_new_tokens=self.max_tokens,
sampling_pipeline=self.sampling_pipeline,
stop_ids=self.stop_ids,
)
gen_ids = gen_out["generated_ids"].reshape(B, G, -1)
gen_mask = gen_out["generated_mask"].reshape(B, G, -1)
gen_logprobs = gen_out["logprobs"].reshape(B, G, -1)
response_texts = self._decode(gen_ids, gen_mask)
reward_tensor = self.reward_model.score(prompt_texts, response_texts)
rewards = reward_tensor.to(device=device)
return RolloutResult(
prompts=prompt_ids,
responses=gen_ids,
response_mask=gen_mask,
rewards=rewards,
logprobs_old=gen_logprobs,
prompt_texts=prompt_texts,
response_texts=response_texts,
)
def __call__(self, batch: Dict[str, Tensor]) -> RolloutResult:
"""Return cached or fresh :class:`RolloutResult`.
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
or when the cache is empty.
"""
if self._cache is None or self._steps_since_rollout >= self.rollout_interval:
self._cache = self._run(batch)
self._steps_since_rollout = 0
return self._cache
+101 -3
View File
@@ -9,6 +9,7 @@ import torch.nn.functional as F
from torch import Tensor
from astrai.factory import BaseFactory
from astrai.trainer.rollout import RolloutResult
def create_ref_model(
@@ -87,7 +88,15 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
class BaseStrategy(ABC):
"""Abstract base class for training strategies."""
"""Abstract base class for training strategies.
When a :class:`~astrai.trainer.rollout.RolloutRunner` is injected via
:meth:`set_rollout_runner`, the strategy transparently switches to
online mode: each ``__call__`` produces a :class:`RolloutResult`,
converts it to a training batch via :meth:`prepare_from_rollout`, and
then computes the loss. Without a runner the strategy runs in
offline mode and consumes the batch directly.
"""
def __init__(
self,
@@ -99,6 +108,8 @@ class BaseStrategy(ABC):
self.device = device
self.executor = kwargs.pop("executor", None)
self.extra_kwargs = kwargs
self._rollout_runner = None
self._prev_rollout_result = None
@abstractmethod
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
@@ -112,9 +123,51 @@ class BaseStrategy(ABC):
"""
raise NotImplementedError
def supports_online(self) -> bool:
"""Whether this strategy can operate with a rollout runner.
Base implementation returns ``False``; strategies that implement
:meth:`prepare_from_rollout` should override to return ``True``.
"""
return False
def set_rollout_runner(self, runner):
"""Inject a :class:`RolloutRunner` to enable online rollout mode."""
self._rollout_runner = runner
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
"""Map a :class:`RolloutResult` to the batch layout expected by
:meth:`compute_loss`.
Strategies that return ``True`` from :meth:`supports_online` must
override this. Default raises :class:`NotImplementedError`.
"""
raise NotImplementedError(
f"{type(self).__name__} does not support online rollout"
)
def _on_rollout_refresh(self):
"""Hook fired when a fresh rollout result is produced.
Override to refresh stale state (e.g. syncing the behaviour
policy). Default is a no-op.
"""
pass
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
"""Allow calling strategy directly as a callable."""
return self.compute_loss(batch)
"""Run offline or online forward depending on runner injection."""
if self._rollout_runner is None:
return self.compute_loss(batch)
result = self._rollout_runner(batch)
if result is not self._prev_rollout_result:
self._on_rollout_refresh()
self._prev_rollout_result = result
if self.executor and self.executor.sync_gradients:
self._rollout_runner.step()
train_batch = self.prepare_from_rollout(result)
return self.compute_loss(train_batch)
class StrategyFactory(BaseFactory["BaseStrategy"]):
@@ -260,6 +313,29 @@ class DPOStrategy(BaseStrategy):
return dpo_loss
def supports_online(self) -> bool:
return True
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
"""Pick best/worst response per prompt by reward as chosen/rejected."""
rewards = result.rewards
responses = result.responses
masks = result.response_mask
best = rewards.argmax(dim=-1)
worst = rewards.argmin(dim=-1)
B = responses.shape[0]
idx = torch.arange(B, device=responses.device)
chosen = responses[idx, best]
chosen_mask = masks[idx, best].float()
rejected = responses[idx, worst]
rejected_mask = masks[idx, worst].float()
return {
"chosen": chosen,
"chosen_mask": chosen_mask,
"rejected": rejected,
"rejected_mask": rejected_mask,
}
@StrategyFactory.register("grpo")
class GRPOStrategy(BaseStrategy):
@@ -371,3 +447,25 @@ class GRPOStrategy(BaseStrategy):
total_loss = policy_loss + kl_penalty
return total_loss
def supports_online(self) -> bool:
return True
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
return {
"prompts": result.prompts,
"responses": result.responses,
"masks": result.response_mask,
"rewards": result.rewards,
}
def _on_rollout_refresh(self):
"""Sync the behaviour policy whenever a fresh rollout arrives."""
self.sync_old_model()
# Factory aliases: online variants use the same strategy class; the
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
# online mode, so no separate subclass is needed.
StrategyFactory._entries["online_grpo"] = GRPOStrategy
StrategyFactory._entries["online_dpo"] = DPOStrategy
+52 -3
View File
@@ -8,11 +8,19 @@ from torch.utils.data import DataLoader, random_split
from astrai.config.train_config import TrainConfig
from astrai.dataset import RDSampler
from astrai.inference.sample import (
SamplingPipeline,
TemperatureStrategy,
TopKStrategy,
TopPStrategy,
)
from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
from astrai.serialization import Checkpoint, load_json
from astrai.tokenize import AutoTokenizer
from astrai.trainer.rollout import RolloutRunner
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
@@ -27,7 +35,6 @@ class TrainContext:
config: TrainConfig = field(default=None)
model_config: dict = field(default_factory=dict)
executor: BaseExecutor = field(default=None)
epoch: int = field(default=0)
consumed_samples: int = field(default=0)
loss: float = field(default=0.0)
@@ -194,13 +201,22 @@ class TrainContextBuilder:
strategy_kwargs = dict(cfg.extra_kwargs)
if cfg.strategy in ("dpo", "grpo"):
needs_ref = cfg.strategy in (
"dpo",
"grpo",
"online_grpo",
"online_dpo",
)
needs_old = cfg.strategy in ("grpo", "online_grpo")
if needs_ref:
ref_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device)
strategy_kwargs["ref_model"] = ref_model
if cfg.strategy == "grpo":
old_model = None
if needs_old:
old_model = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model)
).to(device=device)
@@ -214,4 +230,37 @@ class TrainContextBuilder:
**strategy_kwargs,
)
# Enable online rollout when the train_type is an ``online_*`` variant.
is_online = cfg.strategy.startswith("online_")
if is_online:
if not context.strategy.supports_online():
raise ValueError(
f"Strategy '{cfg.strategy}' does not support online rollout"
)
if cfg.reward_model_fn is None:
raise ValueError("reward_model_fn is required for online RL strategies")
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
reward_model = cfg.reward_model_fn()
pipeline = SamplingPipeline(
[
TemperatureStrategy(cfg.rollout_temperature),
TopKStrategy(cfg.rollout_top_k),
TopPStrategy(cfg.rollout_top_p),
]
)
runner = RolloutRunner(
policy_model=context.model,
old_model=old_model,
tokenizer=tokenizer,
reward_model=reward_model,
sampling_pipeline=pipeline,
max_tokens=cfg.rollout_max_tokens,
group_size=strategy_kwargs.get("group_size", 8),
rollout_interval=cfg.rollout_interval,
)
context.strategy.set_rollout_runner(runner)
return context