#!/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")