wave2notes / models /model_loader.py
Razvanix's picture
Upload 13 files
e6ed91e verified
Raw
History Blame Contribute Delete
2.51 kB
import os
import tensorflow as tf
import threading
class ModelLoader:
def __init__(self):
# Model is stored directly in the Space
self.local_model_path = "models/saved_model"
self.model = None
self._model_ready = False
self._lock = threading.Lock() # Add thread safety for web/mobile concurrent requests
print("ModelLoader initialized - ready for lazy loading")
def get_model(self):
"""Get model, loading it if necessary (lazy loading)"""
with self._lock: # Prevent race conditions between web/mobile requests
if self.model is None or not self._model_ready:
print("Model not loaded yet - loading now...")
self.load_model()
return self.model
def load_model(self):
"""Load model from local files"""
try:
print(f"Loading SavedModel from: {self.local_model_path}")
# Check if model exists
if not os.path.exists(self.local_model_path):
raise FileNotFoundError(f"Model not found at: {self.local_model_path}")
# Clear any existing model first (important for resets)
self.model = None
self._model_ready = False
# Load the model
loaded_model = tf.saved_model.load(self.local_model_path)
# Add predict wrapper
def predict_wrapper(input_data):
return loaded_model(input_data)
loaded_model.predict = predict_wrapper
self.model = loaded_model
self._model_ready = True
print("βœ… Model loaded successfully from local files")
except Exception as e:
print(f"❌ Error loading model: {e}")
self.model = None
self._model_ready = False
raise
def reset(self):
"""Reset the model loader state - fixes corruption issues"""
with self._lock:
print("πŸ”„ Resetting ModelLoader state...")
self.model = None
self._model_ready = False
# Force garbage collection to free memory
try:
import gc
gc.collect()
except:
pass
print("βœ… ModelLoader state reset complete")
def is_model_ready(self):
"""Check if model is loaded and ready"""
with self._lock:
return self.model is not None and self._model_ready