debpc
Deploy claims-first v97 extraction and validation
b2e9550
Raw
History Blame Contribute Delete
3.9 kB
#!/usr/bin/env python3
"""Model loading and lifecycle management for MedCPT and PubMedBERT-NLI."""
from typing import Dict, Any
import torch
from transformers import AutoModel, AutoModelForSequenceClassification, AutoTokenizer
from config import (
DEVICE,
LOAD_CROSS_ENCODER,
MEDCPT_ARTICLE_MODEL,
MEDCPT_CROSS_MODEL,
MEDCPT_QUERY_MODEL,
NLI_MODEL,
)
from logger import setup_logger
logger = setup_logger("models.loader")
_CACHE: Dict[str, Any] = {}
def _place_model(model):
model.to(DEVICE)
model.eval()
return model
def load_medcpt_query_tokenizer():
if "medcpt_query_tokenizer" not in _CACHE:
_CACHE["medcpt_query_tokenizer"] = AutoTokenizer.from_pretrained(
MEDCPT_QUERY_MODEL
)
return _CACHE["medcpt_query_tokenizer"]
def load_medcpt_query_model():
if "medcpt_query_model" not in _CACHE:
logger.info("Loading MedCPT query encoder from %s", MEDCPT_QUERY_MODEL)
_CACHE["medcpt_query_model"] = _place_model(
AutoModel.from_pretrained(MEDCPT_QUERY_MODEL)
)
return _CACHE["medcpt_query_model"]
def load_medcpt_article_tokenizer():
if "medcpt_article_tokenizer" not in _CACHE:
_CACHE["medcpt_article_tokenizer"] = AutoTokenizer.from_pretrained(
MEDCPT_ARTICLE_MODEL
)
return _CACHE["medcpt_article_tokenizer"]
def load_medcpt_article_model():
if "medcpt_article_model" not in _CACHE:
logger.info("Loading MedCPT article encoder from %s", MEDCPT_ARTICLE_MODEL)
_CACHE["medcpt_article_model"] = _place_model(
AutoModel.from_pretrained(MEDCPT_ARTICLE_MODEL)
)
return _CACHE["medcpt_article_model"]
def load_medcpt_cross_tokenizer():
if "medcpt_cross_tokenizer" not in _CACHE:
_CACHE["medcpt_cross_tokenizer"] = AutoTokenizer.from_pretrained(
MEDCPT_CROSS_MODEL
)
return _CACHE["medcpt_cross_tokenizer"]
def load_medcpt_cross_model():
if "medcpt_cross_model" not in _CACHE:
logger.info("Loading MedCPT cross encoder from %s", MEDCPT_CROSS_MODEL)
_CACHE["medcpt_cross_model"] = _place_model(
AutoModelForSequenceClassification.from_pretrained(MEDCPT_CROSS_MODEL)
)
return _CACHE["medcpt_cross_model"]
def load_nli_model():
if "nli_model" not in _CACHE:
logger.info("Loading PubMedBERT-NLI from %s", NLI_MODEL)
_CACHE["nli_model"] = _place_model(
AutoModelForSequenceClassification.from_pretrained(NLI_MODEL)
)
return _CACHE["nli_model"]
def load_nli_tokenizer():
if "nli_tokenizer" not in _CACHE:
_CACHE["nli_tokenizer"] = AutoTokenizer.from_pretrained(NLI_MODEL)
return _CACHE["nli_tokenizer"]
def load_all():
"""Eagerly load the dual encoders and NLI model; cross encoder stays optional."""
load_medcpt_query_tokenizer()
load_medcpt_query_model()
load_medcpt_article_tokenizer()
load_medcpt_article_model()
load_nli_tokenizer()
load_nli_model()
if LOAD_CROSS_ENCODER:
load_medcpt_cross_tokenizer()
load_medcpt_cross_model()
logger.info("Core models loaded")
def models_ready() -> bool:
required = {
"medcpt_query_tokenizer",
"medcpt_query_model",
"medcpt_article_tokenizer",
"medcpt_article_model",
"nli_tokenizer",
"nli_model",
}
return required.issubset(_CACHE)
def get_info() -> Dict[str, bool]:
return {
"medcpt_query": "medcpt_query_model" in _CACHE,
"medcpt_article": "medcpt_article_model" in _CACHE,
"medcpt_cross": "medcpt_cross_model" in _CACHE,
"nli_model": "nli_model" in _CACHE,
"nli_tokenizer": "nli_tokenizer" in _CACHE,
}
def cleanup():
_CACHE.clear()
if torch.cuda.is_available():
torch.cuda.empty_cache()
logger.info("Models unloaded")