Spaces:
Runtime error
Runtime error
| """ | |
| Carga y gestión de los 3 modelos (OCR, Traducción, Descripciones). | |
| """ | |
| import torch | |
| from transformers import ( | |
| LightOnOcrForConditionalGeneration, | |
| LightOnOcrProcessor, | |
| AutoModelForCausalLM, | |
| AutoTokenizer, | |
| AutoModelForSeq2SeqLM, | |
| ) | |
| from config import ( | |
| DEVICE, DTYPE, | |
| OCR_MODEL_ID, TRANSLATION_MODEL_ID, DESCRIPTION_MODEL_ID, | |
| LANG_CODES, | |
| ) | |
| class ModelManager: | |
| """Gestiona la carga y acceso a los 3 modelos.""" | |
| def __init__(self): | |
| self.device = DEVICE | |
| self.dtype = DTYPE | |
| self.models_loaded = False | |
| # OCR | |
| self.ocr_processor = None | |
| self.ocr_model = None | |
| # Traducción | |
| self.translation_tokenizer = None | |
| self.translation_model = None | |
| # Descripciones | |
| self.description_tokenizer = None | |
| self.description_model = None | |
| def load_all(self): | |
| """Carga los 3 modelos en memoria.""" | |
| print("\n" + "=" * 70) | |
| print("CARGANDO MODELOS") | |
| print("=" * 70) | |
| try: | |
| self._load_ocr() | |
| self._load_translation() | |
| self._load_description() | |
| self.models_loaded = True | |
| print("\n" + "=" * 70) | |
| print("✅ TODOS LOS MODELOS CARGADOS EXITOSAMENTE") | |
| print("=" * 70 + "\n") | |
| except Exception as e: | |
| print(f"\n❌ ERROR cargando modelos: {e}") | |
| raise | |
| def _load_ocr(self): | |
| print(f"\n📷 [1/3] Cargando modelo OCR: {OCR_MODEL_ID}") | |
| self.ocr_processor = LightOnOcrProcessor.from_pretrained(OCR_MODEL_ID) | |
| self.ocr_model = LightOnOcrForConditionalGeneration.from_pretrained( | |
| OCR_MODEL_ID, | |
| torch_dtype=self.dtype, | |
| device_map="auto", | |
| low_cpu_mem_usage=True, | |
| ) | |
| print(" ✅ OCR cargado correctamente") | |
| def _load_translation(self): | |
| print(f"\n🌐 [2/3] Cargando modelo de traducción: {TRANSLATION_MODEL_ID}") | |
| self.translation_tokenizer = AutoTokenizer.from_pretrained( | |
| TRANSLATION_MODEL_ID, | |
| src_lang=LANG_CODES["Spanish"], | |
| ) | |
| self.translation_model = AutoModelForSeq2SeqLM.from_pretrained( | |
| TRANSLATION_MODEL_ID, | |
| torch_dtype=self.dtype, | |
| device_map="auto", | |
| low_cpu_mem_usage=True, | |
| ) | |
| print(" ✅ Traducción cargada correctamente") | |
| def _load_description(self): | |
| print(f"\n📝 [3/3] Cargando modelo de descripciones: {DESCRIPTION_MODEL_ID}") | |
| self.description_tokenizer = AutoTokenizer.from_pretrained(DESCRIPTION_MODEL_ID) | |
| self.description_model = AutoModelForCausalLM.from_pretrained( | |
| DESCRIPTION_MODEL_ID, | |
| torch_dtype=self.dtype, | |
| device_map="auto", | |
| low_cpu_mem_usage=True, | |
| ) | |
| if self.description_tokenizer.pad_token is None: | |
| self.description_tokenizer.pad_token = self.description_tokenizer.eos_token | |
| print(" ✅ Descripciones cargadas correctamente") | |