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