fugu-lite / src /fugu_lite /generate.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
5.03 kB
from __future__ import annotations
import asyncio
from collections import Counter
from pathlib import Path
from rich.console import Console
from .config import WorkerPoolConfig
from .io import assigned_split, load_rewards
from .providers import WorkerRunner
from .rewards import compute_utility, deterministic_quality
from .schemas import RewardRecord, TaskRecord, WorkerResult
console = Console()
def _append_record(destination: Path, record: RewardRecord) -> None:
with destination.open("a", encoding="utf-8") as handle:
handle.write(record.model_dump_json() + "\n")
def _winner_label(record: RewardRecord) -> str:
best_reward = max(record.rewards)
winners = [
worker_id
for worker_id, reward in zip(record.worker_ids, record.rewards, strict=True)
if abs(reward - best_reward) <= 1e-12
]
if len(winners) == 1:
return winners[0]
return f"tie({','.join(winners)})"
async def _score_result(
runner: WorkerRunner,
task: TaskRecord,
result: WorkerResult,
config: WorkerPoolConfig,
) -> WorkerResult:
if task.grader.type == "llm_judge":
quality = await runner.judge(task, result.response)
else:
quality = deterministic_quality(task, result.response)
scored = result.model_copy(update={"quality": quality})
return scored.model_copy(
update={"utility": compute_utility(scored, config.reward)}
)
async def generate_reward_dataset(
tasks: list[TaskRecord],
config: WorkerPoolConfig,
output_path: str | Path,
limit: int | None = None,
resume: bool = False,
repetitions: int = 1,
) -> list[RewardRecord]:
if repetitions < 1:
raise ValueError("repetitions must be >= 1")
selected_tasks = tasks[:limit] if limit else tasks
destination = Path(output_path)
destination.parent.mkdir(parents=True, exist_ok=True)
records: list[RewardRecord] = []
if resume and destination.exists() and destination.stat().st_size:
records = load_rewards(destination)
if records[0].worker_ids != config.worker_ids:
raise ValueError(
"Cannot resume: existing reward file has a different worker order"
)
if any(record.repetitions != repetitions for record in records):
raise ValueError(
"Cannot resume: existing reward file uses a different repetition count"
)
elif not resume:
destination.write_text("", encoding="utf-8")
completed_ids = {record.task_id for record in records}
remaining_tasks = [
task for task in selected_tasks
if task.task_id not in completed_ids
]
runner = WorkerRunner(config)
try:
for index, task in enumerate(remaining_tasks, start=1):
# Worker-major ordering:
# qwen run1, run2, run3,
# deepseek run1, run2, run3,
# gemini run1, run2, run3
raw_results = await asyncio.gather(
*(
runner.complete(worker, task)
for worker in config.workers
for _ in range(repetitions)
)
)
scored_results = await asyncio.gather(
*(
_score_result(runner, task, result, config)
for result in raw_results
)
)
result_groups = [
scored_results[
worker_index * repetitions:
(worker_index + 1) * repetitions
]
for worker_index in range(len(config.workers))
]
averaged_rewards = [
sum(result.utility for result in group) / len(group)
for group in result_groups
]
record = RewardRecord(
task_id=task.task_id,
prompt=task.prompt,
domain=task.domain,
split=assigned_split(task.task_id, task.split),
tags=task.tags,
metadata=task.metadata,
repetitions=repetitions,
worker_ids=config.worker_ids,
rewards=averaged_rewards,
results=list(scored_results),
)
records.append(record)
await asyncio.to_thread(
_append_record,
destination,
record,
)
console.print(
f"[{index}/{len(remaining_tasks)}] "
f"{task.task_id}: "
f"best={_winner_label(record)} "
f"avg_rewards="
f"{[round(value, 4) for value in record.rewards]}"
)
finally:
await runner.close()
winners = Counter(
_winner_label(record)
for record in records
)
console.print(
f"Dataset now has {len(records)} records at {output_path}; "
f"oracle winners={dict(winners)}"
)
return records