Files

451 lines
18 KiB
Python

#!/usr/bin/env python3
"""Export turn-ending assistant reply prose from SWE-chat for the prose slice.
`docs/CONTEXT_SWITCH.md` §3.2 classifies the agent's settled-turn reply as well as the
user's prompt, but `purpose-lite` is trained only on user prompts, so reply prose is
out-of-domain. This exporter builds the missing slice: the same gated, revision-pinned
SWE-chat snapshot as `export_swe_chat.py`, restricted to conversational
`assistant_response` rows that **end a turn**, with the student text produced by
`prose_extract.reply_prose` — a port of the runtime's own extraction, so the classifier
trains on exactly the text it is asked to classify.
Like its sibling this importer never downloads anything: callers supply an already
accepted, revision-pinned Parquet snapshot below the gitignored `.artifacts/`. It reads in
record batches and takes two passes, the first of which deliberately does **not** project
`content` — turn-end structure is decided from small metadata columns alone, and reply text
is read only for the rows that survive.
The preceding user prompt travels in the candidate as teacher-only context. It never
reaches the student record: `label_swe_chat_prose.py` writes prose only.
"""
from __future__ import annotations
import argparse
import json
import sys
from collections import Counter, defaultdict
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Iterable, Sequence
from export_swe_chat import LICENSE, REPOSITORY_ID, parquet_rows, session_metadata
from prose_extract import (
DEFAULT_CHARACTER_LIMIT, ends_in_question, graphemes, reply_prose, strip_fenced_code, trim
)
from purpose_data import DataError, file_sha256, prompt_hash, write_json, write_jsonl
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_RAW_DIR = SCRIPT_DIR / ".artifacts" / "swe-chat" / "raw"
DEFAULT_OUTPUT = SCRIPT_DIR / ".artifacts" / "swe-chat" / "prose-candidates.jsonl"
DEFAULT_MANIFEST = SCRIPT_DIR / ".artifacts" / "swe-chat" / "prose-export-manifest.json"
SCHEMA_VERSION = 1
MARKER_COLUMNS = (
"session_id", "turn_id", "conversation_turn_number", "turn_number",
"turn_type", "role", "is_conversational", "is_continuation",
)
TEXT_COLUMNS = ("session_id", "turn_id", "content")
@dataclass(frozen=True)
class Marker:
"""One conversational row's ordering metadata, without its text."""
session_id: str
turn_id: str
conversation_turn_number: int
turn_number: int
kind: str # "user" | "assistant"
@property
def order(self) -> tuple[int, int, str]:
return (self.conversation_turn_number, self.turn_number, self.turn_id)
@dataclass(frozen=True)
class Pair:
"""A turn-ending assistant response and the user prompt that opened its turn."""
session_id: str
prompt_turn_id: str
response_turn_id: str
conversation_turn_number: int
@dataclass(frozen=True)
class Candidate:
session_id: str
repo_id: str | None
user_id: str | None
pair: Pair
teacher_prompt: str
prose: str
tail_biased: bool
@property
def prose_hash(self) -> str:
return prompt_hash(self.prose)
def json(self, revision: str, character_limit: int) -> dict[str, Any]:
return {
"schemaVersion": SCHEMA_VERSION,
"repoID": self.repo_id,
"userID": self.user_id,
"sessionID": self.session_id,
"sourceTurnIDs": [self.pair.prompt_turn_id, self.pair.response_turn_id],
"sourceRevision": revision,
"conversationTurnNumber": self.pair.conversation_turn_number,
"proseHash": self.prose_hash,
"teacherPromptHash": prompt_hash(self.teacher_prompt),
"replyCharacterLimit": character_limit,
"tailBiased": self.tail_biased,
"endsInQuestion": ends_in_question(self.prose),
# Student text: exactly what the runtime would classify.
"prose": self.prose,
# Teacher-only context; it lives solely in this ignored pre-labeling artifact.
"teacherPrompt": self.teacher_prompt,
}
def _as_text(value: Any) -> str | None:
return value if isinstance(value, str) and value.strip() else None
def _as_int(value: Any) -> int | None:
if isinstance(value, bool):
return None
if isinstance(value, int):
return value
if isinstance(value, float) and value.is_integer():
return int(value)
return None
def _paths(raw_dir: Path, stem: str) -> list[Path]:
candidates = sorted(raw_dir.rglob(f"*{stem}*.parquet")) if raw_dir.is_dir() else []
if not candidates:
raise DataError(
f"{raw_dir}: no {stem} Parquet files found; place the accepted pinned "
f"SWE-chat snapshot under this directory or pass --{stem}"
)
return candidates
def conversational_markers(
rows: Iterable[dict[str, Any]]
) -> tuple[dict[str, list[Marker]], Counter[str]]:
"""First pass: order every conversational row per session, carrying no text."""
funnel: Counter[str] = Counter()
markers: dict[str, list[Marker]] = defaultdict(list)
seen: set[tuple[str, str]] = set()
for row in rows:
funnel["conversationRows"] += 1
turn_type, role = row.get("turn_type"), row.get("role")
if turn_type == "user_prompt" and role == "user":
kind = "user"
elif turn_type == "assistant_response" and role == "assistant":
kind = "assistant"
else:
funnel["rejectedTurnTypeOrRole"] += 1
continue
if row.get("is_conversational") is not True:
funnel["rejectedNonConversational"] += 1
continue
if row.get("is_continuation") is True:
funnel["rejectedContinuation"] += 1
continue
session_id = _as_text(row.get("session_id"))
turn_id = _as_text(row.get("turn_id"))
ordinal = _as_int(row.get("conversation_turn_number"))
turn_number = _as_int(row.get("turn_number"))
if session_id is None or turn_id is None or ordinal is None or turn_number is None:
funnel["rejectedMalformedMarker"] += 1
continue
if (session_id, turn_id) in seen:
raise DataError(f"{session_id}: duplicate turn_id {turn_id!r} in conversations")
seen.add((session_id, turn_id))
markers[session_id].append(Marker(session_id, turn_id, ordinal, turn_number, kind))
funnel[f"conversational{kind.capitalize()}Rows"] += 1
funnel["sessionsWithConversationalRows"] = len(markers)
return dict(markers), funnel
def select_turn_ends(
markers: dict[str, list[Marker]], funnel: Counter[str], *, max_per_session: int
) -> list[Pair]:
"""Assistant responses whose next conversational event is a user turn, or nothing.
A multi-message assistant run contributes only its final message: the earlier ones are
mid-turn, and the runtime watcher only ever sees a settled turn's last reply.
"""
pairs: list[Pair] = []
for session_id in sorted(markers):
ordered = sorted(markers[session_id], key=lambda marker: marker.order)
kept = 0
latest_prompt: Marker | None = None
for index, marker in enumerate(ordered):
if marker.kind == "user":
latest_prompt = marker
continue
following = ordered[index + 1] if index + 1 < len(ordered) else None
if following is not None and following.kind != "user":
funnel["rejectedMidTurnAssistantResponse"] += 1
continue
if latest_prompt is None:
# An assistant response with no preceding user prompt has no teacher
# context; the prose alone cannot be adjudicated against a request.
funnel["rejectedNoPrecedingUserPrompt"] += 1
continue
if kept >= max_per_session:
funnel["rejectedSessionCap"] += 1
continue
pairs.append(
Pair(session_id, latest_prompt.turn_id, marker.turn_id,
marker.conversation_turn_number)
)
kept += 1
funnel["turnEndingAssistantResponses"] = len(pairs)
return pairs
def attach_text(
rows: Iterable[dict[str, Any]],
pairs: Sequence[Pair],
*,
sessions: dict[str, tuple[str | None, str | None]],
funnel: Counter[str],
character_limit: int,
) -> list[Candidate]:
"""Second pass: read content only for the selected prompt and response rows."""
wanted: dict[tuple[str, str], None] = {}
for pair in pairs:
wanted[(pair.session_id, pair.prompt_turn_id)] = None
wanted[(pair.session_id, pair.response_turn_id)] = None
text: dict[tuple[str, str], str] = {}
for row in rows:
session_id = _as_text(row.get("session_id"))
turn_id = _as_text(row.get("turn_id"))
key = (session_id or "", turn_id or "")
if key not in wanted:
continue
content = row.get("content")
if not isinstance(content, str) or not content.strip() or "\x00" in content:
funnel["rejectedMalformedOrEmptyContent"] += 1
continue
text[key] = content
candidates: list[Candidate] = []
for pair in pairs:
prompt = text.get((pair.session_id, pair.prompt_turn_id))
reply = text.get((pair.session_id, pair.response_turn_id))
if prompt is None or reply is None:
funnel["rejectedMissingText"] += 1
continue
prose = reply_prose(reply, character_limit)
if prose is None:
# A reply that is entirely fenced code or whitespace carries no language
# signal. The runtime skips it too, so it is not training data.
funnel["rejectedNoProseAfterExtraction"] += 1
continue
# Whether the limit truncated a long reply, rather than merely whether fences were
# stripped: a labeler reading a tail-biased record is not seeing the whole reply.
tail_biased = len(graphemes(trim(strip_fenced_code(reply)))) > character_limit
if tail_biased:
funnel["tailBiasedProse"] += 1
repo_id, user_id = sessions.get(pair.session_id, (None, None))
candidates.append(
Candidate(
session_id=pair.session_id, repo_id=repo_id, user_id=user_id, pair=pair,
teacher_prompt=prompt, prose=prose, tail_biased=tail_biased,
)
)
funnel["candidatesWithProse"] = len(candidates)
return candidates
def cap_and_dedupe(
preliminary: Sequence[Candidate],
funnel: Counter[str],
*,
max_per_repo: int,
max_per_user: int,
) -> tuple[list[Candidate], dict[str, int]]:
"""Dedupe normalized prose and cap repository/user concentration.
Agent sign-off prose is far more repetitive than human prompts ("All tests pass. Let me
know if you want anything else."), so the duplicate count here is expected to be large.
It is reported rather than smoothed over.
"""
seen: set[str] = set()
deduped: list[Candidate] = []
for candidate in preliminary:
digest = candidate.prose_hash
if digest in seen:
funnel["rejectedDuplicateProse"] += 1
continue
seen.add(digest)
deduped.append(candidate)
accepted: list[Candidate] = []
repo_counts: Counter[str] = Counter()
user_counts: Counter[str] = Counter()
for candidate in deduped:
repo_key = candidate.repo_id or "<unknown>"
user_key = candidate.user_id or "<unknown>"
if repo_counts[repo_key] >= max_per_repo:
funnel["rejectedRepoCap"] += 1
continue
if user_counts[user_key] >= max_per_user:
funnel["rejectedUserCap"] += 1
continue
accepted.append(candidate)
repo_counts[repo_key] += 1
user_counts[user_key] += 1
funnel["exportedCandidates"] = len(accepted)
funnel["exportedEndingInQuestion"] = sum(
1 for candidate in accepted if ends_in_question(candidate.prose)
)
return accepted, dict(sorted(funnel.items()))
def export(
*,
conversations: Sequence[Path],
sessions_path: Sequence[Path],
revision: str,
output: Path,
manifest_path: Path,
max_per_repo: int,
max_per_user: int,
max_per_session: int,
character_limit: int,
) -> dict[str, Any]:
if not revision or revision == "main":
raise DataError("--revision must be an accepted immutable SWE-chat commit, never main")
if max_per_repo <= 0 or max_per_user <= 0 or max_per_session <= 0:
raise DataError("source concentration caps must be positive")
if character_limit <= 0:
raise DataError("--reply-character-limit must be positive")
sessions = session_metadata(parquet_rows(sessions_path, ("session_id", "repo_id", "user_id")))
markers, funnel = conversational_markers(parquet_rows(conversations, MARKER_COLUMNS))
pairs = select_turn_ends(markers, funnel, max_per_session=max_per_session)
preliminary = attach_text(
parquet_rows(conversations, TEXT_COLUMNS), pairs,
sessions=sessions, funnel=funnel, character_limit=character_limit,
)
candidates, funnel = cap_and_dedupe(
preliminary, funnel, max_per_repo=max_per_repo, max_per_user=max_per_user
)
write_jsonl(output, [c.json(revision, character_limit) for c in candidates])
manifest = {
"schemaVersion": SCHEMA_VERSION,
"generatedAt": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"),
"slice": "assistant-prose",
"source": {
"repoID": REPOSITORY_ID,
"revision": revision,
"license": LICENSE,
"attribution": "SALT-NLP/SWE-chat; SWE-chat paper arXiv:2604.20779",
"rawFiles": [
{"path": str(path), "sha256": file_sha256(path)}
for path in sorted([*conversations, *sessions_path])
],
},
"selection": {
"rowFilter": (
"conversational assistant_response with role=assistant, "
"not is_continuation"
),
"perSession": (
"assistant responses whose next conversational event is a user prompt or "
"the end of the session, capped at maxPerSession"
),
"studentText": (
"reply prose from prose_extract.reply_prose — a port of "
"HeuristicSummary.contextSwitchReplyProse, pinned by "
"Tests/NucleicCoreTests/Fixtures/context-switch-prose.json"
),
"teacherContext": (
"the user prompt that opened the turn, retained only in the ignored "
"candidate JSONL until labeling"
),
"dedupe": "exact normalized reply prose",
"replyCharacterLimit": character_limit,
"maxPerRepo": max_per_repo,
"maxPerUser": max_per_user,
"maxPerSession": max_per_session,
},
"funnel": funnel,
"output": {
"path": str(output), "records": len(candidates), "sha256": file_sha256(output)
},
"removalLineage": (
"sourceTurnIDs and prose hashes are retained in ignored audit sidecars; "
"rebuild the next dataset version after a source tombstone."
),
}
write_json(manifest_path, manifest)
return manifest
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--raw-dir", type=Path, default=DEFAULT_RAW_DIR)
parser.add_argument("--conversations", action="append", type=Path)
parser.add_argument("--sessions", action="append", type=Path)
parser.add_argument("--revision", required=True, help="accepted immutable Hugging Face revision")
parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument("--manifest", type=Path, default=DEFAULT_MANIFEST)
parser.add_argument("--max-per-repo", type=int, default=100)
parser.add_argument("--max-per-user", type=int, default=50)
parser.add_argument(
"--max-per-session", type=int, default=3,
help="turn-ending replies taken from one session (sign-offs repeat within a chat)",
)
parser.add_argument(
"--reply-character-limit", type=int, default=DEFAULT_CHARACTER_LIMIT,
help="must match the runtime watcher's limit; changing it changes the student text",
)
return parser
def main(argv: Sequence[str] | None = None) -> int:
args = build_parser().parse_args(argv)
try:
raw_dir = args.raw_dir.expanduser().resolve()
conversations = (
[path.expanduser().resolve() for path in args.conversations]
if args.conversations else _paths(raw_dir, "conversations")
)
sessions = (
[path.expanduser().resolve() for path in args.sessions]
if args.sessions else _paths(raw_dir, "sessions")
)
manifest = export(
conversations=conversations, sessions_path=sessions, revision=args.revision,
output=args.output.expanduser().resolve(),
manifest_path=args.manifest.expanduser().resolve(),
max_per_repo=args.max_per_repo, max_per_user=args.max_per_user,
max_per_session=args.max_per_session,
character_limit=args.reply_character_limit,
)
except (DataError, OSError, ValueError) as error:
print(f"error: {error}", file=sys.stderr)
return 1
print(json.dumps(manifest["funnel"], sort_keys=True))
print(f"Candidates: {manifest['output']['path']}")
print(f"Manifest: {args.manifest}")
return 0
if __name__ == "__main__":
raise SystemExit(main())