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"