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")