Files

368 lines
15 KiB
Python

#!/usr/bin/env python3
"""Export first-prompt/first-response SWE-chat candidates without later text.
The source snapshot is gated and deliberately stays below ``.artifacts/``. This
importer does not download it: callers supply an already accepted, revision-pinned
Parquet snapshot. It reads Parquet in record batches, joins the small sessions table
only for repository/user grouping, and writes an unlabeled JSONL that contains the
first prompt plus a teacher-response hash for later labeling.
"""
from __future__ import annotations
import argparse
import json
import sys
from collections import Counter
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Iterable, Iterator, Sequence
from purpose_data import DataError, canonical_json, file_sha256, prompt_hash, write_json
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_RAW_DIR = SCRIPT_DIR / ".artifacts" / "swe-chat" / "raw"
DEFAULT_OUTPUT = SCRIPT_DIR / ".artifacts" / "swe-chat" / "candidates.jsonl"
DEFAULT_MANIFEST = SCRIPT_DIR / ".artifacts" / "swe-chat" / "export-manifest.json"
SCHEMA_VERSION = 1
REPOSITORY_ID = "SALT-NLP/SWE-chat"
LICENSE = "ODC-By-1.0"
@dataclass(frozen=True)
class Turn:
session_id: str
turn_id: str
conversation_turn_number: int
turn_number: int
prompt: str
@dataclass(frozen=True)
class Candidate:
session_id: str
repo_id: str | None
user_id: str | None
prompt_turn: Turn
response_turn: Turn
def json(self, revision: str) -> dict[str, Any]:
first, response = self.prompt_turn, self.response_turn
return {
"schemaVersion": SCHEMA_VERSION,
"repoID": self.repo_id,
"userID": self.user_id,
"sessionID": self.session_id,
"sourceTurnIDs": [first.turn_id, response.turn_id],
"sourceRevision": revision,
"promptHash": prompt_hash(first.prompt),
"teacherResponseHash": prompt_hash(response.prompt),
"prompt": first.prompt,
# Response text lives only in this ignored pre-labeling artifact.
"teacherResponse": response.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 parquet_rows(paths: Sequence[Path], columns: Sequence[str]) -> Iterator[dict[str, Any]]:
try:
import pyarrow.parquet as pq
except ImportError as error:
raise DataError(
"Parquet import requires pyarrow; install it in the purpose-classifier "
"environment (the raw gated snapshot is not read otherwise)"
) from error
for path in paths:
try:
parquet = pq.ParquetFile(path)
except Exception as error:
raise DataError(f"{path}: cannot open Parquet: {error}") from error
available = set(parquet.schema_arrow.names)
missing = sorted(set(columns) - available)
if missing:
raise DataError(
f"{path}: missing required columns {missing}; available columns are "
f"{sorted(available)}"
)
for batch in parquet.iter_batches(columns=list(columns), batch_size=16_384):
values = batch.to_pydict()
for index in range(batch.num_rows):
yield {column: values[column][index] for column in columns}
def session_metadata(rows: Iterable[dict[str, Any]]) -> dict[str, tuple[str | None, str | None]]:
result: dict[str, tuple[str | None, str | None]] = {}
for row in rows:
session_id = _as_text(row.get("session_id"))
if session_id is None:
continue
metadata = (_as_text(row.get("repo_id")), _as_text(row.get("user_id")))
previous = result.get(session_id)
if previous is not None and previous != metadata:
raise DataError(f"sessions config assigns conflicting repository/user to {session_id!r}")
result[session_id] = metadata
return result
def select_candidates(
conversation_rows: Iterable[dict[str, Any]],
*,
sessions: dict[str, tuple[str | None, str | None]],
max_per_repo: int,
max_per_user: int,
) -> tuple[list[Candidate], dict[str, int]]:
"""Test-friendly pair selection. Production uses two streaming passes below."""
rows = list(conversation_rows)
prompts, funnel = first_user_prompts(rows)
preliminary = attach_first_responses(rows, prompts, sessions=sessions, funnel=funnel)
return cap_and_dedupe(preliminary, funnel, max_per_repo=max_per_repo, max_per_user=max_per_user)
def _turn(row: dict[str, Any], *, funnel: Counter[str], prefix: str) -> Turn | None:
session_id = _as_text(row.get("session_id"))
turn_id = _as_text(row.get("turn_id"))
content = row.get("content")
conversation_turn_number = _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 not isinstance(content, str)
or not content.strip() or "\x00" in content
or conversation_turn_number is None or turn_number is None
):
funnel[f"rejected{prefix}MalformedOrEmpty"] += 1
return None
return Turn(session_id, turn_id, conversation_turn_number, turn_number, content)
def first_user_prompts(rows: Iterable[dict[str, Any]]) -> tuple[dict[str, Turn], Counter[str]]:
"""First streaming pass: retain the earliest eligible user prompt per session."""
funnel: Counter[str] = Counter()
prompts: dict[str, Turn] = {}
for row in rows:
funnel["conversationRows"] += 1
if row.get("turn_type") != "user_prompt":
funnel["rejectedTurnType"] += 1
continue
if row.get("role") != "user":
funnel["rejectedRole"] += 1
continue
if row.get("is_conversational") is not True:
funnel["rejectedUserNonConversational"] += 1
continue
if row.get("is_continuation") is True:
funnel["rejectedContinuation"] += 1
continue
turn = _turn(row, funnel=funnel, prefix="User")
if turn is None:
continue
previous = prompts.get(turn.session_id)
if previous is None or (turn.conversation_turn_number, turn.turn_number, turn.turn_id) < (
previous.conversation_turn_number, previous.turn_number, previous.turn_id
):
prompts[turn.session_id] = turn
funnel["eligibleUserPrompts"] += 1
funnel["sessionsWithEligibleFirstPrompt"] = len(prompts)
return prompts, funnel
def attach_first_responses(
rows: Iterable[dict[str, Any]],
prompts: dict[str, Turn],
*,
sessions: dict[str, tuple[str | None, str | None]],
funnel: Counter[str],
) -> list[Candidate]:
"""Second pass: find the immediately following conversational assistant response."""
responses: dict[str, Turn] = {}
for row in rows:
if row.get("turn_type") != "assistant_response" or row.get("role") != "assistant":
continue
if row.get("is_conversational") is not True:
funnel["rejectedAssistantNonConversational"] += 1
continue
session_id = _as_text(row.get("session_id"))
prompt = prompts.get(session_id or "")
if prompt is None:
continue
turn = _turn(row, funnel=funnel, prefix="Assistant")
if turn is None or turn.conversation_turn_number != prompt.conversation_turn_number + 1:
continue
previous = responses.get(turn.session_id)
if previous is not None:
# The dataset can contain multiple transcript rows at one conversational
# ordinal. `turn_number` is the documented checked tie-breaker; retain the
# earliest actual response and make the condition visible in the funnel.
funnel["multipleAssistantResponsesAtFirstResponseOrdinal"] += 1
if (turn.turn_number, turn.turn_id) >= (previous.turn_number, previous.turn_id):
continue
responses[turn.session_id] = turn
preliminary: list[Candidate] = []
for session_id, prompt in sorted(prompts.items()):
response = responses.get(session_id)
if response is None:
funnel["sessionsWithoutFirstAssistantResponse"] += 1
continue
repo_id, user_id = sessions.get(session_id, (None, None))
preliminary.append(Candidate(session_id, repo_id, user_id, prompt, response))
funnel["sessionsWithPromptAndFirstAssistantResponse"] = len(preliminary)
return preliminary
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 first prompts and cap repository/user concentration."""
first_by_hash: dict[str, str] = {}
deduped: list[Candidate] = []
for candidate in preliminary:
digest = prompt_hash(candidate.prompt_turn.prompt)
if digest in first_by_hash:
funnel["rejectedDuplicateFirstPrompt"] += 1
continue
first_by_hash[digest] = candidate.session_id
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)
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,
) -> 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:
raise DataError("source concentration caps must be positive")
sessions = session_metadata(parquet_rows(sessions_path, ("session_id", "repo_id", "user_id")))
columns = (
"session_id", "turn_id", "conversation_turn_number", "turn_number",
"turn_type", "role", "is_conversational", "is_continuation", "content",
)
prompts, funnel = first_user_prompts(parquet_rows(conversations, columns))
preliminary = attach_first_responses(
parquet_rows(conversations, columns), prompts, sessions=sessions, funnel=funnel
)
candidates, funnel = cap_and_dedupe(
preliminary, funnel, max_per_repo=max_per_repo, max_per_user=max_per_user
)
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text("".join(f"{canonical_json(candidate.json(revision))}\n" for candidate in candidates), encoding="utf-8")
manifest = {
"schemaVersion": SCHEMA_VERSION,
"generatedAt": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"),
"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": "first user_prompt with role=user, is_conversational=true, not is_continuation",
"perSession": "first eligible user prompt plus the immediately following conversational assistant_response",
"studentText": "first prompt only",
"teacherContext": "first assistant response retained only in ignored candidate JSONL until labeling",
"dedupe": "exact normalized first prompt",
"maxPerRepo": max_per_repo,
"maxPerUser": max_per_user,
},
"funnel": funnel,
"output": {"path": str(output), "records": len(candidates), "sha256": file_sha256(output)},
"removalLineage": "sourceTurnIDs and prompt 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)
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,
)
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())