362 lines
12 KiB
Python
Executable File
362 lines
12 KiB
Python
Executable File
"""
|
|
Model Manager - Hệ thống quản lý và vận hành tất cả các loại models
|
|
Hỗ trợ: XGBoost, Random Forest, Decision Tree, SVM, CNN, và các model khác
|
|
"""
|
|
|
|
import joblib
|
|
import json
|
|
from pathlib import Path
|
|
from typing import Optional, Dict, List, Any, Tuple
|
|
from datetime import datetime
|
|
import numpy as np
|
|
import warnings
|
|
|
|
# PyTorch for CNN models
|
|
try:
|
|
import torch
|
|
PYTORCH_AVAILABLE = True
|
|
except ImportError:
|
|
PYTORCH_AVAILABLE = False
|
|
|
|
warnings.filterwarnings('ignore')
|
|
|
|
|
|
class ModelManager:
|
|
"""Quản lý tất cả các models: load, save, list, validate"""
|
|
|
|
def __init__(self, models_dir: str = "model_train"):
|
|
self.models_dir = Path(models_dir)
|
|
self.models_dir.mkdir(exist_ok=True)
|
|
self.current_model = None
|
|
self.current_metadata = None
|
|
|
|
def list_models(self) -> List[Dict[str, Any]]:
|
|
"""
|
|
Liệt kê tất cả models có sẵn với metadata
|
|
|
|
Returns:
|
|
List of dicts containing model info
|
|
"""
|
|
models = []
|
|
|
|
# Tìm tất cả file .joblib
|
|
for model_file in self.models_dir.glob("*.joblib"):
|
|
# Skip Zone.Identifier files
|
|
if "Zone.Identifier" in model_file.name:
|
|
continue
|
|
|
|
model_info = {
|
|
"filename": model_file.name,
|
|
"path": str(model_file),
|
|
"size_mb": model_file.stat().st_size / (1024 * 1024),
|
|
"modified": datetime.fromtimestamp(model_file.stat().st_mtime).isoformat(),
|
|
}
|
|
|
|
# Tìm metadata file tương ứng
|
|
metadata_file = model_file.with_suffix('.json')
|
|
if not metadata_file.exists():
|
|
# Try with _info.json suffix
|
|
metadata_file = model_file.parent / (model_file.stem + "_info.json")
|
|
|
|
if metadata_file.exists():
|
|
try:
|
|
with open(metadata_file, 'r') as f:
|
|
metadata = json.load(f)
|
|
model_info["metadata"] = metadata
|
|
model_info["has_metadata"] = True
|
|
|
|
# Extract key info
|
|
model_info["model_type"] = metadata.get("model_type", "unknown")
|
|
model_info["features"] = metadata.get("features", [])
|
|
model_info["n_features"] = metadata.get("n_features", 0)
|
|
model_info["n_classes"] = metadata.get("n_classes", 0)
|
|
model_info["test_accuracy"] = metadata.get("test_accuracy", None)
|
|
model_info["timestamp"] = metadata.get("timestamp", None)
|
|
model_info["data_source"] = metadata.get("data_source", "unknown")
|
|
|
|
except Exception as e:
|
|
model_info["has_metadata"] = False
|
|
model_info["metadata_error"] = str(e)
|
|
else:
|
|
model_info["has_metadata"] = False
|
|
|
|
models.append(model_info)
|
|
|
|
# Sort by modified time (newest first)
|
|
models.sort(key=lambda x: x["modified"], reverse=True)
|
|
|
|
return models
|
|
|
|
def load_model(self, model_filename: str) -> Tuple[Any, Optional[Any], Dict[str, Any]]:
|
|
"""
|
|
Load model từ file
|
|
|
|
Args:
|
|
model_filename: Tên file model (ví dụ: "model_odc.joblib")
|
|
|
|
Returns:
|
|
Tuple of (model, label_encoder, metadata)
|
|
"""
|
|
model_path = self.models_dir / model_filename
|
|
|
|
if not model_path.exists():
|
|
raise FileNotFoundError(f"Model không tồn tại: {model_filename}")
|
|
|
|
# Load model
|
|
print(f"[MODEL MANAGER] Loading model: {model_filename}")
|
|
model_data = joblib.load(model_path)
|
|
|
|
# Extract model and encoder
|
|
if isinstance(model_data, dict):
|
|
model = model_data.get('model')
|
|
label_encoder = model_data.get('label_encoder')
|
|
else:
|
|
# Old format: model only
|
|
model = model_data
|
|
label_encoder = None
|
|
|
|
# Load metadata
|
|
metadata = self._load_metadata(model_filename)
|
|
|
|
# Store current model
|
|
self.current_model = model
|
|
self.current_metadata = metadata
|
|
|
|
# Check if CNN model and set to eval mode
|
|
if PYTORCH_AVAILABLE and hasattr(model, '__class__') and 'CNN' in model.__class__.__name__:
|
|
model.eval()
|
|
print(f"[MODEL MANAGER] PyTorch CNN model detected and set to eval mode")
|
|
|
|
print(f"[MODEL MANAGER] Model loaded successfully")
|
|
print(f" - Type: {metadata.get('model_type', 'unknown')}")
|
|
print(f" - Features: {metadata.get('n_features', 'N/A')}")
|
|
print(f" - Classes: {metadata.get('n_classes', 'N/A')}")
|
|
print(f" - Accuracy: {metadata.get('test_accuracy', 'N/A')}")
|
|
|
|
return model, label_encoder, metadata
|
|
|
|
def _load_metadata(self, model_filename: str) -> Dict[str, Any]:
|
|
"""Load metadata cho model"""
|
|
model_path = self.models_dir / model_filename
|
|
|
|
# Try multiple metadata file patterns
|
|
metadata_files = [
|
|
model_path.with_suffix('.json'),
|
|
model_path.parent / (model_path.stem + "_info.json"),
|
|
]
|
|
|
|
for metadata_file in metadata_files:
|
|
if metadata_file.exists():
|
|
try:
|
|
with open(metadata_file, 'r') as f:
|
|
return json.load(f)
|
|
except Exception as e:
|
|
print(f"[MODEL MANAGER] Warning: Could not load metadata from {metadata_file}: {e}")
|
|
|
|
# Return default metadata if not found
|
|
print(f"[MODEL MANAGER] Warning: No metadata found for {model_filename}")
|
|
return {
|
|
"model_type": "unknown",
|
|
"features": [],
|
|
"n_features": 0,
|
|
"n_classes": 0,
|
|
"timestamp": None
|
|
}
|
|
|
|
def save_model(self, model: Any, metadata: Dict[str, Any],
|
|
model_filename: Optional[str] = None,
|
|
label_encoder: Optional[Any] = None) -> str:
|
|
"""
|
|
Save model với metadata
|
|
|
|
Args:
|
|
model: Model object
|
|
metadata: Dict chứa thông tin về model
|
|
model_filename: Tên file (optional, sẽ auto-generate nếu không có)
|
|
label_encoder: Label encoder (optional)
|
|
|
|
Returns:
|
|
Path to saved model file
|
|
"""
|
|
# Generate filename if not provided
|
|
if model_filename is None:
|
|
model_type = metadata.get("model_type", "model")
|
|
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
|
model_filename = f"model_{model_type}_{timestamp}.joblib"
|
|
|
|
model_path = self.models_dir / model_filename
|
|
metadata_path = model_path.parent / (model_path.stem + "_info.json")
|
|
|
|
# Prepare model data
|
|
if label_encoder is not None:
|
|
model_data = {
|
|
'model': model,
|
|
'label_encoder': label_encoder
|
|
}
|
|
else:
|
|
model_data = {
|
|
'model': model
|
|
}
|
|
|
|
# Save model
|
|
print(f"[MODEL MANAGER] Saving model to: {model_path}")
|
|
joblib.dump(model_data, model_path)
|
|
|
|
# Save metadata
|
|
print(f"[MODEL MANAGER] Saving metadata to: {metadata_path}")
|
|
with open(metadata_path, 'w') as f:
|
|
json.dump(metadata, f, indent=2)
|
|
|
|
print(f"[MODEL MANAGER] Model saved successfully!")
|
|
|
|
return str(model_path)
|
|
|
|
def validate_model(self, model_filename: str) -> Dict[str, Any]:
|
|
"""
|
|
Validate model file và kiểm tra integrity
|
|
|
|
Returns:
|
|
Dict with validation results
|
|
"""
|
|
result = {
|
|
"valid": False,
|
|
"errors": [],
|
|
"warnings": []
|
|
}
|
|
|
|
model_path = self.models_dir / model_filename
|
|
|
|
# Check file exists
|
|
if not model_path.exists():
|
|
result["errors"].append(f"File không tồn tại: {model_filename}")
|
|
return result
|
|
|
|
# Try to load model
|
|
try:
|
|
model, encoder, metadata = self.load_model(model_filename)
|
|
result["valid"] = True
|
|
|
|
# Check metadata
|
|
if not metadata or metadata.get("model_type") == "unknown":
|
|
result["warnings"].append("Không có metadata hoặc metadata không đầy đủ")
|
|
|
|
# Check required features
|
|
if not metadata.get("features"):
|
|
result["warnings"].append("Danh sách features không có trong metadata")
|
|
|
|
# Check model object
|
|
if model is None:
|
|
result["errors"].append("Model object is None")
|
|
result["valid"] = False
|
|
|
|
except Exception as e:
|
|
result["errors"].append(f"Lỗi khi load model: {str(e)}")
|
|
result["valid"] = False
|
|
|
|
return result
|
|
|
|
def get_required_features(self, model_filename: str) -> List[str]:
|
|
"""
|
|
Lấy danh sách features cần thiết cho model
|
|
|
|
Returns:
|
|
List of feature names
|
|
"""
|
|
metadata = self._load_metadata(model_filename)
|
|
return metadata.get("features", [])
|
|
|
|
def predict(self, model_filename: str, X: np.ndarray) -> np.ndarray:
|
|
"""
|
|
Predict using specified model
|
|
|
|
Args:
|
|
model_filename: Model file name
|
|
X: Features array (n_samples, n_features)
|
|
|
|
Returns:
|
|
Predictions array
|
|
"""
|
|
if self.current_model is None or model_filename != getattr(self, '_current_model_filename', None):
|
|
model, encoder, metadata = self.load_model(model_filename)
|
|
self._current_model_filename = model_filename
|
|
else:
|
|
model = self.current_model
|
|
metadata = self.current_metadata
|
|
|
|
# Validate input features
|
|
expected_features = metadata.get("n_features", 0)
|
|
if X.shape[1] != expected_features:
|
|
raise ValueError(f"Expected {expected_features} features, got {X.shape[1]}")
|
|
|
|
# Predict
|
|
predictions = model.predict(X)
|
|
|
|
return predictions
|
|
|
|
def get_model_info(self, model_filename: str) -> Dict[str, Any]:
|
|
"""Get detailed info about a model"""
|
|
models = self.list_models()
|
|
for model in models:
|
|
if model["filename"] == model_filename:
|
|
return model
|
|
return None
|
|
|
|
def delete_model(self, model_filename: str) -> bool:
|
|
"""
|
|
Xóa model và metadata
|
|
|
|
Returns:
|
|
True if successful
|
|
"""
|
|
model_path = self.models_dir / model_filename
|
|
|
|
if not model_path.exists():
|
|
return False
|
|
|
|
# Delete model file
|
|
model_path.unlink()
|
|
|
|
# Delete metadata file if exists
|
|
metadata_file = model_path.with_suffix('.json')
|
|
if metadata_file.exists():
|
|
metadata_file.unlink()
|
|
|
|
# Try alternative metadata file name
|
|
metadata_file_alt = model_path.parent / (model_path.stem + "_info.json")
|
|
if metadata_file_alt.exists():
|
|
metadata_file_alt.unlink()
|
|
|
|
return True
|
|
|
|
def get_latest_model(self, model_type: Optional[str] = None) -> Optional[str]:
|
|
"""
|
|
Lấy model mới nhất (theo thời gian modified)
|
|
|
|
Args:
|
|
model_type: Filter by model type (xgboost, cnn, etc.), None for any
|
|
|
|
Returns:
|
|
Model filename or None
|
|
"""
|
|
models = self.list_models()
|
|
|
|
if model_type:
|
|
models = [m for m in models if m.get("model_type") == model_type]
|
|
|
|
if not models:
|
|
return None
|
|
|
|
# Already sorted by modified time
|
|
return models[0]["filename"]
|
|
|
|
|
|
# Singleton instance
|
|
_model_manager = None
|
|
|
|
def get_model_manager() -> ModelManager:
|
|
"""Get singleton ModelManager instance"""
|
|
global _model_manager
|
|
if _model_manager is None:
|
|
_model_manager = ModelManager()
|
|
return _model_manager
|