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" @staticmethod def header(text: str, color: str = "cyan") -> str: return f"[bold {color}]{text}[/bold {color}]" @staticmethod def success(text: str) -> str: return f"[green]{text}[/green]" @staticmethod def warning(text: str) -> str: return f"[yellow]{text}[/yellow]" @staticmethod def error(text: str) -> str: return f"[red]{text}[/red]" @staticmethod def info(text: str) -> str: return f"[blue]{text}[/blue]" @staticmethod def dim(text: str) -> str: return f"[dim]{text}[/dim]" @staticmethod def gold(text: str) -> str: return f"[gold1]{text}[/gold1]" @staticmethod def primary(text: str) -> str: return f"[cyan]{text}[/cyan]" @staticmethod def secondary(text: str) -> str: return f"[green]{text}[/green]" @staticmethod def magenta(text: str) -> str: return f"[magenta]{text}[/magenta]" @staticmethod def purple(text: str) -> str: return f"[purple]{text}[/purple]" @staticmethod def orange(text: str) -> str: return f"[orange1]{text}[/orange1]" @staticmethod def teal(text: str) -> str: return f"[teal]{text}[/teal]" @staticmethod def lime(text: str) -> str: return f"[lime]{text}[/lime]" @staticmethod def coral(text: str) -> str: return f"[coral]{text}[/coral]" @staticmethod def box(text: str, color: str = "cyan") -> str: return f"[{color}]{text}[/{color}]" @staticmethod def metric(name: str, value: Any) -> str: return f"[cyan]{name}:[/cyan] [green]{value}[/green]" @staticmethod 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() @contextmanager 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) @contextmanager 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)" ) @contextmanager 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 @contextmanager 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) @contextmanager 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: @wraps(func) 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: @wraps(func) 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: @wraps(func) 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 {}