56 lines
1.3 KiB
Python
56 lines
1.3 KiB
Python
from transformers import AutoTokenizer
|
|
from transformers import AutoModelForSequenceClassification
|
|
|
|
from app.config.settings import settings
|
|
from app.config.logging import logger
|
|
|
|
import torch
|
|
|
|
|
|
class ModelManager:
|
|
|
|
def __init__(self):
|
|
self.device = None
|
|
self.text_model = None
|
|
self.text_tokenizer = None
|
|
|
|
def load_device(self):
|
|
if settings.DEVICE == "auto":
|
|
self.device = torch.device(
|
|
"cuda" if torch.cuda.is_available() else "cpu"
|
|
)
|
|
else:
|
|
self.device = torch.device(settings.DEVICE)
|
|
|
|
def load_tokenizer(self):
|
|
print("Loading tokenizer...")
|
|
|
|
self.text_tokenizer = AutoTokenizer.from_pretrained(
|
|
settings.TEXT_MODEL,
|
|
cache_dir=settings.MODEL_CACHE_DIR
|
|
)
|
|
|
|
print("Tokenizer loaded.")
|
|
|
|
def load_model(self):
|
|
print("Loading model...")
|
|
|
|
self.text_model = AutoModelForSequenceClassification.from_pretrained(
|
|
settings.TEXT_MODEL,
|
|
cache_dir=settings.MODEL_CACHE_DIR
|
|
)
|
|
|
|
self.text_model.to(self.device)
|
|
self.text_model.eval()
|
|
print(self.text_model.config.id2label)
|
|
|
|
print("Text model loaded.")
|
|
|
|
def initialize(self):
|
|
self.load_device()
|
|
print(f"Using device: {self.device}")
|
|
self.load_tokenizer()
|
|
self.load_model()
|
|
|
|
|
|
model_manager = ModelManager() |