From 4a40596d02854f71ae7d771ab6ed5d0d95048fd1 Mon Sep 17 00:00:00 2001 From: Chuming Yao <1416004356qq@gmail.com> Date: Mon, 14 Sep 2026 19:29:02 +0800 Subject: [PATCH] [feature]reward_loop --- .../rl/reward_loop/bench_streaming_local.py | 1069 +++++++++++++++++ src/twinkle/reward_loop/__init__.py | 11 + src/twinkle/reward_loop/config.py | 24 + src/twinkle/reward_loop/data.py | 50 + src/twinkle/reward_loop/default_score.py | 26 + src/twinkle/reward_loop/metrics.py | 23 + src/twinkle/reward_loop/pipeline.py | 140 +++ .../reward_loop/reward_manager/__init__.py | 13 + .../reward_loop/reward_manager/base.py | 76 ++ .../reward_loop/reward_manager/dapo.py | 19 + .../reward_loop/reward_manager/gdpo.py | 15 + .../reward_loop/reward_manager/limited.py | 68 ++ .../reward_loop/reward_manager/naive.py | 7 + .../reward_loop/reward_manager/registry.py | 35 + .../reward_loop/reward_manager/remote.py | 8 + src/twinkle/reward_loop/worker.py | 56 + .../sampler/vllm_sampler/vllm_sampler.py | 66 + .../server/sampler/twinkle_handlers.py | 256 ++++ src/twinkle_client/types/__init__.py | 2 + src/twinkle_client/types/component.py | 16 + 20 files changed, 1980 insertions(+) create mode 100644 cookbook/rl/reward_loop/bench_streaming_local.py create mode 100644 src/twinkle/reward_loop/__init__.py create mode 100644 src/twinkle/reward_loop/config.py create mode 100644 src/twinkle/reward_loop/data.py create mode 100644 src/twinkle/reward_loop/default_score.py create mode 100644 src/twinkle/reward_loop/metrics.py create mode 100644 src/twinkle/reward_loop/pipeline.py create mode 100644 src/twinkle/reward_loop/reward_manager/__init__.py create mode 100644 src/twinkle/reward_loop/reward_manager/base.py create mode 100644 src/twinkle/reward_loop/reward_manager/dapo.py create mode 100644 src/twinkle/reward_loop/reward_manager/gdpo.py create mode 100644 src/twinkle/reward_loop/reward_manager/limited.py create mode 100644 src/twinkle/reward_loop/reward_manager/naive.py create mode 100644 src/twinkle/reward_loop/reward_manager/registry.py create mode 100644 src/twinkle/reward_loop/reward_manager/remote.py create mode 100644 src/twinkle/reward_loop/worker.py diff --git a/cookbook/rl/reward_loop/bench_streaming_local.py b/cookbook/rl/reward_loop/bench_streaming_local.py new file mode 100644 index 00000000..4d71a21a --- /dev/null +++ b/cookbook/rl/reward_loop/bench_streaming_local.py @@ -0,0 +1,1069 @@ +"""Benchmark: streaming sampling benefit for GRPO + reward loop (local ray mode). + +Runs in the same local process two sampling paths over the *same* prompts and +weights, then compares them: + +- Path A (batch baseline): ``sampler.sample`` returns the whole batch, then all + RewardItems are submitted at once; rewards are collected when sampling ends. +- Path B (streaming): one concurrent remote ``sample([input])`` call per + sequence; as each sequence finishes its RewardItem is submitted immediately, + so reward computation overlaps with the remaining sequences' generation + (local-mode equivalent of ``stream_sample_to_data_plane`` per-sample events). + +Correctness is verified on two levels before / during the sweep: +- Level 1 (deterministic): greedy + fixed seed, both paths must produce + identical tokens / logprobs / rewards on the same inputs. +- Level 2 (semantic): under random sampling, no dropped / duplicated / shuffled + items, reward function determinism, and group-mean-zero advantages. + +Outputs (JSONL timeline + CSV per-step summary) go to ``BENCH_OUT_DIR``. + +Environment knobs: TWINKLE_MODEL_ID / TWINKLE_DATASET_ID / TWINKLE_MODEL_GPUS / +TWINKLE_SAMPLER_GPUS / TWINKLE_LEARNING_RATE / TWINKLE_ADAPTER_NAME / +TWINKLE_REWARD_NUM_WORKERS / TWINKLE_REWARD_DELAY_MS / BENCH_RUNS (all|smoke) / +BENCH_OUT_DIR. +""" +from __future__ import annotations + +import asyncio +import csv +import itertools +import json +import os +import re +import threading +import time +from concurrent.futures import ThreadPoolExecutor, as_completed +from typing import Any, Dict, List, Optional + +from peft import LoraConfig + +import twinkle +from twinkle import DeviceMesh, DeviceGroup, get_device_placement, get_logger +from twinkle.advantage import GRPOAdvantage +from twinkle.checkpoint_engine import CheckpointEngineManager +from twinkle.data_format import SamplingParams, user_data_get +from twinkle.dataloader import DataLoader +from twinkle.dataset import Dataset, DatasetMeta +from twinkle.metric import CompletionRewardMetric +from twinkle.model import TransformersModel +from twinkle.processor import InputProcessor +from twinkle.preprocessor.base import Preprocessor +from twinkle.reward import GSM8KAccuracyReward +from twinkle.reward_loop import (AsyncRewardPipeline, RewardItem, RewardResult, + register) +from twinkle.reward_loop.reward_manager import RewardManagerBase +from twinkle.sampler import vLLMSampler + +logger = get_logger() + +# --------------------------------------------------------------------------- +# Configuration (env-overridable) +# --------------------------------------------------------------------------- +MODEL_ID = os.environ.get('TWINKLE_MODEL_ID', 'ms://Qwen/Qwen3.5-4B') +DATASET_ID = os.environ.get('TWINKLE_DATASET_ID', 'ms://modelscope/gsm8k') +# 奖励模型(judge)可独立于训练模型配置;默认跟随训练模型。 +REWARD_MODEL_ID = os.environ.get('TWINKLE_REWARD_MODEL_ID', MODEL_ID) +# Qwen3.5/3.6 are multimodal (vision tower); other models use the plain chat +# template. Override explicitly with TWINKLE_TEMPLATE_CLS when needed. +_IS_MULTIMODAL_QWEN = 'Qwen3.5' in MODEL_ID or 'Qwen3.6' in MODEL_ID +TEMPLATE_CLS = os.environ.get( + 'TWINKLE_TEMPLATE_CLS', + 'Qwen3_5Template' if _IS_MULTIMODAL_QWEN else 'Template', +) +_IS_MULTIMODAL_REWARD = 'Qwen3.5' in REWARD_MODEL_ID or 'Qwen3.6' in REWARD_MODEL_ID +REWARD_TEMPLATE_CLS = os.environ.get( + 'TWINKLE_REWARD_TEMPLATE_CLS', + 'Qwen3_5Template' if _IS_MULTIMODAL_REWARD else 'Template', +) + +MODEL_GPUS = int(os.environ.get('TWINKLE_MODEL_GPUS', '1')) +SAMPLER_GPUS = int(os.environ.get('TWINKLE_SAMPLER_GPUS', '1')) +REWARD_GPUS = int(os.environ.get('TWINKLE_REWARD_GPUS', '1')) +BENCH_RM = os.environ.get('BENCH_RM', '0') == '1' +NUM_GPUS = MODEL_GPUS + SAMPLER_GPUS + (REWARD_GPUS if BENCH_RM else 0) + +LEARNING_RATE = float(os.environ.get('TWINKLE_LEARNING_RATE', '1e-5')) +ADAPTER_NAME = os.environ.get('TWINKLE_ADAPTER_NAME', 'bench-streaming-grpo') +REWARD_NUM_WORKERS = int(os.environ.get('TWINKLE_REWARD_NUM_WORKERS', '2')) +REWARD_DELAY_MS = float(os.environ.get('TWINKLE_REWARD_DELAY_MS', '0')) +# Absolute default: repo-root/results, independent of the working directory. +_SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) +_REPO_ROOT = os.path.dirname(os.path.dirname(os.path.dirname(_SCRIPT_DIR))) +BENCH_OUT_DIR = os.environ.get('BENCH_OUT_DIR', os.path.join(_REPO_ROOT, 'results')) +TIMELINE_PATH = os.path.join(BENCH_OUT_DIR, 'bench_timeline.jsonl') +SUMMARY_PATH = os.path.join(BENCH_OUT_DIR, 'bench_summary.csv') +LEVEL1_TEXTS_PATH = os.path.join(BENCH_OUT_DIR, 'level1_texts.json') +BENCH_RUNS = os.environ.get('BENCH_RUNS', 'all') +# 'engine' = one remote call (vLLMSampler.sample_sequences_to_queue): all +# sequences scheduled concurrently in the sampler actor (vLLM keeps batching), +# completion events stream back through a Ray queue. +# 'legacy' = N concurrent per-input remote calls (serialized by the actor). +BENCH_PATH_B_STREAM = os.environ.get('BENCH_PATH_B_STREAM', 'engine') +# RM 模式:奖励提交粒度(whole=整批 / mini=每 K 条一批 / per-item=逐条)。 +BENCH_SUBMIT_GRANULARITY = os.environ.get('BENCH_SUBMIT_GRANULARITY', 'per-item') +MINI_SUBMIT_SIZE = int(os.environ.get('BENCH_MINI_BATCH_SIZE', '2')) + +BASE_STEPS = 6 +SWEEP_STEPS = 4 + +# Backlog must cover one streaming step's in-flight per-sequence handles. +MAX_TOTAL_PER_STEP = 4 * 8 # batch=4, gen=8 +REWARD_BACKLOG = MAX_TOTAL_PER_STEP + 2 + + +def build_runs() -> List[Dict[str, Any]]: + """Run matrix: one base config plus single-variable sweep points. + + ``BENCH_RUNS`` accepts 'all', 'smoke', or a comma-separated run-name list + (e.g. 'base,gen8,d1000') for targeted reruns. + """ + runs = [ + dict(name='base', batch=4, gen=4, max_tokens=1024, delay_ms=0, steps=BASE_STEPS), + dict(name='gen2', batch=4, gen=2, max_tokens=1024, delay_ms=0, steps=SWEEP_STEPS), + dict(name='gen8', batch=4, gen=8, max_tokens=1024, delay_ms=0, steps=SWEEP_STEPS), + dict(name='tok512', batch=4, gen=4, max_tokens=512, delay_ms=0, steps=SWEEP_STEPS), + dict(name='tok2048', batch=4, gen=4, max_tokens=2048, delay_ms=0, steps=SWEEP_STEPS), + dict(name='d200', batch=4, gen=4, max_tokens=1024, delay_ms=200, steps=SWEEP_STEPS), + dict(name='d1000', batch=4, gen=4, max_tokens=1024, delay_ms=1000, steps=SWEEP_STEPS), + ] + if BENCH_RM: + # RM 场景专用矩阵:路径 × 提交粒度(小规模,每步 4 条序列)。 + base = dict(batch=2, gen=2, max_tokens=BENCH_RM_MAX_TOKENS, delay_ms=0, steps=3) + return [ + dict(name='rm-whole', path='A', granularity='whole', **base), + dict(name='rm-b-whole', path='B', granularity='whole', **base), + dict(name='rm-b-mini', path='B', granularity='mini', **base), + dict(name='rm-b-per', path='B', granularity='per-item', **base), + ] + if BENCH_RUNS == 'smoke': + return [dict(name='smoke', batch=1, gen=2, max_tokens=64, delay_ms=0, steps=1)] + if BENCH_RUNS == 'all': + return runs + names = [r['name'] for r in runs] + selected = [name.strip() for name in BENCH_RUNS.split(',') if name.strip()] + unknown = [name for name in selected if name not in names] + if unknown or not selected: + raise ValueError( + f"BENCH_RUNS must be 'all', 'smoke', or a comma-separated subset of " + f"{names}; got {BENCH_RUNS!r}") + return [r for r in runs if r['name'] in selected] + + +# --------------------------------------------------------------------------- +# Timeline instrumentation (thread-safe; reward workers append from threads) +# --------------------------------------------------------------------------- +class Timeline: + """Monotonic-clock event log shared by the main thread and reward workers.""" + + def __init__(self) -> None: + self._events: List[Dict[str, Any]] = [] + self._lock = threading.Lock() + + def record(self, run: str, path: str, kind: str, step: Optional[int] = None, + idx: Optional[int] = None, item_id: Optional[str] = None, + value: Optional[float] = None) -> None: + event = { + 'ts': time.perf_counter(), 'run': run, 'path': path, 'kind': kind, + 'step': step, 'idx': idx, 'item_id': item_id, 'value': value, + } + with self._lock: + self._events.append(event) + + def finalize(self) -> None: + """Patch reward-worker events (which don't know run/path) from item_id. + + item_id format: ``{run}/{path}/step-{step}/sample-{idx}``. + """ + for event in self._events: + if event['run'] == '__run__' and event['item_id']: + parts = event['item_id'].split('/') + if len(parts) >= 4: + event['run'] = parts[0] + event['path'] = parts[1] + event['step'] = int(parts[2].split('-')[1]) if parts[2].startswith('step-') else None + event['idx'] = int(parts[3].split('-')[1]) if parts[3].startswith('sample-') else None + + def dump(self, path: str) -> None: + self.finalize() + with open(path, 'w', encoding='utf-8') as fh: + for event in self._events: + fh.write(json.dumps(event, ensure_ascii=False) + '\n') + + +TIMELINE = Timeline() + +# Current run's reward delay (ms), read by gsm8k_score running in worker threads. +_current_delay_ms: float = REWARD_DELAY_MS +_current_delay_lock = threading.Lock() + + +def set_reward_delay(delay_ms: float) -> None: + global _current_delay_ms + with _current_delay_lock: + _current_delay_ms = delay_ms + + +def get_reward_delay() -> float: + with _current_delay_lock: + return _current_delay_ms + + +# --------------------------------------------------------------------------- +# Dataset + reward scoring (same contract as minimal_grpo_local.py) +# --------------------------------------------------------------------------- +def create_dataset() -> Dataset: + """GSM8K dataset adapted from its message-format rows. + + Local jsonl files (``TWINKLE_DATASET_ID`` pointing at a file) are loaded + through the in-memory path (``DatasetMeta(data=rows)``) so loading never + touches the modelscope hub loader — guarantees offline operation even on + hosts where modelscope's loader is unavailable or network-restricted. + + Otherwise (``ms://...``) rows come from the hub and are already + ``messages`` (user question + assistant reference solution ending with + ``#### ``). ``_MessagesGSM8KPreprocessor`` extracts the ground truth + into ``user_data``, drops the reference message so it never leaks into + the sampled prompt, and prepends the ``\\boxed{}`` system prompt. Rows + stay trajectories (not ``encode()``-d). + """ + dataset = Dataset(DatasetMeta(DATASET_ID, subset_name='main', split='train')) + dataset.map(_MessagesGSM8KPreprocessor(system='Put the final answer within \\boxed{}.')) + dataset.set_template(TEMPLATE_CLS, model_id=MODEL_ID, max_length=400) + return dataset + + +class _MessagesGSM8KPreprocessor(Preprocessor): + """Adapt message-format GSM8K rows to (messages, user_data) trajectories.""" + + def __init__(self, system: str = None): + self.system = system + + def __call__(self, rows: Dict[str, List[Any]]) -> Dict[str, List[Any]]: + rows = self.map_col_to_row(rows) + rows = [self.preprocess(row) for row in rows] + return self.map_row_to_col(rows) + + def preprocess(self, row: Dict[str, Any]) -> Dict[str, Any]: + messages = list(row.get('messages') or []) + ground_truth = '' + kept = [] + for msg in messages: + if msg.get('role') == 'assistant': + # Reference solution embedded in messages; use for GT, drop + # from the prompt so it never leaks into sampling. + if not ground_truth: + ground_truth = _extract_ground_truth_from_answer(msg.get('content', '')) + else: + kept.append(msg) + if not ground_truth: + # modelscope rows carry the reference solution in a separate + # 'gold_answer' column instead of an assistant message. + ground_truth = _extract_ground_truth_from_answer(row.get('gold_answer', '')) + if self.system: + kept = [{'role': 'system', 'content': self.system}] + kept + return {'messages': kept, 'user_data': [('ground_truth', ground_truth)]} + + +def _extract_predicted_answer(completion: str) -> str: + """Extract the model's answer: \\boxed{} > #### > last number. + + ``GSM8KAccuracyReward.extract_answer`` only recognizes \\boxed{} and ####; + models without the boxed instruction emit plain text like + ``**Final Answer:** 72 clips.``, so fall back to the last number. + """ + predicted = GSM8KAccuracyReward.extract_answer(completion) + if predicted: + return predicted + tail = completion[-200:] if len(completion) > 200 else completion + numbers = re.findall(r'-?\d+(?:[.,]\d+)?', tail) + return numbers[-1].replace(',', '') if numbers else '' + + +def _numerically_equal(predicted: str, ground_truth: str) -> bool: + try: + return abs(float(predicted) - float(str(ground_truth).strip())) < 1e-5 + except (ValueError, OverflowError): + return predicted == str(ground_truth).strip() + + +def _score(data_source: str, solution_str: str, ground_truth: str, extra_info: dict): + """Pure scalar reward: 1.0 iff the extracted answer matches ground truth.""" + prompt = extra_info.get('prompt') if isinstance(extra_info, dict) else {} + prompt = dict(prompt) if isinstance(prompt, dict) else {} + messages = list(prompt.get('messages') or []) + messages.append({'role': 'assistant', 'content': solution_str}) + trajectory = {**prompt, 'messages': messages} + user_data = list(trajectory.get('user_data') or []) + if user_data_get(user_data, 'ground_truth', None) in (None, ''): + user_data.append(('ground_truth', str(ground_truth))) + trajectory['user_data'] = user_data + reward = GSM8KAccuracyReward()([trajectory])[0] + if reward == 0.0 and str(ground_truth).strip(): + predicted = _extract_predicted_answer(solution_str) + if predicted and _numerically_equal(predicted, str(ground_truth).strip()): + reward = 1.0 + return reward, {'data_source': data_source} + + +def gsm8k_score(data_source: str, solution_str: str, ground_truth: str, extra_info: dict): + """reward_loop worker entry: optional artificial delay + timeline events. + + Workers don't know run/path; Timeline.finalize() patches them from item_id. + """ + item_id = extra_info.get('_bench_item_id') if isinstance(extra_info, dict) else None + TIMELINE.record('__run__', '__path__', 'reward_start', item_id=item_id) + delay_s = get_reward_delay() / 1000.0 + if delay_s > 0: + time.sleep(delay_s) + score, meta = _score(data_source, solution_str, ground_truth, extra_info) + TIMELINE.record('__run__', '__path__', 'reward_end', item_id=item_id) + return score, meta + + +# --------------------------------------------------------------------------- +# RM(奖励模型)模式:生成式 judge 打分(独立 GPU 批级调用) +# --------------------------------------------------------------------------- +# 边界验证(长尾生成 + 昂贵 judge):扩大判词长度让 judge 先推理再判定, +# 单条延迟升至秒级;RM 采样 max_tokens 由 BENCH_RM_MAX_TOKENS 控制。 +_JUDGE_MAX_TOKENS = int(os.environ.get('TWINKLE_REWARD_JUDGE_MAX_TOKENS', '8')) +BENCH_RM_MAX_TOKENS = int(os.environ.get('BENCH_RM_MAX_TOKENS', '512')) +_JUDGE_SYSTEM = ('You are a strict math answer verifier. Reason briefly about ' + 'whether the model answer matches the ground truth, then end ' + 'your response with exactly one word on the last line: Correct ' + 'or Incorrect.') +_JUDGE_PARAMS = SamplingParams(max_tokens=_JUDGE_MAX_TOKENS, num_samples=1, + temperature=0.0) + + +_JUDGE_MAX_SOLUTION_CHARS = 3000 + + +def judge_prompt_for(item: RewardItem) -> Dict[str, Any]: + """Verification trajectory (dict with ``messages``): question + model + answer + ground truth — the shape ``sampler.sample`` expects. + + Long model answers are truncated to their tail (the final answer region), + so judge prompts stay within the judge engine's context window. + """ + prompt = item.extra_info.get('prompt') if isinstance(item.extra_info, dict) else {} + question = '' + for msg in reversed(list((prompt or {}).get('messages') or [])): + if msg.get('role') == 'user': + question = msg.get('content', '') + break + solution = item.solution_str or '' + if len(solution) > _JUDGE_MAX_SOLUTION_CHARS: + solution = '...[truncated, showing tail]...\n' + solution[-_JUDGE_MAX_SOLUTION_CHARS:] + return {'messages': [ + {'role': 'system', 'content': _JUDGE_SYSTEM}, + {'role': 'user', 'content': + f'Question: {question}\n\nModel answer: {solution}\n\n' + f'Ground truth: {item.ground_truth}\n\nIs the model answer correct?'}, + ]} + + +def parse_judge_verdict(text: str) -> Optional[float]: + if re.search(r'\bcorrect\b', text, re.IGNORECASE): + return 1.0 + if re.search(r'\bincorrect\b', text, re.IGNORECASE): + return 0.0 + return None + + +@register('batch_judge') +class BatchJudgeRewardManager(RewardManagerBase): + """Reward manager that scores a whole chunk with one judge engine call. + + Items in a chunk are packed into one ``judge_sampler.sample(prompts)`` + call so the judge's vLLM batches them (real RM batching), then each + verdict is parsed per item. Chunk size = reward submission granularity + (whole / mini / per-item), so the granularity experiment controls exactly + how many trajectories the judge sees per engine call. + """ + + def __init__(self, compute_score=None, judge_sampler=None, **kwargs): + super().__init__(compute_score=compute_score, **kwargs) + self.judge_sampler = judge_sampler + + async def run_batch(self, items): + if not items: + return [] + prompts = [judge_prompt_for(item) for item in items] + responses = await asyncio.to_thread( + self.judge_sampler.sample, prompts, _JUDGE_PARAMS, '') + results = [] + for item, response in zip(items, responses): + text = response.sequences[0].decoded or '' + score = parse_judge_verdict(text) + if score is None: + score = 0.0 + logger.warning(f'[judge] unparsed verdict {text[:80]!r} for {item.item_id}') + results.append(RewardResult( + item.item_id, score, {'judge_verdict': text.strip()[:60]})) + return results + + +_ANS_RE = re.compile(r'####\s*(-?\d+(?:[.,]\d+)?)') + + +def _extract_ground_truth_from_answer(answer: Any) -> str: + """Extract the final numeric answer from a raw GSM8K row's ``answer`` field. + + Raw rows (no GSM8KProcessor) carry the full solution in ``answer`` ending + with ``#### `` (huggingface-style); fall back to the last number. + """ + if not answer: + return '' + text = str(answer) + match = _ANS_RE.search(text) + if match: + return match.group(1).replace(',', '') + numbers = re.findall(r'-?\d+(?:[.,]\d+)?', text) + return numbers[-1].replace(',', '') if numbers else '' + + +def _ground_truth(prompt: Dict[str, Any]) -> str: + gt = user_data_get(prompt.get('user_data'), 'ground_truth', '') + if gt in (None, ''): + gt = _extract_ground_truth_from_answer(prompt.get('answer', '')) + return str(gt) + + +def make_reward_items(run: str, path: str, prompts: List[Dict[str, Any]], + sequences: List[Any], step: int, num_generations: int) -> List[RewardItem]: + """One RewardItem per completed sequence, in prompt-major index order. + + ``sequences`` is a list of (index, SampledSequence) aligned to the expanded + prompt copies (index // num_generations -> prompt). + """ + items: List[RewardItem] = [] + for idx, sequence in sequences: + prompt = prompts[idx // num_generations] + item_id = f'{run}/{path}/step-{step}/sample-{idx}' + items.append(RewardItem( + item_id=item_id, + data_source='gsm8k', + solution_str=sequence.decoded or '', + ground_truth=_ground_truth(prompt), + extra_info={'prompt': prompt, '_bench_item_id': item_id}, + )) + return items + + +# --------------------------------------------------------------------------- +# Sampling paths +# --------------------------------------------------------------------------- +def _expand(prompts: List[Dict[str, Any]], num_generations: int) -> List[Dict[str, Any]]: + return [prompt for prompt in prompts for _ in range(num_generations)] + + +def _seq_tokens(sequence) -> List[int]: + return list(sequence.tokens) + + +def _seq_logprobs(sequence) -> List[float]: + return [logprob[0][1] for logprob in sequence.logprobs] + + +def _seq_input_feature(sequence): + return sequence.new_input_feature + + +def _collect_payload(sequence) -> tuple: + return (_seq_input_feature(sequence), _seq_logprobs(sequence), len(sequence.tokens)) + + +def sample_batch(sampler, prompts: List[Dict[str, Any]], params: SamplingParams, + num_generations: int) -> List[Any]: + """Path A: single batch call; returns one SampledSequence per copy.""" + responses = sampler.sample(_expand(prompts, num_generations), params, ADAPTER_NAME) + return [resp.sequences[0] for resp in responses] + + +def sample_stream(sampler, prompts: List[Dict[str, Any]], params: SamplingParams, + num_generations: int): + """Path B: per-sequence completion events with engine-level concurrency. + + 'engine' (default): one remote call + (``vLLMSampler.sample_sequences_to_queue``) schedules ALL sequences in the + sampler actor's event loop — vLLM keeps batching the whole batch, so + t_sample stays at the batch level — and streams ``(index, SampleResponse)`` + events back through a Ray queue in completion order (local-mode counterpart + of the server's ``stream_sample_to_data_plane``). + + 'legacy': N concurrent per-input remote calls, serialized by the actor + (~N x single-sequence time; kept for comparison). + """ + expanded = _expand(prompts, num_generations) + if BENCH_PATH_B_STREAM == 'legacy': + with ThreadPoolExecutor(max_workers=len(expanded)) as pool: + futures = {pool.submit(sampler.sample, [traj], params, ADAPTER_NAME): idx + for idx, traj in enumerate(expanded)} + for future in as_completed(futures): + idx = futures[future] + response = future.result()[0] + yield idx, response.sequences[0] + return + if BENCH_PATH_B_STREAM != 'engine': + raise ValueError(f"BENCH_PATH_B_STREAM must be 'engine' or 'legacy', got {BENCH_PATH_B_STREAM!r}") + + import queue as stdlib_queue + import ray + from ray.util.queue import Queue + queue = Queue() + with ThreadPoolExecutor(max_workers=1) as pool: + remote = pool.submit( + sampler.sample_sequences_to_queue, queue, expanded, params, ADAPTER_NAME) + try: + expected = len(expanded) + received = 0 + while True: + try: + idx, response = queue.get(timeout=1.0) + except stdlib_queue.Empty: + if remote.done(): + # Engine side finished (or failed) without a sentinel. + remote.result() # re-raises engine-side errors + raise RuntimeError( + f'stream ended early: {received}/{expected} events, no sentinel') + continue + if idx is None: + break + received += 1 + yield idx, response.sequences[0] + finally: + # Drain complete; propagate any engine-side error. + remote.result() + + +# --------------------------------------------------------------------------- +# Training step (same as minimal_grpo_local.py) +# --------------------------------------------------------------------------- +def train_batch(*, model, advantage_fn, metrics, input_data, old_logps, + completion_lengths, rewards, num_generations, micro_batch_size=2) -> None: + advantages = advantage_fn(rewards, num_generations=num_generations, scale='group').tolist() + metrics.accumulate(completion_lengths=completion_lengths, rewards={'total': rewards}) + total = len(input_data) + for mb_start in range(0, total, micro_batch_size): + mb_end = min(mb_start + micro_batch_size, total) + model.forward_backward( + inputs=input_data[mb_start:mb_end], + old_logps=old_logps[mb_start:mb_end], + advantages=advantages[mb_start:mb_end], + micro_batch_size=micro_batch_size, + ) + model.clip_grad_and_step() + log_dict = metrics.calculate() + log_dict.update(model.calculate_metric(is_training=True)) + return advantages, log_dict + + +# --------------------------------------------------------------------------- +# Correctness checks +# --------------------------------------------------------------------------- +def check_semantics(run: str, path: str, step: int, prompts, items, results, + rewards, advantages, num_generations) -> Dict[str, Any]: + """Level-2 structural checks; returns a dict of pass/fail booleans.""" + total = len(prompts) * num_generations + checks: Dict[str, Any] = {} + checks['item_count'] = len(items) == total + ids = [item.item_id for item in items] + checks['no_duplicate_ids'] = len(set(ids)) == len(ids) + checks['results_aligned'] = len(results) == len(items) and {r.item_id for r in results} == set(ids) + checks['reward_count'] = len(rewards) == total + checks['advantage_count'] = len(advantages) == total + # Group-mean-zero: GRPOAdvantage normalizes each group of num_generations. + groups_ok = True + for g in range(len(prompts)): + group = advantages[g * num_generations:(g + 1) * num_generations] + if abs(sum(group)) > 1e-4: + groups_ok = False + checks['advantage_group_mean_zero'] = groups_ok + # Reward-function determinism: re-score the first item without the pipeline. + if items: + first = items[0] + direct, _ = _score('gsm8k', first.solution_str, first.ground_truth, first.extra_info) + checks['reward_deterministic'] = abs(direct - results[0].reward_score) < 1e-9 + ok = all(checks.values()) + checks['all_ok'] = ok + TIMELINE.record(run, path, 'semantic_checks', step=step, + value=1.0 if ok else 0.0, item_id=json.dumps(checks)) + if not ok: + logger.warning(f'[semantic checks failed] run={run} path={path} step={step}: {checks}') + return checks + + +def run_level1(sampler, pipeline, prompts, num_generations, max_tokens, run='level1'): + """Determinism: greedy + fixed seed, both paths on identical inputs, no + training in between. Returns a summary dict of per-check results.""" + params = SamplingParams(max_tokens=max_tokens, num_samples=1, logprobs=1, + temperature=0.0, seed=0) + checks: Dict[str, Any] = {'run': run} + path_a = list(sample_batch(sampler, prompts, params, num_generations)) + # Streaming yields in completion order; sort by input index for comparison. + streamed = sorted(sample_stream(sampler, prompts, params, num_generations), key=lambda p: p[0]) + path_b = [seq for _, seq in streamed] + checks['sample_count_equal'] = len(path_a) == len(path_b) == len(prompts) * num_generations + n = min(len(path_a), len(path_b)) + + max_token_diff = 0 + token_mismatches = 0 + max_logprob_diff = 0.0 + logprob_mismatches = 0 + decoded_mismatches = 0 + text_pairs = [] + for i in range(n): + ta, tb = _seq_tokens(path_a[i]), _seq_tokens(path_b[i]) + first_diff = None + if ta != tb: + token_mismatches += 1 + shorter = min(len(ta), len(tb)) + for j in range(shorter): + if ta[j] != tb[j]: + first_diff = j + break + if first_diff is None: + first_diff = shorter # prefix identical, lengths differ + max_token_diff = max(max_token_diff, abs(len(ta) - len(tb))) + text_pairs.append({ + 'idx': i, + 'tokens_identical': ta == tb, + 'first_diff_pos': first_diff, + 'len_a': len(ta), + 'len_b': len(tb), + 'a_text': path_a[i].decoded or '', + 'b_text': path_b[i].decoded or '', + }) + la, lb = _seq_logprobs(path_a[i]), _seq_logprobs(path_b[i]) + if la != lb: + logprob_mismatches += 1 + shorter = min(len(la), len(lb)) + if shorter: + diffs = [abs(x - y) for x, y in zip(la[:shorter], lb[:shorter])] + max_logprob_diff = max(max_logprob_diff, max(diffs)) + max_logprob_diff = max(max_logprob_diff, abs(len(la) - len(lb))) + if (path_a[i].decoded or '') != (path_b[i].decoded or ''): + decoded_mismatches += 1 + checks['tokens_identical'] = token_mismatches == 0 + checks['token_mismatches'] = token_mismatches + checks['max_token_len_diff'] = max_token_diff + checks['logprobs_identical'] = logprob_mismatches == 0 + checks['logprob_mismatches'] = logprob_mismatches + checks['max_logprob_diff'] = max_logprob_diff + checks['decoded_identical'] = decoded_mismatches == 0 + checks['decoded_mismatches'] = decoded_mismatches + + # Reward path must also agree deterministically through the pipeline. + def _scored_items(sequences, path_label): + seqs = [(i, s) for i, s in enumerate(sequences)] + return make_reward_items(run, path_label, prompts, seqs, 0, num_generations) + items_a = _scored_items(path_a, 'A') + items_b = _scored_items(path_b, 'B') + results_a = pipeline.collect(pipeline.submit(items_a)) + results_b = pipeline.collect(pipeline.submit(items_b)) + rewards_a = [r.reward_score for r in results_a] + rewards_b = [r.reward_score for r in results_b] + checks['reward_count_equal'] = len(rewards_a) == len(rewards_b) == n + checks['rewards_identical'] = rewards_a == rewards_b + if rewards_a == rewards_b and rewards_a: + adv_a = GRPOAdvantage()(rewards_a, num_generations=num_generations, scale='group').tolist() + adv_b = GRPOAdvantage()(rewards_b, num_generations=num_generations, scale='group').tolist() + checks['advantages_identical'] = adv_a == adv_b + else: + checks['advantages_identical'] = False + checks['all_ok'] = all( + v is True for k, v in checks.items() if k not in ('run',) and isinstance(v, bool)) + TIMELINE.record(run, 'both', 'level1_checks', value=1.0 if checks['all_ok'] else 0.0, + item_id=json.dumps(checks)) + logger.info(f'[Level-1 determinism] {json.dumps(checks, ensure_ascii=False)}') + + # Dump per-pair texts so divergent pairs can be inspected by hand. + _write_level1_texts(prompts, text_pairs, rewards_a, rewards_b, num_generations, checks) + return checks + + +def _write_level1_texts(prompts, text_pairs, rewards_a, rewards_b, num_generations, checks) -> None: + """Write A/B text pairs, rewards and first divergence position per index.""" + for i, pair in enumerate(text_pairs): + prompt = prompts[i // num_generations] + pair['ground_truth'] = _ground_truth(prompt) + pair['reward_a'] = rewards_a[i] if i < len(rewards_a) else None + pair['reward_b'] = rewards_b[i] if i < len(rewards_b) else None + record = { + 'num_pairs': len(text_pairs), + 'num_generations': num_generations, + 'level1_checks': {k: v for k, v in checks.items() if k != 'run'}, + 'pairs': text_pairs, + } + with open(LEVEL1_TEXTS_PATH, 'w', encoding='utf-8') as fh: + json.dump(record, fh, ensure_ascii=False, indent=2) + logger.info(f'[Level-1] per-pair texts written to {LEVEL1_TEXTS_PATH}') + + +# --------------------------------------------------------------------------- +# Per-path step loop +# --------------------------------------------------------------------------- +def run_path(run_cfg: Dict[str, Any], path: str, sampler, pipeline, model, + advantage_fn, metrics, batches: List[List[Dict[str, Any]]], + sync_weights) -> List[Dict[str, Any]]: + """Run one path over the given batches; returns per-step summary rows.""" + run = run_cfg['name'] + num_generations = run_cfg['gen'] + # 提交粒度:RM 模式由 run 矩阵决定;普通模式 B 保持逐条(流式语义)。 + granularity = run_cfg.get('granularity') or ( + 'per-item' if path == 'B' else 'whole') + params = SamplingParams(max_tokens=run_cfg['max_tokens'], num_samples=1, + logprobs=1, temperature=1.0, top_p=0.95) + rows: List[Dict[str, Any]] = [] + metrics.reset() + for step, prompts in enumerate(batches): + TIMELINE.record(run, path, 'step_start', step=step) + t_step0 = time.perf_counter() + sync_weights() + sampler.reset_prefix_cache() + + if path == 'A': + t0 = time.perf_counter() + TIMELINE.record(run, path, 'sample_start', step=step) + sequences = sample_batch(sampler, prompts, params, num_generations) + t_sample = time.perf_counter() - t0 + TIMELINE.record(run, path, 'sample_end', step=step, value=t_sample) + seqs = [(i, s) for i, s in enumerate(sequences)] + else: + # Streaming: submit rewards as sequences complete; the submission + # granularity controls how many rewards ride in one handle + # (per-item = immediate, mini = every K, whole = after sampling). + t0 = time.perf_counter() + TIMELINE.record(run, path, 'sample_start', step=step) + seqs, items, handles = [], [], [] + pending: List[RewardItem] = [] + submit_threshold = 1 if granularity == 'per-item' else ( + MINI_SUBMIT_SIZE if granularity == 'mini' else len(prompts) * num_generations) + + def _flush_pending(): + if not pending: + return + if len(pending) < submit_threshold: + return + _items, pending[:] = pending[:], [] + handles.append(pipeline.submit(_items)) + + t_submit = 0.0 + for idx, sequence in sample_stream(sampler, prompts, params, num_generations): + TIMELINE.record(run, path, 'sample_done', step=step, idx=idx) + seqs.append((idx, sequence)) + item = make_reward_items(run, path, prompts, [(idx, sequence)], + step, num_generations)[0] + items.append(item) + pending.append(item) + t1 = time.perf_counter() + _flush_pending() + t_submit += time.perf_counter() - t1 + if pending: + handles.append(pipeline.submit(pending)) + t_sample = time.perf_counter() - t0 + TIMELINE.record(run, path, 'sample_end', step=step, value=t_sample) + + # Submit remaining rewards (A: one batch call after sampling). + if path == 'A': + t0 = time.perf_counter() + items = make_reward_items(run, path, prompts, seqs, step, num_generations) + TIMELINE.record(run, path, 'submit_start', step=step) + handle = pipeline.submit(items) + TIMELINE.record(run, path, 'submit_end', step=step) + t_submit = time.perf_counter() - t0 + handles = [handle] + + # Collect all rewards for this step (single-buffer schedule). + t0 = time.perf_counter() + TIMELINE.record(run, path, 'collect_start', step=step) + results = [] + for handle in handles: + results.extend(pipeline.collect(handle)) + t_collect = time.perf_counter() - t0 + TIMELINE.record(run, path, 'collect_end', step=step, value=t_collect) + by_id = {r.item_id: r for r in results} + rewards = [by_id[item.item_id].reward_score for item in items] + + # Train. + t0 = time.perf_counter() + TIMELINE.record(run, path, 'train_start', step=step) + payloads = [_collect_payload(seq) for _, seq in seqs] + advantages, log_dict = train_batch( + model=model, advantage_fn=advantage_fn, metrics=metrics, + input_data=[p[0] for p in payloads], old_logps=[p[1] for p in payloads], + completion_lengths=[p[2] for p in payloads], rewards=rewards, + num_generations=num_generations, + ) + t_train = time.perf_counter() - t0 + TIMELINE.record(run, path, 'train_end', step=step, value=t_train) + + checks = check_semantics(run, path, step, prompts, items, results, rewards, + advantages, num_generations) + + t_total = time.perf_counter() - t_step0 + TIMELINE.record(run, path, 'step_end', step=step, value=t_total) + row = dict(run=run, path=path, step=step, batch=len(prompts), + gen=num_generations, max_tokens=run_cfg['max_tokens'], + delay_ms=run_cfg['delay_ms'], + t_sample=t_sample, t_submit=t_submit, t_collect=t_collect, + t_train=t_train, t_total=t_total, + seq_per_s=(len(prompts) * num_generations) / t_total, + semantic_ok=checks['all_ok']) + rows.append(row) + logger.info(f"[{run}/{path}] step={step} t_sample={t_sample:.2f}s " + f"t_submit={t_submit:.2f}s t_collect={t_collect:.2f}s " + f"t_train={t_train:.2f}s t_total={t_total:.2f}s " + f"reward_mean={sum(rewards) / len(rewards):.4f} {log_dict}") + return rows + + +# --------------------------------------------------------------------------- +# Main +# --------------------------------------------------------------------------- +def main() -> None: + runs = build_runs() + os.makedirs(BENCH_OUT_DIR, exist_ok=True) + + _sampler_start = MODEL_GPUS + _reward_start = MODEL_GPUS + SAMPLER_GPUS + device_groups = [ + DeviceGroup(name='model', ranks=list(range(MODEL_GPUS)), device_type='GPU'), + DeviceGroup(name='sampler', ranks=list(range(_sampler_start, _reward_start)), + device_type='GPU'), + ] + if BENCH_RM: + device_groups.append(DeviceGroup( + name='reward', ranks=list(range(_reward_start, NUM_GPUS)), device_type='GPU')) + model_mesh = DeviceMesh.from_sizes(world_size=MODEL_GPUS, dp_size=MODEL_GPUS) + sampler_mesh = DeviceMesh.from_sizes(world_size=SAMPLER_GPUS, dp_size=SAMPLER_GPUS) + twinkle.initialize(mode='ray', nproc_per_node=NUM_GPUS, groups=device_groups, + lazy_collect=False) + + lora_config = LoraConfig( + target_modules=[ + 'q_proj', 'k_proj', 'v_proj', 'o_proj', + 'gate_proj', 'up_proj', 'down_proj', + 'in_proj_qkv', 'in_proj_z', 'in_proj_a', 'in_proj_b', 'out_proj', + ], + r=32, lora_alpha=64, lora_dropout=0.05, + ) + model = TransformersModel(model_id=MODEL_ID, device_mesh=model_mesh, remote_group='model') + model.add_adapter_to_model(ADAPTER_NAME, lora_config, gradient_accumulation_steps=1) + model.set_optimizer('AdamW', lr=LEARNING_RATE) + model.set_lr_scheduler('CosineAnnealingLR', T_max=200, eta_min=0) + model.set_loss('GRPOLoss', epsilon=0.2) + model.set_processor(InputProcessor) + model.set_template(TEMPLATE_CLS, model_id=MODEL_ID) + + sampler = vLLMSampler( + model_id=MODEL_ID, + engine_args={ + 'gpu_memory_utilization': 0.8, + 'max_model_len': 4496, + 'max_lora_rank': 32, + 'enable_lora': True, + }, + device_mesh=sampler_mesh, + remote_group='sampler', + ) + sampler.set_template(TEMPLATE_CLS, model_id=MODEL_ID) + + judge_sampler = None + if BENCH_RM: + # 独立 GPU 上的 judge 引擎(冻结权重,不参与训练更新)。 + reward_mesh = DeviceMesh.from_sizes(world_size=REWARD_GPUS, dp_size=REWARD_GPUS) + judge_sampler = vLLMSampler( + model_id=REWARD_MODEL_ID, + engine_args={ + 'gpu_memory_utilization': 0.8, + 'max_model_len': 8192, + }, + device_mesh=reward_mesh, + remote_group='reward', + ) + judge_sampler.set_template(REWARD_TEMPLATE_CLS, model_id=REWARD_MODEL_ID) + + ckpt_manager = CheckpointEngineManager(model=model, sampler=sampler) + advantage_fn = GRPOAdvantage() + metrics = CompletionRewardMetric() + if BENCH_RM: + pipeline = AsyncRewardPipeline( + num_workers=REWARD_NUM_WORKERS, + mode='async', + backlog=REWARD_BACKLOG, + worker_kwargs={ + 'compute_score': gsm8k_score, + 'manager_name': 'batch_judge', + 'reward_kwargs': {'judge_sampler': judge_sampler}, + }, + ) + else: + pipeline = AsyncRewardPipeline( + num_workers=REWARD_NUM_WORKERS, + mode='async', + backlog=REWARD_BACKLOG, + worker_kwargs={'compute_score': gsm8k_score}, + ) + + def sync_weights(): + ckpt_manager.sync_weights(merge_and_sync=False) + + logger.info(get_device_placement()) + if BENCH_RM: + logger.info(f'[bench] ** RM MODE ENABLED ** judge={REWARD_MODEL_ID} ' + f'reward_gpus={REWARD_GPUS} granularity-matrix=' + f'{BENCH_SUBMIT_GRANULARITY}') + else: + logger.info(f'[bench] RM mode disabled (add BENCH_RM=1 for reward-model runs)') + logger.info(f'[bench] outputs -> {BENCH_OUT_DIR}') + summary_rows: List[Dict[str, Any]] = [] + try: + if not BENCH_RM: + first = runs[0] + base_batch = max(1, first['batch']) + # Level-1 determinism check on one batch, before any training. + # instance_id keeps Ray actor names unique: remote_class derives the + # actor name from the caller's source line, so two DataLoaders created + # on the same line would collide with ActorAlreadyExistsError. + dataloader = DataLoader( + dataset=create_dataset, batch_size=base_batch, min_batch_size=base_batch, + device_mesh=model_mesh, remote_group='model', instance_id='bench-level1', + ) + one_batch = next(iter(dataloader)) + prompts = list(one_batch) if isinstance(one_batch, list) else [one_batch] + # Diagnostic: confirm the row schema and that ground truths resolve. + row0 = prompts[0] + preview = json.dumps(row0, ensure_ascii=False, default=str) + logger.info(f'[level1 row0] keys={sorted(row0.keys())} ' + f'user_data={row0.get("user_data")!r} preview={preview[:300]!r}') + for i, p in enumerate(prompts[:3]): + logger.info(f'[level1 prompt {i}] gt={_ground_truth(p)!r} ' + f'answer={str(p.get("answer", ""))[:60]!r}') + sync_weights() + sampler.reset_prefix_cache() + level1 = run_level1(sampler, pipeline, prompts, first['gen'], first['max_tokens']) + summary_rows.append(dict(run='level1', path='both', step=-1, batch=len(prompts), + gen=first['gen'], max_tokens=first['max_tokens'], + delay_ms=0, t_sample=0.0, t_submit=0.0, t_collect=0.0, + t_train=0.0, t_total=0.0, seq_per_s=0.0, + semantic_ok=level1['all_ok'])) + # Persist level-1 results even if a sweep run fails afterwards. + flush_outputs(summary_rows) + + for run_cfg in runs: + run = run_cfg['name'] + set_reward_delay(run_cfg['delay_ms']) + dataloader = DataLoader( + dataset=create_dataset, batch_size=run_cfg['batch'], + min_batch_size=run_cfg['batch'], + device_mesh=model_mesh, remote_group='model', + instance_id=f'bench-{run}', + ) + batches = list(itertools.islice(iter(dataloader), run_cfg['steps'])) + batches = [list(b) if isinstance(b, list) else [b] for b in batches] + logger.info(f'[{run}] materialized {len(batches)} batches x {run_cfg["batch"]} prompts ' + f'(gen={run_cfg["gen"]}, max_tokens={run_cfg["max_tokens"]}, ' + f'delay={run_cfg["delay_ms"]}ms, granularity={run_cfg.get("granularity", "-")})') + paths = [run_cfg['path']] if run_cfg.get('path') else ('A', 'B') + for path in paths: + summary_rows.extend(run_path( + run_cfg, path, sampler, pipeline, model, advantage_fn, metrics, + batches, sync_weights)) + # Persist after every run so a crash keeps all completed runs. + flush_outputs(summary_rows) + except BaseException: + logger.exception('[bench] run failed; flushing partial results before re-raising') + flush_outputs(summary_rows) + raise + finally: + pipeline.close() + + logger.info(f'[bench] final outputs: {TIMELINE_PATH}, {SUMMARY_PATH}') + + # Compact console summary: per-run-path means over steps. + logger.info('=== benchmark summary (per-run path means) ===') + means: Dict[str, Dict[str, float]] = {} + for row in summary_rows: + if row['step'] < 0: + continue + key = f"{row['run']}/{row['path']}" + bucket = means.setdefault(key, {k: 0.0 for k in + ('t_sample', 't_submit', 't_collect', 't_train', + 't_total', 'seq_per_s', 'reward_head_start', + 'reward_tail_after_sample', 'n')}) + for k in ('t_sample', 't_submit', 't_collect', 't_train', 't_total', + 'seq_per_s', 'reward_head_start', 'reward_tail_after_sample'): + if isinstance(row.get(k), (int, float)): + bucket[k] += row[k] + bucket['n'] += 1 + for key, bucket in means.items(): + n = bucket.pop('n') + if n: + means[key] = {k: v / n for k, v in bucket.items()} + logger.info(f"{key}: " + ' '.join(f'{k}={v:.3f}' for k, v in means[key].items())) + + +def flush_outputs(summary_rows: List[Dict[str, Any]]) -> None: + """Write timeline JSONL + summary CSV with data collected so far. + + Safe to call repeatedly (after each run) and from the failure handler: + ``Timeline.dump`` patches reward-worker events in place, and the CSV is + fully rewritten from ``summary_rows`` each time. + """ + TIMELINE.dump(TIMELINE_PATH) + overlap = _aggregate_reward_overlap() + with open(SUMMARY_PATH, 'w', newline='', encoding='utf-8') as fh: + fieldnames = ['run', 'path', 'step', 'batch', 'gen', 'max_tokens', 'delay_ms', + 't_sample', 't_submit', 't_collect', 't_train', 't_total', + 'seq_per_s', 'reward_head_start', 'reward_tail_after_sample', + 'semantic_ok'] + writer = csv.DictWriter(fh, fieldnames=fieldnames) + writer.writeheader() + for row in summary_rows: + key = (row['run'], row['path'], row['step']) + if key in overlap: + row['reward_head_start'] = overlap[key][0] + row['reward_tail_after_sample'] = overlap[key][1] + else: + row['reward_head_start'] = '' + row['reward_tail_after_sample'] = '' + writer.writerow(row) + logger.info(f'[bench] flushed {len(summary_rows)} summary rows, ' + f'{len(TIMELINE._events)} timeline events -> {BENCH_OUT_DIR}') + + +def _aggregate_reward_overlap() -> Dict[tuple, tuple]: + """Per (run, path, step) -> (reward_head_start, reward_tail_after_sample). + + reward_head_start: time from sampling start to the first reward computation. + reward_tail_after_sample: time from sampling end to the last reward ready. + Both in seconds; for streaming both shrink as rewards overlap with sampling. + """ + starts: Dict[tuple, float] = {} + ends: Dict[tuple, float] = {} + rewards: Dict[tuple, List[float]] = {} + for event in TIMELINE._events: + key = (event['run'], event['path'], event['step']) + if key[0] == '__run__' or key[0] == 'level1': + continue + if event['kind'] == 'sample_start': + starts.setdefault(key, event['ts']) + elif event['kind'] == 'sample_end': + ends.setdefault(key, event['ts']) + elif event['kind'] == 'reward_start': + rewards.setdefault(key, []).append(event['ts']) + elif event['kind'] == 'reward_end': + rewards.setdefault(key, []).append(event['ts']) + result: Dict[tuple, tuple] = {} + for key in set(starts) & set(rewards): + r_ts = sorted(rewards[key]) + head = (r_ts[0] - starts[key]) if len(r_ts) >= 2 else 0.0 + tail = (r_ts[-1] - ends.get(key, r_ts[-1])) if len(r_ts) >= 2 else 0.0 + result[key] = (head, tail) + return result + + +if __name__ == '__main__': + main() diff --git a/src/twinkle/reward_loop/__init__.py b/src/twinkle/reward_loop/__init__.py new file mode 100644 index 00000000..3359d3ed --- /dev/null +++ b/src/twinkle/reward_loop/__init__.py @@ -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"] diff --git a/src/twinkle/reward_loop/config.py b/src/twinkle/reward_loop/config.py new file mode 100644 index 00000000..fd380e34 --- /dev/null +++ b/src/twinkle/reward_loop/config.py @@ -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) diff --git a/src/twinkle/reward_loop/data.py b/src/twinkle/reward_loop/data.py new file mode 100644 index 00000000..cabea6ee --- /dev/null +++ b/src/twinkle/reward_loop/data.py @@ -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] diff --git a/src/twinkle/reward_loop/default_score.py b/src/twinkle/reward_loop/default_score.py new file mode 100644 index 00000000..edd71923 --- /dev/null +++ b/src/twinkle/reward_loop/default_score.py @@ -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 {}) diff --git a/src/twinkle/reward_loop/metrics.py b/src/twinkle/reward_loop/metrics.py new file mode 100644 index 00000000..5d8a6f4d --- /dev/null +++ b/src/twinkle/reward_loop/metrics.py @@ -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) diff --git a/src/twinkle/reward_loop/pipeline.py b/src/twinkle/reward_loop/pipeline.py new file mode 100644 index 00000000..64d0c5d0 --- /dev/null +++ b/src/twinkle/reward_loop/pipeline.py @@ -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) diff --git a/src/twinkle/reward_loop/reward_manager/__init__.py b/src/twinkle/reward_loop/reward_manager/__init__.py new file mode 100644 index 00000000..82079ab8 --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/__init__.py @@ -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"] diff --git a/src/twinkle/reward_loop/reward_manager/base.py b/src/twinkle/reward_loop/reward_manager/base.py new file mode 100644 index 00000000..909c1b9c --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/base.py @@ -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 diff --git a/src/twinkle/reward_loop/reward_manager/dapo.py b/src/twinkle/reward_loop/reward_manager/dapo.py new file mode 100644 index 00000000..f98d6884 --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/dapo.py @@ -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 diff --git a/src/twinkle/reward_loop/reward_manager/gdpo.py b/src/twinkle/reward_loop/reward_manager/gdpo.py new file mode 100644 index 00000000..5359c091 --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/gdpo.py @@ -0,0 +1,15 @@ +from .base import RewardManagerBase +from .registry import register + + +@register("gdpo") +class GDPORewardManager(RewardManagerBase): + def __init__(self, *args, experiment_name=None, **kwargs): + super().__init__(*args, **kwargs) + self.experiment_name = experiment_name + + async def run_single(self, item): + result = await super().run_single(item) + if self.experiment_name: + result.reward_extra_info.setdefault("experiment_name", self.experiment_name) + return result diff --git a/src/twinkle/reward_loop/reward_manager/limited.py b/src/twinkle/reward_loop/reward_manager/limited.py new file mode 100644 index 00000000..a648312e --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/limited.py @@ -0,0 +1,68 @@ +import asyncio +import time +from collections.abc import Callable +from typing import Optional + +from .base import RewardManagerBase +from .registry import register + + +class AsyncTokenBucket: + def __init__(self, rate=None, capacity=None): + self.rate = float(rate or 0) + self.capacity = float(capacity if capacity is not None else (rate or 1)) + if self.rate < 0 or self.capacity <= 0: + raise ValueError("token bucket rate must be non-negative and capacity positive") + self.tokens = self.capacity + self.updated = time.monotonic() + self._lock = asyncio.Lock() + + async def acquire(self, amount=1): + amount = float(amount) + if amount < 0: + raise ValueError("token amount must be non-negative") + if not self.rate or amount == 0: + return + if amount > self.capacity: + raise ValueError(f"token request {amount:g} exceeds bucket capacity {self.capacity:g}") + while True: + async with self._lock: + now = time.monotonic() + self.tokens = min(self.capacity, self.tokens + (now - self.updated) * self.rate) + self.updated = now + if self.tokens >= amount: + self.tokens -= amount + return + wait = (amount - self.tokens) / self.rate + await asyncio.sleep(wait) + + +@register("rate_limited") +class RateLimitedRewardManager(RewardManagerBase): + def __init__(self, *args, max_rpm=None, max_tpm=None, max_concurrent=1, timeout=300.0, + token_counter: Optional[Callable] = None, fallback_on_error=False, **kwargs): + if max_concurrent < 1: + raise ValueError("max_concurrent must be positive") + if timeout <= 0: + raise ValueError("timeout must be positive") + super().__init__(*args, max_concurrent=max_concurrent, **kwargs) + self.rpm = AsyncTokenBucket((max_rpm or 0) / 60.0, capacity=max_rpm or 1) + self.tpm = AsyncTokenBucket((max_tpm or 0) / 60.0, capacity=max_tpm or 1) + self.timeout = timeout + self.token_counter = token_counter or (lambda item: len(item.solution_str)) + self.fallback_on_error = fallback_on_error + + async def run_single(self, item): + try: + await self.rpm.acquire() + await self.tpm.acquire(self.token_counter(item)) + return await asyncio.wait_for(super().run_single(item), self.timeout) + except (asyncio.TimeoutError, ValueError) as exc: + if not self.fallback_on_error: + raise + return self._zero_result(item, exc) + + @staticmethod + def _zero_result(item, exc): + from ..data import RewardResult + return RewardResult(item.item_id, 0.0, {"error": str(exc)}) diff --git a/src/twinkle/reward_loop/reward_manager/naive.py b/src/twinkle/reward_loop/reward_manager/naive.py new file mode 100644 index 00000000..9133a5fc --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/naive.py @@ -0,0 +1,7 @@ +from .base import RewardManagerBase +from .registry import register + + +@register("naive") +class NaiveRewardManager(RewardManagerBase): + pass diff --git a/src/twinkle/reward_loop/reward_manager/registry.py b/src/twinkle/reward_loop/reward_manager/registry.py new file mode 100644 index 00000000..6159b073 --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/registry.py @@ -0,0 +1,35 @@ +"""Reward manager registry.""" +from typing import Dict, Type + +_REGISTRY: Dict[str, Type] = {} + + +def _normalize(name: str) -> str: + if not isinstance(name, str) or not name.strip(): + raise ValueError("reward manager name must be a non-empty string") + return name.strip().lower().replace("-", "_") + + +def register(name: str): + key = _normalize(name) + + def decorator(cls): + if key in _REGISTRY and _REGISTRY[key] is not cls: + raise ValueError(f"reward manager already registered: {key}") + _REGISTRY[key] = cls + return cls + + return decorator + + +def get_reward_manager_cls(name: str): + key = _normalize(name) + try: + return _REGISTRY[key] + except KeyError as exc: + available = ", ".join(sorted(_REGISTRY)) or "none" + raise KeyError(f"unknown reward manager {name!r}; available: {available}") from exc + + +def registered_managers(): + return dict(_REGISTRY) diff --git a/src/twinkle/reward_loop/reward_manager/remote.py b/src/twinkle/reward_loop/reward_manager/remote.py new file mode 100644 index 00000000..7ee6957b --- /dev/null +++ b/src/twinkle/reward_loop/reward_manager/remote.py @@ -0,0 +1,8 @@ +from .naive import NaiveRewardManager +from .registry import register + + +@register("remote") +class RemoteRewardManager(NaiveRewardManager): + """Manager hook for CPU-isolated execution.""" + pass diff --git a/src/twinkle/reward_loop/worker.py b/src/twinkle/reward_loop/worker.py new file mode 100644 index 00000000..0abb734f --- /dev/null +++ b/src/twinkle/reward_loop/worker.py @@ -0,0 +1,56 @@ +import asyncio +import importlib +import threading +from typing import Optional + +from .data import RewardItem +from .default_score import compute_score as default_compute_score +from .reward_manager import get_reward_manager_cls + +try: + from twinkle.infra import remote_class +except ImportError: # pragma: no cover + remote_class = lambda **kwargs: (lambda cls: cls) + + +@remote_class(execute="all") +class RewardLoopWorker: + def __init__(self, manager_name="naive", compute_score=None, custom_reward_function_path=None, + custom_reward_function_name="compute_score", unknown_rewards="warn", reward_kwargs=None, + max_rpm=None, max_tpm=None, max_concurrent=1, timeout=300.0, **kwargs): + if compute_score is None and custom_reward_function_path: + module = importlib.import_module(custom_reward_function_path) + compute_score = getattr(module, custom_reward_function_name) + self.unknown_rewards = unknown_rewards + self.compute_score = compute_score or default_compute_score + manager_cls = get_reward_manager_cls(manager_name) + options = dict(reward_kwargs or {}) + if manager_name == "rate_limited": + options.update(max_rpm=max_rpm, max_tpm=max_tpm, max_concurrent=max_concurrent, timeout=timeout) + self.manager = manager_cls(compute_score=self.compute_score, **options) + # Managers build asyncio primitives in __init__ (the base class semaphore, + # the token bucket locks), and a primitive binds itself to the first loop + # it is awaited on. Serving every batch from one long-lived loop keeps them + # reusable across batches; a per-batch asyncio.run() breaks on the second. + self._loop = asyncio.new_event_loop() + self._loop_thread = threading.Thread(target=self._run_event_loop, daemon=True, + name="RewardLoopWorker-EventLoop") + self._loop_thread.start() + + def _run_event_loop(self): + asyncio.set_event_loop(self._loop) + self._loop.run_forever() + + def compute_score_batch(self, items): + future = asyncio.run_coroutine_threadsafe(self.manager.run_batch(items), self._loop) + return future.result() + + def close(self): + """Stop the worker's event loop and release its thread pool. Idempotent.""" + if self._loop.is_running(): + self._loop.call_soon_threadsafe(self._loop.stop) + if self._loop_thread.is_alive(): + self._loop_thread.join(timeout=5) + if not self._loop.is_closed(): + # Also shuts down the default executor that call_score() runs in. + self._loop.close() diff --git a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py index 7877ddb5..de6a65c1 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py @@ -413,6 +413,72 @@ def sample_stream_to_queue(self, queue, inputs, sampling_params=None, adapter_na from twinkle.server.sampler.backends import stream_to_queue stream_to_queue(self, queue, inputs, sampling_params, adapter_name, adapter_path) + @remote_function(dispatch='all', execute='first', collect='first', lazy_collect=False) + def sample_sequences_to_queue(self, queue, inputs, sampling_params=None, adapter_name='', adapter_path=None): + """Engine-level per-sequence streaming (local mode). + + Schedules ALL inputs concurrently in this actor's event loop — the + same batching behaviour as ``sample``, so the vLLM engine keeps + batching the whole batch — and pushes one ``(index, SampleResponse)`` + event onto ``queue`` as each sequence finishes (completion order), + followed by a ``(None, None)`` sentinel. + + Local-mode counterpart of the server's ``stream_sample_to_data_plane``: + call it from the driver on a worker thread while draining the queue on + the main path. ``index`` is the global input index, so per-sequence + reward submission and index alignment work like the server's + StreamSample events. + + Note: with multiple sampler DP workers only the first actor executes + (``execute='first'``); this targets the default SAMPLER_GPUS=1 + deployment. + """ + if sampling_params is None: + sampling_params = SamplingParams() + elif isinstance(sampling_params, dict): + sampling_params = SamplingParams.from_dict(sampling_params) + + inputs_list = self._normalize_inputs(inputs) + logprobs_only = False + if sampling_params.max_tokens == 0: + sampling_params = copy(sampling_params) + sampling_params.max_tokens = 1 + logprobs_only = True + + encoded_inputs = [] + for feat in inputs_list: + if 'input_ids' not in feat: + encoded_inputs.append(self.encode_trajectory_for_vllm(feat, adapter_name)) + else: + encoded_inputs.append(feat) + multi_modal_data_list = [self._extract_multi_modal_data(f) for f in encoded_inputs] + + lora_request = None + if adapter_path is not None: + logger.info(f'Loading LoRA from {adapter_path}') + adapter_path = HubOperation.download_model(model_id_or_path=adapter_path) + lora_request = self._run_in_loop(self.engine._get_or_load_lora(adapter_path)) + if lora_request is None: + logger.warning(f'Failed to pre-load LoRA from {adapter_path}, ' + 'sampling will proceed without LoRA') + + async def runner(): + async def generate(idx, feat, multi_modal_data): + response = await self._sample_single( + feat, sampling_params, lora_request=lora_request, + multi_modal_data=multi_modal_data, logprobs_only=logprobs_only) + # ray queue puts can block; keep the actor event loop free. + await asyncio.to_thread(queue.put, (idx, response)) + + tasks = [ + asyncio.ensure_future(generate(idx, feat, mm)) + for idx, (feat, mm) in enumerate(zip(encoded_inputs, multi_modal_data_list)) + ] + await asyncio.gather(*tasks) + await asyncio.to_thread(queue.put, (None, None)) + + self._run_in_loop(runner()) + @remote_function(dispatch='all', collect='first') def sleep(self, level: int = 1) -> None: """ diff --git a/src/twinkle/server/sampler/twinkle_handlers.py b/src/twinkle/server/sampler/twinkle_handlers.py index 609c25f2..3b534376 100644 --- a/src/twinkle/server/sampler/twinkle_handlers.py +++ b/src/twinkle/server/sampler/twinkle_handlers.py @@ -99,6 +99,47 @@ def _build_rollout_rows_and_tags( return rows, tags +def _sample_model_to_row( + model: types.SampleResponseModel, + *, + group_id: str, + prompt_index: int, + generation_idx: int, + policy_version: int | None, + adapter_uri: str | None, +) -> tuple[dict, dict]: + """Convert one generated sequence into a single TQ row and its tag. + + The row layout must stay identical to ``_build_rollout_rows_and_tags`` so + incremental writes land on the same global keys, where the global row index + is ``prompt_index * num_samples + generation_idx``. ``model`` carries a + single sequence (``collect_ready_samples`` pops per-sequence responses). + """ + sequence = model.sequences[0] + sampled_logprobs = [ + 0.0 if not position else float(position[0][1]) for position in (sequence.logprobs or []) + ] + row = { + 'train_input': sequence.new_input_feature, + 'sampled_logprobs': sampled_logprobs, + 'tokens': sequence.tokens, + 'decoded': sequence.decoded, + 'stop_reason': sequence.stop_reason, + 'prompt_logprobs': model.prompt_logprobs, + 'topk_prompt_logprobs': model.topk_prompt_logprobs, + } + tag = { + 'record_type': 'sample', + 'group_id': group_id, + 'prompt_index': prompt_index, + 'generation_idx': generation_idx, + 'rollout_status': 'ROLLOUT_DONE', + 'rollout_policy_version': policy_version, + 'rollout_adapter_uri': adapter_uri, + } + return row, tag + + def _to_sample_response_models(responses) -> list[types.SampleResponseModel]: """Convert internal sampler responses to the HTTP response schema.""" sample_models = [] @@ -180,6 +221,123 @@ async def _await_generation( except Exception: logger.warning('Failed to cancel generation %s', submission_id, exc_info=True) +async def _poll_generation_events( + sampler, + submission_id: str, +): + """Poll generation status, yielding each round's flattened DP worker states. + + The final round (all workers completed) is yielded as well, then the + generator returns. Poll cancellations raised by Ray are retried; worker + failures propagate immediately so the caller can cancel the submission. + """ + poll_interval = 0.01 + while True: + try: + states = _submission_states( + await asyncio.to_thread(sampler.get_generation_status, submission_id)) + except Exception as error: + # A pending read-only actor call can be cancelled by Ray while + # the generation submitted just above remains alive. Treating + # that as a generation failure makes the finally block discard + # otherwise valid rollout work. Retry only Ray's explicit task + # cancellation; actor death and application errors must still + # propagate immediately. + from ray.exceptions import TaskCancelledError + if not isinstance(error, TaskCancelledError): + raise + logger.warning( + 'Generation status poll was cancelled; retrying submission %s', + submission_id, + ) + await asyncio.sleep(poll_interval) + poll_interval = min(poll_interval * 1.5, 0.25) + continue + failed = next( + (state for state in states if state.get('status') not in ('running', 'completed')), + None, + ) + if failed is not None: + error = failed.get('error') or failed.get('status', 'unknown failure') + raise RuntimeError(f'generation {submission_id} failed: {error}') + yield states + if states and all(state.get('status') == 'completed' for state in states): + return + await asyncio.sleep(poll_interval) + poll_interval = min(poll_interval * 1.5, 0.25) + +async def _stream_sample_events( + sampler, + data_plane, + *, + submission_id: str, + ref: types.DataRef, + num_samples: int, + total: int, + group_ids: list[str], + policy_version: int | None, + adapter_uri: str | None, +): + """Yield NDJSON lines for the incremental sampling endpoint. + + Emits one ``progress`` line per completed sequence (row layout matches + ``_build_rollout_rows_and_tags`` and is written in place via ``put_rows``), + then a ``ref`` line with the fully-filled DataRef. On failure it cancels + the generation, releases the ref and yields an ``error`` line. + """ + written: set[int] = set() + try: + async for states in _poll_generation_events(sampler, submission_id): + new_indices: list[int] = [] + for state in states: + if not isinstance(state, dict): + continue + for index in state.get('completed_indices', []): + if index not in written: + written.add(index) + new_indices.append(index) + if not new_indices: + continue + responses = await asyncio.to_thread( + sampler.collect_ready_samples, submission_id, new_indices) + models = _to_sample_response_models([response for _, response in responses]) + for (index, _response), model in zip(responses, models): + prompt_index, generation_idx = divmod(index, num_samples) + row, tag = _sample_model_to_row( + model, + group_id=group_ids[prompt_index], + prompt_index=prompt_index, + generation_idx=generation_idx, + policy_version=policy_version, + adapter_uri=adapter_uri, + ) + ref = await data_plane.put_rows( + ref, + [json_safe(row)], + [index], + tags=[json_safe(tag)], + ) + yield json.dumps({ + 'event': 'progress', + 'index': index, + 'row': json_safe(row), + 'ref': ref.model_dump(), + 'done': len(written), + 'total': total, + }) + '\n' + yield json.dumps({'event': 'ref', 'ref': ref.model_dump()}) + '\n' + except Exception as error: + logger.error(traceback.format_exc()) + try: + await asyncio.to_thread(sampler.cancel_generation, submission_id) + except Exception: + logger.warning('Failed to cancel generation %s', submission_id, exc_info=True) + try: + await data_plane.release(ref) + except Exception: + logger.warning('Failed to release streaming ref %s', ref.ref_id, exc_info=True) + yield json.dumps({'event': 'error', 'error': str(error)}) + '\n' + def _register_twinkle_sampler_routes(app: FastAPI, self_fn: Callable[[], SamplerManagement]) -> None: """Register all /twinkle/* sampler routes on the given FastAPI app. @@ -350,6 +508,104 @@ async def _admit(): tags=tags, ) + @app.post('/twinkle/sample_to_data_plane_stream') + async def sample_to_data_plane_stream( + request: Request, + body: types.DataPlaneSampleRequest, + self: SamplerManagement = Depends(self_fn), + ) -> StreamingResponse: + """Generate completions and stream each finished sequence as NDJSON. + + Lines: + ``{"event":"progress","index":..,"row":{..},"done":N,"total":M}`` + per completed sequence (row layout matches ``_build_rollout_rows_and_tags`` + so the pre-allocated ref is filled incrementally in place), + then ``{"event":"ref","ref":{..}}`` once every sequence has been + written. Failures surface as ``{"event":"error","error":".."}`` and + release the pre-allocated ref. + """ + token = await self._on_request_start(request) + if not self.data_plane.enabled: + raise HTTPException(status_code=503, detail='sample_to_data_plane requires data_plane_url') + if not callable(getattr(self.sampler, 'submit_generation', None)) or not callable( + getattr(self.sampler, 'collect_ready_samples', None)): + raise HTTPException(status_code=503, detail='sampler_type must be vllm_async') + + adapter_path = None + full_adapter_name = _get_twinkle_sampler_adapter_name(request, body.adapter_name) or '' + if body.adapter_uri: + from twinkle.server.checkpoint import create_checkpoint_manager + checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') + _, adapter_path = checkpoint_manager.parse_adapter_uri(body.adapter_uri) + + inputs = ( + await self.data_plane.get(body.input_ref) + if body.input_ref is not None else body.inputs + ) + if isinstance(inputs, list) and inputs: + first = inputs[0] + if isinstance(first, dict) and 'input_ids' in first: + inputs = [InputFeature(**item) for item in inputs] + else: + inputs = [Trajectory(**item) for item in inputs] + elif isinstance(inputs, dict): + inputs = [InputFeature(**inputs)] if 'input_ids' in inputs else [Trajectory(**inputs)] + + input_count = len(inputs) if isinstance(inputs, list) else 1 + params_dict = dict(body.sampling_params or {}) + params_dict['num_samples'] = body.num_samples + params = SamplingParams.from_dict(params_dict) + submission_id = uuid.uuid4().hex + total = input_count * body.num_samples + resolved_group_ids = body.group_ids or [uuid.uuid4().hex for _ in range(input_count)] + if len(resolved_group_ids) != input_count: + raise ValueError( + f'group_ids contains {len(resolved_group_ids)} values for {input_count} sampler inputs') + + async def _admit(): + await asyncio.to_thread( + self.sampler.submit_generation, + submission_id, + inputs, + params, + adapter_name=full_adapter_name, + adapter_path=adapter_path, + ) + return submission_id + + inline_inputs = body.inputs if isinstance(body.inputs, list) else [body.inputs] + input_tokens = ( + body.input_ref.num_tokens + if body.input_ref is not None else + sum(len(item.get('input_ids', [])) for item in inline_inputs if isinstance(item, dict)) + ) + await run_task( + self.schedule_task_and_wait( + _admit, + model_id=full_adapter_name or None, + token=token, + input_tokens=input_tokens, + task_type='sample_admission', + )) + + ref = await self.data_plane.create(total, kind='rollout') + + async def _stream(): + async for line in _stream_sample_events( + self.sampler, + self.data_plane, + submission_id=submission_id, + ref=ref, + num_samples=body.num_samples, + total=total, + group_ids=resolved_group_ids, + policy_version=body.policy_version, + adapter_uri=body.adapter_uri, + ): + yield line + + return StreamingResponse(_stream(), media_type='application/x-ndjson') + @app.post('/twinkle/unload_adapter_paths') async def unload_adapter_paths( request: Request, diff --git a/src/twinkle_client/types/__init__.py b/src/twinkle_client/types/__init__.py index 1c25324a..e770c104 100644 --- a/src/twinkle_client/types/__init__.py +++ b/src/twinkle_client/types/__init__.py @@ -96,9 +96,11 @@ from .checkpoint import ResolvedLoadPath from .component import ( DataAppendRequest, + DataCreateRequest, DataGetRequest, DataPlaneSampleRequest, DataPutRequest, + DataPutRowsRequest, DataRef, DataReleaseRequest, DataRowsResponse, diff --git a/src/twinkle_client/types/component.py b/src/twinkle_client/types/component.py index d7e9ec2a..dbb9609c 100644 --- a/src/twinkle_client/types/component.py +++ b/src/twinkle_client/types/component.py @@ -35,6 +35,22 @@ class DataAppendRequest(BaseModel): tags: list[dict[str, Any]] | None = None +class DataCreateRequest(BaseModel): + """Pre-allocate a DataRef whose row keys are filled incrementally later.""" + + size: int + kind: str = 'data' + + +class DataPutRowsRequest(BaseModel): + """Overwrite specific row keys of an existing ref in place.""" + + ref: DataRef + rows: list[dict[str, Any]] + indices: list[int] + tags: list[dict[str, Any]] | None = None + + class DataReleaseRequest(BaseModel): ref: DataRef