dora7's picture
Albedo SN97 workspace v16: RECORD, merged v16, adapters v13/v15/v16-dpo, packs, harness, evals
2abcc30 verified
Raw
History Blame Contribute Delete
5.31 kB
from __future__ import annotations
import json
import random
from pathlib import Path
from typing import Any
import pyarrow.parquet as pq
from albedo_eval_service.remote.dataset import EvalSample, apply_submit_protocol, load_manifest_samples
from albedo_eval_service.shared.dataset_manifest import load_manifest_file
from albedo_eval_service.shared.sampling import multi_source_manifest_sample_ids
from .constants import TOKENIZER_DIR
def load_samples(
dataset_root: Path,
*,
sample_ids: list[str] | None,
sample_count: int,
seed: str,
tokenizer_path: Path | None = None,
) -> list[EvalSample]:
tokenizer = str(tokenizer_path or TOKENIZER_DIR)
if not sample_ids:
sample_ids = pick_sample_ids(dataset_root, sample_count=sample_count, seed=seed)
samples = load_manifest_samples(
dataset_root=dataset_root,
sample_ids=sample_ids,
tokenizer_path=tokenizer,
enable_thinking=True,
)
return apply_submit_protocol(
samples,
salt=seed,
keep_original_ratio=0.25,
tokenizer_path=tokenizer,
enable_thinking=True,
)
def pick_sample_ids(dataset_root: Path, *, sample_count: int, seed: str) -> list[str]:
manifest_path = dataset_root / "manifest.json"
if manifest_path.is_file():
try:
manifest = load_manifest_file(manifest_path, expected_sha256="")
if "sources" in manifest:
return multi_source_manifest_sample_ids(
manifest, block_hash=seed, sample_count=sample_count
)
except Exception as exc:
print(f"official sampler unavailable ({exc}); falling back to lite sampler", flush=True)
return lite_sample_ids(dataset_root, sample_count=sample_count, seed=seed)
def lite_sample_ids(dataset_root: Path, *, sample_count: int, seed: str) -> list[str]:
"""Prefix sampler for a partial local tree (no pinned multi-source manifest)."""
shards = sorted(dataset_root.glob("*/data/train-*.parquet"))
if not shards:
raise FileNotFoundError(
f"no parquet shards under {dataset_root}/*/data/ — run: python -m local_eval download-lite-data"
)
rng = random.Random(seed)
candidates: list[str] = []
for shard in shards:
rel = shard.relative_to(dataset_root).as_posix()
table = pq.read_table(shard, columns=_message_columns(shard))
for row_idx, raw in enumerate(table.to_pydict()[table.column_names[0]]):
turns = _as_turns(raw)
assistant = [i for i, turn in enumerate(turns) if _role(turn) == "assistant"]
if len(assistant) < 3:
continue
for turn_idx in (1, 2, min(len(assistant) - 1, 4)):
if turn_idx < len(assistant):
candidates.append(f"{rel}:{row_idx}:{turn_idx}")
rng.shuffle(candidates)
if len(candidates) < sample_count:
raise ValueError(f"only {len(candidates)} lite prefixes available, need {sample_count}")
return candidates[:sample_count]
def leftover_observations(dataset_root: Path, sample_id: str) -> list[str]:
"""Gold user turns after the sampled cut — offline stand-in for repo-context."""
shard_name, row_idx, turn_idx = sample_id.rsplit(":", 2)
row = _read_row(dataset_root / shard_name, int(row_idx))
turns = _as_turns(row.get("messages") or row.get("turns") or row.get("conversation"))
assistant = [i for i, turn in enumerate(turns) if _role(turn) == "assistant"]
if int(turn_idx) >= len(assistant):
return []
start = assistant[int(turn_idx)] + 1
return [
_content(turn)
for turn in turns[start:]
if _role(turn) in {"user", "tool"} and _content(turn)
]
def _message_columns(shard: Path) -> list[str]:
schema = pq.read_schema(shard)
for name in ("messages", "turns", "conversation", "trajectory"):
if name in schema.names:
return [name]
return [schema.names[0]]
def _read_row(path: Path, row_idx: int) -> dict[str, Any]:
parquet = pq.ParquetFile(path)
seen = 0
for batch in parquet.iter_batches(batch_size=512):
if seen + batch.num_rows <= row_idx:
seen += batch.num_rows
continue
return {k: (v[0] if isinstance(v, list) and len(v) == 1 else v) for k, v in batch.slice(row_idx - seen, 1).to_pydict().items()}
raise IndexError(row_idx)
def _as_turns(value: Any) -> list[Any]:
parsed = value
if isinstance(value, str):
try:
parsed = json.loads(value)
except json.JSONDecodeError:
return []
if isinstance(parsed, dict):
for key in ("messages", "turns", "conversation"):
if isinstance(parsed.get(key), list):
return parsed[key]
return []
return parsed if isinstance(parsed, list) else []
def _role(turn: Any) -> str:
if not isinstance(turn, dict):
return ""
return str(turn.get("role") or turn.get("speaker") or turn.get("from") or "").lower()
def _content(turn: Any) -> str:
if not isinstance(turn, dict):
return str(turn or "")
for key in ("content", "text", "value", "message"):
value = turn.get(key)
if value:
return str(value)
return ""