Merge nucleic/plucky-north-vole-sdna into dev

This commit is contained in:
2026-08-01 16:18:23 -07:00
parent 9903296da6
commit 97d63c7714
2 changed files with 18 additions and 4 deletions
+11 -4
View File
@@ -306,6 +306,16 @@ def _append(path: Path, states: Sequence[dict[str, Any]]) -> None:
os.fsync(handle.fileno()) os.fsync(handle.fileno())
def record_from_decision(candidate: Candidate, decision: dict[str, Any]) -> dict[str, Any]:
"""Construct the exact seven-field student record from a teacher decision."""
record = {field: decision[field] for field in SOURCE_FIELDS - {"prompt"}}
record["prompt"] = candidate.prompt
if not decision["recoverableFromFirst"]:
record["slice"] = "vague-eval"
validate_source_record(record, candidate.id)
return record
def run(args: argparse.Namespace) -> dict[str, int]: def run(args: argparse.Namespace) -> dict[str, int]:
lines = base.source_lines(args.input) lines = base.source_lines(args.input)
source = candidates(lines) source = candidates(lines)
@@ -335,10 +345,7 @@ def run(args: argparse.Namespace) -> dict[str, int]:
newly: list[dict[str, Any]] = [] newly: list[dict[str, Any]] = []
for candidate, decision in decisions: for candidate, decision in decisions:
if decision["keep"]: if decision["keep"]:
record = {field: decision[field] for field in SOURCE_FIELDS} record = record_from_decision(candidate, decision)
if not decision["recoverableFromFirst"]:
record["slice"] = "vague-eval"
validate_source_record(record, candidate.id)
newly.append(_state(candidate, status="labeled", record=record, reason=None, recoverable=decision["recoverableFromFirst"])) newly.append(_state(candidate, status="labeled", record=record, reason=None, recoverable=decision["recoverableFromFirst"]))
else: else:
newly.append(_state(candidate, status="rejected", record=None, reason=f"semantic_junk:{decision['junkReason'].strip()}", recoverable=None)) newly.append(_state(candidate, status="rejected", record=None, reason=f"semantic_junk:{decision['junkReason'].strip()}", recoverable=None))
+7
View File
@@ -83,6 +83,13 @@ class LabelSWEChatPromptsTests(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "rather than truncating"): with self.assertRaisesRegex(ValueError, "rather than truncating"):
label_swe_chat_prompts.response_for_labeling("x" * 1_001, 1_000) label_swe_chat_prompts.response_for_labeling("x" * 1_001, 1_000)
def test_record_uses_only_the_first_user_message_as_student_text(self):
candidate = self.candidate()
decision = self.decision(candidate, True)["items"][0]
record = label_swe_chat_prompts.record_from_decision(candidate, decision)
self.assertEqual(candidate.prompt, record["prompt"])
self.assertNotIn(candidate.response, record.values())
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()