"""G6 attention queue, calibration dataset, version drift와 manifest 코어.""" from __future__ import annotations import hashlib import json from collections import defaultdict from pathlib import Path from typing import Iterable from ..contracts.supervision_research import ( AttentionQueueItem, AttentionQueueReason, CalibrationDatasetRow, EvaluationVersionBatch, EvaluationVersionDriftReport, LearnerAttentionSignal, Phase3EvidenceArtifact, Phase3OutcomeEvidenceManifest, SubgroupVersionMetric, SupervisionResearchBenchmarkPack, TeacherAiDisagreement, ) _SIGNAL_PRIORITY = { "safety_boundary": 0, "deterioration": 1, "unresolved_rupture": 2, "persistent_overconfidence": 3, "growth_stagnation": 4, "transfer_failure": 5, } _SEVERITY_PRIORITY = {"high": 0, "moderate": 1, "low": 2} MIN_DRIFT_MATCHES = 4 MIN_SUBGROUP_MATCHES = 2 ACCURACY_DROP_THRESHOLD = 0.05 SUBGROUP_DROP_THRESHOLD = 0.10 def build_attention_queue( signals: Iterable[LearnerAttentionSignal], ) -> tuple[AttentionQueueItem, ...]: """활성 신호를 계획의 임상 워크벤치 우선순위로 정렬한다.""" grouped: dict[str, list[LearnerAttentionSignal]] = defaultdict(list) seen: set[str] = set() for signal in signals: if signal.signal_id in seen: raise ValueError(f"duplicate attention signal id: {signal.signal_id}") seen.add(signal.signal_id) if signal.state in {"active", "monitoring"}: grouped[signal.learner_ref].append(signal) ranked: list[ tuple[tuple[int, int, int, str], str, list[LearnerAttentionSignal]] ] = [] for learner_ref, items in grouped.items(): ordered = sorted( items, key=lambda item: ( _SIGNAL_PRIORITY[item.signal_type], _SEVERITY_PRIORITY[item.severity], item.observed_sequence, item.signal_id, ), ) primary = ordered[0] ranked.append( ( ( _SIGNAL_PRIORITY[primary.signal_type], _SEVERITY_PRIORITY[primary.severity], min(item.observed_sequence for item in ordered), learner_ref, ), learner_ref, ordered, ) ) output: list[AttentionQueueItem] = [] for position, (_, learner_ref, items) in enumerate(sorted(ranked), start=1): routes = tuple( dict.fromkeys( evidence.route_hint for item in items for evidence in item.evidence ) )[:3] output.append( AttentionQueueItem( learner_ref=learner_ref, queue_position=position, primary_signal=items[0].signal_type, oldest_active_sequence=min(item.observed_sequence for item in items), reasons=tuple( AttentionQueueReason( signal_id=item.signal_id, signal_type=item.signal_type, severity=item.severity, uncertainty=item.uncertainty, evidence=item.evidence, ) for item in items ), drilldown_routes=routes, ) ) return tuple(output) def _dataset_row_id(item: TeacherAiDisagreement) -> str: payload = { "disagreement_id": item.disagreement_id, "case_ref": item.case_ref, "competency_id": item.competency_id, "ai_label": item.ai_label, "teacher_label": item.teacher_label, "ai_model": item.ai_model, "prompt_version": item.prompt_version, "instrument_id": item.instrument_id, "instrument_version": item.instrument_version, "ai_evidence": sorted(value.event_id for value in item.ai_evidence), "teacher_evidence": sorted( value.event_id for value in item.teacher_correction_evidence ), "correction_reason_code": item.correction_reason_code, } canonical = json.dumps(payload, sort_keys=True, separators=(",", ":")) return hashlib.sha256(canonical.encode("utf-8")).hexdigest() def build_calibration_dataset( disagreements: Iterable[TeacherAiDisagreement], ) -> tuple[CalibrationDatasetRow, ...]: rows: list[CalibrationDatasetRow] = [] seen: set[str] = set() for item in disagreements: if item.disagreement_id in seen: raise ValueError(f"duplicate disagreement id: {item.disagreement_id}") seen.add(item.disagreement_id) evidence_ids = tuple( dict.fromkeys( value.event_id for value in (*item.ai_evidence, *item.teacher_correction_evidence) ) ) rows.append( CalibrationDatasetRow( row_id=_dataset_row_id(item), disagreement_id=item.disagreement_id, case_ref=item.case_ref, competency_id=item.competency_id, ai_label=item.ai_label, teacher_label=item.teacher_label, ai_model=item.ai_model, prompt_version=item.prompt_version, instrument_id=item.instrument_id, instrument_version=item.instrument_version, evidence_event_ids=evidence_ids, correction_reason_code=item.correction_reason_code, ) ) return tuple(rows) def compare_evaluation_versions( baseline: EvaluationVersionBatch, candidate: EvaluationVersionBatch, ) -> EvaluationVersionDriftReport: baseline_by_key = { (item.case_ref, item.competency_id): item for item in baseline.observations } candidate_by_key = { (item.case_ref, item.competency_id): item for item in candidate.observations } keys = sorted(set(baseline_by_key) & set(candidate_by_key)) if len(keys) < MIN_DRIFT_MATCHES: return EvaluationVersionDriftReport( baseline_batch_id=baseline.batch_id, candidate_batch_id=candidate.batch_id, matched_count=len(keys), status="insufficient_evidence", disagreement_case_refs=(), subgroup_metrics=(), alerts=(f"matched_cases_below_minimum:{len(keys)}/{MIN_DRIFT_MATCHES}",), ) baseline_correct = [ baseline_by_key[key].predicted_label == baseline_by_key[key].gold_label for key in keys ] candidate_correct = [ candidate_by_key[key].predicted_label == candidate_by_key[key].gold_label for key in keys ] baseline_accuracy = sum(baseline_correct) / len(keys) candidate_accuracy = sum(candidate_correct) / len(keys) accuracy_delta = candidate_accuracy - baseline_accuracy alerts: list[str] = [] if accuracy_delta < -ACCURACY_DROP_THRESHOLD: alerts.append("overall_accuracy_regression") subgroup_metrics: list[SubgroupVersionMetric] = [] subgroups = sorted( {baseline_by_key[key].synthetic_subgroup for key in keys} | {candidate_by_key[key].synthetic_subgroup for key in keys} ) for subgroup in subgroups: subgroup_keys = [ key for key in keys if baseline_by_key[key].synthetic_subgroup == subgroup and candidate_by_key[key].synthetic_subgroup == subgroup ] if len(subgroup_keys) < MIN_SUBGROUP_MATCHES: subgroup_metrics.append( SubgroupVersionMetric( subgroup=subgroup, matched_count=len(subgroup_keys) ) ) continue baseline_rate = sum( baseline_by_key[key].predicted_label == baseline_by_key[key].gold_label for key in subgroup_keys ) / len(subgroup_keys) candidate_rate = sum( candidate_by_key[key].predicted_label == candidate_by_key[key].gold_label for key in subgroup_keys ) / len(subgroup_keys) delta = candidate_rate - baseline_rate subgroup_metrics.append( SubgroupVersionMetric( subgroup=subgroup, matched_count=len(subgroup_keys), baseline_accuracy=baseline_rate, candidate_accuracy=candidate_rate, accuracy_delta=delta, ) ) if delta < -SUBGROUP_DROP_THRESHOLD: alerts.append(f"synthetic_subgroup_regression:{subgroup}") disagreements = tuple( sorted( { key[0] for key in keys if baseline_by_key[key].predicted_label != candidate_by_key[key].predicted_label } ) ) return EvaluationVersionDriftReport( baseline_batch_id=baseline.batch_id, candidate_batch_id=candidate.batch_id, matched_count=len(keys), status="drift_flagged" if alerts else "stable", baseline_accuracy=baseline_accuracy, candidate_accuracy=candidate_accuracy, accuracy_delta=accuracy_delta, disagreement_case_refs=disagreements, subgroup_metrics=tuple(subgroup_metrics), alerts=tuple(dict.fromkeys(alerts)), ) def build_phase3_outcome_manifest( artifacts: Iterable[Phase3EvidenceArtifact], ) -> Phase3OutcomeEvidenceManifest: return Phase3OutcomeEvidenceManifest( schema_version="vignette.phase3-outcome-evidence-manifest.v1", artifacts=tuple(artifacts), ) def load_supervision_research_benchmark( path: str | Path, ) -> SupervisionResearchBenchmarkPack: return SupervisionResearchBenchmarkPack.model_validate_json( Path(path).read_text(encoding="utf-8") ) def evaluate_supervision_research_benchmark( pack: SupervisionResearchBenchmarkPack, ) -> dict[str, object]: queue = build_attention_queue(pack.attention_signals) dataset = build_calibration_dataset(pack.disagreements) drift = compare_evaluation_versions(pack.baseline_batch, pack.candidate_batch) manifest = build_phase3_outcome_manifest(pack.phase3_artifacts) return { "schema_version": "vignette.supervision-research-benchmark-report.v1", "data_classification": pack.data_classification, "clinical_claim_allowed": pack.clinical_claim_allowed, "queue_order_correct": tuple(item.learner_ref for item in queue) == pack.expected_queue_order, "queue_items": [item.model_dump(mode="json") for item in queue], "calibration_dataset_rows": len(dataset), "raw_transcript_rows": sum(item.raw_transcript_included for item in dataset), "drift_status": drift.status, "drift_status_correct": drift.status == pack.expected_drift_status, "manifest_domains": sorted(item.domain for item in manifest.artifacts), } __all__ = [ "ACCURACY_DROP_THRESHOLD", "MIN_DRIFT_MATCHES", "MIN_SUBGROUP_MATCHES", "SUBGROUP_DROP_THRESHOLD", "build_attention_queue", "build_calibration_dataset", "build_phase3_outcome_manifest", "compare_evaluation_versions", "evaluate_supervision_research_benchmark", "load_supervision_research_benchmark", ]