fixed code
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user