451 lines
18 KiB
Python
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())
|