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
7.72 kB
from __future__ import annotations
from dataclasses import replace
from pathlib import Path
from albedo_eval_service.evaluator.shared.questions import assign_horizons
from albedo_eval_service.remote.dataset import EvalSample, format_messages
from albedo_eval_service.remote.generation import GenerationResult, VllmProcessGenerator, format_scored_trajectory
from albedo_eval_service.shared.observation_format import detect_format, first_bash_block, truncation_notice, wrap
from albedo_eval_service.simulator.prompt_simulator import COMPLETE_MARKER, missing_command_output
from .live_protocol import generate_retrying_bad_turns, is_live_submit
from .constants import (
DEFAULT_GPU_MEMORY_UTILIZATION,
DEFAULT_LOCAL_TURNS,
DEFAULT_MAX_MODEL_LEN,
MAX_NEW_TOKENS,
TEMPERATURE,
TOKENIZER_DIR,
TOP_K,
TOP_P,
)
from .samples import leftover_observations
def build_generator(
model: str,
gpu_ids: list[str],
*,
max_model_len: int = DEFAULT_MAX_MODEL_LEN,
max_new_tokens: int = MAX_NEW_TOKENS,
gpu_memory_utilization: float = DEFAULT_GPU_MEMORY_UTILIZATION,
enforce_eager: bool = True,
) -> VllmProcessGenerator:
return VllmProcessGenerator(
model=model,
gpu_ids=gpu_ids,
max_new_tokens=max_new_tokens,
temperature=TEMPERATURE,
top_p=TOP_P,
top_k=TOP_K,
max_model_len=max_model_len,
enforce_eager=enforce_eager,
gpu_memory_utilization=gpu_memory_utilization,
# Triton GDN — this box has no system CUDA toolkit for flashinfer JIT.
gdn_prefill_backend="triton",
)
def generate_side(
*,
generator: VllmProcessGenerator,
samples: list[EvalSample],
dataset_root: Path,
max_turns: int = DEFAULT_LOCAL_TURNS,
max_new_tokens: int = MAX_NEW_TOKENS,
use_gold_env: bool = True,
) -> list[GenerationResult]:
official = assign_horizons(samples)
horizons = {
sample.sample_id: min(official.get(sample.sample_id, max_turns), max_turns)
for sample in samples
}
leftover = {
sample.sample_id: leftover_observations(dataset_root, sample.sample_id) if use_gold_env else []
for sample in samples
}
leftover_idx = {sample.sample_id: 0 for sample in samples}
current = list(samples)
turn_results: list[list[GenerationResult]] = []
turn_observations: list[dict[tuple[str, str], _Obs]] = []
try:
for turn_index in range(max(horizons.values(), default=max_turns)):
alive = [s for s in current if horizons.get(s.sample_id, max_turns) > turn_index]
if not alive:
break
results = generate_retrying_bad_turns(generator, alive)
turn_results.append(results)
if turn_index + 1 >= max(horizons.values(), default=max_turns):
break
observations = {}
next_samples: list[EvalSample] = []
result_by_id = {r.sample_id: r for r in results}
for sample in alive:
if horizons.get(sample.sample_id, max_turns) <= turn_index + 1:
continue
result = result_by_id.get(sample.sample_id)
if result is None or result.error or result.truncated:
continue
if _submitted(sample, result.text):
continue
obs = _observation(sample, result.text, leftover, leftover_idx)
observations[("local", sample.sample_id)] = obs
messages = list(sample.messages or []) + [
{"role": "assistant", "content": result.text},
{"role": "user", "content": obs.text},
]
next_samples.append(
replace(
sample,
prompt=format_messages(
messages,
tokenizer_path=str(TOKENIZER_DIR),
enable_thinking=True,
),
messages=messages,
)
)
turn_observations.append(observations)
current = next_samples
finally:
generator.close()
return _merge(samples, turn_results, turn_observations, max_new_tokens, horizons)
class _Obs:
def __init__(self, text: str, error: str | None = None):
self.observation = text
self.text = text
self.error = error
def _submitted(sample: EvalSample, text: str) -> bool:
return is_live_submit(
text,
command=sample.submit_command or "",
marker=sample.submit_marker or "",
)
def _observation(
sample: EvalSample,
text: str,
leftover: dict[str, list[str]],
leftover_idx: dict[str, int],
) -> _Obs:
fmt = detect_format(sample.sample_id, sample.messages)
if _submitted(sample, text):
return _Obs(wrap(sample.submit_marker or COMPLETE_MARKER, fmt))
if not first_bash_block(text):
return _Obs(missing_command_output(fmt))
gold = leftover.get(sample.sample_id) or []
index = leftover_idx.get(sample.sample_id, 0)
if index < len(gold):
leftover_idx[sample.sample_id] = index + 1
return _Obs(gold[index])
return _Obs(wrap("command completed with no captured output", fmt))
def _merge(
samples: list[EvalSample],
turn_results: list[list[GenerationResult]],
turn_observations: list[dict[tuple[str, str], _Obs]],
token_limit: int,
horizons: dict[str, int],
) -> list[GenerationResult]:
maps = [{r.sample_id: r for r in results} for results in turn_results]
merged: list[GenerationResult] = []
for sample in samples:
turns: list[dict] = [
{"role": m.get("role", "user"), "content": m.get("content", "")}
for m in (sample.messages or [])
]
error = None
truncated = False
for index, result_map in enumerate(maps):
result = result_map.get(sample.sample_id)
if result is None:
error = f"missing_generation_turn_{index + 1}"
break
if result.error:
error = result.error
break
if result.truncated:
truncated = True
turns.append(
{
"role": "assistant",
"content": truncation_notice(token_limit),
"score_target": True,
"truncated": True,
}
)
break
turns.append({"role": "assistant", "content": result.text, "score_target": True})
if index + 1 >= horizons.get(sample.sample_id, len(maps)):
break
if _submitted(sample, result.text):
break
if index >= len(turn_observations):
continue
obs = turn_observations[index].get(("local", sample.sample_id))
if obs is None or obs.error:
error = obs.error if obs else f"missing_observation_turn_{index + 1}"
break
turns.append(
{"role": "user", "content": obs.observation, "environment_observation": True}
)
if error:
merged.append(GenerationResult(sample.sample_id, "", error))
else:
merged.append(
GenerationResult(
sample_id=sample.sample_id,
text=format_scored_trajectory(turns),
turns=turns,
truncated=truncated,
)
)
return merged