|
| 1 | +"""CIT- calibration improvement tracker schema. |
| 2 | +
|
| 3 | +Aggregates MBL- records across batches and computes whether the pipeline's |
| 4 | +prediction quality is trending in the right direction. Requires at least 2 |
| 5 | +batches with wet-lab data to compute a trend. Flags when calibration is not |
| 6 | +producing measurable improvement. |
| 7 | +""" |
| 8 | + |
| 9 | +from __future__ import annotations |
| 10 | + |
| 11 | +from dataclasses import dataclass |
| 12 | + |
| 13 | +VALID_CIT_TREND_DIRECTIONS: frozenset[str] = frozenset({ |
| 14 | + "improving", |
| 15 | + "stable", |
| 16 | + "degrading", |
| 17 | + "insufficient_data", |
| 18 | +}) |
| 19 | + |
| 20 | +VALID_CIT_SUMMARY_GRADES: frozenset[str] = frozenset({ |
| 21 | + "A", |
| 22 | + "B", |
| 23 | + "C", |
| 24 | + "D", |
| 25 | + "N/A", |
| 26 | +}) |
| 27 | + |
| 28 | +MIN_BATCHES_FOR_TREND: int = 2 |
| 29 | +IMPROVEMENT_THRESHOLD: float = 0.05 |
| 30 | +DEGRADATION_THRESHOLD: float = -0.05 |
| 31 | + |
| 32 | + |
| 33 | +@dataclass |
| 34 | +class BatchHitRateEntry: |
| 35 | + batch_id: str |
| 36 | + mbl_id: str |
| 37 | + hit_rate: float |
| 38 | + quality_grade: str |
| 39 | + |
| 40 | + |
| 41 | +@dataclass |
| 42 | +class CalibrationImprovementTracker: |
| 43 | + cit_id: str |
| 44 | + pipeline_version: str |
| 45 | + batch_entries: list[BatchHitRateEntry] |
| 46 | + n_batches: int |
| 47 | + n_batches_with_data: int |
| 48 | + first_batch_hit_rate: float |
| 49 | + latest_batch_hit_rate: float |
| 50 | + hit_rate_delta: float |
| 51 | + trend_direction: str |
| 52 | + summary_grade: str |
| 53 | + dry_lab_only: bool |
| 54 | + limitations: list[str] |
| 55 | + created_at: str |
| 56 | + |
| 57 | + |
| 58 | +def validate_calibration_improvement_tracker(cit: CalibrationImprovementTracker) -> None: |
| 59 | + if not cit.cit_id.startswith("CIT-"): |
| 60 | + raise ValueError(f"cit_id must start with 'CIT-': {cit.cit_id!r}") |
| 61 | + if not cit.pipeline_version: |
| 62 | + raise ValueError("pipeline_version must be non-empty") |
| 63 | + for entry in cit.batch_entries: |
| 64 | + if not (0.0 <= entry.hit_rate <= 1.0): |
| 65 | + raise ValueError( |
| 66 | + f"hit_rate must be in [0, 1] for batch {entry.batch_id!r}: {entry.hit_rate}" |
| 67 | + ) |
| 68 | + if cit.n_batches != len(cit.batch_entries): |
| 69 | + raise ValueError("n_batches must equal len(batch_entries)") |
| 70 | + entries_with_data = [e for e in cit.batch_entries if e.quality_grade != "N/A"] |
| 71 | + if cit.n_batches_with_data != len(entries_with_data): |
| 72 | + raise ValueError("n_batches_with_data mismatch") |
| 73 | + if cit.trend_direction not in VALID_CIT_TREND_DIRECTIONS: |
| 74 | + raise ValueError( |
| 75 | + f"trend_direction {cit.trend_direction!r} not in VALID_CIT_TREND_DIRECTIONS" |
| 76 | + ) |
| 77 | + if cit.summary_grade not in VALID_CIT_SUMMARY_GRADES: |
| 78 | + raise ValueError( |
| 79 | + f"summary_grade {cit.summary_grade!r} not in VALID_CIT_SUMMARY_GRADES" |
| 80 | + ) |
| 81 | + if cit.trend_direction == "insufficient_data" and cit.summary_grade != "N/A": |
| 82 | + raise ValueError( |
| 83 | + "summary_grade must be 'N/A' when trend_direction='insufficient_data'" |
| 84 | + ) |
| 85 | + if not cit.dry_lab_only: |
| 86 | + raise ValueError("dry_lab_only must be True") |
| 87 | + if not cit.limitations: |
| 88 | + raise ValueError("limitations must be non-empty") |
| 89 | + if not cit.created_at: |
| 90 | + raise ValueError("created_at must be non-empty") |
| 91 | + |
| 92 | + |
| 93 | +def _compute_trend( |
| 94 | + entries_with_data: list[BatchHitRateEntry], |
| 95 | +) -> tuple[float, float, float, str]: |
| 96 | + if len(entries_with_data) < MIN_BATCHES_FOR_TREND: |
| 97 | + first = entries_with_data[0].hit_rate if entries_with_data else 0.0 |
| 98 | + latest = entries_with_data[-1].hit_rate if entries_with_data else 0.0 |
| 99 | + return first, latest, 0.0, "insufficient_data" |
| 100 | + first = entries_with_data[0].hit_rate |
| 101 | + latest = entries_with_data[-1].hit_rate |
| 102 | + delta = round(latest - first, 6) |
| 103 | + if delta >= IMPROVEMENT_THRESHOLD: |
| 104 | + direction = "improving" |
| 105 | + elif delta <= DEGRADATION_THRESHOLD: |
| 106 | + direction = "degrading" |
| 107 | + else: |
| 108 | + direction = "stable" |
| 109 | + return first, latest, delta, direction |
| 110 | + |
| 111 | + |
| 112 | +def _compute_summary_grade( |
| 113 | + trend: str, |
| 114 | + latest_hit_rate: float, |
| 115 | +) -> str: |
| 116 | + if trend == "insufficient_data": |
| 117 | + return "N/A" |
| 118 | + if trend == "improving" and latest_hit_rate >= 0.25: |
| 119 | + return "A" |
| 120 | + if trend in ("improving", "stable") and latest_hit_rate >= 0.10: |
| 121 | + return "B" |
| 122 | + if trend == "stable" or latest_hit_rate >= 0.05: |
| 123 | + return "C" |
| 124 | + return "D" |
| 125 | + |
| 126 | + |
| 127 | +def build_calibration_improvement_tracker( |
| 128 | + *, |
| 129 | + cit_id: str, |
| 130 | + pipeline_version: str, |
| 131 | + batch_entry_dicts: list[dict], |
| 132 | + limitations: list[str], |
| 133 | + created_at: str, |
| 134 | +) -> CalibrationImprovementTracker: |
| 135 | + """Build a CalibrationImprovementTracker. |
| 136 | +
|
| 137 | + batch_entry_dicts: list of dicts with keys: |
| 138 | + batch_id, mbl_id, hit_rate, quality_grade |
| 139 | + """ |
| 140 | + entries = [ |
| 141 | + BatchHitRateEntry( |
| 142 | + batch_id=d["batch_id"], |
| 143 | + mbl_id=d["mbl_id"], |
| 144 | + hit_rate=float(d["hit_rate"]), |
| 145 | + quality_grade=d["quality_grade"], |
| 146 | + ) |
| 147 | + for d in batch_entry_dicts |
| 148 | + ] |
| 149 | + entries_with_data = [e for e in entries if e.quality_grade != "N/A"] |
| 150 | + first_hr, latest_hr, delta, trend = _compute_trend(entries_with_data) |
| 151 | + grade = _compute_summary_grade(trend, latest_hr) |
| 152 | + cit = CalibrationImprovementTracker( |
| 153 | + cit_id=cit_id, |
| 154 | + pipeline_version=pipeline_version, |
| 155 | + batch_entries=entries, |
| 156 | + n_batches=len(entries), |
| 157 | + n_batches_with_data=len(entries_with_data), |
| 158 | + first_batch_hit_rate=first_hr, |
| 159 | + latest_batch_hit_rate=latest_hr, |
| 160 | + hit_rate_delta=delta, |
| 161 | + trend_direction=trend, |
| 162 | + summary_grade=grade, |
| 163 | + dry_lab_only=True, |
| 164 | + limitations=limitations, |
| 165 | + created_at=created_at, |
| 166 | + ) |
| 167 | + validate_calibration_improvement_tracker(cit) |
| 168 | + return cit |
| 169 | + |
| 170 | + |
| 171 | +def format_calibration_improvement_tracker(cit: CalibrationImprovementTracker) -> str: |
| 172 | + lines = [ |
| 173 | + f"Calibration Improvement Tracker — {cit.cit_id}", |
| 174 | + f"Pipeline: {cit.pipeline_version}", |
| 175 | + f"Trend: {cit.trend_direction} | Grade: {cit.summary_grade}", |
| 176 | + f"Batches: {cit.n_batches} total, {cit.n_batches_with_data} with data", |
| 177 | + f"Hit rate: first={cit.first_batch_hit_rate:.1%} " |
| 178 | + f"latest={cit.latest_batch_hit_rate:.1%} " |
| 179 | + f"delta={cit.hit_rate_delta:+.1%}", |
| 180 | + ] |
| 181 | + if cit.batch_entries: |
| 182 | + lines.append("Batch history:") |
| 183 | + for entry in cit.batch_entries: |
| 184 | + lines.append( |
| 185 | + f" {entry.batch_id} ({entry.mbl_id}): " |
| 186 | + f"hr={entry.hit_rate:.1%} grade={entry.quality_grade}" |
| 187 | + ) |
| 188 | + lines.append(f"Created: {cit.created_at}") |
| 189 | + lines.append(f"Limitations: {'; '.join(cit.limitations)}") |
| 190 | + lines.append(f"dry_lab_only: {cit.dry_lab_only}") |
| 191 | + return "\n".join(lines) |
0 commit comments