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 ""