File size: 8,746 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 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 | import json
from collections import Counter
from pathlib import Path
def _accepted_v5_row():
from figment.observation_targets import required_observation_targets
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
from scripts.generate_finetune_data import v5_required_selected_observation_ids
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_v5")
prepared = prepare_case(spec, cards_by_id)
candidate = assemble_teacher_navigator_output(
prepared,
{
"facts": ["confirmed field concern"],
"missing": ["highest-value observation pending"],
"observe": ["highest-value observation pending"],
"checklist": ["cite deterministic rule cards"],
"uncertain": ["some vitals remain incomplete"],
"sbar": {
"situation": "confirmed handoff concern",
"background": "field workflow setting",
"assessment_observations_only": "observations only from confirmed intake",
"handoff_request": "request protocol review",
},
"script": "I am checking protocol observations.",
},
)
selected_ids = v5_required_selected_observation_ids(
source_card_ids=[str(card_id) for card_id in candidate.get("source_cards", [])],
retrieved_cards=prepared.retrieved_cards,
)
required_targets_by_id = {str(target["id"]): target for target in required_observation_targets(prepared.retrieved_cards)}
required_observation_text = [
str(required_targets_by_id[selected_id]["display_text"])
for selected_id in selected_ids
if selected_id in required_targets_by_id
]
candidate["selected_required_observation_ids"] = selected_ids
candidate["missing_info_to_collect"] = required_observation_text + list(candidate["missing_info_to_collect"])
candidate["next_observations_to_collect"] = required_observation_text
result = score_candidate(candidate, prepared)
assert result.passed is True
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_v5_failure_distribution_matches_focused_plan():
from scripts.generate_finetune_data import V5_FOCUSED_COUNTS
from scripts.generate_finetune_data import _failure_class_for_index
categories = Counter(_failure_class_for_index(index, dataset_version="figment_sft_v5") for index in range(1100))
assert categories == V5_FOCUSED_COUNTS
def test_v5_full_corpus_wrapper_pins_v5_defaults():
from scripts.generate_v5_full_corpus import DEFAULT_ARGS
from scripts.generate_v5_full_corpus import DEFAULT_COUNTS
from scripts.generate_v5_full_corpus import DEFAULT_NAVIGATOR_COUNT
from scripts.generate_v5_full_corpus import DEFAULT_OUTPUT_VERSION
from scripts.generate_v5_full_corpus import DEFAULT_TEACHER_MODEL_ID
from scripts.generate_v5_full_corpus import build_corpus_args
assert DEFAULT_OUTPUT_VERSION == "figment_sft_v5"
assert DEFAULT_TEACHER_MODEL_ID == "nvidia/nemotron-3-ultra-550b-a55b:free"
assert DEFAULT_COUNTS == {
"sbar_observation_ownership": 350,
"required_observation_id_selection": 250,
"source_card_invariant": 150,
"noisy_field_audio_style": 100,
"general_regression": 250,
}
assert DEFAULT_NAVIGATOR_COUNT == 1100
assert DEFAULT_ARGS[DEFAULT_ARGS.index("--dataset-version") + 1] == "figment_sft_v5"
assert DEFAULT_ARGS[DEFAULT_ARGS.index("--teacher-model-id") + 1] == "nvidia/nemotron-3-ultra-550b-a55b:free"
assert DEFAULT_ARGS[DEFAULT_ARGS.index("--navigator-count") + 1] == "1100"
assert DEFAULT_ARGS[DEFAULT_ARGS.index("--repair-count") + 1] == "200"
assert DEFAULT_ARGS[DEFAULT_ARGS.index("--output") + 1] == "data/finetune/figment_sft_v5.jsonl"
assert DEFAULT_ARGS[DEFAULT_ARGS.index("--modal-output-dir") + 1] == "data/finetune/modal/figment_sft_v5"
args = build_corpus_args(["--navigator-count", "2", "--output", "tmp/v5_smoke.jsonl"])
assert args[-4:] == ["--navigator-count", "2", "--output", "tmp/v5_smoke.jsonl"]
dry_run_args = build_corpus_args(["--navigator-count", "2", "--dry-run"])
assert dry_run_args[-3:] == ["--navigator-count", "2", "--dry-run"]
def test_v5_sft_row_records_training_focus_and_required_observation_ids():
row, spec_record = _accepted_v5_row()
output = json.loads(row["messages"][1]["content"])
metadata = row["metadata"]
assert row["version"] == "figment_sft_v5"
assert row["category"] == "sbar_observation_ownership"
assert metadata["training_focus"] == "sbar_observation_ownership"
assert metadata["excluded_eval_case_ids"] == [
"field_workflow_holdout_v1-000054",
"field_workflow_holdout_v1-000099",
]
assert metadata["must_include_source_cards"]
assert set(metadata["must_include_source_cards"]) <= set(output["source_cards"])
assert output["selected_required_observation_ids"]
assert set(metadata["must_include_selected_required_observation_ids"]) <= set(
output["selected_required_observation_ids"]
)
assert spec_record["dataset_version"] == "figment_sft_v5"
assert spec_record["workflow_category"] == "sbar_observation_ownership"
def test_v5_policy_rejects_missing_fired_card_selected_ids_and_generic_observations():
from scripts.generate_finetune_data import v5_policy_issues
output = {
"source_cards": ["SAFETY-BOUNDARIES-v1"],
"missing_info_to_collect": ["repeat vitals"],
"next_observations_to_collect": ["monitor closely"],
"handoff_note_sbar": {
"situation": "",
"background": "",
"assessment_observations_only": "",
"handoff_request": "",
},
}
retrieved_cards = [
{
"card_id": "STROKE-SIGNS-v1",
"card": {
"card_id": "STROKE-SIGNS-v1",
"required_observations": ["time last known well"],
},
}
]
issues = v5_policy_issues(
output,
failure_class="source_card_invariant",
expected_red_flag_rule_ids=["STROKE-001"],
expected_candidate_pathway_card_ids=["STROKE-SIGNS-v1"],
structured_intake={},
rule_results=[{"rule_id": "STROKE-001", "card_id": "STROKE-SIGNS-v1"}],
retrieved_cards=retrieved_cards,
target_protocol_card_id="STROKE-SIGNS-v1",
)
assert "fired_rule_source_card_missing:STROKE-SIGNS-v1" in issues
assert "generic_observation_phrase:repeat_vitals" in issues
assert "generic_observation_phrase:monitor_closely" in issues
def test_verify_v5_rejects_rows_without_selected_ids(tmp_path):
from scripts.verify_finetune_harness_alignment import verify_rows
row, spec_record = _accepted_v5_row()
output = json.loads(row["messages"][1]["content"])
output.pop("selected_required_observation_ids", None)
output["missing_info_to_collect"] = ["repeat vitals"]
output["next_observations_to_collect"] = ["monitor closely"]
row["messages"][1]["content"] = json.dumps(output, sort_keys=True)
dataset = tmp_path / "rows.jsonl"
case_specs = tmp_path / "specs.jsonl"
dataset.write_text(json.dumps(row, sort_keys=True) + "\n", encoding="utf-8")
case_specs.write_text(json.dumps(spec_record, sort_keys=True) + "\n", encoding="utf-8")
summary = verify_rows(dataset_path=dataset, case_specs_path=case_specs)
assert summary["passed"] is False
assert summary["issue_types"]["v5_selected_required_observation_ids_missing"] >= 1
assert summary["issue_types"]["v5_generic_observation_phrase:repeat_vitals"] >= 1
def test_v5_repair_scope_schedule_targets_observation_and_handoff_repairs():
from scripts.augment_finetune_repair_rows import _scope_schedule
counts = Counter(_scope_schedule(200, dataset_version="figment_sft_v5"))
assert counts == {
"missing_observations": 55,
"handoff_note_sbar": 45,
"citations_and_pathways": 35,
"forbidden_clinical_language": 25,
"protocol_urgency": 20,
"schema": 20,
}
|