Download staplebridge/utils/profiling.py from ChatterjeeLab/StapleBridge: direct link, hf CLI and curl.
- Browser
- Download file 3.05 kB
-
https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/utils/profiling.py
- Command line
-
hf download hf://ChatterjeeLab/StapleBridge/staplebridge/utils/profiling.py
-
curl -L -o profiling.py https://huggingface.co/ChatterjeeLab/StapleBridge/resolve/main/staplebridge/utils/profiling.py
3.05 kB
| """Lightweight per-stage timing accumulator. | |
| A single global ``TimingAccumulator`` collects wall-clock time and call counts | |
| for named stages (e.g. ``reference_logits_time``, ``ESM2_prior_time``, | |
| ``PeptiVerse_time``, ``backward_time``, ...). Callers wrap their code in | |
| ``with STAGE_TIMER.section("stage_name"):`` and periodically call | |
| ``STAGE_TIMER.report_and_reset()`` to print/log a table. | |
| Cache hit/miss counters (``bump("esm2_cache_hit")``) live on the same object. | |
| The whole module is process-local and thread-safe enough for our single-process | |
| training loop; no cross-process aggregation is attempted. | |
| """ | |
| from __future__ import annotations | |
| import contextlib | |
| import threading | |
| import time | |
| from collections import defaultdict | |
| class TimingAccumulator: | |
| def __init__(self) -> None: | |
| self._lock = threading.Lock() | |
| self._time: dict[str, float] = defaultdict(float) | |
| self._calls: dict[str, int] = defaultdict(int) | |
| self._counters: dict[str, int] = defaultdict(int) | |
| def section(self, name: str): | |
| t0 = time.perf_counter() | |
| try: | |
| yield | |
| finally: | |
| dt = time.perf_counter() - t0 | |
| with self._lock: | |
| self._time[name] += dt | |
| self._calls[name] += 1 | |
| def add(self, name: str, seconds: float) -> None: | |
| with self._lock: | |
| self._time[name] += float(seconds) | |
| self._calls[name] += 1 | |
| def bump(self, name: str, amount: int = 1) -> None: | |
| with self._lock: | |
| self._counters[name] += int(amount) | |
| def snapshot(self) -> dict[str, float]: | |
| with self._lock: | |
| snap: dict[str, float] = {} | |
| for k, v in self._time.items(): | |
| snap[k] = float(v) | |
| snap[f"{k}_calls"] = int(self._calls.get(k, 0)) | |
| for k, v in self._counters.items(): | |
| snap[k] = int(v) | |
| return snap | |
| def reset(self) -> None: | |
| with self._lock: | |
| self._time.clear() | |
| self._calls.clear() | |
| self._counters.clear() | |
| def format_table(self, title: str = "timings") -> str: | |
| with self._lock: | |
| rows: list[tuple[str, float, int]] = [] | |
| for k in sorted(self._time.keys()): | |
| rows.append((k, float(self._time[k]), int(self._calls.get(k, 0)))) | |
| counters = dict(self._counters) | |
| lines = [f"[{title}]"] | |
| for name, secs, calls in rows: | |
| per = (secs / calls) if calls else 0.0 | |
| lines.append( | |
| f" {name:<28s} total={secs:>8.3f}s calls={calls:>8d} avg={per*1000:>8.3f}ms" | |
| ) | |
| if counters: | |
| lines.append(" -- counters --") | |
| for k in sorted(counters.keys()): | |
| lines.append(f" {k:<28s} {counters[k]}") | |
| return "\n".join(lines) | |
| def report_and_reset(self, title: str = "timings") -> str: | |
| s = self.format_table(title) | |
| self.reset() | |
| return s | |
| STAGE_TIMER = TimingAccumulator() | |