Merge nucleic/plucky-north-vole-sdna into dev
This commit is contained in:
@@ -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))
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user