Download utils.py from FannyFa/FannyFa-Model-V1: direct link, hf CLI and curl.
- Browser
- Download file 27.5 kB
-
https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/utils.py
- Command line
-
hf download hf://FannyFa/FannyFa-Model-V1/utils.py
-
curl -L -o utils.py https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/utils.py
27.5 kB
| import os | |
| import sys | |
| import json | |
| import logging | |
| import time | |
| import gc | |
| import resource | |
| import platform | |
| import threading | |
| import signal | |
| import hashlib | |
| import pickle | |
| import shutil | |
| from pathlib import Path | |
| from typing import Dict, List, Any, Optional, Callable, Generator, Tuple, Type | |
| from contextlib import contextmanager | |
| from functools import wraps | |
| from collections import deque | |
| import torch | |
| from rich.console import Console | |
| from rich.theme import Theme as RichTheme | |
| from rich.panel import Panel | |
| from rich.prompt import Prompt, Confirm | |
| from rich.table import Table | |
| from rich import box | |
| from rich.markup import escape as rich_escape | |
| from config import ( | |
| LOG_DIR, | |
| CACHE_DIR, | |
| ERROR_LOG_FILE, | |
| TRAINING_LOG_FILE, | |
| CHAT_LOG_FILE, | |
| SYSTEM_LOG_FILE, | |
| EVAL_LOG_FILE, | |
| TEST_LOG_FILE, | |
| DEBUG_LOG_FILE, | |
| PERFORMANCE_FILE, | |
| DEFAULT_MEMORY_LIMIT_GB, | |
| ) | |
| console = Console() | |
| class Theme: | |
| PRIMARY = "cyan" | |
| SECONDARY = "green" | |
| SUCCESS = "green" | |
| WARNING = "yellow" | |
| ERROR = "red" | |
| INFO = "blue" | |
| DIM = "dim" | |
| WHITE = "white" | |
| MAGENTA = "magenta" | |
| GOLD = "gold1" | |
| PURPLE = "purple" | |
| ORANGE = "orange1" | |
| PINK = "pink1" | |
| TEAL = "teal" | |
| VIOLET = "violet" | |
| CRIMSON = "crimson" | |
| LIME = "lime" | |
| OLIVE = "olive" | |
| INDIGO = "indigo" | |
| CORAL = "coral" | |
| SALMON = "salmon" | |
| TURQUOISE = "turquoise2" | |
| SKY_BLUE = "sky_blue1" | |
| STEEL_BLUE = "steel_blue" | |
| def header(text: str, color: str = "cyan") -> str: | |
| return f"[bold {color}]{text}[/bold {color}]" | |
| def success(text: str) -> str: | |
| return f"[green]{text}[/green]" | |
| def warning(text: str) -> str: | |
| return f"[yellow]{text}[/yellow]" | |
| def error(text: str) -> str: | |
| return f"[red]{text}[/red]" | |
| def info(text: str) -> str: | |
| return f"[blue]{text}[/blue]" | |
| def dim(text: str) -> str: | |
| return f"[dim]{text}[/dim]" | |
| def gold(text: str) -> str: | |
| return f"[gold1]{text}[/gold1]" | |
| def primary(text: str) -> str: | |
| return f"[cyan]{text}[/cyan]" | |
| def secondary(text: str) -> str: | |
| return f"[green]{text}[/green]" | |
| def magenta(text: str) -> str: | |
| return f"[magenta]{text}[/magenta]" | |
| def purple(text: str) -> str: | |
| return f"[purple]{text}[/purple]" | |
| def orange(text: str) -> str: | |
| return f"[orange1]{text}[/orange1]" | |
| def teal(text: str) -> str: | |
| return f"[teal]{text}[/teal]" | |
| def lime(text: str) -> str: | |
| return f"[lime]{text}[/lime]" | |
| def coral(text: str) -> str: | |
| return f"[coral]{text}[/coral]" | |
| def box(text: str, color: str = "cyan") -> str: | |
| return f"[{color}]{text}[/{color}]" | |
| def metric(name: str, value: Any) -> str: | |
| return f"[cyan]{name}:[/cyan] [green]{value}[/green]" | |
| def progress_bar(percent: float) -> str: | |
| bar_len = 20 | |
| filled = int(bar_len * percent / 100) | |
| empty = bar_len - filled | |
| return f"[green]{'█' * filled}[/green][dim]{'░' * empty}[/dim]" | |
| class EnhancedLogger: | |
| def __init__(self): | |
| self.log_dir = Path(LOG_DIR) | |
| self.log_dir.mkdir(exist_ok=True) | |
| self.formatter = logging.Formatter( | |
| "%(asctime)s - %(name)s - %(levelname)s - %(message)s", | |
| datefmt="%Y-%m-%d %H:%M:%S", | |
| ) | |
| self.json_formatter = logging.Formatter( | |
| '{"time": "%(asctime)s", "name": "%(name)s", "level": "%(levelname)s", "message": "%(message)s"}', | |
| datefmt="%Y-%m-%dT%H:%M:%S", | |
| ) | |
| self.loggers: Dict[str, logging.Logger] = {} | |
| self.handlers: Dict[str, logging.FileHandler] = {} | |
| self._create_logger("training", TRAINING_LOG_FILE, logging.INFO) | |
| self._create_logger("error", ERROR_LOG_FILE, logging.ERROR) | |
| self._create_logger("debug", DEBUG_LOG_FILE, logging.DEBUG) | |
| self._create_logger("chat", CHAT_LOG_FILE, logging.INFO) | |
| self._create_logger("system", SYSTEM_LOG_FILE, logging.INFO) | |
| self._create_logger("eval", EVAL_LOG_FILE, logging.INFO) | |
| self._create_logger("test", TEST_LOG_FILE, logging.INFO) | |
| self._create_logger("performance", PERFORMANCE_FILE, logging.INFO) | |
| console_handler = logging.StreamHandler(sys.stdout) | |
| console_handler.setFormatter(self.formatter) | |
| console_handler.setLevel(logging.INFO) | |
| for logger in self.loggers.values(): | |
| logger.addHandler(console_handler) | |
| def _create_logger(self, name: str, filename: str, level: int) -> logging.Logger: | |
| logger = logging.getLogger(name) | |
| logger.setLevel(level) | |
| file_handler = logging.FileHandler(self.log_dir / filename) | |
| file_handler.setFormatter( | |
| self.json_formatter if name == "performance" else self.formatter | |
| ) | |
| file_handler.setLevel(level) | |
| if not logger.handlers: | |
| logger.addHandler(file_handler) | |
| self.handlers[name] = file_handler | |
| self.loggers[name] = logger | |
| return logger | |
| def get(self, name: str) -> logging.Logger: | |
| if name not in self.loggers: | |
| return self._create_logger(name, f"{name}.log", logging.INFO) | |
| return self.loggers[name] | |
| def log_performance(self, event: str, **kwargs) -> None: | |
| perf_logger = self.get("performance") | |
| metrics = {"event": event, "timestamp": time.time(), **kwargs} | |
| perf_logger.info(json.dumps(metrics)) | |
| log_system = EnhancedLogger() | |
| train_logger = log_system.get("training") | |
| error_logger = log_system.get("error") | |
| debug_logger = log_system.get("debug") | |
| chat_logger = log_system.get("chat") | |
| system_logger = log_system.get("system") | |
| eval_logger = log_system.get("eval") | |
| test_logger = log_system.get("test") | |
| perf_logger = log_system.get("performance") | |
| class GracefulExit: | |
| _instance: Optional["GracefulExit"] = None | |
| def __new__(cls) -> "GracefulExit": | |
| if cls._instance is None: | |
| cls._instance = super().__new__(cls) | |
| cls._instance._initialized = False | |
| return cls._instance | |
| def __init__(self) -> None: | |
| if self._initialized: | |
| return | |
| self._initialized = True | |
| self.should_exit = False | |
| self.exit_code = 0 | |
| self._cleanup_hooks: List[Callable] = [] | |
| self._context_managers: List[Any] = [] | |
| self._start_time = time.time() | |
| self._total_exits = 0 | |
| signal.signal(signal.SIGINT, self._signal_handler) | |
| signal.signal(signal.SIGTERM, self._signal_handler) | |
| import atexit | |
| atexit.register(self._atexit_cleanup) | |
| system_logger.info("Graceful exit handler initialized") | |
| def _signal_handler(self, signum: int, frame: Any) -> None: | |
| elapsed = time.time() - self._start_time | |
| self._total_exits += 1 | |
| if self.should_exit or self._total_exits > 2: | |
| print(f"\n[red] Force exit (signal {signum}, elapsed {elapsed:.1f}s)[/red]") | |
| self._force_exit() | |
| return | |
| self.should_exit = True | |
| print( | |
| f"\n[yellow] Received interrupt signal (signal {signum}). Cleaning up...[/yellow]" | |
| ) | |
| for hook in self._cleanup_hooks: | |
| try: | |
| hook() | |
| except Exception as e: | |
| print(f"[red]Cleanup hook error: {e}[/red]") | |
| error_logger.error(f"Cleanup hook error: {e}") | |
| def _force_exit(self) -> None: | |
| try: | |
| for hook in self._cleanup_hooks: | |
| try: | |
| hook() | |
| except Exception: | |
| pass | |
| except Exception: | |
| pass | |
| finally: | |
| sys.exit(1) | |
| def _atexit_cleanup(self) -> None: | |
| if not self.should_exit: | |
| return | |
| for hook in self._cleanup_hooks: | |
| try: | |
| hook() | |
| except Exception: | |
| pass | |
| def register_cleanup(self, hook: Callable) -> None: | |
| self._cleanup_hooks.append(hook) | |
| def register_context(self, ctx: Any) -> None: | |
| self._context_managers.append(ctx) | |
| def __enter__(self) -> "GracefulExit": | |
| return self | |
| def __exit__( | |
| self, | |
| exc_type: Optional[Type[BaseException]], | |
| exc_val: Optional[BaseException], | |
| exc_tb: Optional[Any], | |
| ) -> None: | |
| for hook in self._cleanup_hooks: | |
| try: | |
| hook() | |
| except Exception as e: | |
| error_logger.error(f"Cleanup hook error: {e}") | |
| for ctx in self._context_managers: | |
| try: | |
| ctx.__exit__(exc_type, exc_val, exc_tb) | |
| except Exception: | |
| pass | |
| if exc_type is not None and exc_type is not KeyboardInterrupt: | |
| error_logger.error(f"Exit with exception: {exc_type.__name__}: {exc_val}") | |
| graceful_exit = GracefulExit() | |
| def timer(name: str, logger: logging.Logger = None) -> Generator: | |
| start = time.perf_counter() | |
| try: | |
| yield | |
| finally: | |
| elapsed = time.perf_counter() - start | |
| msg = f"{name} completed in {elapsed:.2f}s" | |
| if logger: | |
| logger.info(msg) | |
| else: | |
| debug_logger.debug(msg) | |
| def memory_tracker(threshold_gb: float = 1.0) -> Generator: | |
| try: | |
| import psutil | |
| HAS_PSUTIL = True | |
| except ImportError: | |
| HAS_PSUTIL = False | |
| if not HAS_PSUTIL: | |
| yield | |
| return | |
| process = psutil.Process(os.getpid()) | |
| start_mem = process.memory_info().rss / (1024**3) | |
| try: | |
| yield | |
| finally: | |
| end_mem = process.memory_info().rss / (1024**3) | |
| diff = end_mem - start_mem | |
| if diff > threshold_gb: | |
| debug_logger.warning(f"Memory increased by {diff:.2f} GB") | |
| debug_logger.debug( | |
| f"Memory: {start_mem:.2f} GB -> {end_mem:.2f} GB ( {diff:.2f} GB)" | |
| ) | |
| def safe_file_operation( | |
| filepath: str, mode: str = "r", backup: bool = True | |
| ) -> Generator: | |
| if os.path.exists(filepath) and backup: | |
| backup_path = f"{filepath}.bak" | |
| try: | |
| shutil.copy2(filepath, backup_path) | |
| debug_logger.debug(f"Backup created: {backup_path}") | |
| except Exception as e: | |
| debug_logger.debug(f"Backup failed: {e}") | |
| try: | |
| with open(filepath, mode, encoding="utf-8") as f: | |
| yield f | |
| except Exception as e: | |
| error_logger.error(f"File operation failed: {e}") | |
| raise | |
| def device_context(device: str = None) -> Generator: | |
| if device is None: | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| original_device = torch.cuda.current_device() if torch.cuda.is_available() else None | |
| try: | |
| yield device | |
| finally: | |
| if original_device is not None and torch.cuda.is_available(): | |
| torch.cuda.set_device(original_device) | |
| def temporary_seed(seed: int) -> Generator: | |
| import random | |
| import numpy as np | |
| from transformers import set_seed | |
| orig_py_state = random.getstate() | |
| orig_np_state = np.random.get_state() | |
| orig_torch_state = torch.get_rng_state() | |
| orig_torch_cuda_state = ( | |
| torch.cuda.get_rng_state() if torch.cuda.is_available() else None | |
| ) | |
| try: | |
| set_seed(seed) | |
| yield | |
| finally: | |
| random.setstate(orig_py_state) | |
| np.random.set_state(orig_np_state) | |
| torch.set_rng_state(orig_torch_state) | |
| if orig_torch_cuda_state is not None and torch.cuda.is_available(): | |
| torch.cuda.set_rng_state(orig_torch_cuda_state) | |
| def timed(func: Callable) -> Callable: | |
| def wrapper(*args, **kwargs): | |
| start = time.perf_counter() | |
| result = func(*args, **kwargs) | |
| elapsed = time.perf_counter() - start | |
| debug_logger.debug(f"{func.__name__} took {elapsed:.2f}s") | |
| return result | |
| return wrapper | |
| def retry( | |
| max_attempts: int = 3, | |
| delay: float = 1.0, | |
| backoff: float = 2.0, | |
| exceptions: Tuple[Type[Exception]] = (Exception,), | |
| ) -> Callable: | |
| def decorator(func: Callable) -> Callable: | |
| def wrapper(*args, **kwargs): | |
| last_exception = None | |
| current_delay = delay | |
| for attempt in range(max_attempts): | |
| try: | |
| return func(*args, **kwargs) | |
| except exceptions as e: | |
| last_exception = e | |
| if attempt < max_attempts - 1: | |
| debug_logger.debug( | |
| f"Retry {attempt + 1}/{max_attempts} for {func.__name__}: {e}" | |
| ) | |
| time.sleep(current_delay) | |
| current_delay *= backoff | |
| else: | |
| debug_logger.error( | |
| f"All {max_attempts} attempts failed for {func.__name__}: {e}" | |
| ) | |
| raise last_exception | |
| return wrapper | |
| return decorator | |
| def suppress_errors(logger: logging.Logger = None, fallback: Any = None) -> Callable: | |
| def decorator(func: Callable) -> Callable: | |
| def wrapper(*args, **kwargs): | |
| try: | |
| return func(*args, **kwargs) | |
| except Exception as e: | |
| if logger: | |
| logger.error(f"Error in {func.__name__}: {e}") | |
| debug_logger.debug(f"Suppressed error in {func.__name__}: {e}") | |
| return fallback | |
| return wrapper | |
| return decorator | |
| def set_memory_hard_limit(limit_gb: float) -> bool: | |
| if platform.system().lower() != "linux": | |
| debug_logger.debug("RLIMIT_AS not supported on non-Linux systems") | |
| return False | |
| if limit_gb <= 0: | |
| return False | |
| limit_bytes = int(limit_gb * (1024**3)) | |
| try: | |
| soft, hard = resource.getrlimit(resource.RLIMIT_AS) | |
| new_hard = hard if hard == resource.RLIM_INFINITY else min(hard, limit_bytes) | |
| resource.setrlimit(resource.RLIMIT_AS, (limit_bytes, new_hard)) | |
| system_logger.info(f"Memory limit set to {limit_gb} GB") | |
| return True | |
| except (ValueError, OSError) as e: | |
| error_logger.error(f"Failed to set RLIMIT_AS: {e}") | |
| return False | |
| def get_memory_usage() -> Dict[str, float]: | |
| try: | |
| import psutil | |
| except ImportError: | |
| return {} | |
| try: | |
| process = psutil.Process(os.getpid()) | |
| mem = process.memory_info() | |
| return { | |
| "rss_gb": mem.rss / (1024**3), | |
| "vms_gb": mem.vms / (1024**3), | |
| "shared_gb": getattr(mem, "shared", 0) / (1024**3), | |
| "text_gb": getattr(mem, "text", 0) / (1024**3), | |
| "data_gb": getattr(mem, "data", 0) / (1024**3), | |
| "lib_gb": getattr(mem, "lib", 0) / (1024**3), | |
| "dirty_gb": getattr(mem, "dirty", 0) / (1024**3), | |
| "percent": process.memory_percent(), | |
| "num_threads": process.num_threads(), | |
| } | |
| except Exception as e: | |
| error_logger.error(f"Memory usage error: {e}") | |
| return {} | |
| class MemoryWatchdog: | |
| def __init__( | |
| self, | |
| limit_gb: float = DEFAULT_MEMORY_LIMIT_GB, | |
| warn_threshold_ratio: float = 0.80, | |
| critical_threshold_ratio: float = 0.95, | |
| check_interval_sec: float = 2.0, | |
| step_getter: Optional[Callable[[], int]] = None, | |
| on_warning: Optional[Callable] = None, | |
| on_critical: Optional[Callable] = None, | |
| ): | |
| self.limit_bytes = limit_gb * (1024**3) | |
| self.warn_threshold_bytes = self.limit_bytes * warn_threshold_ratio | |
| self.critical_threshold_bytes = self.limit_bytes * critical_threshold_ratio | |
| self.check_interval_sec = check_interval_sec | |
| self.step_getter = step_getter or (lambda: -1) | |
| self.on_warning = on_warning | |
| self.on_critical = on_critical | |
| self._stop_event = threading.Event() | |
| self._thread: Optional[threading.Thread] = None | |
| self._warning_count = 0 | |
| self._critical_count = 0 | |
| self._max_warnings = 5 | |
| self._history: List[Dict] = [] | |
| self.is_running = False | |
| self._last_stats = {} | |
| self._peak_rss_gb = 0.0 | |
| self._peak_vms_gb = 0.0 | |
| self.stats = { | |
| "rss_samples": deque(maxlen=1000), | |
| "vms_samples": deque(maxlen=1000), | |
| "cpu_samples": deque(maxlen=1000), | |
| "timestamps": deque(maxlen=1000), | |
| "steps": deque(maxlen=1000), | |
| } | |
| def _monitor_loop(self) -> None: | |
| try: | |
| import psutil | |
| except ImportError: | |
| return | |
| process = psutil.Process(os.getpid()) | |
| while not self._stop_event.is_set(): | |
| try: | |
| mem_info = process.memory_info() | |
| rss_bytes = mem_info.rss | |
| rss_gb = rss_bytes / (1024**3) | |
| vms_gb = mem_info.vms / (1024**3) | |
| if rss_gb > self._peak_rss_gb: | |
| self._peak_rss_gb = rss_gb | |
| if vms_gb > self._peak_vms_gb: | |
| self._peak_vms_gb = vms_gb | |
| cpu_percent = process.cpu_percent(interval=0.1) | |
| mem_percent = process.memory_percent() | |
| self._last_stats = { | |
| "rss_gb": rss_gb, | |
| "vms_gb": vms_gb, | |
| "cpu_percent": cpu_percent, | |
| "memory_percent": mem_percent, | |
| "num_threads": process.num_threads(), | |
| "step": self.step_getter(), | |
| } | |
| self.stats["rss_samples"].append(rss_gb) | |
| self.stats["vms_samples"].append(vms_gb) | |
| self.stats["cpu_samples"].append(cpu_percent) | |
| self.stats["timestamps"].append(time.time()) | |
| self.stats["steps"].append(self.step_getter()) | |
| if rss_bytes >= self.critical_threshold_bytes: | |
| self._critical_count += 1 | |
| if self.on_critical: | |
| self.on_critical(self._last_stats) | |
| critical_msg = ( | |
| f"\n{'=' * 70}\n" | |
| f" CRITICAL MEMORY (Level {self._critical_count})\n" | |
| f" RSS: {rss_gb:.2f} GB / {self.limit_bytes / (1024**3):.0f} GB\n" | |
| f" Step: {self.step_getter()}\n" | |
| f"{'=' * 70}" | |
| ) | |
| console.print(Theme.error(critical_msg)) | |
| system_logger.critical(f"Critical memory: {rss_gb:.2f} GB") | |
| self._force_cleanup() | |
| elif rss_bytes >= self.warn_threshold_bytes: | |
| self._warning_count += 1 | |
| if self.on_warning: | |
| self.on_warning(self._last_stats) | |
| if self._warning_count <= self._max_warnings: | |
| warning_msg = ( | |
| f"\n{'=' * 60}\n" | |
| f" MEMORY WARNING (Level {self._warning_count})\n" | |
| f" RSS: {rss_gb:.2f} GB / {self.limit_bytes / (1024**3):.0f} GB\n" | |
| f" Step: {self.step_getter()}\n" | |
| f"{'=' * 60}" | |
| ) | |
| console.print(Theme.warning(warning_msg)) | |
| system_logger.warning(f"Memory warning: {rss_gb:.2f} GB") | |
| if self._warning_count >= self._max_warnings: | |
| console.print( | |
| Theme.error( | |
| f" MULTIPLE MEMORY WARNINGS ({self._warning_count}). " | |
| f"Consider reducing batch size." | |
| ) | |
| ) | |
| self._warning_count = 0 | |
| elif rss_bytes < self.warn_threshold_bytes * 0.6: | |
| if self._warning_count > 0: | |
| console.print(Theme.success(" Memory recovered to safe levels")) | |
| self._warning_count = 0 | |
| if self._critical_count > 0: | |
| self._critical_count = 0 | |
| if self.stats["rss_samples"] and len(self.stats["rss_samples"]) > 100: | |
| if ( | |
| rss_gb | |
| > sum(self.stats["rss_samples"]) | |
| / len(self.stats["rss_samples"]) | |
| * 1.5 | |
| ): | |
| self._force_cleanup() | |
| except Exception as e: | |
| debug_logger.debug(f"Watchdog error: {e}") | |
| time.sleep(self.check_interval_sec) | |
| def _force_cleanup(self) -> None: | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| torch.cuda.synchronize() | |
| gc.collect() | |
| debug_logger.debug("Forced memory cleanup") | |
| def start(self) -> None: | |
| try: | |
| import psutil | |
| except ImportError: | |
| console.print(Theme.warning(" psutil not installed - watchdog disabled")) | |
| return | |
| if self.is_running: | |
| return | |
| self._thread = threading.Thread(target=self._monitor_loop, daemon=True) | |
| self._thread.start() | |
| self.is_running = True | |
| console.print( | |
| Theme.dim( | |
| f" Memory watchdog active (check every {self.check_interval_sec}s)" | |
| ) | |
| ) | |
| system_logger.info("Memory watchdog started") | |
| def stop(self) -> None: | |
| if not self.is_running: | |
| return | |
| self._stop_event.set() | |
| if self._thread is not None: | |
| self._thread.join(timeout=self.check_interval_sec + 2.0) | |
| self.is_running = False | |
| system_logger.info("Memory watchdog stopped") | |
| def get_stats(self) -> Dict: | |
| return { | |
| **self._last_stats, | |
| "peak_rss_gb": self._peak_rss_gb, | |
| "peak_vms_gb": self._peak_vms_gb, | |
| "warning_count": self._warning_count, | |
| "critical_count": self._critical_count, | |
| "is_running": self.is_running, | |
| "history_len": len(self.stats["rss_samples"]), | |
| "avg_rss_gb": sum(self.stats["rss_samples"]) | |
| / len(self.stats["rss_samples"]) | |
| if self.stats["rss_samples"] | |
| else 0, | |
| "avg_cpu": sum(self.stats["cpu_samples"]) / len(self.stats["cpu_samples"]) | |
| if self.stats["cpu_samples"] | |
| else 0, | |
| } | |
| def __enter__(self) -> "MemoryWatchdog": | |
| self.start() | |
| return self | |
| def __exit__(self, exc_type, exc_val, exc_tb) -> None: | |
| self.stop() | |
| class CacheManager: | |
| def __init__(self, cache_dir: str = CACHE_DIR, max_size_mb: int = 1024): | |
| self.cache_dir = Path(cache_dir) | |
| self.cache_dir.mkdir(exist_ok=True) | |
| self.max_size_bytes = max_size_mb * (1024**2) | |
| self._memory_cache: Dict[str, Any] = {} | |
| self._memory_size = 0 | |
| self._cache_stats = { | |
| "hits": 0, | |
| "misses": 0, | |
| "evictions": 0, | |
| "disk_writes": 0, | |
| "disk_reads": 0, | |
| } | |
| def _get_cache_path(self, key: str) -> Path: | |
| key_hash = hashlib.sha256(key.encode()).hexdigest() | |
| return self.cache_dir / key_hash | |
| def get(self, key: str, default: Any = None) -> Optional[Any]: | |
| if key in self._memory_cache: | |
| self._cache_stats["hits"] += 1 | |
| return self._memory_cache[key] | |
| cache_path = self._get_cache_path(key) | |
| if cache_path.exists(): | |
| try: | |
| with open(cache_path, "rb") as f: | |
| value = pickle.load(f) | |
| self._cache_stats["disk_reads"] += 1 | |
| self._set_memory(key, value) | |
| return value | |
| except Exception: | |
| pass | |
| self._cache_stats["misses"] += 1 | |
| return default | |
| def set(self, key: str, value: Any) -> None: | |
| self._set_memory(key, value) | |
| try: | |
| cache_path = self._get_cache_path(key) | |
| with open(cache_path, "wb") as f: | |
| pickle.dump(value, f) | |
| self._cache_stats["disk_writes"] += 1 | |
| except Exception as e: | |
| debug_logger.debug(f"Disk cache write failed: {e}") | |
| self._cleanup() | |
| def _set_memory(self, key: str, value: Any) -> None: | |
| value_size = sys.getsizeof(value) | |
| while self._memory_size + value_size > self.max_size_bytes: | |
| if not self._memory_cache: | |
| break | |
| oldest_key = next(iter(self._memory_cache)) | |
| old_value = self._memory_cache.pop(oldest_key) | |
| self._memory_size -= sys.getsizeof(old_value) | |
| self._cache_stats["evictions"] += 1 | |
| self._memory_cache[key] = value | |
| self._memory_size += value_size | |
| def _cleanup(self) -> None: | |
| try: | |
| files = list(self.cache_dir.glob("*")) | |
| total_size = sum(f.stat().st_size for f in files) | |
| if total_size > self.max_size_bytes: | |
| files.sort(key=lambda f: f.stat().st_mtime) | |
| for f in files: | |
| if total_size <= self.max_size_bytes * 0.8: | |
| break | |
| f.unlink() | |
| total_size -= f.stat().st_size | |
| self._cache_stats["evictions"] += 1 | |
| except Exception: | |
| pass | |
| def clear(self) -> None: | |
| self._memory_cache.clear() | |
| self._memory_size = 0 | |
| for f in self.cache_dir.glob("*"): | |
| try: | |
| f.unlink() | |
| except Exception: | |
| pass | |
| def get_stats(self) -> Dict: | |
| return { | |
| **self._cache_stats, | |
| "memory_entries": len(self._memory_cache), | |
| "memory_size_mb": self._memory_size / (1024**2), | |
| "disk_files": len(list(self.cache_dir.glob("*"))), | |
| "hit_rate": self._cache_stats["hits"] | |
| / (self._cache_stats["hits"] + self._cache_stats["misses"]) | |
| if self._cache_stats["hits"] + self._cache_stats["misses"] > 0 | |
| else 0, | |
| } | |
| cache_manager = CacheManager() | |
| def get_device() -> str: | |
| """Return best available device: cuda / mps / cpu""" | |
| if torch.cuda.is_available(): | |
| return "cuda" | |
| elif hasattr(torch, "mps") and torch.mps.is_available(): | |
| return "mps" | |
| else: | |
| return "cpu" | |
| def get_gpu_info() -> Dict[str, Any]: | |
| """Return GPU info dict, atau {} kalau tidak ada CUDA""" | |
| if not torch.cuda.is_available(): | |
| return {} | |
| try: | |
| return { | |
| "name": torch.cuda.get_device_name(0), | |
| "memory_total_gb": torch.cuda.get_device_properties(0).total_memory | |
| / (1024**3), | |
| "memory_allocated_gb": torch.cuda.memory_allocated() / (1024**3), | |
| "memory_reserved_gb": torch.cuda.memory_reserved() / (1024**3), | |
| "max_memory_allocated_gb": torch.cuda.max_memory_allocated() / (1024**3), | |
| "device_count": torch.cuda.device_count(), | |
| "is_available": True, | |
| } | |
| except Exception: | |
| return {} | |