| from __future__ import annotations |
|
|
| import importlib.util |
| import threading |
| import time |
| from contextlib import AbstractContextManager |
| from dataclasses import dataclass, field |
| from typing import Callable |
|
|
| import torch |
|
|
|
|
| class ProfilingUnavailable(RuntimeError): |
| pass |
|
|
|
|
| def count_parameters(model: torch.nn.Module) -> dict[str, float]: |
| total = sum(p.numel() for p in model.parameters()) |
| trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) |
| return { |
| "params": float(total), |
| "params_m": total / 1_000_000.0, |
| "trainable_params": float(trainable), |
| "trainable_params_m": trainable / 1_000_000.0, |
| } |
|
|
|
|
| def count_flops(model: torch.nn.Module, input_factory: Callable[[], tuple], device: torch.device) -> dict[str, object]: |
| model_was_training = model.training |
| model.eval() |
| inputs = tuple(x.to(device) if torch.is_tensor(x) else x for x in input_factory()) |
| shape = [list(x.shape) for x in inputs if torch.is_tensor(x)] |
| try: |
| if importlib.util.find_spec("thop") is not None: |
| from thop import profile |
|
|
| flops, _ = profile(model, inputs=inputs, verbose=False) |
| return { |
| "flops": float(flops), |
| "flops_g": float(flops) / 1_000_000_000.0, |
| "flops_input_shape": shape, |
| "flops_library": "thop", |
| } |
| if importlib.util.find_spec("fvcore") is not None: |
| from fvcore.nn import FlopCountAnalysis |
|
|
| flops = FlopCountAnalysis(model, inputs).total() |
| return { |
| "flops": float(flops), |
| "flops_g": float(flops) / 1_000_000_000.0, |
| "flops_input_shape": shape, |
| "flops_library": "fvcore", |
| } |
| if importlib.util.find_spec("ptflops") is not None: |
| raise ProfilingUnavailable( |
| "ptflops is installed, but this repository needs a two-input adapter before it can be used safely." |
| ) |
| raise ProfilingUnavailable("Install one FLOPs backend: thop or fvcore.") |
| finally: |
| model.train(model_was_training) |
|
|
|
|
| @dataclass |
| class GpuProfiler(AbstractContextManager): |
| device: torch.device |
| interval_s: float = 0.1 |
| required: bool = False |
| _stop: threading.Event = field(default_factory=threading.Event, init=False) |
| _thread: threading.Thread | None = field(default=None, init=False) |
| _samples_util: list[float] = field(default_factory=list, init=False) |
| _samples_mem: list[float] = field(default_factory=list, init=False) |
| _nvml: object | None = field(default=None, init=False) |
| _handle: object | None = field(default=None, init=False) |
| _error: str | None = field(default=None, init=False) |
|
|
| def __enter__(self) -> "GpuProfiler": |
| if self.device.type != "cuda": |
| self._error = "CUDA is not available; GPU utilization profiling was not run." |
| if self.required: |
| raise ProfilingUnavailable(self._error) |
| return self |
| if importlib.util.find_spec("pynvml") is None: |
| self._error = "pynvml is not installed; install pynvml for GPU utilization profiling." |
| if self.required: |
| raise ProfilingUnavailable(self._error) |
| return self |
| import pynvml |
|
|
| self._nvml = pynvml |
| pynvml.nvmlInit() |
| index = self.device.index if self.device.index is not None else torch.cuda.current_device() |
| self._handle = pynvml.nvmlDeviceGetHandleByIndex(index) |
| torch.cuda.reset_peak_memory_stats(self.device) |
| self._thread = threading.Thread(target=self._sample_loop, daemon=True) |
| self._thread.start() |
| return self |
|
|
| def _sample_loop(self) -> None: |
| assert self._nvml is not None and self._handle is not None |
| while not self._stop.is_set(): |
| util = self._nvml.nvmlDeviceGetUtilizationRates(self._handle) |
| mem = self._nvml.nvmlDeviceGetMemoryInfo(self._handle) |
| self._samples_util.append(float(util.gpu)) |
| self._samples_mem.append(float(mem.used) / (1024.0**3)) |
| time.sleep(self.interval_s) |
|
|
| def __exit__(self, exc_type, exc, tb) -> bool: |
| self._stop.set() |
| if self._thread is not None: |
| self._thread.join(timeout=2.0) |
| if self._nvml is not None: |
| self._nvml.nvmlShutdown() |
| return False |
|
|
| def summary(self) -> dict[str, object]: |
| if self.device.type == "cuda": |
| index = self.device.index if self.device.index is not None else torch.cuda.current_device() |
| base = { |
| "gpu_index": int(index), |
| "gpu_name": torch.cuda.get_device_name(self.device), |
| "gpu_mem_allocated_peak_gb": torch.cuda.max_memory_allocated(self.device) / (1024.0**3), |
| "gpu_mem_reserved_peak_gb": torch.cuda.max_memory_reserved(self.device) / (1024.0**3), |
| } |
| else: |
| base = { |
| "gpu_index": None, |
| "gpu_name": None, |
| "gpu_mem_allocated_peak_gb": None, |
| "gpu_mem_reserved_peak_gb": None, |
| } |
| if not self._samples_util: |
| base.update({ |
| "gpu_util_mean": None, |
| "gpu_util_max": None, |
| "gpu_util_min": None, |
| "gpu_mem_used_mean_gb": None, |
| "gpu_mem_used_max_gb": None, |
| "gpu_profiling_error": self._error, |
| }) |
| return base |
| base.update({ |
| "gpu_util_mean": sum(self._samples_util) / len(self._samples_util), |
| "gpu_util_max": max(self._samples_util), |
| "gpu_util_min": min(self._samples_util), |
| "gpu_mem_used_mean_gb": sum(self._samples_mem) / len(self._samples_mem), |
| "gpu_mem_used_max_gb": max(self._samples_mem), |
| "gpu_profiling_error": None, |
| }) |
| return base |
|
|