Merge nucleic/plucky-north-vole-sdna into dev
This commit is contained in:
@@ -26,16 +26,15 @@ def row(session, turn, prompt, **overrides):
|
||||
|
||||
|
||||
class ExportSWEChatTests(unittest.TestCase):
|
||||
def test_selects_first_three_orders_dedupes_and_caps_sources(self):
|
||||
def test_selects_first_prompt_and_its_response_dedupes_and_caps_sources(self):
|
||||
rows = [
|
||||
row("s1", "t3", "third", conversation_turn_number=3),
|
||||
row("s1", "t1", "first", conversation_turn_number=1),
|
||||
row("s1", "t2", "second", conversation_turn_number=2),
|
||||
row("s2", "t1", " first ", conversation_turn_number=1),
|
||||
row("s2", "t2", "later", conversation_turn_number=2),
|
||||
row("s2", "t3", "later again", conversation_turn_number=3),
|
||||
row("s3", "t1", "one"), row("s3", "t2", "two"),
|
||||
row("s4", "t1", "another one"), row("s4", "t2", "another two"), row("s4", "t3", "another three"),
|
||||
row("s1", "t2", "answer", conversation_turn_number=1, turn_type="assistant_response", role="assistant"),
|
||||
row("s1", "t1", "first", conversation_turn_number=0),
|
||||
row("s2", "t1", " first ", conversation_turn_number=0),
|
||||
row("s2", "t2", "answer", conversation_turn_number=1, turn_type="assistant_response", role="assistant"),
|
||||
row("s3", "t1", "one"),
|
||||
row("s4", "t1", "another one", conversation_turn_number=0),
|
||||
row("s4", "t2", "answer", conversation_turn_number=1, turn_type="assistant_response", role="assistant"),
|
||||
]
|
||||
candidates, funnel = export_swe_chat.select_candidates(
|
||||
rows,
|
||||
@@ -43,30 +42,31 @@ class ExportSWEChatTests(unittest.TestCase):
|
||||
max_per_repo=1, max_per_user=1,
|
||||
)
|
||||
self.assertEqual(1, len(candidates))
|
||||
self.assertEqual(["t1", "t2", "t3"], [turn.turn_id for turn in candidates[0].turns])
|
||||
self.assertEqual("first", candidates[0].turns[0].prompt)
|
||||
self.assertEqual(["t1", "t2"], [candidates[0].prompt_turn.turn_id, candidates[0].response_turn.turn_id])
|
||||
self.assertEqual("first", candidates[0].prompt_turn.prompt)
|
||||
self.assertEqual(1, funnel["rejectedDuplicateFirstPrompt"])
|
||||
self.assertEqual(1, funnel["sessionsFewerThanThreeEligiblePrompts"])
|
||||
self.assertEqual(1, funnel["sessionsWithoutFirstAssistantResponse"])
|
||||
self.assertEqual(1, funnel["rejectedRepoCap"])
|
||||
|
||||
def test_filters_non_user_continuation_empty_and_ambiguous_ordinals(self):
|
||||
def test_filters_non_user_continuation_empty_and_missing_response(self):
|
||||
rows = [
|
||||
row("s1", "t1", "x", role="assistant"), row("s1", "t2", "x", is_continuation=True), row("s1", "t3", ""),
|
||||
row("s2", "t1", "a", conversation_turn_number=1, turn_number=1), row("s2", "t2", "b", conversation_turn_number=1, turn_number=1), row("s2", "t3", "c", conversation_turn_number=3),
|
||||
row("s2", "t1", "a", conversation_turn_number=0), row("s2", "t2", "wrong ordinal", conversation_turn_number=2, turn_type="assistant_response", role="assistant"),
|
||||
]
|
||||
candidates, funnel = export_swe_chat.select_candidates(rows, sessions={}, max_per_repo=10, max_per_user=10)
|
||||
self.assertEqual([], candidates)
|
||||
self.assertEqual(1, funnel["rejectedRole"])
|
||||
self.assertEqual(1, funnel["rejectedContinuation"])
|
||||
self.assertEqual(1, funnel["rejectedMalformedOrEmpty"])
|
||||
self.assertEqual(1, funnel["sessionsAmbiguousTurnOrder"])
|
||||
self.assertEqual(1, funnel["rejectedUserMalformedOrEmpty"])
|
||||
self.assertEqual(1, funnel["sessionsWithoutFirstAssistantResponse"])
|
||||
|
||||
def test_candidate_only_retains_first_prompt_as_student_text(self):
|
||||
turns = tuple(export_swe_chat.Turn("s", f"t{index}", index, index, prompt) for index, prompt in enumerate(("first", "second", "third"), start=1))
|
||||
value = export_swe_chat.Candidate("s", "r", "u", turns).json("a" * 40)
|
||||
prompt = export_swe_chat.Turn("s", "t1", 0, 0, "first")
|
||||
response = export_swe_chat.Turn("s", "t2", 1, 1, "answer")
|
||||
value = export_swe_chat.Candidate("s", "r", "u", prompt, response).json("a" * 40)
|
||||
self.assertEqual("first", value["prompt"])
|
||||
self.assertEqual(["second", "third"], value["teacherContext"])
|
||||
self.assertNotIn("second", value["contextPromptHashes"])
|
||||
self.assertEqual("answer", value["teacherResponse"])
|
||||
self.assertNotIn("answer", value["teacherResponseHash"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -20,12 +20,12 @@ class LabelSWEChatPromptsTests(unittest.TestCase):
|
||||
"repoID": "repo",
|
||||
"userID": "user",
|
||||
"sessionID": "session",
|
||||
"sourceTurnIDs": ["one", "two", "three"],
|
||||
"sourceTurnIDs": ["one", "two"],
|
||||
"prompt": "What is making this test fail?",
|
||||
"teacherContext": ["It fails only on CI.", "Please diagnose it."],
|
||||
"teacherResponse": "It fails only on CI.",
|
||||
}
|
||||
value["promptHash"] = prompt_hash(value["prompt"])
|
||||
value["contextPromptHashes"] = [prompt_hash(item) for item in value["teacherContext"]]
|
||||
value["teacherResponseHash"] = prompt_hash(value["teacherResponse"])
|
||||
line = base.SourceLine(1, json.dumps(value), "line-hash")
|
||||
return label_swe_chat_prompts.candidates([line])[0]
|
||||
|
||||
@@ -39,7 +39,7 @@ class LabelSWEChatPromptsTests(unittest.TestCase):
|
||||
}]
|
||||
}
|
||||
|
||||
def test_context_hashes_are_checked_and_context_dependent_labels_become_vague_eval(self):
|
||||
def test_response_hash_is_checked_and_context_dependent_labels_become_vague_eval(self):
|
||||
candidate = self.candidate()
|
||||
decisions = label_swe_chat_prompts.validate_decisions(
|
||||
[candidate], self.decision(candidate, False)
|
||||
@@ -50,8 +50,8 @@ class LabelSWEChatPromptsTests(unittest.TestCase):
|
||||
def test_candidate_rejects_context_hash_mismatch(self):
|
||||
candidate = self.candidate()
|
||||
value = json.loads(candidate.line.raw)
|
||||
value["contextPromptHashes"][0] = "bad"
|
||||
with self.assertRaisesRegex(ValueError, "contextPromptHashes"):
|
||||
value["teacherResponseHash"] = "bad"
|
||||
with self.assertRaisesRegex(ValueError, "teacherResponseHash"):
|
||||
label_swe_chat_prompts.candidates(
|
||||
[base.SourceLine(1, json.dumps(value), "line-hash")]
|
||||
)
|
||||
@@ -63,7 +63,6 @@ class LabelSWEChatPromptsTests(unittest.TestCase):
|
||||
)
|
||||
encoded = json.dumps(state)
|
||||
self.assertNotIn("It fails only on CI.", encoded)
|
||||
self.assertNotIn("Please diagnose it.", encoded)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user