trainmeter 0.0.2__py3-none-any.whl

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.
Files changed (46) hide show
  1. trainmeter/__init__.py +17 -0
  2. trainmeter/__main__.py +5 -0
  3. trainmeter/agent/__init__.py +1 -0
  4. trainmeter/agent/bootstrap/sitecustomize.py +28 -0
  5. trainmeter/agent/core.py +578 -0
  6. trainmeter/agent/sender.py +68 -0
  7. trainmeter/agent/taps.py +144 -0
  8. trainmeter/agent/wrap.py +23 -0
  9. trainmeter/catalog.py +100 -0
  10. trainmeter/cli.py +382 -0
  11. trainmeter/commands.py +99 -0
  12. trainmeter/config.py +65 -0
  13. trainmeter/doctor.py +118 -0
  14. trainmeter/emit.py +64 -0
  15. trainmeter/engine.py +621 -0
  16. trainmeter/export/__init__.py +0 -0
  17. trainmeter/export/files.py +24 -0
  18. trainmeter/export/wandb.py +87 -0
  19. trainmeter/facts.py +50 -0
  20. trainmeter/flops.py +52 -0
  21. trainmeter/metrics.py +68 -0
  22. trainmeter/passport.py +327 -0
  23. trainmeter/peaks.py +92 -0
  24. trainmeter/records.py +97 -0
  25. trainmeter/replay.py +110 -0
  26. trainmeter/report.py +47 -0
  27. trainmeter/sources/__init__.py +1 -0
  28. trainmeter/sources/gpu.py +458 -0
  29. trainmeter/sources/host.py +135 -0
  30. trainmeter/supervisor/__init__.py +1 -0
  31. trainmeter/supervisor/ingest.py +86 -0
  32. trainmeter/supervisor/launcher.py +90 -0
  33. trainmeter/supervisor/live.py +133 -0
  34. trainmeter/supervisor/store.py +118 -0
  35. trainmeter/timeline.py +102 -0
  36. trainmeter/viewer.py +273 -0
  37. trainmeter/web/__init__.py +0 -0
  38. trainmeter/web/server.py +192 -0
  39. trainmeter/web/static/app.js +571 -0
  40. trainmeter/web/static/index.html +43 -0
  41. trainmeter/web/static/style.css +157 -0
  42. trainmeter-0.0.2.dist-info/METADATA +109 -0
  43. trainmeter-0.0.2.dist-info/RECORD +46 -0
  44. trainmeter-0.0.2.dist-info/WHEEL +4 -0
  45. trainmeter-0.0.2.dist-info/entry_points.txt +3 -0
  46. trainmeter-0.0.2.dist-info/licenses/LICENSE +202 -0
trainmeter/cli.py ADDED
@@ -0,0 +1,382 @@
1
+ """The `tm` command: `tm [options] -- <command>` or `tm script.py [args]`."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ import re
7
+ import socket
8
+ import sys
9
+ import time
10
+ import warnings
11
+ from collections.abc import Mapping, Sequence
12
+ from dataclasses import dataclass, field
13
+ from pathlib import Path
14
+ from typing import Any
15
+
16
+ from . import __version__
17
+ from .config import VALUE_FLAGS, declared_facts
18
+ from .engine import Engine
19
+ from .records import make_record, redact_argv
20
+ from .report import format_elapsed, format_summary
21
+ from .sources.gpu import GpuSampler, valid_gpus_spec
22
+ from .sources.host import HostSampler, read_uptime
23
+ from .supervisor.ingest import Pipeline
24
+ from .supervisor.launcher import launch
25
+ from .supervisor.live import LiveState
26
+ from .supervisor.store import RunStore
27
+ from .web.server import DEFAULT_PORT, WebServer
28
+
29
+ USAGE = """\
30
+ usage: tm [options] -- <command> [args...]
31
+ tm [options] script.py [args...]
32
+ tm doctor [--gpus LIST] what tm can see on this node
33
+ tm passport [RUN] the run passport of a finished run, for a model card
34
+ tm export wandb [RUN] --project P upload a finished run to Weights & Biases
35
+ tm view [RUN] a finished run in the dashboard
36
+ tm attach PID watch a program that is already running, from outside
37
+
38
+ Runs the command untouched and records the run in .trainmeter/runs/.
39
+
40
+ options:
41
+ --name NAME label added to the run directory name
42
+ --gpus LIST GPUs to sample, e.g. 0,1 (default: CUDA_VISIBLE_DEVICES, else all)
43
+ --port N first port to try for the dashboard (default 8765)
44
+ --no-web do not serve the dashboard
45
+ -q, --quiet print no banner and no summary line
46
+ --no-agent do not inject the agent into the training process
47
+ --no-tap do not pick up series from wandb, TensorBoard and MLflow
48
+ --version print the version and exit
49
+ -h, --help print this help and exit
50
+
51
+ declare what only you know (20B, 1.5e9 and 300M are accepted):
52
+ --budget-tokens N --budget-steps N --budget-flops N
53
+ --tokens-per-step N --flops-per-token N --convention NAME
54
+ --params N --embedding-params N --layers N
55
+ --d-model N --seq-len N --dp N
56
+ --plan-mfu X --plan-hours H what you planned, for the Plan vs actual view
57
+ """
58
+
59
+ RUN_ID_ENV = "TRAINMETER_RUN_ID"
60
+ SOCK_ENV = "TRAINMETER_SOCK"
61
+ AGENT_ENV = "TRAINMETER_AGENT"
62
+
63
+
64
+ @dataclass(frozen=True)
65
+ class Parsed:
66
+ command: list[str]
67
+ name: str | None = None
68
+ quiet: bool = False
69
+ agent: bool = True
70
+ declared: dict[str, Any] = field(default_factory=dict)
71
+ gpus: str | None = None
72
+ port: int = DEFAULT_PORT
73
+ web: bool = True
74
+ tap: bool = True
75
+
76
+
77
+ def parse_args(argv: Sequence[str]) -> Parsed | int:
78
+ """Parse the options before the command. Returns an exit status for help, version, errors."""
79
+ opts: dict[str, Any] = {
80
+ "name": None,
81
+ "gpus": None,
82
+ "port": DEFAULT_PORT,
83
+ "web": True,
84
+ "tap": True,
85
+ "quiet": False,
86
+ "agent": True,
87
+ "values": {},
88
+ }
89
+ i = 0
90
+ while i < len(argv):
91
+ token = argv[i]
92
+ flag, eq, inline = token.partition("=")
93
+ if token == "--":
94
+ return _finish(list(argv[i + 1 :]), opts)
95
+ if token in ("-q", "--quiet"):
96
+ opts["quiet"] = True
97
+ elif token == "--no-agent":
98
+ opts["agent"] = False
99
+ elif token == "--no-tap":
100
+ opts["tap"] = False
101
+ elif token == "--no-web":
102
+ opts["web"] = False
103
+ elif flag in ("--name", "--gpus", "--port", *VALUE_FLAGS):
104
+ if eq:
105
+ value = inline
106
+ else:
107
+ i += 1
108
+ if i >= len(argv):
109
+ return _fail(f"{flag} needs a value")
110
+ value = argv[i]
111
+ if flag == "--name":
112
+ opts["name"] = value
113
+ elif flag == "--port":
114
+ if not value.isdigit() or not 0 <= int(value) <= 65535:
115
+ return _fail("--port: expected a port number, 0 picks any free port")
116
+ opts["port"] = int(value)
117
+ elif flag == "--gpus":
118
+ if not valid_gpus_spec(value):
119
+ return _fail("--gpus: expected indices such as 0,1 or GPU UUIDs")
120
+ opts["gpus"] = value
121
+ else:
122
+ opts["values"][flag] = value
123
+ elif token in ("-h", "--help"):
124
+ print(USAGE, end="")
125
+ return 0
126
+ elif token == "--version":
127
+ print(f"trainmeter {__version__}")
128
+ return 0
129
+ elif token.startswith("-"):
130
+ return _fail(f"unknown option {token}")
131
+ else:
132
+ return _finish(list(argv[i:]), opts)
133
+ i += 1
134
+ return _finish([], opts)
135
+
136
+
137
+ def _fail(message: str) -> int:
138
+ print(f"tm: {message}\n\n{USAGE}", end="", file=sys.stderr)
139
+ return 2
140
+
141
+
142
+ def _finish(command: list[str], opts: dict[str, Any]) -> Parsed | int:
143
+ if not command:
144
+ return _fail("no command given")
145
+ try:
146
+ declared = declared_facts(opts["values"])
147
+ except ValueError as exc:
148
+ return _fail(str(exc))
149
+ if command[0].endswith(".py") and Path(command[0]).is_file():
150
+ command = [sys.executable, *command] # shorthand: run the script with this interpreter
151
+ return Parsed(
152
+ command,
153
+ opts["name"],
154
+ opts["quiet"],
155
+ opts["agent"],
156
+ declared,
157
+ opts["gpus"],
158
+ opts["port"],
159
+ opts["web"],
160
+ opts["tap"],
161
+ )
162
+
163
+
164
+ def default_name(command: list[str]) -> str:
165
+ """Label for the run directory: the script for `python train.py`, else the program."""
166
+ first = Path(command[0]).name
167
+ if (
168
+ re.fullmatch(r"python[0-9.]*", first)
169
+ and len(command) > 1
170
+ and not command[1].startswith("-")
171
+ ):
172
+ return Path(command[1]).stem
173
+ return Path(first).stem
174
+
175
+
176
+ def _display(path: Path) -> str:
177
+ try:
178
+ return str(path.relative_to(Path.cwd()))
179
+ except ValueError:
180
+ return str(path)
181
+
182
+
183
+ def _bootstrap_dir() -> str:
184
+ return str(Path(__file__).parent / "agent" / "bootstrap")
185
+
186
+
187
+ def main(argv: Sequence[str] | None = None, env: Mapping[str, str] | None = None) -> int:
188
+ args = list(sys.argv[1:] if argv is None else argv)
189
+ if args[:1] == ["passport"]:
190
+ from .commands import passport_main
191
+
192
+ return passport_main(args[1:])
193
+ if args[:1] == ["export"]:
194
+ from .commands import export_main
195
+
196
+ return export_main(args[1:])
197
+ if args[:1] == ["view"]:
198
+ from .viewer import view_main
199
+
200
+ return view_main(args[1:])
201
+ if args[:1] == ["attach"]:
202
+ from .viewer import attach_main
203
+
204
+ return attach_main(args[1:])
205
+ if args[:1] == ["doctor"]:
206
+ from .doctor import main as doctor_main
207
+
208
+ return doctor_main(args[1:], env)
209
+ parsed = parse_args(args)
210
+ if isinstance(parsed, int):
211
+ return parsed
212
+ environ = dict(os.environ if env is None else env)
213
+
214
+ if RUN_ID_ENV in environ:
215
+ print("tm: already running under tm, running the command untouched", file=sys.stderr)
216
+ try:
217
+ os.execvpe(parsed.command[0], parsed.command, environ)
218
+ except OSError as exc:
219
+ print(f"tm: cannot run {parsed.command[0]!r}: {exc}", file=sys.stderr)
220
+ return 127 if isinstance(exc, FileNotFoundError) else 126
221
+
222
+ store = _open_store(parsed.name or default_name(parsed.command))
223
+ engine = Engine()
224
+ live = LiveState(engine, _header(parsed)) # feeds the engine under a lock, for every thread
225
+
226
+ def handle(record: dict[str, Any]) -> None:
227
+ if store is not None:
228
+ store.append_record(record)
229
+ live.feed(record)
230
+
231
+ pipeline = _open_pipeline(handle)
232
+ if store is not None:
233
+ environ[RUN_ID_ENV] = store.run_id
234
+ if not parsed.quiet:
235
+ print(f"trainmeter: run {_display(store.path)}", file=sys.stderr)
236
+ _write_meta(store, live.header)
237
+ handle(make_record("run.start", "supervisor", time.monotonic(), live.header, pid=os.getpid()))
238
+ _declare(handle, parsed.declared)
239
+ if pipeline is not None:
240
+ environ[SOCK_ENV] = pipeline.path
241
+ if parsed.agent:
242
+ environ["PYTHONPATH"] = os.pathsep.join(
243
+ p for p in (_bootstrap_dir(), environ.get("PYTHONPATH", "")) if p
244
+ )
245
+ if not parsed.agent:
246
+ environ[AGENT_ENV] = "0"
247
+ if not parsed.tap:
248
+ environ["TRAINMETER_TAP"] = "0"
249
+
250
+ web = WebServer(live, parsed.port) if parsed.web else None
251
+ url = web.start() if web is not None else None
252
+ if url is not None:
253
+ live.start()
254
+ if not parsed.quiet:
255
+ print(f"trainmeter: dashboard {url}", file=sys.stderr)
256
+ gpu_sampler = GpuSampler(
257
+ lambda d: handle(make_record("sample.gpu", "gpu", time.monotonic(), d)),
258
+ gpus=parsed.gpus,
259
+ cuda_visible=environ.get("CUDA_VISIBLE_DEVICES"),
260
+ )
261
+ gpu_sampler.start()
262
+ sampler: HostSampler | None = None
263
+
264
+ def on_spawn(pid: int) -> None:
265
+ nonlocal sampler
266
+ sampler = HostSampler(
267
+ pid, lambda d: handle(make_record("sample.host", "host", time.monotonic(), d))
268
+ )
269
+ sampler.start()
270
+
271
+ try:
272
+ outcome = launch(parsed.command, env=environ, on_spawn=on_spawn)
273
+ finally:
274
+ gpu_sampler.stop()
275
+ if sampler is not None:
276
+ sampler.stop()
277
+ if pipeline is not None:
278
+ pipeline.close()
279
+
280
+ end = {"elapsed_s": outcome.elapsed_s, "exit_status": outcome.status}
281
+ if outcome.signal is not None:
282
+ end["signal"] = outcome.signal
283
+ handle(make_record("run.end", "supervisor", time.monotonic(), end, pid=os.getpid()))
284
+ live.finish({"exit_status": outcome.status, "signal": outcome.signal})
285
+ live.stop()
286
+ if web is not None:
287
+ web.stop()
288
+ if outcome.error is not None:
289
+ print(f"tm: cannot run {parsed.command[0]!r}: {outcome.error}", file=sys.stderr)
290
+ snap = _snapshot(engine)
291
+ _finish_run(store, end, snap, pipeline.dropped if pipeline is not None else 0)
292
+ _write_passport(store, snap, live, engine, end)
293
+ if not parsed.quiet:
294
+ print(format_summary(outcome.elapsed_s, snap), file=sys.stderr)
295
+ return outcome.status
296
+
297
+
298
+ def _snapshot(engine: Engine) -> dict[str, Any]:
299
+ try:
300
+ return engine.snapshot()
301
+ except Exception as exc: # noqa: BLE001 - the exit status is the child's, whatever happens here
302
+ warnings.warn(f"trainmeter: could not compute the summary: {exc!r}", stacklevel=2)
303
+ return {}
304
+
305
+
306
+ def _open_store(name: str) -> RunStore | None:
307
+ try:
308
+ return RunStore.create(Path.cwd(), name)
309
+ except OSError as exc:
310
+ warnings.warn(f"trainmeter: no run log, the command runs anyway: {exc}", stacklevel=2)
311
+ return None
312
+
313
+
314
+ def _open_pipeline(handle) -> Pipeline | None:
315
+ try:
316
+ pipeline = Pipeline(handle)
317
+ pipeline.start()
318
+ return pipeline
319
+ except OSError as exc:
320
+ warnings.warn(f"trainmeter: no event socket, the command runs anyway: {exc}", stacklevel=2)
321
+ return None
322
+
323
+
324
+ def _declare(handle, declared: dict[str, Any]) -> None:
325
+ for key, value in declared.items():
326
+ handle(
327
+ make_record(
328
+ "fact", "supervisor", 0.0, {"key": key, "value": value, "source": "declared"}
329
+ )
330
+ )
331
+
332
+
333
+ def _header(parsed: Parsed) -> dict[str, Any]:
334
+ return {
335
+ "argv": redact_argv(parsed.command),
336
+ "cwd": os.getcwd(),
337
+ "host": socket.gethostname(),
338
+ "tm_version": __version__,
339
+ "declared": parsed.declared,
340
+ "node_uptime_s": read_uptime(),
341
+ }
342
+
343
+
344
+ def _write_meta(store: RunStore, header: dict[str, Any]) -> None:
345
+ try:
346
+ store.write_json("meta.json", {**header, "mono0": store.mono0, "wall0": store.wall0})
347
+ except Exception as exc: # noqa: BLE001 - the run must go on
348
+ warnings.warn(f"trainmeter: could not record the start: {exc!r}", stacklevel=2)
349
+
350
+
351
+ def _finish_run(
352
+ store: RunStore | None, end: dict[str, Any], snap: dict[str, Any], dropped: int
353
+ ) -> None:
354
+ if store is None:
355
+ return
356
+ try:
357
+ store.write_json("summary.json", {**end, "dropped_datagrams": dropped, **snap})
358
+ except Exception as exc: # noqa: BLE001
359
+ warnings.warn(f"trainmeter: could not record the end: {exc!r}", stacklevel=2)
360
+ finally:
361
+ store.close()
362
+
363
+
364
+ def _write_passport(
365
+ store: RunStore | None,
366
+ snap: dict[str, Any],
367
+ live: LiveState,
368
+ engine: Engine,
369
+ end: dict[str, Any],
370
+ ) -> None:
371
+ if store is None or not snap:
372
+ return
373
+ try:
374
+ from .export.files import write_passport
375
+ from .passport import build
376
+
377
+ write_passport(store.path, build(snap, live.header, live.facts(), engine.hello, end))
378
+ except Exception as exc: # noqa: BLE001 - a passport must never cost the exit status
379
+ warnings.warn(f"trainmeter: could not write the passport: {exc!r}", stacklevel=2)
380
+
381
+
382
+ __all__ = ["main", "parse_args", "default_name", "format_elapsed"]
trainmeter/commands.py ADDED
@@ -0,0 +1,99 @@
1
+ """Subcommands that read a finished run: `tm passport` and `tm export wandb`."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import sys
7
+ from pathlib import Path
8
+
9
+ from .export.files import write_passport
10
+ from .passport import to_markdown
11
+ from .replay import find_run, replay
12
+
13
+ PASSPORT_USAGE = """\
14
+ usage: tm passport [RUN] [--json] [--write]
15
+
16
+ Prints the run passport of RUN (a run directory or its name; default: the latest run) as Markdown.
17
+ --json print JSON instead
18
+ --write write passport.json and passport.md into the run directory
19
+ """
20
+
21
+ EXPORT_USAGE = """\
22
+ usage: tm export wandb [RUN] --project PROJECT [--entity ENTITY] [--name NAME]
23
+
24
+ Uploads a finished run to Weights & Biases. Nothing is sent unless you run this command.
25
+ Needs `pip install wandb` and a logged-in wandb, or WANDB_MODE=offline.
26
+ """
27
+
28
+
29
+ def _fail(message: str, usage: str) -> int:
30
+ print(f"tm: {message}\n\n{usage}", end="", file=sys.stderr)
31
+ return 2
32
+
33
+
34
+ def passport_main(args: list[str]) -> int:
35
+ as_json = write = False
36
+ name = None
37
+ for arg in args:
38
+ if arg == "--json":
39
+ as_json = True
40
+ elif arg == "--write":
41
+ write = True
42
+ elif arg in ("-h", "--help"):
43
+ print(PASSPORT_USAGE, end="")
44
+ return 0
45
+ elif arg.startswith("-") or name is not None:
46
+ return _fail(f"unexpected argument {arg!r}", PASSPORT_USAGE)
47
+ else:
48
+ name = arg
49
+ try:
50
+ replayed = replay(find_run(Path.cwd(), name))
51
+ except FileNotFoundError as exc:
52
+ print(f"tm: {exc}", file=sys.stderr)
53
+ return 1
54
+ data = replayed.passport()
55
+ if write:
56
+ for path in write_passport(replayed.path, data):
57
+ print(f"wrote {path}", file=sys.stderr)
58
+ sys.stdout.write(json.dumps(data, indent=2) + "\n" if as_json else to_markdown(data))
59
+ return 0
60
+
61
+
62
+ def export_main(args: list[str]) -> int:
63
+ if args[:1] != ["wandb"]:
64
+ return _fail("expected `tm export wandb`", EXPORT_USAGE)
65
+ opts: dict[str, str] = {}
66
+ name = None
67
+ rest = args[1:]
68
+ i = 0
69
+ while i < len(rest):
70
+ arg = rest[i]
71
+ if arg in ("--project", "--entity", "--name"):
72
+ i += 1
73
+ if i >= len(rest):
74
+ return _fail(f"{arg} needs a value", EXPORT_USAGE)
75
+ opts[arg] = rest[i]
76
+ elif arg in ("-h", "--help"):
77
+ print(EXPORT_USAGE, end="")
78
+ return 0
79
+ elif arg.startswith("-") or name is not None:
80
+ return _fail(f"unexpected argument {arg!r}", EXPORT_USAGE)
81
+ else:
82
+ name = arg
83
+ i += 1
84
+ if "--project" not in opts:
85
+ return _fail("--project is required", EXPORT_USAGE)
86
+ try:
87
+ replayed = replay(find_run(Path.cwd(), name))
88
+ write_passport(replayed.path, replayed.passport())
89
+ from .export.wandb import export
90
+
91
+ where = export(replayed, opts["--project"], opts.get("--entity"), opts.get("--name"))
92
+ except FileNotFoundError as exc:
93
+ print(f"tm: {exc}", file=sys.stderr)
94
+ return 1
95
+ except ImportError:
96
+ print("tm: wandb is not installed here (pip install wandb)", file=sys.stderr)
97
+ return 1
98
+ print(f"exported {replayed.path.name} to {where}", file=sys.stderr)
99
+ return 0
trainmeter/config.py ADDED
@@ -0,0 +1,65 @@
1
+ """Declared facts: the command-line flags that describe the run, turned into facts.
2
+
3
+ Pure. `tm.toml` is not read yet: it needs `tomllib`, which Python 3.10 lacks.
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ from typing import Any
9
+
10
+ from .flops import CONVENTIONS
11
+
12
+ _SUFFIX = {"k": 1e3, "m": 1e6, "b": 1e9, "t": 1e12}
13
+
14
+ # flag -> (fact key, kind). "count" accepts 20B, 1.5e9, 300M; "int" a whole number.
15
+ VALUE_FLAGS: dict[str, tuple[str, str]] = {
16
+ "--budget-tokens": ("budget_tokens", "count"),
17
+ "--budget-steps": ("budget_steps", "count"),
18
+ "--budget-flops": ("budget_flops", "count"),
19
+ "--tokens-per-step": ("tokens_per_step", "count"),
20
+ "--flops-per-token": ("flops_per_token", "count"),
21
+ "--convention": ("convention", "convention"),
22
+ "--params": ("n_params", "count"),
23
+ "--embedding-params": ("n_embedding_params", "count"),
24
+ "--layers": ("n_layers", "int"),
25
+ "--d-model": ("d_model", "int"),
26
+ "--seq-len": ("seq_len", "int"),
27
+ "--dp": ("dp", "int"),
28
+ "--plan-mfu": ("plan_mfu", "fraction"),
29
+ "--plan-hours": ("plan_hours", "positive"),
30
+ }
31
+
32
+
33
+ def parse_count(text: str) -> float:
34
+ """`20B` -> 2e10, `1.5e9` -> 1.5e9, `300M` -> 3e8. Raises ValueError otherwise."""
35
+ raw = text.strip().lower().replace("_", "")
36
+ if raw and raw[-1] in _SUFFIX:
37
+ return float(raw[:-1]) * _SUFFIX[raw[-1]]
38
+ return float(raw)
39
+
40
+
41
+ def declared_facts(options: dict[str, str]) -> dict[str, Any]:
42
+ """Fact key -> value for the value flags in `options`. Raises ValueError with a message."""
43
+ facts: dict[str, Any] = {}
44
+ for flag, raw in options.items():
45
+ key, kind = VALUE_FLAGS[flag]
46
+ try:
47
+ if kind == "convention":
48
+ if raw not in CONVENTIONS:
49
+ raise ValueError(f"expected one of {', '.join(CONVENTIONS)}")
50
+ facts[key] = raw
51
+ continue
52
+ if kind == "fraction": # 0.4 or 40%
53
+ text = raw.strip()
54
+ value = float(text[:-1]) / 100 if text.endswith("%") else float(text)
55
+ if not 0 < value <= 1:
56
+ raise ValueError("expected a fraction in (0, 1] such as 0.4 or 40%")
57
+ facts[key] = value
58
+ continue
59
+ number = parse_count(raw)
60
+ except ValueError as exc:
61
+ raise ValueError(f"{flag}: {exc}") from exc
62
+ if number <= 0:
63
+ raise ValueError(f"{flag}: must be positive")
64
+ facts[key] = int(number) if kind == "int" or number == int(number) else number
65
+ return facts
trainmeter/doctor.py ADDED
@@ -0,0 +1,118 @@
1
+ """`tm doctor`: what trainmeter can see on this node, and why anything is missing.
2
+
3
+ `run_checks` is pure over a backend and an environment, so it is tested with fakes. The live GPM
4
+ read at the end is the one line that proves the counters work; on real hardware only.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import importlib.util
10
+ import os
11
+ import sys
12
+ import time
13
+ from collections.abc import Callable, Mapping
14
+ from dataclasses import dataclass
15
+ from pathlib import Path
16
+
17
+ from . import __version__
18
+ from .peaks import lookup
19
+ from .sources.gpu import Backend, GpuUnavailable, select
20
+
21
+ OK, WARN, FAIL = "ok", "warn", "FAIL"
22
+
23
+
24
+ @dataclass(frozen=True)
25
+ class Check:
26
+ status: str
27
+ name: str
28
+ detail: str
29
+
30
+
31
+ def run_checks(
32
+ backend: Backend,
33
+ env: Mapping[str, str],
34
+ *,
35
+ gpus: str | None = None,
36
+ proc_root: Path = Path("/proc"),
37
+ sleep: Callable[[float], None] = time.sleep,
38
+ ) -> list[Check]:
39
+ checks = [Check(OK, "trainmeter", f"{__version__}, python {sys.version.split()[0]}")]
40
+ checks.append(
41
+ Check(OK if (proc_root / "uptime").exists() else WARN, "host stats", "read from /proc")
42
+ )
43
+ torch_found = importlib.util.find_spec("torch") is not None
44
+ checks.append(
45
+ Check(
46
+ OK if torch_found else WARN,
47
+ "torch",
48
+ "importable, the agent can attach" if torch_found else "not importable here",
49
+ )
50
+ )
51
+ try:
52
+ backend.open()
53
+ except GpuUnavailable as exc:
54
+ return [*checks, Check(FAIL, "NVML", f"{exc}; OFU and observed FLOPs will be absent")]
55
+ try:
56
+ versions = backend.versions()
57
+ detail = ", ".join(f"{k} {v}" for k, v in versions.items()) or "initialized"
58
+ checks.append(Check(OK, "NVML", detail))
59
+ idents = backend.identities()
60
+ indices, how, warning = select(idents, gpus, env.get("CUDA_VISIBLE_DEVICES"))
61
+ label = {"flag": "--gpus", "env": "CUDA_VISIBLE_DEVICES", "all": "all GPUs"}[how]
62
+ picked = ", ".join(map(str, indices)) or "none"
63
+ checks.append(
64
+ Check(WARN if warning or not indices else OK, "selection", f"{label}: {picked}")
65
+ )
66
+ if warning:
67
+ checks.append(Check(WARN, "selection", warning))
68
+ for i in indices:
69
+ checks.extend(_device(backend, i, idents[i], sleep))
70
+ finally:
71
+ backend.close()
72
+ return checks
73
+
74
+
75
+ def _device(backend: Backend, i: int, ident: dict, sleep) -> list[Check]:
76
+ name = ident.get("name", "?")
77
+ peak = lookup(name)
78
+ out = [
79
+ Check(
80
+ OK if peak else WARN,
81
+ f"GPU {i}",
82
+ f"{name}, dense BF16 peak {peak.bf16_dense_flops / 1e12:.1f} TFLOP/s ({peak.name})"
83
+ if peak
84
+ else f"{name}: not in the peak table, MFU and observed FLOPs will be absent",
85
+ )
86
+ ]
87
+ reason = backend.gpm_support(i)
88
+ if reason is not None:
89
+ return [*out, Check(WARN, f"GPU {i} GPM", f"{reason}; OFU and arithmetic intensity absent")]
90
+ backend.gpm(i) # primes the first sample
91
+ sleep(1.0)
92
+ counters = backend.gpm(i)
93
+ if not counters or "tensor_active" not in counters:
94
+ return [
95
+ *out,
96
+ Check(WARN, f"GPU {i} GPM", "supported, but a live read returned no counters"),
97
+ ]
98
+ shown = ", ".join(f"{k} {v:.3g}" for k, v in sorted(counters.items()))
99
+ return [*out, Check(OK, f"GPU {i} GPM", f"live read: {shown} (idle GPU: zeros are normal)")]
100
+
101
+
102
+ def format_checks(checks: list[Check]) -> str:
103
+ width = max(len(c.name) for c in checks)
104
+ return "\n".join(f"[{c.status:>4}] {c.name:<{width}} {c.detail}" for c in checks) + "\n"
105
+
106
+
107
+ def main(argv: list[str], env: Mapping[str, str] | None = None) -> int:
108
+ from .sources.gpu import default_backend
109
+
110
+ gpus = None
111
+ if argv[:1] == ["--gpus"] and len(argv) == 2:
112
+ gpus = argv[1]
113
+ elif argv:
114
+ print("usage: tm doctor [--gpus LIST]", file=sys.stderr)
115
+ return 2
116
+ checks = run_checks(default_backend(), os.environ if env is None else env, gpus=gpus)
117
+ sys.stdout.write(format_checks(checks))
118
+ return 1 if any(c.status == FAIL for c in checks) else 0