Files
whetstone_RSA/tools/gen_binary_gate_tsv.py

216 lines
6.6 KiB
Python
Raw Permalink Normal View History

2026-04-30 19:39:01 +00:00
#!/usr/bin/env python3
"""Generate Fabricate-format TSVs for simple binary gates from extracted rows.
These datasets are intentionally pragmatic. The original extraction marked
target-file, acceptance-command, and required-tool rows as rejected for strict
gate-supervision purposes when labels were empty or sibling-uniform. For a first
small-transformer pass, those same fields can still support presence/absence
gates:
target_files_present
acceptance_commands_present
required_tools_present
manual_approval_present
Each output file uses Fabricate TSV format:
label<TAB>0<TAB>text
Usage:
python3 tools/gen_binary_gate_tsv.py \\
--input semantic/extracted_gate_rows_1000.ndjson \\
--output-dir semantic/binary_gate_tsv \\
--eval-frac 0.15 \\
--seed 42
"""
import argparse
import json
import random
from pathlib import Path
from typing import Callable
REPO_ROOT = Path(__file__).resolve().parent.parent
def compose_text(ctx: dict) -> str:
parts: list[str] = []
title = (ctx.get("title") or "").strip()
if title:
parts.append(f"title: {title}")
reasons = ctx.get("reasons") or []
if reasons:
clean = [str(r).strip() for r in reasons if str(r).strip()]
if clean:
parts.append("reasons: " + "; ".join(clean))
deps = ctx.get("dependency_task_ids") or []
if deps:
parts.append(f"dependency_count: {len(deps)}")
confidence = ctx.get("confidence")
if confidence is not None:
parts.append(f"confidence: {confidence}")
ambiguity_count = ctx.get("ambiguity_count")
if ambiguity_count is not None:
parts.append(f"ambiguity_count: {ambiguity_count}")
reqs = ctx.get("normalized_requirements") or []
goals = [r["normalizedText"] for r in reqs if r.get("kind") == "goal" and r.get("normalizedText")]
accepts = [r["normalizedText"] for r in reqs if r.get("kind") == "acceptance" and r.get("normalizedText")]
constraints = [r["normalizedText"] for r in reqs if r.get("kind") == "constraint" and r.get("normalizedText")]
dependencies = [r["normalizedText"] for r in reqs if r.get("kind") == "dependency" and r.get("normalizedText")]
if goals:
parts.append("goals: " + "; ".join(goals))
if accepts:
parts.append("accepts: " + "; ".join(accepts))
if constraints:
parts.append("constraints: " + "; ".join(constraints))
if dependencies:
parts.append("dependencies: " + "; ".join(dependencies))
prereqs = ctx.get("prerequisite_ops") or []
if prereqs:
parts.append("prerequisite_ops: " + "; ".join(str(p) for p in prereqs))
targets = ctx.get("target_files") or []
if targets:
parts.append("target_files: " + "; ".join(str(t) for t in targets))
return " | ".join(parts)
def load_rows(path: Path) -> list[dict]:
rows: list[dict] = []
with path.open("r", encoding="utf-8") as fh:
for line in fh:
line = line.strip()
if line:
rows.append(json.loads(line))
return rows
def write_split(
rows: list[tuple[int, str]],
output_dir: Path,
stem: str,
eval_frac: float,
seed: int,
) -> dict:
rng = random.Random(seed)
shuffled = list(rows)
rng.shuffle(shuffled)
n_eval = max(1, int(len(shuffled) * eval_frac))
eval_rows = shuffled[:n_eval]
train_rows = shuffled[n_eval:]
output_dir.mkdir(parents=True, exist_ok=True)
for split_name, split_rows in (("train", train_rows), ("eval", eval_rows)):
path = output_dir / f"{stem}_{split_name}.tsv"
with path.open("w", encoding="utf-8") as fh:
for label, text in split_rows:
if text:
fh.write(f"{label}\t0\t{text}\n")
def summarize(split_rows: list[tuple[int, str]]) -> dict:
pos = sum(1 for label, _ in split_rows if label == 1)
neg = len(split_rows) - pos
return {
"total": len(split_rows),
"pos": pos,
"neg": neg,
"pos_pct": round(pos / len(split_rows) * 100, 1) if split_rows else 0,
}
return {"train": summarize(train_rows), "eval": summarize(eval_rows)}
def build_gate(
rows: list[dict],
gate_id: str,
label_fn: Callable[[dict], int],
) -> list[tuple[int, str]]:
result: list[tuple[int, str]] = []
for row in rows:
if row.get("gate_id") != gate_id:
continue
text = compose_text(row.get("decision_context") or {})
if not text:
continue
result.append((label_fn(row), text))
return result
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--input", default="semantic/extracted_gate_rows_1000.ndjson")
parser.add_argument("--output-dir", default="semantic/binary_gate_tsv")
parser.add_argument("--eval-frac", type=float, default=0.15)
parser.add_argument("--seed", type=int, default=42)
return parser.parse_args()
def main() -> int:
args = parse_args()
input_path = Path(args.input)
if not input_path.is_absolute():
input_path = (REPO_ROOT / input_path).resolve()
output_dir = Path(args.output_dir)
if not output_dir.is_absolute():
output_dir = (REPO_ROOT / output_dir).resolve()
rows = load_rows(input_path)
gates = {
"target_files_present": build_gate(
rows,
"target_file_selection",
lambda row: 1 if row.get("decision_label", {}).get("target_files") else 0,
),
"acceptance_commands_present": build_gate(
rows,
"acceptance_command_selection",
lambda row: 1 if row.get("decision_label", {}).get("acceptance_commands") else 0,
),
"required_tools_present": build_gate(
rows,
"required_tool_selection",
lambda row: 1 if row.get("decision_label", {}).get("required_tools") else 0,
),
"manual_approval_present": build_gate(
rows,
"prereq_op_selector",
lambda row: 1
if "manual-approval" in (row.get("decision_label", {}).get("source_values") or [])
else 0,
),
}
summary = {
"input": str(input_path),
"output_dir": str(output_dir),
"gates": {},
}
for stem, gate_rows in gates.items():
if not gate_rows:
continue
summary["gates"][stem] = write_split(
gate_rows,
output_dir,
stem,
args.eval_frac,
args.seed,
)
print(json.dumps(summary, indent=2, sort_keys=True))
return 0
if __name__ == "__main__":
raise SystemExit(main())