refactor: map instruction/input/output to chat roles
- RolloutGenerator._instruction_to_messages builds system/user/assistant list (instruction->system, input->user, output->assistant), replacing single-user-turn concatenation - Remove _iter_samples helper; _prepare_prompts zips parallel list-of-strings fields directly per the collate_fn contract - Tests adopt a system-aware chat template and pin the three-field role mapping - Drop unused imports caught by ruff F401 (torch.Tensor in scheduler.py, iter_raw_records in pipeline.py, Tuple in evaluate_rouge.py)
This commit is contained in:
@@ -4,7 +4,6 @@ import uuid
|
|||||||
from typing import Any, Dict, List, Optional, Tuple
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
from astrai.inference.core.cache import ContiguousCache, KVCache
|
from astrai.inference.core.cache import ContiguousCache, KVCache
|
||||||
from astrai.inference.core.executor import Executor
|
from astrai.inference.core.executor import Executor
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ import tqdm
|
|||||||
from astrai.config.preprocess_config import PipelineConfig
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.preprocessing.core import (
|
from astrai.preprocessing.core import (
|
||||||
build_preprocessing_components,
|
build_preprocessing_components,
|
||||||
iter_raw_records,
|
|
||||||
primary_ids,
|
primary_ids,
|
||||||
)
|
)
|
||||||
from astrai.preprocessing.packing import PackingStrategyFactory
|
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||||
|
|||||||
+112
-32
@@ -18,7 +18,6 @@ from dataclasses import dataclass, field
|
|||||||
from typing import Dict, List, Optional, Tuple
|
from typing import Dict, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
@@ -33,25 +32,26 @@ class RawRollout:
|
|||||||
|
|
||||||
Fields are designed to cover all common RL algorithms:
|
Fields are designed to cover all common RL algorithms:
|
||||||
GRPO, PPO, Online DPO, Rejection Sampling, etc.
|
GRPO, PPO, Online DPO, Rejection Sampling, etc.
|
||||||
|
|
||||||
|
Fields:
|
||||||
|
prompts: Tokenized prompts, shape ``[B, P_len]``.
|
||||||
|
responses: Generated response token IDs, shape ``[B, G, R_max]``.
|
||||||
|
response_mask: Boolean mask for real (non-pad) response tokens,
|
||||||
|
shape ``[B, G, R_max]``.
|
||||||
|
logprobs_old: Per-token log-probs under the behaviour policy,
|
||||||
|
shape ``[B, G, R_max]``.
|
||||||
|
prompt_texts: Decoded prompt strings (for reward models that
|
||||||
|
need text).
|
||||||
|
response_texts: Decoded response strings, shape ``[B, G]``
|
||||||
|
(for reward models).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
prompts: Tensor
|
prompts: Tensor
|
||||||
"""Tokenized prompts, shape ``[B, P_len]``."""
|
|
||||||
|
|
||||||
responses: Tensor
|
responses: Tensor
|
||||||
"""Generated response token IDs, shape ``[B, G, R_max]``."""
|
|
||||||
|
|
||||||
response_mask: Tensor
|
response_mask: Tensor
|
||||||
"""Boolean mask for real (non-pad) response tokens, shape ``[B, G, R_max]``."""
|
|
||||||
|
|
||||||
logprobs_old: Tensor
|
logprobs_old: Tensor
|
||||||
"""Per-token log-probs under the behaviour policy, shape ``[B, G, R_max]``."""
|
|
||||||
|
|
||||||
prompt_texts: List[str] = field(default_factory=list)
|
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)
|
response_texts: List[List[str]] = field(default_factory=list)
|
||||||
"""Decoded response strings, shape ``[B, G]`` (for reward models)."""
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(kw_only=True)
|
@dataclass(kw_only=True)
|
||||||
@@ -60,10 +60,12 @@ class RolloutResult(RawRollout):
|
|||||||
|
|
||||||
Produced by :class:`RolloutRunner` once the :class:`BaseRewardModel`
|
Produced by :class:`RolloutRunner` once the :class:`BaseRewardModel`
|
||||||
has scored the decoded responses.
|
has scored the decoded responses.
|
||||||
|
|
||||||
|
Fields:
|
||||||
|
rewards: Reward per response, shape ``[B, G]``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
rewards: Tensor
|
rewards: Tensor
|
||||||
"""Reward per response, shape ``[B, G]``."""
|
|
||||||
|
|
||||||
|
|
||||||
class BaseRewardModel(ABC):
|
class BaseRewardModel(ABC):
|
||||||
@@ -126,26 +128,31 @@ class RolloutGenerator:
|
|||||||
self.rep_window = rep_window
|
self.rep_window = rep_window
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def generate(self, batch: Dict[str, Tensor]) -> RawRollout:
|
def generate(self, batch: Dict) -> RawRollout:
|
||||||
"""Expand prompts by ``group_size`` and generate one response each."""
|
"""Expand prompts by ``group_size`` and generate one response each.
|
||||||
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, _ = prompt_ids.shape
|
|
||||||
G = self.group_size
|
|
||||||
|
|
||||||
prompt_texts: List[str] = []
|
Accepted batch formats (per sample, repeated B times):
|
||||||
flat_prompt_ids: List[List[int]] = []
|
|
||||||
for i in range(B):
|
- **messages**: ``{"messages": [{"role": "user", "content": "..."}, ...]}``
|
||||||
ids = prompt_ids[i, prompt_mask[i]].tolist()
|
- **instruction + input + output**: ``{"instruction": "...",
|
||||||
text = self.tokenizer.decode(ids, skip_special_tokens=True)
|
"input": "...", "output": "..."}`` — mapped to ``system`` /
|
||||||
for _ in range(G):
|
``user`` / ``assistant`` messages; ``input`` and ``output``
|
||||||
flat_prompt_ids.append(list(ids))
|
are optional and skipped when empty.
|
||||||
prompt_texts.append(text)
|
|
||||||
|
Both are rendered through the tokenizer's chat template with
|
||||||
|
``add_generation_prompt=True`` so rollout prompts match the
|
||||||
|
format the policy was SFT-trained on.
|
||||||
|
"""
|
||||||
|
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
|
||||||
|
B = len(prompt_texts)
|
||||||
|
G = self.group_size
|
||||||
|
# Re-expand flat list to G copies per prompt for run_batch.
|
||||||
|
expanded_prompt_ids: List[List[int]] = []
|
||||||
|
for ids in flat_prompt_ids:
|
||||||
|
expanded_prompt_ids.extend([list(ids)] * G)
|
||||||
|
|
||||||
results = self.scheduler.run_batch(
|
results = self.scheduler.run_batch(
|
||||||
flat_prompt_ids,
|
expanded_prompt_ids,
|
||||||
max_tokens=self.max_tokens,
|
max_tokens=self.max_tokens,
|
||||||
temperature=self.temperature,
|
temperature=self.temperature,
|
||||||
top_k=self.top_k,
|
top_k=self.top_k,
|
||||||
@@ -161,7 +168,14 @@ class RolloutGenerator:
|
|||||||
max_len = max(max_len, len(token_ids))
|
max_len = max(max_len, len(token_ids))
|
||||||
max_len = max(max_len, 1)
|
max_len = max(max_len, 1)
|
||||||
|
|
||||||
device = prompt_ids.device
|
device = self.scheduler.device
|
||||||
|
P_len = max(len(ids) for ids in flat_prompt_ids)
|
||||||
|
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
|
||||||
|
for i, ids in enumerate(flat_prompt_ids):
|
||||||
|
prompts_tensor[i, : len(ids)] = torch.tensor(
|
||||||
|
ids, dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
|
||||||
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
|
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
|
||||||
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
|
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
|
||||||
logprobs_old = torch.zeros((B, G, max_len), dtype=torch.float, device=device)
|
logprobs_old = torch.zeros((B, G, max_len), dtype=torch.float, device=device)
|
||||||
@@ -186,7 +200,7 @@ class RolloutGenerator:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return RawRollout(
|
return RawRollout(
|
||||||
prompts=prompt_ids,
|
prompts=prompts_tensor,
|
||||||
responses=responses,
|
responses=responses,
|
||||||
response_mask=response_mask,
|
response_mask=response_mask,
|
||||||
logprobs_old=logprobs_old,
|
logprobs_old=logprobs_old,
|
||||||
@@ -194,6 +208,72 @@ class RolloutGenerator:
|
|||||||
response_texts=response_texts,
|
response_texts=response_texts,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _prepare_prompts(self, batch: Dict) -> Tuple[List[str], List[List[int]]]:
|
||||||
|
"""Render batch prompts to ``(texts, token_id_lists)``.
|
||||||
|
|
||||||
|
Returns two parallel lists of length B (number of prompts in
|
||||||
|
the batch). Dispatches by batch keys:
|
||||||
|
|
||||||
|
- ``"messages"``: treated as a pre-built message list per sample.
|
||||||
|
- ``"instruction"`` (optionally ``"input"`` and ``"output"``): mapped
|
||||||
|
to ``system`` / ``user`` / ``assistant`` messages respectively.
|
||||||
|
|
||||||
|
Both paths go through the tokenizer's chat template with
|
||||||
|
``add_generation_prompt=True``.
|
||||||
|
"""
|
||||||
|
if "messages" in batch:
|
||||||
|
messages_list = batch["messages"]
|
||||||
|
elif "instruction" in batch:
|
||||||
|
instructions = batch["instruction"]
|
||||||
|
B = len(instructions)
|
||||||
|
inputs = batch.get("input") or [""] * B
|
||||||
|
outputs = batch.get("output") or [""] * B
|
||||||
|
messages_list = [
|
||||||
|
self._instruction_to_messages(i, u, o)
|
||||||
|
for i, u, o in zip(instructions, inputs, outputs)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"Rollout batch must contain either 'messages' or "
|
||||||
|
"'instruction' (optionally 'input'/'output'); got keys: "
|
||||||
|
f"{list(batch.keys())}"
|
||||||
|
)
|
||||||
|
|
||||||
|
prompt_texts: List[str] = []
|
||||||
|
flat_prompt_ids: List[List[int]] = []
|
||||||
|
for messages in messages_list:
|
||||||
|
text = self.tokenizer.apply_chat_template(
|
||||||
|
messages, tokenize=False, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
ids = self.tokenizer.apply_chat_template(
|
||||||
|
messages, tokenize=True, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
prompt_texts.append(text)
|
||||||
|
flat_prompt_ids.append(list(ids))
|
||||||
|
return prompt_texts, flat_prompt_ids
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _instruction_to_messages(
|
||||||
|
instruction: str, inp: str = "", output: str = ""
|
||||||
|
) -> List[Dict[str, str]]:
|
||||||
|
"""Map instruction/input/output to chat messages.
|
||||||
|
|
||||||
|
Role mapping follows the convention used throughout the
|
||||||
|
preprocessing pipeline: ``instruction`` → system, ``input`` →
|
||||||
|
user, ``output`` → assistant. Empty fields are skipped so a
|
||||||
|
bare instruction produces a ``[system]`` list and the chat
|
||||||
|
template's ``add_generation_prompt`` adds the assistant header
|
||||||
|
for sampling.
|
||||||
|
"""
|
||||||
|
messages: List[Dict[str, str]] = []
|
||||||
|
if instruction:
|
||||||
|
messages.append({"role": "system", "content": instruction})
|
||||||
|
if inp:
|
||||||
|
messages.append({"role": "user", "content": inp})
|
||||||
|
if output:
|
||||||
|
messages.append({"role": "assistant", "content": output})
|
||||||
|
return messages
|
||||||
|
|
||||||
|
|
||||||
class RolloutRunner:
|
class RolloutRunner:
|
||||||
"""Produces :class:`RolloutResult` from a prompt batch.
|
"""Produces :class:`RolloutResult` from a prompt batch.
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ Usage::
|
|||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
from collections import Counter
|
from collections import Counter
|
||||||
from typing import Dict, List, Tuple
|
from typing import Dict, List
|
||||||
|
|
||||||
|
|
||||||
def _tokenize(text: str) -> List[str]:
|
def _tokenize(text: str) -> List[str]:
|
||||||
|
|||||||
@@ -13,23 +13,40 @@ from astrai.trainer.rollout import BaseRewardModel
|
|||||||
from astrai.trainer.schedule import SchedulerFactory
|
from astrai.trainer.schedule import SchedulerFactory
|
||||||
from astrai.trainer.trainer import Trainer
|
from astrai.trainer.trainer import Trainer
|
||||||
|
|
||||||
|
_CHAT_TEMPLATE = (
|
||||||
|
"{% for message in messages %}"
|
||||||
|
"{% if message['role'] == 'system' %}"
|
||||||
|
"SYSTEM: {{ message['content'] }}\n"
|
||||||
|
"{% elif message['role'] == 'user' %}"
|
||||||
|
"USER: {{ message['content'] }}\n"
|
||||||
|
"{% elif message['role'] == 'assistant' %}"
|
||||||
|
"ASSISTANT: {{ message['content'] }}\n"
|
||||||
|
"{% endif %}"
|
||||||
|
"{% endfor %}"
|
||||||
|
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
|
||||||
|
)
|
||||||
|
|
||||||
class PromptDataset(Dataset):
|
|
||||||
"""Toy prompt-only dataset for online RL rollout."""
|
|
||||||
|
|
||||||
def __init__(self, n=4, seq_len=8, vocab_size=1000):
|
class InstructionDataset(Dataset):
|
||||||
self.n = n
|
"""Toy instruction/input dataset for online RL rollout.
|
||||||
self.seq_len = seq_len
|
|
||||||
self.vocab_size = vocab_size
|
Each sample has an ``instruction`` and an optional ``input``; the
|
||||||
|
RolloutGenerator renders both through the tokenizer's chat template
|
||||||
|
so the prompt matches the SFT-trained format.
|
||||||
|
"""
|
||||||
|
|
||||||
|
_SAMPLES = [
|
||||||
|
{"instruction": "Hello", "input": ""},
|
||||||
|
{"instruction": "Tell me a story", "input": "about dragons"},
|
||||||
|
{"instruction": "Summarize", "input": "the article"},
|
||||||
|
{"instruction": "Translate", "input": "to French: hi"},
|
||||||
|
]
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
return self.n
|
return len(self._SAMPLES)
|
||||||
|
|
||||||
def __getitem__(self, idx):
|
def __getitem__(self, idx):
|
||||||
return {
|
return dict(self._SAMPLES[idx])
|
||||||
"input_ids": torch.randint(3, self.vocab_size, (self.seq_len,)),
|
|
||||||
"attention_mask": torch.ones(self.seq_len, dtype=torch.bool),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class LengthRewardModel(BaseRewardModel):
|
class LengthRewardModel(BaseRewardModel):
|
||||||
@@ -48,6 +65,14 @@ class LengthRewardModel(BaseRewardModel):
|
|||||||
return rewards
|
return rewards
|
||||||
|
|
||||||
|
|
||||||
|
def instruction_collate_fn(batch):
|
||||||
|
"""Stack a list of instruction/input dicts into a batch dict of lists."""
|
||||||
|
return {
|
||||||
|
"instruction": [b["instruction"] for b in batch],
|
||||||
|
"input": [b.get("input", "") for b in batch],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def _model_fn(model_config):
|
def _model_fn(model_config):
|
||||||
return AutoRegressiveLM(model_config).to(dtype=torch.float32)
|
return AutoRegressiveLM(model_config).to(dtype=torch.float32)
|
||||||
|
|
||||||
@@ -70,15 +95,16 @@ def test_online_dpo_end_to_end(base_test_env):
|
|||||||
tokenizer = base_test_env["tokenizer"]
|
tokenizer = base_test_env["tokenizer"]
|
||||||
model_config = base_test_env["transformer_config"]
|
model_config = base_test_env["transformer_config"]
|
||||||
|
|
||||||
# base_test_env already wrote config.json into test_dir; we only need
|
# Equip tokenizer with a chat template so RolloutGenerator can
|
||||||
# to drop the tokenizer files so AutoTokenizer.from_pretrained works.
|
# render instruction/input via apply_chat_template.
|
||||||
|
tokenizer.set_chat_template(_CHAT_TEMPLATE)
|
||||||
tokenizer.save_pretrained(test_dir)
|
tokenizer.save_pretrained(test_dir)
|
||||||
|
|
||||||
model_fn = partial(_model_fn, model_config)
|
model_fn = partial(_model_fn, model_config)
|
||||||
optimizer_fn = _optimizer_fn
|
optimizer_fn = _optimizer_fn
|
||||||
scheduler_fn = _scheduler_fn
|
scheduler_fn = _scheduler_fn
|
||||||
|
|
||||||
dataset = PromptDataset(n=4, seq_len=8, vocab_size=model_config.vocab_size)
|
dataset = InstructionDataset()
|
||||||
|
|
||||||
train_config = TrainConfig(
|
train_config = TrainConfig(
|
||||||
strategy="online_dpo",
|
strategy="online_dpo",
|
||||||
@@ -103,7 +129,7 @@ def test_online_dpo_end_to_end(base_test_env):
|
|||||||
rollout_top_p=1.0,
|
rollout_top_p=1.0,
|
||||||
rollout_max_tokens=4,
|
rollout_max_tokens=4,
|
||||||
reward_model_fn=LengthRewardModel,
|
reward_model_fn=LengthRewardModel,
|
||||||
collate_fn=None,
|
collate_fn=instruction_collate_fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
trainer = Trainer(train_config)
|
trainer = Trainer(train_config)
|
||||||
|
|||||||
@@ -1,10 +1,4 @@
|
|||||||
"""Unit tests for the online rollout module.
|
"""Unit tests for the online rollout module."""
|
||||||
|
|
||||||
Covers :class:`RolloutResult` / :class:`RawRollout`, :class:`BaseRewardModel`,
|
|
||||||
:class:`RolloutGenerator` (KV-cache-backed via :class:`InferenceScheduler.run_batch`)
|
|
||||||
and :class:`RolloutRunner` including its internal cache and rollout-interval
|
|
||||||
trigger logic.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
@@ -20,26 +14,49 @@ from astrai.trainer.rollout import (
|
|||||||
RolloutRunner,
|
RolloutRunner,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
_CHAT_TEMPLATE = (
|
||||||
|
"{% for message in messages %}"
|
||||||
|
"{% if message['role'] == 'system' %}SYSTEM: {{ message['content'] }}\n{% endif %}"
|
||||||
|
"{% if message['role'] == 'user' %}USER: {{ message['content'] }}\n{% endif %}"
|
||||||
|
"{% if message['role'] == 'assistant' %}ASSISTANT: {{ message['content'] }}\n{% endif %}"
|
||||||
|
"{% endfor %}"
|
||||||
|
"{% if add_generation_prompt %}ASSISTANT: {% endif %}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class FakeTokenizer:
|
class FakeTokenizer:
|
||||||
"""Minimal char-level tokenizer stub for rollout tests.
|
"""Minimal stub tokenizer with a chat template for rollout tests."""
|
||||||
|
|
||||||
Vocab: 0 = pad, 1..255 = byte values. ``stop_ids = [2]`` (a fake
|
|
||||||
EOS) so tests can verify early-stopping behaviour.
|
|
||||||
"""
|
|
||||||
|
|
||||||
stop_ids = [2]
|
stop_ids = [2]
|
||||||
|
|
||||||
def encode(self, texts, out_ids=True, **_):
|
def __init__(self):
|
||||||
|
from astrai.tokenize.chat_template import ChatTemplate
|
||||||
|
|
||||||
|
self._chat_template = ChatTemplate.from_string(_CHAT_TEMPLATE)
|
||||||
|
|
||||||
|
def encode(self, texts, **_):
|
||||||
if isinstance(texts, str):
|
if isinstance(texts, str):
|
||||||
texts = [texts]
|
texts = [texts]
|
||||||
return [[b for b in t.encode("utf-8")] for t in texts]
|
return [[b for b in t.encode("utf-8")] for t in texts]
|
||||||
|
|
||||||
def decode(self, ids, skip_special_tokens=True):
|
def decode(self, ids, skip_special_tokens=True):
|
||||||
out = bytes(b for b in ids if b > 2 or not skip_special_tokens).decode(
|
if isinstance(ids, list):
|
||||||
"utf-8", errors="ignore"
|
return bytes(b for b in ids if b > 2).decode("utf-8", errors="ignore")
|
||||||
|
return str(ids)
|
||||||
|
|
||||||
|
def apply_chat_template(
|
||||||
|
self, messages, tokenize=True, add_generation_prompt=True, **_
|
||||||
|
):
|
||||||
|
rendered = self._chat_template.render(
|
||||||
|
messages=messages, add_generation_prompt=add_generation_prompt
|
||||||
)
|
)
|
||||||
return out
|
if tokenize:
|
||||||
|
return (
|
||||||
|
self.encode(rendered)[0]
|
||||||
|
if isinstance(rendered, str)
|
||||||
|
else [self.encode(t)[0] for t in rendered]
|
||||||
|
)
|
||||||
|
return rendered
|
||||||
|
|
||||||
|
|
||||||
class ConstantRewardModel(BaseRewardModel):
|
class ConstantRewardModel(BaseRewardModel):
|
||||||
@@ -84,10 +101,11 @@ def _make_scheduler(model, tokenizer, max_batch_size=8, max_len=128):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _make_prompt_batch(batch_size=2, prompt_len=6, device="cpu"):
|
def _make_instruction_batch(n=2):
|
||||||
ids = torch.randint(3, 200, (batch_size, prompt_len), device=device)
|
"""Build a batch of instruction+input prompts as lists of strings."""
|
||||||
mask = torch.ones(batch_size, prompt_len, dtype=torch.bool, device=device)
|
instructions = [f"Tell me about topic {i}" for i in range(n)]
|
||||||
return {"input_ids": ids, "attention_mask": mask}
|
inputs = [f"context {i}" for i in range(n)]
|
||||||
|
return {"instruction": instructions, "input": inputs}
|
||||||
|
|
||||||
|
|
||||||
def test_raw_rollout_fields():
|
def test_raw_rollout_fields():
|
||||||
@@ -114,8 +132,6 @@ def test_rollout_result_inherits_raw_rollout_fields():
|
|||||||
assert r.rewards.shape == (2, 3)
|
assert r.rewards.shape == (2, 3)
|
||||||
assert r.prompts.shape == (2, 4)
|
assert r.prompts.shape == (2, 4)
|
||||||
assert r.responses.shape == (2, 3, 5)
|
assert r.responses.shape == (2, 3, 5)
|
||||||
assert r.prompt_texts == []
|
|
||||||
assert r.response_texts == []
|
|
||||||
|
|
||||||
|
|
||||||
def test_base_reward_model_is_abstract():
|
def test_base_reward_model_is_abstract():
|
||||||
@@ -158,9 +174,8 @@ def _make_generator(device, **kw):
|
|||||||
|
|
||||||
def test_rollout_generator_shapes(device):
|
def test_rollout_generator_shapes(device):
|
||||||
gen, _ = _make_generator(device, group_size=3, max_tokens=5)
|
gen, _ = _make_generator(device, group_size=3, max_tokens=5)
|
||||||
batch = _make_prompt_batch(batch_size=2, prompt_len=4, device=device)
|
batch = _make_instruction_batch(n=2)
|
||||||
r = gen.generate(batch)
|
r = gen.generate(batch)
|
||||||
assert r.prompts.shape == (2, 4)
|
|
||||||
assert r.responses.shape == (2, 3, 5)
|
assert r.responses.shape == (2, 3, 5)
|
||||||
assert r.response_mask.shape == (2, 3, 5)
|
assert r.response_mask.shape == (2, 3, 5)
|
||||||
assert r.logprobs_old.shape == (2, 3, 5)
|
assert r.logprobs_old.shape == (2, 3, 5)
|
||||||
@@ -172,14 +187,12 @@ def test_rollout_generator_shapes(device):
|
|||||||
def test_rollout_generator_mask_matches_responses(device):
|
def test_rollout_generator_mask_matches_responses(device):
|
||||||
"""Positions beyond a response's length are pad (mask False)."""
|
"""Positions beyond a response's length are pad (mask False)."""
|
||||||
gen, _ = _make_generator(device, group_size=2, max_tokens=6)
|
gen, _ = _make_generator(device, group_size=2, max_tokens=6)
|
||||||
batch = _make_prompt_batch(batch_size=2, prompt_len=4, device=device)
|
batch = _make_instruction_batch(n=2)
|
||||||
r = gen.generate(batch)
|
r = gen.generate(batch)
|
||||||
for i in range(2):
|
for i in range(2):
|
||||||
for g in range(2):
|
for g in range(2):
|
||||||
real = r.response_mask[i, g].sum().item()
|
real = r.response_mask[i, g].sum().item()
|
||||||
# Pad positions should be 0
|
|
||||||
assert r.responses[i, g, real:].sum() == 0
|
assert r.responses[i, g, real:].sum() == 0
|
||||||
# logprobs after the real tokens are 0 (padding)
|
|
||||||
if real < r.logprobs_old.size(-1):
|
if real < r.logprobs_old.size(-1):
|
||||||
assert torch.all(r.logprobs_old[i, g, real:] == 0)
|
assert torch.all(r.logprobs_old[i, g, real:] == 0)
|
||||||
|
|
||||||
@@ -187,7 +200,7 @@ def test_rollout_generator_mask_matches_responses(device):
|
|||||||
def test_rollout_generator_logprobs_are_nonpositive(device):
|
def test_rollout_generator_logprobs_are_nonpositive(device):
|
||||||
"""Behaviour-policy logprobs of sampled tokens should be ≤ 0."""
|
"""Behaviour-policy logprobs of sampled tokens should be ≤ 0."""
|
||||||
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
|
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
|
||||||
batch = _make_prompt_batch(batch_size=1, prompt_len=3, device=device)
|
batch = _make_instruction_batch(n=1)
|
||||||
r = gen.generate(batch)
|
r = gen.generate(batch)
|
||||||
for i in range(1):
|
for i in range(1):
|
||||||
for g in range(2):
|
for g in range(2):
|
||||||
@@ -196,6 +209,45 @@ def test_rollout_generator_logprobs_are_nonpositive(device):
|
|||||||
assert torch.all(lp <= 1e-5)
|
assert torch.all(lp <= 1e-5)
|
||||||
|
|
||||||
|
|
||||||
|
def test_rollout_generator_instruction_role_mapping(device):
|
||||||
|
"""instruction → system, input → user, output → assistant."""
|
||||||
|
gen, _ = _make_generator(device, group_size=1, max_tokens=4)
|
||||||
|
batch = {
|
||||||
|
"instruction": ["Be helpful"],
|
||||||
|
"input": ["What is 2+2?"],
|
||||||
|
"output": ["Four"],
|
||||||
|
}
|
||||||
|
r = gen.generate(batch)
|
||||||
|
text = r.prompt_texts[0]
|
||||||
|
assert "SYSTEM: Be helpful" in text
|
||||||
|
assert "USER: What is 2+2?" in text
|
||||||
|
assert "ASSISTANT: Four" in text
|
||||||
|
|
||||||
|
|
||||||
|
def test_rollout_generator_messages_format(device):
|
||||||
|
"""Rollout also accepts pre-built messages."""
|
||||||
|
gen, _ = _make_generator(device, group_size=2, max_tokens=4)
|
||||||
|
batch = {
|
||||||
|
"messages": [
|
||||||
|
[{"role": "user", "content": "Hello"}],
|
||||||
|
[{"role": "user", "content": "Goodbye"}],
|
||||||
|
]
|
||||||
|
}
|
||||||
|
r = gen.generate(batch)
|
||||||
|
assert r.responses.shape[0] == 2
|
||||||
|
assert len(r.prompt_texts) == 2
|
||||||
|
assert "Hello" in r.prompt_texts[0] or "USER" in r.prompt_texts[0]
|
||||||
|
|
||||||
|
|
||||||
|
def test_rollout_generator_bad_batch_raises(device):
|
||||||
|
"""Batch without messages or instruction raises a clear error."""
|
||||||
|
gen, _ = _make_generator(device)
|
||||||
|
with pytest.raises(
|
||||||
|
ValueError, match="must contain either 'messages' or 'instruction'"
|
||||||
|
):
|
||||||
|
gen.generate({"input_ids": torch.zeros(2, 4, dtype=torch.long)})
|
||||||
|
|
||||||
|
|
||||||
def _make_runner(device, **kw):
|
def _make_runner(device, **kw):
|
||||||
generator, model = _make_generator(
|
generator, model = _make_generator(
|
||||||
device,
|
device,
|
||||||
@@ -217,10 +269,9 @@ def _make_runner(device, **kw):
|
|||||||
|
|
||||||
def test_rollout_runner_shapes(device):
|
def test_rollout_runner_shapes(device):
|
||||||
runner, _ = _make_runner(device, group_size=3, max_tokens=5)
|
runner, _ = _make_runner(device, group_size=3, max_tokens=5)
|
||||||
batch = _make_prompt_batch(batch_size=2, prompt_len=4, device=device)
|
batch = _make_instruction_batch(n=2)
|
||||||
r, is_fresh = runner(batch)
|
r, is_fresh = runner(batch)
|
||||||
assert is_fresh
|
assert is_fresh
|
||||||
assert r.prompts.shape == (2, 4)
|
|
||||||
assert r.responses.shape == (2, 3, 5)
|
assert r.responses.shape == (2, 3, 5)
|
||||||
assert r.response_mask.shape == (2, 3, 5)
|
assert r.response_mask.shape == (2, 3, 5)
|
||||||
assert r.rewards.shape == (2, 3)
|
assert r.rewards.shape == (2, 3)
|
||||||
@@ -232,7 +283,7 @@ def test_rollout_runner_shapes(device):
|
|||||||
|
|
||||||
def test_rollout_runner_cache_returns_stale_flag(device):
|
def test_rollout_runner_cache_returns_stale_flag(device):
|
||||||
runner, _ = _make_runner(device, rollout_interval=10)
|
runner, _ = _make_runner(device, rollout_interval=10)
|
||||||
batch = _make_prompt_batch(device=device)
|
batch = _make_instruction_batch()
|
||||||
r1, fresh1 = runner(batch)
|
r1, fresh1 = runner(batch)
|
||||||
r2, fresh2 = runner(batch)
|
r2, fresh2 = runner(batch)
|
||||||
assert r1 is r2
|
assert r1 is r2
|
||||||
@@ -242,7 +293,7 @@ def test_rollout_runner_cache_returns_stale_flag(device):
|
|||||||
|
|
||||||
def test_rollout_runner_step_triggers_new_rollout(device):
|
def test_rollout_runner_step_triggers_new_rollout(device):
|
||||||
runner, _ = _make_runner(device, rollout_interval=2)
|
runner, _ = _make_runner(device, rollout_interval=2)
|
||||||
batch = _make_prompt_batch(device=device)
|
batch = _make_instruction_batch()
|
||||||
r1, fresh1 = runner(batch)
|
r1, fresh1 = runner(batch)
|
||||||
assert fresh1 is True
|
assert fresh1 is True
|
||||||
runner.step()
|
runner.step()
|
||||||
@@ -259,7 +310,7 @@ def test_rollout_runner_step_triggers_new_rollout(device):
|
|||||||
|
|
||||||
def test_rollout_runner_clear_cache_forces_rerun(device):
|
def test_rollout_runner_clear_cache_forces_rerun(device):
|
||||||
runner, _ = _make_runner(device, rollout_interval=100)
|
runner, _ = _make_runner(device, rollout_interval=100)
|
||||||
batch = _make_prompt_batch(device=device)
|
batch = _make_instruction_batch()
|
||||||
r1, _ = runner(batch)
|
r1, _ = runner(batch)
|
||||||
runner.clear_cache()
|
runner.clear_cache()
|
||||||
r2, fresh2 = runner(batch)
|
r2, fresh2 = runner(batch)
|
||||||
@@ -269,7 +320,7 @@ def test_rollout_runner_clear_cache_forces_rerun(device):
|
|||||||
|
|
||||||
def test_rollout_runner_step_resets_counter(device):
|
def test_rollout_runner_step_resets_counter(device):
|
||||||
runner, _ = _make_runner(device, rollout_interval=1)
|
runner, _ = _make_runner(device, rollout_interval=1)
|
||||||
batch = _make_prompt_batch(device=device)
|
batch = _make_instruction_batch()
|
||||||
r1, _ = runner(batch)
|
r1, _ = runner(batch)
|
||||||
runner.step()
|
runner.step()
|
||||||
r2, fresh2 = runner(batch)
|
r2, fresh2 = runner(batch)
|
||||||
|
|||||||
Reference in New Issue
Block a user