AriaLM / src /s07_utils.py
krishnah27's picture
Upload folder using huggingface_hub
30e9297 verified
Raw
History Blame Contribute Delete
1.41 kB
"""
Utility helpers for logging, memory monitoring, and reproducibility.
"""
import gc
import logging
import os
import random
import sys
import numpy as np
import torch
def setup_logging(level: str = "INFO"):
logging.basicConfig(
level=getattr(logging, level.upper(), logging.INFO),
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
handlers=[logging.StreamHandler(sys.stdout)],
)
def set_seed(seed: int = 42):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def log_memory_usage(tag: str = ""):
"""Log current memory usage — critical for constrained hardware."""
prefix = f"[{tag}] " if tag else ""
if torch.cuda.is_available():
allocated = torch.cuda.memory_allocated() / 1024**2
reserved = torch.cuda.memory_reserved() / 1024**2
logging.info(f"{prefix}GPU Memory: {allocated:.0f}MB allocated, {reserved:.0f}MB reserved")
import psutil
proc = psutil.Process(os.getpid())
ram = proc.memory_info().rss / 1024**2
logging.info(f"{prefix}RAM Usage: {ram:.0f}MB")
def clear_memory():
"""Aggressively free memory — call between major operations."""
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.synchronize()