gpu-energy-profiler 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,14 @@
1
+ Metadata-Version: 2.4
2
+ Name: gpu-energy-profiler
3
+ Version: 0.1.0
4
+ Summary:
5
+ Author: Vincent Eichhorn
6
+ Author-email: v.eichhorn@posteo.de
7
+ Requires-Python: >=3.14
8
+ Classifier: Programming Language :: Python :: 3
9
+ Classifier: Programming Language :: Python :: 3.14
10
+ Classifier: Programming Language :: Python :: 3.15
11
+ Requires-Dist: nvidia-ml-py (>=12.0,<13.0)
12
+ Requires-Dist: pandas (>=3.0.6,<4.0.0)
13
+ Requires-Dist: plotly (>=7.1.0,<8.0.0)
14
+ Requires-Dist: torch (>=2.14.1,<3.0.0)
@@ -0,0 +1,25 @@
1
+ [project]
2
+ name = "gpu-energy-profiler"
3
+ version = "0.1.0"
4
+ description = ""
5
+ authors = [
6
+ {name = "Vincent Eichhorn",email = "v.eichhorn@posteo.de"}
7
+ ]
8
+ requires-python = ">=3.14"
9
+ dependencies = [
10
+ "pandas (>=3.0.6,<4.0.0)",
11
+ "torch (>=2.14.1,<3.0.0)",
12
+ "plotly (>=7.1.0,<8.0.0)",
13
+ "nvidia-ml-py (>=12.0,<13.0)"
14
+ ]
15
+
16
+ [tool.poetry.group.dev.dependencies]
17
+ pytest = "^8.0"
18
+
19
+ [tool.pytest.ini_options]
20
+ testpaths = ["tests"]
21
+
22
+
23
+ [build-system]
24
+ requires = ["poetry-core>=2.0.0,<3.0.0"]
25
+ build-backend = "poetry.core.masonry.api"
@@ -0,0 +1,7 @@
1
+ """Tools for profiling PyTorch workloads and GPU power usage."""
2
+
3
+ from .base import Profiler
4
+ from .nvidia_profiler import NvidiaProfiler
5
+ from .torch_profiler import TorchProfiler
6
+
7
+ __all__ = ["NvidiaProfiler", "Profiler", "TorchProfiler"]
@@ -0,0 +1,35 @@
1
+ """Shared functionality for profiler implementations."""
2
+
3
+ from contextlib import contextmanager
4
+ from datetime import datetime
5
+ from typing import Generator
6
+
7
+
8
+ class Profiler:
9
+ """Record named timestamps while a profiling session is running."""
10
+
11
+ def __init__(self) -> None:
12
+ """Initialize an empty step timeline."""
13
+ self.record_steps: list[tuple[datetime, str]] = []
14
+ self.record_step("__init__")
15
+
16
+ def record_step(self, name: str) -> None:
17
+ """Record a named timestamp.
18
+
19
+ Args:
20
+ name: Label for the recorded step.
21
+ """
22
+ self.record_steps.append((datetime.now(), name))
23
+
24
+ @contextmanager
25
+ def record_context(self, name: str) -> Generator[None, None, None]:
26
+ """Record a context's start and end markers.
27
+
28
+ Args:
29
+ name: Label recorded when entering the context.
30
+ """
31
+ try:
32
+ self.record_step(name)
33
+ yield
34
+ finally:
35
+ self.record_step("__other__")
@@ -0,0 +1,236 @@
1
+ """Result handlers and helpers for profiler worker processes."""
2
+
3
+ from abc import ABC, abstractmethod
4
+ from collections.abc import Callable
5
+ from itertools import zip_longest
6
+ from multiprocessing import Manager, Process, Queue
7
+ from pathlib import Path
8
+ from typing import Any
9
+ import warnings
10
+
11
+ Result = tuple[Any, ...]
12
+ ResultOrSentinel = Result | None
13
+
14
+
15
+ class ResultHandler(ABC):
16
+ """Store worker results behind a queue-like interface."""
17
+
18
+ def __init__(self) -> None:
19
+ """Initialize an unconfigured result handler."""
20
+ self.column_names: tuple[str, ...] = ()
21
+ self.dtypes: tuple[type, ...] = ()
22
+
23
+ def __enter__(self) -> "ResultHandler":
24
+ """Enter the handler context.
25
+
26
+ Returns:
27
+ This result handler.
28
+ """
29
+ return self
30
+
31
+ def __exit__(self, *context_info: Any) -> None:
32
+ """Exit the handler context.
33
+
34
+ Args:
35
+ *context_info: Context-manager exception information.
36
+ """
37
+ del context_info
38
+
39
+ def set_columns(
40
+ self,
41
+ names: tuple[str, ...],
42
+ dtypes: tuple[type, ...] | None = None,
43
+ ) -> None:
44
+ """Set column names and conversion types used by file-backed results.
45
+
46
+ Args:
47
+ names: Column names in the serialized result rows.
48
+ dtypes: Conversion types for serialized values. Defaults to ``str``.
49
+
50
+ Raises:
51
+ ValueError: If the number of names and types differs.
52
+ """
53
+ dtypes = dtypes or (str,) * len(names)
54
+ if len(names) != len(dtypes):
55
+ raise ValueError("names and dtypes must have the same length")
56
+ self.column_names = names
57
+ self.dtypes = dtypes
58
+
59
+ @abstractmethod
60
+ def put(self, data: ResultOrSentinel) -> None:
61
+ """Store one result or ``None`` as the end-of-results sentinel.
62
+
63
+ Args:
64
+ data: Result tuple or ``None`` to mark the end of the stream.
65
+ """
66
+
67
+ @abstractmethod
68
+ def latest(self) -> ResultOrSentinel:
69
+ """Return the latest result, or ``None`` when no result is available."""
70
+
71
+ @abstractmethod
72
+ def all(self) -> list[Result]:
73
+ """Return all results up to the end-of-results sentinel."""
74
+
75
+ def get(self) -> ResultOrSentinel:
76
+ """Return the next result using the legacy method name.
77
+
78
+ Returns:
79
+ The next result or ``None``.
80
+ """
81
+ return self.latest()
82
+
83
+ def get_all(self) -> list[Result]:
84
+ """Return all results using the legacy method name.
85
+
86
+ Returns:
87
+ All results up to the end-of-results sentinel.
88
+ """
89
+ return self.all()
90
+
91
+
92
+ class MPQueueResultHandler(ResultHandler):
93
+ """Store results in a multiprocessing queue."""
94
+
95
+ def __init__(self) -> None:
96
+ """Initialize an in-memory queue result handler."""
97
+ super().__init__()
98
+ self.queue: Queue[ResultOrSentinel] = Queue()
99
+
100
+ def put(self, data: ResultOrSentinel) -> None:
101
+ """Put a result or sentinel into the queue.
102
+
103
+ Args:
104
+ data: Result tuple or ``None`` as the stream sentinel.
105
+ """
106
+ self.queue.put(data)
107
+
108
+ def latest(self) -> ResultOrSentinel:
109
+ """Block until and return the next queued result.
110
+
111
+ Returns:
112
+ The next queued result or ``None``.
113
+ """
114
+ return self.queue.get()
115
+
116
+ def all(self) -> list[Result]:
117
+ """Read results until the sentinel is encountered.
118
+
119
+ Returns:
120
+ All queued result tuples.
121
+ """
122
+ return [result for result in iter(self.queue.get, None)]
123
+
124
+
125
+ class FileCacheResultHandler(ResultHandler):
126
+ """Append typed result rows to a CSV cache file."""
127
+
128
+ def __init__(self, file_path: str, force: bool = False) -> None:
129
+ """Initialize a CSV-backed result handler.
130
+
131
+ Args:
132
+ file_path: Path to the CSV cache.
133
+ force: Whether to replace an existing cache.
134
+ """
135
+ super().__init__()
136
+ self.file_path = Path(file_path)
137
+ self.file_path.parent.mkdir(parents=True, exist_ok=True)
138
+ if force and self.file_path.is_file():
139
+ self.file_path.unlink()
140
+ if self.file_path.is_file() and self.file_path.stat().st_size > 0:
141
+ warnings.warn(
142
+ f"File {self.file_path} already exists and is not empty",
143
+ UserWarning,
144
+ stacklevel=2,
145
+ )
146
+
147
+ def __enter__(self) -> "FileCacheResultHandler":
148
+ """Enter the file-cache context.
149
+
150
+ Returns:
151
+ This file-cache result handler.
152
+ """
153
+ return self
154
+
155
+ def set_columns(
156
+ self,
157
+ names: tuple[str, ...],
158
+ dtypes: tuple[type, ...] | None = None,
159
+ ) -> None:
160
+ """Write the CSV header when the cache is empty."""
161
+ super().set_columns(names, dtypes)
162
+ if self.file_path.exists() and self.file_path.stat().st_size > 0:
163
+ return
164
+ self.file_path.write_text(",".join(names) + "\n", encoding="utf-8")
165
+
166
+ def put(self, data: ResultOrSentinel) -> None:
167
+ """Append a tuple result and ignore the worker sentinel.
168
+
169
+ Args:
170
+ data: Tuple to append, or ``None`` to end a worker stream.
171
+ """
172
+ if not isinstance(data, tuple):
173
+ return
174
+ with self.file_path.open("a", encoding="utf-8") as file:
175
+ file.write(",".join(str(value) for value in data) + "\n")
176
+
177
+ def _parse_line(self, line: str) -> Result:
178
+ """Convert one serialized CSV row to a typed result tuple.
179
+
180
+ Args:
181
+ line: Comma-separated result row.
182
+
183
+ Returns:
184
+ The converted result tuple.
185
+ """
186
+ values = zip_longest(
187
+ line.strip().split(","),
188
+ self.dtypes,
189
+ fillvalue=str,
190
+ )
191
+ return tuple(converter(value) for value, converter in values)
192
+
193
+ def latest(self) -> ResultOrSentinel:
194
+ """Return the last cached row, or ``None`` for an empty cache."""
195
+ rows = self._rows()
196
+ return self._parse_line(rows[-1]) if rows else None
197
+
198
+ def all(self) -> list[Result]:
199
+ """Return every cached row in file order.
200
+
201
+ Returns:
202
+ Typed result tuples in cache order.
203
+ """
204
+ return [self._parse_line(row) for row in self._rows()]
205
+
206
+ def _rows(self) -> list[str]:
207
+ """Read data rows from the cache file.
208
+
209
+ Returns:
210
+ Raw CSV data rows without the optional header.
211
+ """
212
+ with self.file_path.open("r", encoding="utf-8") as file:
213
+ rows = file.readlines()
214
+ return rows[1:] if self.column_names else rows
215
+
216
+
217
+ def start_separate_process(target: Callable[..., Any], arguments: list[Any]) -> Any:
218
+ """Run ``target`` in a spawned process and return its first result.
219
+
220
+ Args:
221
+ target: Worker function receiving a managed queue first.
222
+ arguments: Additional arguments passed to the worker.
223
+
224
+ Returns:
225
+ The first result placed in the managed queue.
226
+ """
227
+ with Manager() as manager:
228
+ queue = manager.Queue()
229
+ process = Process(target=target, args=[queue, *arguments])
230
+ process.start()
231
+ process.join()
232
+ return queue.get()
233
+
234
+
235
+ # Compatibility alias for the original misspelled function name.
236
+ start_seprate_process = start_separate_process
@@ -0,0 +1,458 @@
1
+ """GPU power and memory profiling through NVML or ``nvidia-smi``."""
2
+
3
+ from datetime import datetime
4
+ import subprocess
5
+ from multiprocessing import Array, Event, Process, Value
6
+ import time
7
+ from typing import Any, Literal
8
+
9
+ import pandas as pd
10
+ import plotly.graph_objects as go
11
+ from plotly.subplots import make_subplots
12
+
13
+ from .base import Profiler
14
+ from .multiprocessing_util import FileCacheResultHandler, MPQueueResultHandler, ResultHandler
15
+ from .plotting_util import sample_colors
16
+
17
+
18
+ def _load_pynvml() -> Any:
19
+ """Import and return the optional ``pynvml`` module.
20
+
21
+ Returns:
22
+ The imported NVML Python bindings.
23
+
24
+ Raises:
25
+ RuntimeError: If the optional dependency is not installed.
26
+ """
27
+ try:
28
+ import pynvml
29
+ except ImportError as error:
30
+ raise RuntimeError(
31
+ "The pynvml backend requires the optional 'nvidia-ml-py' dependency."
32
+ ) from error
33
+ return pynvml
34
+
35
+
36
+ class NvidiaProfiler(Profiler):
37
+ """Sample GPU power and memory usage with a selectable backend."""
38
+
39
+ _COLUMNS = ("gpu_id", "timestamp", "power", "memory", "record_step")
40
+ _BACKENDS = ("nvidia-smi", "pynvml")
41
+
42
+ def __init__(
43
+ self,
44
+ interval: int = 1,
45
+ cache_file: str | None = None,
46
+ force_cache: bool = False,
47
+ backend: Literal["nvidia-smi", "pynvml"] = "nvidia-smi",
48
+ ) -> None:
49
+ """Initialize an NVIDIA power and memory profiler.
50
+
51
+ Args:
52
+ interval: Sampling interval in milliseconds.
53
+ cache_file: Optional CSV path for persisted samples.
54
+ force_cache: Whether to replace an existing cache file.
55
+ backend: Sampling implementation, either ``"nvidia-smi"`` or
56
+ ``"pynvml"``.
57
+
58
+ Raises:
59
+ ValueError: If ``backend`` is not supported.
60
+ """
61
+ if backend not in self._BACKENDS:
62
+ raise ValueError(f"Unknown backend {backend!r}; choose from {self._BACKENDS}.")
63
+ self.current_record_step: Any = Array("c", 1000)
64
+ self.interval = interval
65
+ self.backend = backend
66
+ self.data: list[tuple[Any, ...]] = []
67
+ self.should_profiling_run = Value("i", 1)
68
+ self.profiling_started = Event()
69
+ self.profiling_stopped = Event()
70
+ self.result_handler: ResultHandler = (
71
+ MPQueueResultHandler()
72
+ if cache_file is None
73
+ else FileCacheResultHandler(cache_file, force_cache)
74
+ )
75
+ self.result_handler.set_columns(self._COLUMNS, (int, str, float, float, str))
76
+ self.process = Process(
77
+ target=self._profiling_process,
78
+ args=(
79
+ self.should_profiling_run,
80
+ self.profiling_started,
81
+ self.profiling_stopped,
82
+ self.result_handler,
83
+ self.current_record_step,
84
+ self.interval,
85
+ self.backend,
86
+ ),
87
+ )
88
+ super().__init__()
89
+
90
+ def record_step(self, name: str) -> None:
91
+ """Record a step and share its name with the sampling process.
92
+
93
+ Args:
94
+ name: Label for the recorded step.
95
+ """
96
+ super().record_step(name)
97
+ self.current_record_step.value = name.encode("utf-8")
98
+
99
+ @staticmethod
100
+ def _parse_nvidia_smi_row(line: str, current_record_step: Any) -> tuple[Any, ...]:
101
+ """Parse one CSV row emitted by ``nvidia-smi``.
102
+
103
+ Args:
104
+ line: Raw CSV sample line.
105
+ current_record_step: Shared character array containing the step name.
106
+
107
+ Returns:
108
+ A typed sample tuple containing GPU, timestamp, power, memory, and step.
109
+ """
110
+ values = line.strip().split(", ")
111
+ return (
112
+ int(values[0]),
113
+ datetime.strptime(values[1], "%Y/%m/%d %H:%M:%S.%f"),
114
+ float(values[2].split(" ")[0]),
115
+ float(values[3].split(" ")[0]),
116
+ current_record_step.value.decode("utf-8"),
117
+ )
118
+
119
+ @staticmethod
120
+ def _profiling_process(
121
+ should_run: Any,
122
+ started: Any,
123
+ stopped: Any,
124
+ result_handler: ResultHandler,
125
+ current_record_step: Any,
126
+ interval: int,
127
+ backend: Literal["nvidia-smi", "pynvml"],
128
+ ) -> None:
129
+ """Run the selected sampling backend in a worker process.
130
+
131
+ Args:
132
+ should_run: Shared flag controlling the worker lifecycle.
133
+ started: Event set after the sampler starts.
134
+ stopped: Event set after sampling ends.
135
+ result_handler: Destination for sampled rows.
136
+ current_record_step: Shared current step name.
137
+ interval: Sampling interval in milliseconds.
138
+ backend: Backend to run.
139
+ """
140
+ if backend == "nvidia-smi":
141
+ NvidiaProfiler._nvidia_smi_profiling_process(
142
+ should_run,
143
+ started,
144
+ stopped,
145
+ result_handler,
146
+ current_record_step,
147
+ interval,
148
+ )
149
+ return
150
+ NvidiaProfiler._pynvml_profiling_process(
151
+ should_run,
152
+ started,
153
+ stopped,
154
+ result_handler,
155
+ current_record_step,
156
+ interval,
157
+ )
158
+
159
+ @staticmethod
160
+ def _nvidia_smi_profiling_process(
161
+ should_run: Any,
162
+ started: Any,
163
+ stopped: Any,
164
+ result_handler: ResultHandler,
165
+ current_record_step: Any,
166
+ interval: int,
167
+ ) -> None:
168
+ """Read samples until ``should_run`` is cleared.
169
+
170
+ Args:
171
+ should_run: Shared flag controlling the worker lifecycle.
172
+ started: Event set after the sampler starts.
173
+ stopped: Event set after sampling ends.
174
+ result_handler: Destination for sampled rows.
175
+ current_record_step: Shared current step name.
176
+ interval: Sampling interval in milliseconds.
177
+ """
178
+ command = (
179
+ "nvidia-smi --query-gpu=index,timestamp,power.draw,memory.used "
180
+ f"--format=csv -lms {interval}"
181
+ )
182
+ with (
183
+ subprocess.Popen(
184
+ command,
185
+ shell=True,
186
+ text=True,
187
+ stdout=subprocess.PIPE,
188
+ ) as nvidia_smi_process,
189
+ result_handler as result,
190
+ ):
191
+ assert nvidia_smi_process.stdout is not None
192
+ with nvidia_smi_process.stdout as output:
193
+ output.readline()
194
+ if should_run.value:
195
+ started.set()
196
+ while should_run.value:
197
+ result.put(
198
+ NvidiaProfiler._parse_nvidia_smi_row(output.readline(), current_record_step)
199
+ )
200
+ result.put(None) # type: ignore[arg-type]
201
+ stopped.set()
202
+
203
+ @staticmethod
204
+ def _pynvml_profiling_process(
205
+ should_run: Any,
206
+ started: Any,
207
+ stopped: Any,
208
+ result_handler: ResultHandler,
209
+ current_record_step: Any,
210
+ interval: int,
211
+ ) -> None:
212
+ """Read power and memory samples through the optional ``pynvml`` package.
213
+
214
+ Args:
215
+ should_run: Shared flag controlling the worker lifecycle.
216
+ started: Event set after the sampler starts.
217
+ stopped: Event set after sampling ends.
218
+ result_handler: Destination for sampled rows.
219
+ current_record_step: Shared current step name.
220
+ interval: Sampling interval in milliseconds.
221
+ """
222
+ pynvml = _load_pynvml()
223
+ pynvml.nvmlInit()
224
+ try:
225
+ with result_handler as result:
226
+ if should_run.value:
227
+ started.set()
228
+ while should_run.value:
229
+ for gpu_id in range(pynvml.nvmlDeviceGetCount()):
230
+ handle = pynvml.nvmlDeviceGetHandleByIndex(gpu_id)
231
+ memory = pynvml.nvmlDeviceGetMemoryInfo(handle)
232
+ result.put(
233
+ (
234
+ gpu_id,
235
+ datetime.now(),
236
+ pynvml.nvmlDeviceGetPowerUsage(handle) / 1000.0,
237
+ memory.used / (1024 * 1024),
238
+ current_record_step.value.decode("utf-8"),
239
+ )
240
+ )
241
+ time.sleep(interval / 1000.0)
242
+ result.put(None) # type: ignore[arg-type]
243
+ finally:
244
+ pynvml.nvmlShutdown()
245
+ stopped.set()
246
+
247
+ def __enter__(self) -> "NvidiaProfiler":
248
+ """Start sampling and return this profiler.
249
+
250
+ Returns:
251
+ This active NVIDIA profiler.
252
+
253
+ Raises:
254
+ RuntimeError: If the selected backend is unavailable.
255
+ """
256
+ if self.backend == "nvidia-smi":
257
+ status, output = subprocess.getstatusoutput(
258
+ "nvidia-smi --query-gpu=index --format=csv,noheader"
259
+ )
260
+ if status != 0 or not output.strip():
261
+ raise RuntimeError(
262
+ "nvidia-smi is unavailable or could not query an NVIDIA GPU. "
263
+ f"Command output: {output.strip() or '<empty>'}"
264
+ )
265
+ else:
266
+ pynvml = _load_pynvml()
267
+ try:
268
+ pynvml.nvmlInit()
269
+ if pynvml.nvmlDeviceGetCount() == 0:
270
+ raise RuntimeError("pynvml found no NVIDIA GPU devices.")
271
+ except pynvml.NVMLError as error:
272
+ raise RuntimeError(
273
+ "pynvml could not initialize or query an NVIDIA GPU."
274
+ ) from error
275
+ finally:
276
+ pynvml.nvmlShutdown()
277
+ self.process.start()
278
+ if not self.profiling_started.wait(timeout=5):
279
+ if self.process.is_alive():
280
+ self.process.terminate()
281
+ self.process.join(timeout=5)
282
+ raise RuntimeError(
283
+ f"The {self.backend} profiling worker did not start. "
284
+ "Check the NVIDIA driver and backend installation."
285
+ )
286
+ return self
287
+
288
+ def __exit__(self, *context_info: Any) -> None:
289
+ """Stop sampling and collect all recorded rows.
290
+
291
+ Args:
292
+ *context_info: Context-manager exception information.
293
+ """
294
+ del context_info
295
+ self.should_profiling_run.value = 0
296
+ self.profiling_stopped.wait()
297
+ self.data = self.result_handler.all()
298
+ self.process.join()
299
+ self.process.terminate()
300
+
301
+ @staticmethod
302
+ def from_cache(cache_file: str) -> "NvidiaProfiler":
303
+ """Create a profiler populated from an existing cache file.
304
+
305
+ Args:
306
+ cache_file: Path to a previously written sample cache.
307
+
308
+ Returns:
309
+ A profiler containing the cached samples.
310
+ """
311
+ profiler = NvidiaProfiler(cache_file=cache_file)
312
+ profiler.data = profiler.result_handler.all()
313
+ return profiler
314
+
315
+ def to_pandas(self) -> pd.DataFrame:
316
+ """Return samples with parsed timestamps and standard column names.
317
+
318
+ Returns:
319
+ A dataframe containing GPU sample data.
320
+ """
321
+ df = pd.DataFrame(self.data, columns=self._COLUMNS)
322
+ df["timestamp"] = pd.to_datetime(
323
+ df["timestamp"], format="%Y-%m-%d %H:%M:%S.%f", errors="coerce"
324
+ )
325
+ return df
326
+
327
+ def profiled_gpus(self) -> list[int]:
328
+ """Return the IDs of GPUs present in the samples.
329
+
330
+ Returns:
331
+ Unique GPU IDs found in the samples.
332
+ """
333
+ return [int(gpu_id) for gpu_id in self.to_pandas()["gpu_id"].unique()]
334
+
335
+ def total_energy(
336
+ self,
337
+ gpu_ids: list[int] | None = None,
338
+ record_steps: list[str] | None = None,
339
+ return_data: bool = False,
340
+ ) -> float | list[float]:
341
+ """Return sampled energy in watt-seconds.
342
+
343
+ Args:
344
+ gpu_ids: GPUs to include. Defaults to the first sampled GPU.
345
+ record_steps: Step labels to include. Defaults to all steps.
346
+ return_data: Return one energy value per record-step group.
347
+
348
+ Returns:
349
+ Total energy, or grouped energy values when ``return_data`` is true.
350
+ """
351
+ if not self.data:
352
+ return 0.0
353
+ df = self.to_pandas()
354
+ df["record_step_id"] = df["record_step"].ne(df["record_step"].shift()).cumsum()
355
+ gpu_ids = gpu_ids or [df["gpu_id"].unique()[0]]
356
+ df = df[df["gpu_id"].isin(gpu_ids)].copy()
357
+ df["time_interval"] = df["timestamp"].diff().dt.total_seconds().fillna(0)
358
+ df["energy_interval"] = df["power"] * df["time_interval"]
359
+ record_steps = record_steps or list(df["record_step"].unique())
360
+ df = df[df["record_step"].isin(record_steps)]
361
+ if return_data:
362
+ return list(df.groupby("record_step_id")["energy_interval"].sum())
363
+ return df["energy_interval"].sum()
364
+
365
+ def total_time(self) -> float:
366
+ """Return the time between the first and last sample in seconds.
367
+
368
+ Returns:
369
+ Elapsed profiling time in seconds.
370
+ """
371
+ if not self.data:
372
+ return 0.0
373
+ df = self.to_pandas()
374
+ return (df["timestamp"].max() - df["timestamp"].min()).total_seconds()
375
+
376
+ def avg_memory_usage(self, gpu_id: int | None = None) -> float:
377
+ """Return average memory usage in MiB for a GPU.
378
+
379
+ Args:
380
+ gpu_id: GPU to measure. Defaults to the first sampled GPU.
381
+
382
+ Returns:
383
+ Average memory usage in MiB.
384
+ """
385
+ if not self.data:
386
+ return 0.0
387
+ df = self.to_pandas()
388
+ gpu_id = gpu_id or df["gpu_id"].unique()[0]
389
+ return df.loc[df["gpu_id"] == gpu_id, "memory"].mean()
390
+
391
+ def time_series_plot(self) -> go.Figure | None:
392
+ """Return a Plotly chart showing power, memory, and recorded steps.
393
+
394
+ Returns:
395
+ A Plotly figure, or ``None`` when no samples are available.
396
+ """
397
+ if not self.data:
398
+ return None
399
+ df = self.to_pandas()
400
+ figure = make_subplots(specs=[[{"secondary_y": True}]])
401
+ gpu_ids = self.profiled_gpus()
402
+ gpu_colors = sample_colors("Rainbow", len(gpu_ids))
403
+ for index, gpu_id in enumerate(gpu_ids):
404
+ gpu_df = df[df["gpu_id"] == gpu_id]
405
+ for metric, unit in (("power", "W"), ("memory", "MiB")):
406
+ figure.add_trace(
407
+ go.Scatter(
408
+ x=gpu_df["timestamp"],
409
+ y=gpu_df[metric],
410
+ name=f"{metric.capitalize()} ({unit})",
411
+ mode="lines+markers",
412
+ legendgroup=str(gpu_id),
413
+ legendgrouptitle_text=f"GPU #{gpu_id}",
414
+ line=dict(
415
+ color=gpu_colors[index],
416
+ width=4,
417
+ dash="dot" if metric == "memory" else "solid",
418
+ ),
419
+ ),
420
+ secondary_y=metric == "memory",
421
+ )
422
+ _add_record_step_regions(figure, df, self.record_steps)
423
+ figure.update_layout(
424
+ title="GPU Memory and Power Usage", legend=dict(groupclick="toggleitem")
425
+ )
426
+ figure.update_xaxes(title_text="Time")
427
+ figure.update_yaxes(title_text="Power (W)", secondary_y=False)
428
+ figure.update_yaxes(title_text="Memory (MiB)", secondary_y=True)
429
+ return figure
430
+
431
+ def _add_record_step_regions(
432
+ figure: go.Figure,
433
+ samples: pd.DataFrame,
434
+ record_steps: list[tuple[datetime, str]],
435
+ ) -> None:
436
+ """Add shaded regions for recorded steps to a Plotly figure.
437
+
438
+ Args:
439
+ figure: Figure to modify.
440
+ samples: Sample dataframe containing timestamps.
441
+ record_steps: Timestamp and label pairs defining the regions.
442
+ """
443
+ max_timestamp: datetime = samples["timestamp"].max()
444
+ steps = record_steps + [(max_timestamp, ".")]
445
+ step_names = list(dict.fromkeys(name for _, name in steps))
446
+ colors = dict(zip(step_names, sample_colors("viridis", len(step_names))))
447
+ previous_timestamp, previous_name = steps[0]
448
+ for timestamp, name in steps[1:]:
449
+ figure.add_vrect(
450
+ x0=previous_timestamp,
451
+ x1=timestamp,
452
+ annotation_text=previous_name,
453
+ annotation_position="top left",
454
+ line_width=0,
455
+ opacity=0.25,
456
+ fillcolor=colors[previous_name],
457
+ )
458
+ previous_timestamp, previous_name = timestamp, name
@@ -0,0 +1,22 @@
1
+ """Shared plotting helpers for profiler visualizations."""
2
+
3
+ import plotly.express as px
4
+
5
+
6
+ def sample_colors(scale: str, item_count: int) -> list[str]:
7
+ """Return evenly spaced colors for a collection.
8
+
9
+ Args:
10
+ scale: Plotly color scale name.
11
+ item_count: Number of colors required.
12
+
13
+ Returns:
14
+ At least two colors sampled from the requested scale.
15
+ """
16
+ count = max(item_count, 2)
17
+ return [
18
+ str(color)
19
+ for color in px.colors.sample_colorscale(
20
+ scale, [index / (count - 1) for index in range(count)]
21
+ )
22
+ ]
@@ -0,0 +1,206 @@
1
+ """PyTorch profiler integration and event analysis."""
2
+
3
+ from typing import Any
4
+
5
+ import pandas as pd
6
+ from torch.autograd.profiler_util import EventList, FunctionEvent
7
+ from torch.profiler import ProfilerActivity, profile
8
+
9
+ from .base import Profiler
10
+
11
+
12
+ class TorchProfiler(profile, Profiler):
13
+ """Collect and analyse CPU and CUDA events from PyTorch."""
14
+
15
+ def __init__(self, *args: Any, **kwargs: Any) -> None:
16
+ """Initialize a PyTorch profiler with useful profiling defaults.
17
+
18
+ Args:
19
+ *args: Positional arguments forwarded to ``torch.profiler.profile``.
20
+ **kwargs: Keyword arguments overriding the profiler defaults.
21
+ """
22
+ defaults = {
23
+ "with_flops": True,
24
+ "profile_memory": True,
25
+ "activities": [ProfilerActivity.CPU, ProfilerActivity.CUDA],
26
+ }
27
+ profile.__init__(self, *args, **{**defaults, **kwargs})
28
+ Profiler.__init__(self)
29
+ self.numeric_columns = [
30
+ "flops",
31
+ "count",
32
+ "self_device_time_total",
33
+ "self_cpu_time_total",
34
+ "device_time_total",
35
+ "cpu_time_total",
36
+ "self_device_memory_usage",
37
+ "self_cpu_memory_usage",
38
+ "device_memory_usage",
39
+ "cpu_memory_usage",
40
+ ]
41
+
42
+ def _get_profiler_events(self) -> EventList:
43
+ """Return the collected PyTorch events.
44
+
45
+ Returns:
46
+ The finalized PyTorch event list.
47
+
48
+ Raises:
49
+ AssertionError: If profiling has not been stopped correctly.
50
+ """
51
+ profiler = self.profiler
52
+ assert profiler is not None, "Profiling not stopped correctly"
53
+ profiler._ensure_function_events()
54
+ events = profiler._function_events
55
+ assert events is not None, "Profiler events are unavailable"
56
+ return events
57
+
58
+ def _get_profiler_events_by_record_step(self) -> dict[str, list[FunctionEvent]]:
59
+ """Group events by the most recent recorded step.
60
+
61
+ Returns:
62
+ A mapping from recorded step names to their events.
63
+ """
64
+ events = self._get_profiler_events()
65
+ matched_events = {step: [] for _, step in self.record_steps}
66
+ profiler = self.profiler
67
+ assert profiler is not None, "Profiling not stopped correctly"
68
+ kineto_results = profiler.kineto_results
69
+ assert kineto_results is not None, "Profiler trace is unavailable"
70
+ base_timestamp = kineto_results.trace_start_ns() * 1e-3
71
+
72
+ for event in events:
73
+ event_timestamp = base_timestamp + event.time_range.start
74
+ previous_steps = [
75
+ (event_timestamp - timestamp.timestamp() * 1e6, name)
76
+ for timestamp, name in self.record_steps
77
+ if event_timestamp >= timestamp.timestamp() * 1e6
78
+ ]
79
+ matched_events[min(previous_steps)[1]].append(event)
80
+ return matched_events
81
+
82
+ def _event_rows(self) -> list[dict[str, Any]]:
83
+ """Convert grouped PyTorch events into dataframe rows.
84
+
85
+ Returns:
86
+ Rows containing event metrics and their recorded step.
87
+ """
88
+ rows = []
89
+ for step, events in self._get_profiler_events_by_record_step().items():
90
+ for event in events:
91
+ row = {column: getattr(event, column, None) for column in self.numeric_columns}
92
+ row.update(
93
+ name=getattr(event, "name", None),
94
+ is_annotation=getattr(event, "is_user_annotation", None),
95
+ device=getattr(event, "device_type", None).name, # type: ignore[union-attr]
96
+ record_step=step,
97
+ )
98
+ rows.append(row)
99
+ return rows
100
+
101
+ def to_pandas(self) -> pd.DataFrame:
102
+ """Return one row per PyTorch event with timing and memory metrics.
103
+
104
+ Returns:
105
+ A dataframe containing event metrics, devices, and percentages.
106
+ """
107
+ df = pd.DataFrame(self._event_rows())
108
+ df.loc[df["device"] == "CPU", ["self_device_time_total", "device_time_total"]] = 0
109
+ df.loc[df["device"] == "CUDA", ["self_cpu_time_total", "cpu_time_total"]] = 0
110
+ df["self_cpu_time_total_percentage"] = (
111
+ df["self_cpu_time_total"] / df["self_cpu_time_total"].sum() * 100
112
+ )
113
+ df["cpu_time_total_percentage"] = df["cpu_time_total"] / df["cpu_time_total"].sum() * 100
114
+ return df
115
+
116
+ def summary(self) -> pd.DataFrame:
117
+ """Return metrics summed by event name, excluding annotations.
118
+
119
+ Returns:
120
+ A dataframe with metrics grouped and sorted by event name.
121
+ """
122
+ df = self.to_pandas()
123
+ return (
124
+ df[~df["is_annotation"]][["name"] + self.numeric_columns]
125
+ .groupby("name")
126
+ .sum()
127
+ .sort_values(by=["flops", "count"])
128
+ )
129
+
130
+ def totals(self) -> pd.Series:
131
+ """Return total metrics for non-annotation events.
132
+
133
+ Returns:
134
+ A series containing the sum of each numeric metric.
135
+ """
136
+ df = self.to_pandas()
137
+ return df[~df["is_annotation"]][self.numeric_columns].sum(axis=0)
138
+
139
+ def total_time(self, device: str = "CUDA") -> float:
140
+ """Return total self time in microseconds for one device.
141
+
142
+ Args:
143
+ device: Device to measure, either ``"CPU"`` or ``"CUDA"``.
144
+
145
+ Returns:
146
+ Total self time in microseconds.
147
+
148
+ Raises:
149
+ AssertionError: If ``device`` is not ``"CPU"`` or ``"CUDA"``.
150
+ """
151
+ assert device in {"CPU", "CUDA"}, "device must be either CPU or CUDA"
152
+ time_field = "self_cpu_time_total" if device == "CPU" else "self_device_time_total"
153
+ return sum(
154
+ getattr(event, time_field, 0.0)
155
+ for event in self._get_profiler_events()
156
+ if event.device_type.name == device and not event.is_user_annotation
157
+ )
158
+
159
+ def total_flops(self) -> int:
160
+ """Return the total FLOPs recorded by PyTorch.
161
+
162
+ Returns:
163
+ The total number of floating-point operations.
164
+ """
165
+ return int(sum((getattr(event, "flops", 0.0) or 0.0) for event in self._get_profiler_events()))
166
+
167
+ def flops_by_step(self) -> pd.DataFrame:
168
+ """Return FLOPs summed for each recorded step.
169
+
170
+ Returns:
171
+ A dataframe indexed by step name with a ``flops`` column.
172
+ """
173
+ flops_by_step = {
174
+ name: sum((event.flops or 0.0) for event in events)
175
+ for name, events in self._get_profiler_events_by_record_step().items()
176
+ }
177
+ return pd.DataFrame.from_dict(flops_by_step, orient="index", columns=["flops"])
178
+
179
+ def time_by_step(self) -> pd.DataFrame:
180
+ """Return CPU and CUDA self time in microseconds for each step.
181
+
182
+ Returns:
183
+ A dataframe indexed by step with ``cpu_time`` and ``gpu_time``.
184
+ """
185
+ time_by_step = {}
186
+ for step, events in self._get_profiler_events_by_record_step().items():
187
+ cpu_time = sum(
188
+ getattr(event, "self_cpu_time_total", 0.0)
189
+ for event in events
190
+ if event.device_type.name == "CPU" and not event.is_user_annotation
191
+ )
192
+ gpu_time = sum(
193
+ getattr(event, "self_device_time_total", 0.0)
194
+ for event in events
195
+ if event.device_type.name == "CUDA" and not event.is_user_annotation
196
+ )
197
+ time_by_step[step] = (cpu_time, gpu_time)
198
+ return pd.DataFrame.from_dict(
199
+ time_by_step, orient="index", columns=["cpu_time", "gpu_time"]
200
+ )
201
+
202
+ # Compatibility aliases for versions before the API cleanup.
203
+ get_total_time = total_time
204
+ get_total_flops = total_flops
205
+ get_flops_by_step = flops_by_step
206
+ get_time_by_step = time_by_step