Spaces:
Sleeping
Sleeping
File size: 3,895 Bytes
e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b b2e9550 e16db8b | 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 | #!/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")
|