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
14.6 kB
"""Local pre-eval chain: microtask + reject-first-submit, then official heuristics."""
from __future__ import annotations
import json
import subprocess
from datetime import datetime, timezone
from pathlib import Path
from uuid import uuid4
from albedo_eval_service.remote.dataset import EvalSample, format_messages
from albedo_eval_service.remote.generation import VllmProcessGenerator
from albedo_eval_service.shared.observation_format import (
MAX_CONSECUTIVE_BAD_TURNS,
detect_format,
first_bash_block,
unusable_turn,
wrap,
)
from albedo_eval_service.simulator.prompt_simulator import missing_command_output
from local_eval.constants import DEFAULT_DATA_ROOT, DEFAULT_RUNS_DIR, MAX_NEW_TOKENS, TOKENIZER_DIR
from local_eval.live_protocol import generate_retrying_bad_turns
from local_eval.rollout import build_generator
from local_eval.samples import leftover_observations, load_samples
from sanity_service.chain import (
SUBMIT_NUDGE,
followup_instruction,
micro_instruction,
segment_has_edit,
)
from .chain_gold import LIVE_CHAIN_IDS
from .chain_heuristics import ChainState, evaluate_chain, is_submit_turn, mid_roll_fatal
from .chain_pack import _REJECTION, infer_micro
from .constants import DEFAULT_RL_EXPORT_DIR
LIVE_SAMPLES = 3
LIVE_TURNS = 32
_OBS_CHARS = 2500
_CHAIN_GPUS = 4
_MIN_FREE_GIB = 100.0
def pick_free_gpus(n: int = _CHAIN_GPUS, min_free_gib: float = _MIN_FREE_GIB) -> list[str]:
"""Prefer cards that can actually satisfy vLLM's 0.8 utilization check."""
try:
raw = subprocess.check_output(
["nvidia-smi", "--query-gpu=index,memory.free", "--format=csv,noheader,nounits"],
text=True,
)
except (subprocess.CalledProcessError, FileNotFoundError) as exc:
raise RuntimeError(f"nvidia-smi failed: {exc}") from exc
free: list[str] = []
for line in raw.strip().splitlines():
parts = [p.strip() for p in line.split(",")]
if len(parts) < 2:
continue
idx, mem = parts[0], parts[1]
try:
if float(mem) / 1024.0 >= min_free_gib:
free.append(idx)
except ValueError:
continue
if len(free) < n:
raise RuntimeError(
f"need {n} GPUs with >= {min_free_gib:.0f} GiB free, found {free}. "
"kill leftover VLLM::EngineCore / Worker_TP* processes and retry"
)
return free[:n]
def _micro_for(sample) -> dict[str, str]:
hay = "\n".join(m.get("content") or "" for m in (sample.messages or []))
return infer_micro(hay)
def _pick_chain_samples(
dataset_root: Path, *, samples: int, seed: str, live_gate: bool = True
):
"""Use the 3 live-gate files first, then random extras so we test transfer."""
picked: list[tuple] = []
seen: set[str] = set()
forced = []
if live_gate:
try:
forced = load_samples(
dataset_root,
sample_ids=list(LIVE_CHAIN_IDS),
sample_count=len(LIVE_CHAIN_IDS),
seed=seed,
)
except Exception as exc:
print(f"live-gate samples unavailable ({exc}); falling back to random", flush=True)
forced = []
for sample in forced:
micro = _micro_for(sample)
if not sample.submit_command or not micro.get("file"):
continue
picked.append((sample, micro))
seen.add(sample.sample_id)
if len(picked) >= samples:
return picked
extras = load_samples(
dataset_root,
sample_ids=None,
sample_count=max(samples * 6, 18),
seed=seed,
)
for sample in extras:
if sample.sample_id in seen:
continue
micro = _micro_for(sample)
if not sample.submit_command or not micro.get("file"):
continue
picked.append((sample, micro))
seen.add(sample.sample_id)
if len(picked) >= samples:
break
return picked
def run_chain(
*,
challenger: Path = DEFAULT_RL_EXPORT_DIR,
dataset_root: Path = DEFAULT_DATA_ROOT,
samples: int = LIVE_SAMPLES,
turns: int = LIVE_TURNS,
seed: str = "chain-eval",
live_gate: bool = True,
gpu_ids: list[str] | None = None,
runs_dir: Path = DEFAULT_RUNS_DIR,
reject_first_submit: bool = True,
) -> dict:
from local_eval.cuda_env import apply as apply_cuda
apply_cuda()
dataset_root = Path(dataset_root)
picked = _pick_chain_samples(
dataset_root, samples=samples, seed=seed, live_gate=live_gate
)
states: list[ChainState] = []
leftover: dict[str, list[str]] = {}
leftover_idx: dict[str, int] = {}
for sample, micro in picked:
instruction = micro_instruction(micro, sample.submit_command)
messages = list(sample.messages or []) + [{"role": "user", "content": instruction}]
turns_so_far = [
{"role": m.get("role", "user"), "content": m.get("content", "")}
for m in messages
]
turns_so_far[-1] = {**turns_so_far[-1], "segment": "micro", "injected": True}
state = ChainState(
sample_id=sample.sample_id,
prompt=format_messages(
messages, tokenizer_path=str(TOKENIZER_DIR), enable_thinking=True
),
messages=messages,
turns=turns_so_far,
submit_clause=sample.submit_command,
submit_marker=sample.submit_marker,
rewrite_mode=getattr(sample, "rewrite_mode", ""),
micro=micro,
)
states.append(state)
leftover[sample.sample_id] = leftover_observations(dataset_root, sample.sample_id)
leftover_idx[sample.sample_id] = 0
if len(states) < samples:
raise RuntimeError(f"only {len(states)} chain-able samples, need {samples}")
gpu_ids = gpu_ids or pick_free_gpus()
if len(gpu_ids) > _CHAIN_GPUS:
gpu_ids = gpu_ids[:_CHAIN_GPUS]
generator = build_generator(
str(challenger),
gpu_ids,
max_new_tokens=MAX_NEW_TOKENS,
)
run_id = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") + "-" + uuid4().hex[:8]
out = Path(runs_dir) / f"{run_id}-chain"
out.mkdir(parents=True, exist_ok=True)
print(
f"local chain run={run_id} samples={len(states)} turns<={turns} chal={challenger}",
flush=True,
)
try:
_roll(generator, states, leftover, leftover_idx, turns, reject_first_submit)
finally:
generator.close()
evaluate_chain(states, turns)
report = _report(run_id, challenger, states, turns)
(out / "chain-report.json").write_text(json.dumps(report, indent=2) + "\n")
with (out / "chain-samples.jsonl").open("w") as handle:
for state in states:
handle.write(
json.dumps(
{
"sample_id": state.sample_id,
"passed": not state.error and not state.heuristic_reason,
"reason": state.error or state.heuristic_reason,
"n_submits": len(state.submits),
"empty_adjacent": sum(1 for s in state.submits if not s.get("has_edit")),
"micro": state.micro,
"submit_command": state.submit_clause,
"commands": [
first_bash_block(str(t.get("content") or ""))
for t in state.turns
if t.get("role") == "assistant" and t.get("score_target")
],
}
)
+ "\n"
)
print(json.dumps(report, indent=2), flush=True)
print(f"artifacts: {out}", flush=True)
return report
def _roll(
generator: VllmProcessGenerator,
states: list[ChainState],
leftover: dict[str, list[str]],
leftover_idx: dict[str, int],
turn_count: int,
reject_first_submit: bool,
) -> None:
for turn_index in range(turn_count):
active = [s for s in states if not s.stopped and not s.error and not s.heuristic_reason]
if not active:
break
batch = [
EvalSample(
sample_id=state.sample_id,
prompt=state.prompt,
messages=state.messages,
submit_command=state.submit_clause,
submit_marker=state.submit_marker,
)
for state in active
]
results = generate_retrying_bad_turns(generator, batch)
by_id = {r.sample_id: r for r in results}
for state in active:
result = by_id.get(state.sample_id)
if result is None or result.error:
scored = any(
t.get("role") == "assistant" and t.get("score_target") for t in state.turns
)
if scored:
state.stopped = True
else:
state.error = (
(result.error if result else "missing_generation") or "missing_generation"
)
continue
text = result.text or ""
reason = unusable_turn(text, truncated=bool(result.truncated))
if reason:
state.turns.append(
{
"role": "assistant",
"content": text,
"score_target": True,
"segment": state.segment,
}
)
state.heuristic_reason = (
f"{reason} on {MAX_CONSECUTIVE_BAD_TURNS} consecutive turns"
)
state.stopped = True
continue
state.turns.append(
{
"role": "assistant",
"content": text,
"score_target": True,
"segment": state.segment,
}
)
fatal = mid_roll_fatal(state)
if fatal:
state.heuristic_reason = fatal
state.stopped = True
continue
if turn_index == turn_count - 1:
continue
if is_submit_turn(text, state.submit_clause, state.submit_marker):
_handle_submit(state, text, turn_index, reject_first_submit)
elif not first_bash_block(text):
fmt = detect_format(state.sample_id, state.messages)
_append_user(state, missing_command_output(fmt))
else:
_append_user(state, _env_obs(state, leftover, leftover_idx))
state.prompt = format_messages(
state.messages, tokenizer_path=str(TOKENIZER_DIR), enable_thinking=True
)
if turn_index >= turn_count - 8 and (turn_count - turn_index) % 4 == 0:
_nudge([s for s in states if not s.stopped and not s.error and not s.submits])
def _handle_submit(
state: ChainState, text: str, turn_index: int, reject_first: bool
) -> None:
rejected = reject_first and not any(s.get("rejected") for s in state.submits)
record = {
"turn": turn_index,
"segment": state.segment,
"rejected": rejected,
"format_ok": state.submit_clause.split("&&")[0].strip() in text,
"post_nudge": bool(state.nudged_at) and turn_index > state.nudged_at,
"has_edit": segment_has_edit(state.as_ns(), state.segment),
}
state.submits.append(record)
if rejected:
_append_user(state, followup_instruction(_REJECTION, state.submit_clause, first=False))
return
first = state.segment == "micro"
state.segment_index += not first
state.segment = "real" if first else f"followup_{state.segment_index}"
_append_user(
state,
followup_instruction(
"Thanks, received. Please continue with the original issue and any missed call sites.",
state.submit_clause,
first=first,
),
)
def _nudge(states: list[ChainState]) -> None:
for state in states:
text = SUBMIT_NUDGE.format(clause=state.submit_clause)
_append_user(state, text)
state.turns[-1]["injected"] = True
state.nudged_at = state.nudged_at or len(
[t for t in state.turns if t.get("score_target")]
)
state.prompt = format_messages(
state.messages, tokenizer_path=str(TOKENIZER_DIR), enable_thinking=True
)
def _append_user(state: ChainState, observation: str) -> None:
assistant = str(state.turns[-1].get("content") or "")
state.messages.extend(
[
{"role": "assistant", "content": assistant},
{"role": "user", "content": observation},
]
)
state.turns.append(
{"role": "user", "content": observation, "environment_observation": True}
)
def _env_obs(
state: ChainState, leftover: dict[str, list[str]], leftover_idx: dict[str, int]
) -> str:
gold = leftover.get(state.sample_id) or []
index = leftover_idx.get(state.sample_id, 0)
if index < len(gold):
leftover_idx[state.sample_id] = index + 1
text = gold[index]
if len(text) > _OBS_CHARS:
text = text[:_OBS_CHARS] + "\n...[truncated]"
return text
fmt = detect_format(state.sample_id, state.messages)
return wrap("command completed with no captured output", fmt)
def _report(run_id: str, challenger: Path, states: list[ChainState], turns: int) -> dict:
rows = []
for state in states:
rows.append(
{
"sample_id": state.sample_id,
"passed": not state.error and not state.heuristic_reason,
"reason": state.error or state.heuristic_reason,
"n_submits": len(state.submits),
"unprompted": any(not s.get("post_nudge") for s in state.submits),
"micro_file": (state.micro or {}).get("file"),
}
)
passed = all(r["passed"] for r in rows) and bool(rows)
reasons = [r["reason"] for r in rows if r["reason"]]
if passed:
reasons = ["go: every sample passed the live chain heuristics"]
return {
"run_id": run_id,
"go": passed,
"challenger": str(challenger),
"n": len(rows),
"turns": turns,
"pass_rate": round(sum(1 for r in rows if r["passed"]) / max(len(rows), 1), 4),
"reasons": reasons,
"samples": rows,
}