fixed code

This commit is contained in:
SlimusMinus
2026-08-05 01:44:29 +03:00
parent 72159dbe4a
commit d33f0eaa78
30 changed files with 143 additions and 375 deletions

View File

@@ -1,89 +1,69 @@
from transformers import AutoTokenizer
from transformers import AutoModelForSequenceClassification
from app.config.settings import settings
from app.config.logging import logger
import torch
from transformers import AutoImageProcessor
from transformers import AutoModelForImageClassification
from transformers import AutoModelForSequenceClassification
from transformers import AutoTokenizer
import torch
from app.config.logging import logger
from app.config.settings import settings
class ModelManager:
def __init__(self):
self.device = None
self.text_model = None
self.text_tokenizer = None
self.image_model = None
self.image_processor = None
def __init__(self):
self.device = None
self.text_model = None
self.text_tokenizer = None
self.image_model = None
self.image_processor = 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_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...")
def load_tokenizer(self):
logger.info("Loading tokenizer...")
self.text_tokenizer = AutoTokenizer.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
logger.info("Tokenizer loaded.")
self.text_tokenizer = AutoTokenizer.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
def load_model(self):
logger.info("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()
logger.debug("id2label: %s", self.text_model.config.id2label)
logger.info("Text model loaded.")
print("Tokenizer loaded.")
def load_image_model(self):
logger.info("Loading image model...")
self.image_processor = AutoImageProcessor.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.image_model = AutoModelForImageClassification.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.image_model.to(self.device)
self.image_model.eval()
logger.info("Image model 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()
self.load_image_model()
def load_image_model(self):
print("Loading image model...")
def initialize(self):
self.load_device()
logger.info(f"Using device: {self.device}")
self.load_tokenizer()
self.load_model()
self.load_image_model()
self.image_processor = (
AutoImageProcessor.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
)
self.image_model = (
AutoModelForImageClassification.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
)
self.image_model.to(self.device)
self.image_model.eval()
print("Image model loaded.")
model_manager = ModelManager()
model_manager = ModelManager()