Spaces:
Sleeping
Sleeping
File size: 5,385 Bytes
94cbe85 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 | import json
from collections import Counter
def _accepted_v9_row():
from figment.retrieval import load_protocol_cards
from scripts.generate_finetune_data import assemble_teacher_navigator_output
from scripts.generate_finetune_data import build_sft_row
from scripts.generate_finetune_data import case_spec_record
from scripts.generate_finetune_data import generate_case_spec
from scripts.generate_finetune_data import prepare_case
from scripts.generate_finetune_data import score_candidate
cards_by_id = {str(card["card_id"]): card for card in load_protocol_cards()}
spec = generate_case_spec(0, cards_by_id, dataset_version="figment_sft_v9_delta")
prepared = prepare_case(spec, cards_by_id)
candidate = assemble_teacher_navigator_output(
prepared,
{
"facts": ["postpartum two weeks", "fever with chills", "temperature elevated"],
"missing": [
"pregnancy or postpartum status",
"bleeding report",
"abdominal pain report",
"headache or vision symptoms",
"seizure or fainting report",
"fever report",
"temperature if available",
],
"observe": ["temperature if available", "bleeding report", "abdominal pain report"],
"checklist": ["cite fever and pregnancy danger-sign cards"],
"uncertain": ["blood pressure pending"],
"sbar": {
"situation": "postpartum fever",
"background": "two weeks postpartum in field intake",
"assessment_observations_only": "fever and postpartum context confirmed",
"handoff_request": "request protocol review",
},
"script": "I am checking fever and postpartum danger-sign observations.",
},
)
result = score_candidate(candidate, prepared)
assert result.passed is True, result.reward_components
row = build_sft_row(
prepared=prepared,
result=result,
teacher_model_id="teacher-test",
candidate_total=1,
candidate_passed=1,
)
return row, case_spec_record(prepared)
def test_v9_failure_cycle_targets_remaining_v8_required_obs_gap():
from scripts.generate_finetune_data import V9_NAVIGATOR_COUNTS
from scripts.generate_finetune_data import _failure_class_for_index
categories = Counter(
_failure_class_for_index(index, dataset_version="figment_sft_v9_delta")
for index in range(sum(V9_NAVIGATOR_COUNTS.values()))
)
first_twelve = {
_failure_class_for_index(index, dataset_version="figment_sft_v9_delta")
for index in range(12)
}
assert categories == V9_NAVIGATOR_COUNTS
assert {
"postpartum_fever_required_obs_cross_category",
"postpartum_fever_required_obs_candidate_focus",
} <= first_twelve
def test_v9_sft_row_requires_postpartum_fever_and_preg_observation_text():
row, spec_record = _accepted_v9_row()
output = json.loads(row["messages"][1]["content"])
metadata = row["metadata"]
selected_ids = set(output["selected_required_observation_ids"])
assert row["version"] == "figment_sft_v9_delta"
assert spec_record["dataset_version"] == "figment_sft_v9_delta"
assert spec_record["target_protocol_card_id"] == "FEVER-RED-FLAGS-v1"
assert {"PREG-001", "FEVER-001"} <= set(spec_record["expected_red_flag_rule_ids"])
assert "postpartum two weeks" in json.dumps(spec_record["structured_intake"]).lower()
assert "PREG-DANGER-SIGNS-v1" in output["source_cards"]
assert "FEVER-RED-FLAGS-v1" in output["source_cards"]
assert [item["card_id"] for item in output["candidate_protocol_pathways"]] == [
"FEVER-RED-FLAGS-v1",
"PREG-DANGER-SIGNS-v1",
]
assert any(item.startswith("FEVER-RED-FLAGS-v1::required_observation::") for item in selected_ids)
assert any(item.startswith("PREG-DANGER-SIGNS-v1::required_observation::") for item in selected_ids)
assert set(metadata["must_include_selected_required_observation_ids"]) <= selected_ids
observation_text = json.dumps(output["missing_info_to_collect"]).lower()
for cue in (
"temperature if available",
"pregnancy or postpartum status",
"bleeding report",
"abdominal pain report",
"headache or vision symptoms",
"seizure or fainting report",
"fever report",
):
assert cue in observation_text
def test_v9_full_corpus_wrapper_pins_delta_defaults():
from scripts.generate_v9_full_corpus import DEFAULT_NAVIGATOR_COUNT
from scripts.generate_v9_full_corpus import DEFAULT_OUTPUT_VERSION
from scripts.generate_v9_full_corpus import DEFAULT_REPAIR_COUNT
from scripts.generate_v9_full_corpus import DEFAULT_TEACHER_MODEL_ID
from scripts.generate_v9_full_corpus import build_corpus_args
assert DEFAULT_OUTPUT_VERSION == "figment_sft_v9_delta"
assert DEFAULT_TEACHER_MODEL_ID == "nvidia/nemotron-3-ultra-550b-a55b:free"
assert DEFAULT_NAVIGATOR_COUNT == 400
assert DEFAULT_REPAIR_COUNT == 0
args = build_corpus_args(["--navigator-count", "8", "--repair-count", "0", "--dry-run"])
assert args[args.index("--navigator-count") + 1] == "8"
assert args[args.index("--repair-count") + 1] == "0"
assert args[-1] == "--dry-run"
|