Spaces:
Running
Running
File size: 7,131 Bytes
79b0bef | 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 | """
scripts/run_icd10_batch.py
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Batch ICD-10 mapper β completes the src/nlp/icd_mapper.py wiring.
Reads DISEASE/SYMPTOM entities that don't have ICD-10 mappings
yet, maps each *unique* entity text once (exact -> fuzzy ->
embedding, see ICD10Mapper.map), and fans the result out to every
entity row that shares that text. Persists via load_icd10_mappings().
Only DISEASE and SYMPTOM entities are eligible -- medications,
procedures, and anatomy terms use different code systems (see
ICD10Mapper.map_entities). Caching by unique text matters here:
206k eligible entity rows reduce to ~43k unique texts, so the
cache avoids redundant fuzzy/embedding lookups for repeated terms
like "hypertension".
Idempotent / resumable: an entity is only processed if it has
zero rows in icd10_mappings, so re-running after an interruption
never creates duplicates.
icd10_mappings only ever stores *found* matches -- an entity text
that the mapper couldn't confidently map has no row and therefore
still looks "pending" on the next run. Without a separate record
of what's already been tried, every restart would re-run the slow
fuzzy/embedding matching on every previously-checked no-match text.
A small on-disk cache (data/processed/icd10_mapping_cache.json,
keyed by entity text) closes that gap: any text seen before --
matched or not -- is skipped straight from the cache on restart.
Usage
βββββ
python scripts/run_icd10_batch.py # all pending entities
python scripts/run_icd10_batch.py --limit 500 # quick demo subset
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
"""
from __future__ import annotations
import argparse
import json
import time
from sqlalchemy import select
from src.db.connection import get_session
from src.db.models import Entity, ICD10Mapping
from src.etl.load import load_icd10_mappings
from src.nlp.icd_mapper import ICD10Mapper
from src.utils.config import Paths
from src.utils.logger import get_logger
logger = get_logger(__name__)
_ELIGIBLE_LABELS = ("DISEASE", "SYMPTOM")
_CACHE_PATH = Paths.processed / "icd10_mapping_cache.json"
def _load_cache() -> dict[str, list[dict]]:
"""Load the on-disk text -> match-dicts cache, or an empty dict."""
if not _CACHE_PATH.exists():
return {}
with open(_CACHE_PATH, encoding="utf-8") as f:
cache = json.load(f)
logger.info("Loaded ICD-10 mapping cache: %d entity texts already tried", len(cache))
return cache
def _save_cache(cache: dict[str, list[dict]]) -> None:
"""Write the text -> match-dicts cache to disk."""
_CACHE_PATH.parent.mkdir(parents=True, exist_ok=True)
with open(_CACHE_PATH, "w", encoding="utf-8") as f:
json.dump(cache, f)
def _fetch_pending_entities(limit: int | None) -> list[tuple[int, str]]:
"""Return (id, text) for DISEASE/SYMPTOM entities with no ICD-10 mapping yet.
Args:
limit: Maximum number of entity rows to return, or None for
every pending entity.
Returns:
List of (entity_id, text) tuples, ordered by entity id so a
--limit run always covers the same entities on re-run.
"""
with get_session() as session:
stmt = (
select(Entity.id, Entity.text)
.outerjoin(ICD10Mapping, ICD10Mapping.entity_id == Entity.id)
.where(
Entity.label.in_(_ELIGIBLE_LABELS),
ICD10Mapping.id.is_(None),
)
.order_by(Entity.id)
)
if limit:
stmt = stmt.limit(limit)
return [(row.id, row.text) for row in session.execute(stmt)]
def run(limit: int | None, commit_every: int) -> None:
"""Map pending entities to ICD-10 codes and save the results.
Args:
limit: Maximum number of pending entity rows to process, or
None for every DISEASE/SYMPTOM entity without a mapping.
commit_every: Write to the database once this many mapping
rows have been queued in memory.
"""
pending = _fetch_pending_entities(limit)
total_entities = len(pending)
if total_entities == 0:
logger.info(
"No pending entities -- every DISEASE/SYMPTOM entity "
"already has an ICD-10 mapping."
)
return
by_text: dict[str, list[int]] = {}
for entity_id, text in pending:
by_text.setdefault(text, []).append(entity_id)
unique_texts = list(by_text.keys())
logger.info(
"Mapping %d unique entity texts (%d entity rows) to ICD-10 codes...",
len(unique_texts), total_entities,
)
mapper = ICD10Mapper()
cache = _load_cache()
cache_hits = 0
start = time.time()
processed_entities = 0
inserted_total = 0
pending_dicts: list[dict] = []
for i, text in enumerate(unique_texts, start=1):
if text in cache:
match_dicts = cache[text]
cache_hits += 1
else:
match_dicts = [m.to_dict() for m in mapper.map(text)]
cache[text] = match_dicts
entity_ids = by_text[text]
for entity_id in entity_ids:
for m in match_dicts:
d = dict(m)
d["entity_id"] = entity_id
pending_dicts.append(d)
processed_entities += len(entity_ids)
if len(pending_dicts) >= commit_every or i == len(unique_texts):
inserted_total += load_icd10_mappings(pending_dicts)
pending_dicts = []
_save_cache(cache)
elapsed = time.time() - start
rate = processed_entities / elapsed if elapsed > 0 else 0
eta_seconds = (total_entities - processed_entities) / rate if rate > 0 else float("inf")
logger.info(
" %d/%d entities | %d/%d unique texts (%d cache hits) | %d mappings so far | "
"%.0fs elapsed | ETA %.0fs",
processed_entities, total_entities, i, len(unique_texts), cache_hits,
inserted_total, elapsed, eta_seconds,
)
logger.info(
"Done: %d entities processed, %d mappings inserted in %.0fs",
processed_entities, inserted_total, time.time() - start,
)
def main() -> None:
parser = argparse.ArgumentParser(
description="Map DISEASE/SYMPTOM entities to ICD-10 codes and save the results."
)
parser.add_argument(
"--limit", type=int, default=None,
help="Max number of pending entity rows to process (default: all pending entities)",
)
parser.add_argument(
"--commit-every", type=int, default=500,
help="Queue this many mapping rows before writing to the database (default: 500)",
)
args = parser.parse_args()
run(limit=args.limit, commit_every=args.commit_every)
if __name__ == "__main__":
main()
|