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:
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user