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/__init__.py ADDED
@@ -0,0 +1,17 @@
1
+ """trainmeter: where your training FLOPs go, and how well the GPUs are used."""
2
+
3
+ from .emit import emit
4
+ from .flops import CONVENTIONS, ModelShape, all_conventions, flops_per_token
5
+ from .peaks import Peak, lookup
6
+
7
+ __version__ = "0.0.2"
8
+
9
+ __all__ = [
10
+ "CONVENTIONS",
11
+ "ModelShape",
12
+ "Peak",
13
+ "all_conventions",
14
+ "emit",
15
+ "flops_per_token",
16
+ "lookup",
17
+ ]
trainmeter/__main__.py ADDED
@@ -0,0 +1,5 @@
1
+ import sys
2
+
3
+ from .cli import main
4
+
5
+ sys.exit(main())
@@ -0,0 +1 @@
1
+ """The passive agent that `tm` injects into the training process."""
@@ -0,0 +1,28 @@
1
+ """Loaded by Python at startup because `tm` puts this directory first on PYTHONPATH.
2
+
3
+ It starts the agent, then runs any other `sitecustomize` that the environment already had.
4
+ Every failure here is swallowed: the training process must start exactly as it would without tm.
5
+ """
6
+
7
+ import os
8
+ import sys
9
+
10
+ try:
11
+ from trainmeter.agent.core import install
12
+
13
+ install()
14
+ except Exception: # noqa: BLE001
15
+ pass
16
+
17
+ try:
18
+ import importlib.machinery
19
+ import importlib.util
20
+
21
+ _here = os.path.dirname(os.path.abspath(__file__))
22
+ _rest = [p for p in sys.path if os.path.abspath(p or ".") != _here]
23
+ _spec = importlib.machinery.PathFinder.find_spec("sitecustomize", _rest)
24
+ if _spec is not None and _spec.loader is not None:
25
+ _module = importlib.util.module_from_spec(_spec)
26
+ _spec.loader.exec_module(_module)
27
+ except Exception: # noqa: BLE001
28
+ pass
@@ -0,0 +1,578 @@
1
+ """The agent: passive probes that run inside the training process.
2
+
3
+ Standard library only, plus the records codec through the sender. It never imports torch itself:
4
+ it waits until the user's code does, then attaches. It never synchronizes a device, never reads a
5
+ tensor value, never allocates on a GPU, and never raises: a probe that fails disables itself,
6
+ warns once on stderr, and nothing else changes.
7
+
8
+ Verified against real torch 2.14.0 on CPU. The device probe is verified only against a fake
9
+ `torch.cuda`, not on real hardware.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import atexit
15
+ import importlib
16
+ import importlib.machinery
17
+ import os
18
+ import sys
19
+ import time
20
+ from typing import Any
21
+
22
+ from . import taps
23
+ from .sender import Sender, get_sender
24
+ from .wrap import patch
25
+
26
+ PROBES_ENV = "TRAINMETER_PROBES" # comma-separated allow-list; unset means every probe
27
+ TAP_ENV = "TRAINMETER_TAP" # "0" turns the logger taps off
28
+ COUNT_STEP_ENV = "TRAINMETER_COUNT_STEP" # optimizer steps to skip before counting FLOPs
29
+ COUNT_AFTER_STEP = 5 # warmup, compile and allocator effects are over by then
30
+
31
+ # Attributes that name the layer count and the width in the configs of well-known model families.
32
+ LAYER_ATTRS = ("num_hidden_layers", "n_layer", "num_layers", "n_layers")
33
+ WIDTH_ATTRS = ("hidden_size", "n_embd", "d_model", "dim")
34
+ EXPERT_ATTRS = ("num_experts", "num_local_experts", "n_routed_experts", "moe_num_experts")
35
+
36
+ FLUSH_INTERVAL = 0.05 # at most about 20 step records per second per rank
37
+
38
+
39
+ class Agent:
40
+ def __init__(self, sender: Sender) -> None:
41
+ self.sender = sender
42
+ self.rank = int(os.environ["RANK"]) if "RANK" in os.environ else None
43
+ self.world_size = int(os.environ.get("WORLD_SIZE", "1"))
44
+ self.torch: Any = None
45
+ self.disabled: set[str] = set()
46
+ self._seq_len_sent = False
47
+ # step state
48
+ self._saw_forward = False
49
+ self._forward_tracking = False
50
+ self._step_index = 0
51
+ self._micro = 0
52
+ self._tokens: int | None = 0 # None once any micro-batch of the step had unknown tokens
53
+ # batch state
54
+ self._root: Any = None
55
+ self._global_handle: Any = None
56
+ self._post_handle: Any = None
57
+ self._barrier = True
58
+ # aggregation
59
+ self._pending_k = 0
60
+ self._pending_micro = 0
61
+ self._pending_tokens: int | None = 0
62
+ self._last_flush = 0.0
63
+ # phases
64
+ self.attached: set[str] = set()
65
+ self._data_s = 0.0
66
+ self._data_n = 0
67
+ self._eval_start: float | None = None
68
+ self._tapped = False
69
+ # counted FLOPs: idle -> active (one whole optimizer step) -> done
70
+ self._count_state = "idle"
71
+ self._count_mode: Any = None
72
+ self._count_after = int(os.environ.get(COUNT_STEP_ENV, COUNT_AFTER_STEP))
73
+
74
+ # ---- plumbing ----------------------------------------------------------------------
75
+
76
+ def _fail(self, probe: str, exc: BaseException) -> None:
77
+ if probe not in self.disabled:
78
+ self.disabled.add(probe)
79
+ print(f"trainmeter agent: probe {probe!r} disabled: {exc!r}", file=sys.stderr)
80
+
81
+ def _send(self, kind: str, d: dict[str, Any]) -> None:
82
+ self.sender.send(kind, "agent", d, self.rank if self.rank is not None else 0)
83
+
84
+ def _fact(self, key: str, value: Any) -> None:
85
+ self._send("fact", {"key": key, "value": value, "source": "detected"})
86
+
87
+ def _enabled(self, probe: str) -> bool:
88
+ allowed = os.environ.get(PROBES_ENV)
89
+ if allowed is not None and probe not in {p.strip() for p in allowed.split(",")}:
90
+ return False
91
+ return not (probe == "tap" and os.environ.get(TAP_ENV) == "0")
92
+
93
+ # ---- attaching ---------------------------------------------------------------------
94
+
95
+ def attach(self, torch: Any) -> None:
96
+ self.torch = torch
97
+ try:
98
+ optimizer_module = importlib.import_module("torch.optim.optimizer")
99
+ optimizer_module.register_optimizer_step_post_hook(self._on_optimizer_step)
100
+ except Exception as exc: # noqa: BLE001
101
+ self._fail("steps", exc)
102
+ return
103
+ try:
104
+ from torch.nn.modules.module import register_module_forward_pre_hook
105
+
106
+ self._global_handle = register_module_forward_pre_hook(self._on_any_forward)
107
+ self._forward_tracking = True
108
+ from torch.nn.modules.module import register_module_forward_hook
109
+
110
+ # The global pre-hook cannot see keyword arguments; this one can, after forward.
111
+ self._post_handle = register_module_forward_hook(self._on_first_post, with_kwargs=True)
112
+ except Exception as exc: # noqa: BLE001
113
+ self._fail("batch", exc)
114
+ atexit.register(self.flush)
115
+ self._attach_phases(torch)
116
+ if self._enabled("flops"):
117
+ self.attached.add("flops")
118
+
119
+ # ---- phases: data wait, checkpoints and eval spans ----------------------------------
120
+
121
+ def _attach_phases(self, torch: Any) -> None:
122
+ for probe, attach in (
123
+ ("dataloader", self._attach_dataloader),
124
+ ("checkpoint", self._attach_checkpoint),
125
+ ("eval", self._attach_eval),
126
+ ):
127
+ if not self._enabled(probe):
128
+ continue
129
+ try:
130
+ attach(torch)
131
+ self.attached.add(probe)
132
+ except Exception as exc: # noqa: BLE001
133
+ self._fail(probe, exc)
134
+
135
+ def _attach_dataloader(self, torch: Any) -> None:
136
+ iterator = importlib.import_module("torch.utils.data.dataloader")._BaseDataLoaderIter
137
+
138
+ def factory(orig: Any) -> Any:
139
+ def __next__(it: Any) -> Any: # noqa: N807
140
+ start = time.monotonic()
141
+ try:
142
+ return orig(it)
143
+ finally:
144
+ self._add_data_wait(time.monotonic() - start)
145
+
146
+ return __next__
147
+
148
+ if patch(iterator, "__next__", factory) is None:
149
+ raise RuntimeError("DataLoader iterator not found or already wrapped")
150
+
151
+ def _add_data_wait(self, seconds: float) -> None:
152
+ try:
153
+ if "dataloader" not in self.disabled and self._eval_start is None:
154
+ self._data_s += seconds # waits inside an eval span belong to the eval span
155
+ self._data_n += 1
156
+ except Exception as exc: # noqa: BLE001
157
+ self._fail("dataloader", exc)
158
+
159
+ def _attach_checkpoint(self, torch: Any) -> None:
160
+ def factory(orig: Any) -> Any:
161
+ def save(*args: Any, **kwargs: Any) -> Any:
162
+ start = time.monotonic()
163
+ try:
164
+ return orig(*args, **kwargs)
165
+ finally:
166
+ self._span("checkpoint", time.monotonic() - start, "checkpoint")
167
+
168
+ return save
169
+
170
+ original = patch(torch, "save", factory)
171
+ if original is None:
172
+ raise RuntimeError("torch.save not found or already wrapped")
173
+ serialization = getattr(torch, "serialization", None)
174
+ if getattr(serialization, "save", None) is original:
175
+ serialization.save = torch.save # one wrapper, so a save is timed once
176
+
177
+ def _attach_eval(self, torch: Any) -> None:
178
+ def factory(orig: Any) -> Any:
179
+ def train(module: Any, mode: bool = True) -> Any:
180
+ self._on_train_toggle(module, mode)
181
+ return orig(module, mode)
182
+
183
+ return train
184
+
185
+ if patch(torch.nn.Module, "train", factory) is None:
186
+ raise RuntimeError("Module.train not found or already wrapped")
187
+
188
+ def _on_train_toggle(self, module: Any, mode: Any) -> None:
189
+ """An eval span runs from `eval()` on the root module to the next `train()` on it."""
190
+ try:
191
+ if "eval" in self.disabled or module is not self._root:
192
+ return
193
+ if not mode:
194
+ if self._eval_start is None:
195
+ self._eval_start = time.monotonic()
196
+ if self._count_state == "active":
197
+ self._cancel_count() # eval FLOPs are not training FLOPs
198
+ elif self._eval_start is not None:
199
+ start, self._eval_start = self._eval_start, None
200
+ self._span("eval", time.monotonic() - start, "eval")
201
+ except Exception as exc: # noqa: BLE001
202
+ self._fail("eval", exc)
203
+
204
+ def _span(self, name: str, seconds: float, probe: str) -> None:
205
+ try:
206
+ if probe not in self.disabled:
207
+ self._send("phase", {"name": name, "s": seconds, "n": 1})
208
+ except Exception as exc: # noqa: BLE001
209
+ self._fail(probe, exc)
210
+
211
+ # ---- logger taps ------------------------------------------------------------------
212
+
213
+ def _report(self, key: str, value: float, step: int | None, via: str) -> None:
214
+ if self.rank not in (None, 0):
215
+ return # one copy of a series, whichever rank the loop logs from
216
+ d: dict[str, Any] = {"key": key, "value": value, "via": via}
217
+ if step is not None:
218
+ d["step"] = step
219
+ self.sender.send("metric", "tap", d, self.rank if self.rank is not None else 0)
220
+
221
+ def _tap_fail(self, via: str, exc: BaseException) -> None:
222
+ self._fail(f"tap:{via}", exc)
223
+
224
+ def attach_tracker(self, module_name: str) -> None:
225
+ """Called once the tracker's module has finished importing."""
226
+ if not self._enabled("tap"):
227
+ return
228
+ try:
229
+ taps.ADAPTERS[module_name](self._report, self._tap_fail, self._tap_disabled)
230
+ self.attached.add("tap")
231
+ except Exception as exc: # noqa: BLE001
232
+ self._fail(f"tap:{module_name}", exc)
233
+
234
+ def on_import(self, name: str, module: Any) -> None:
235
+ if name == "torch":
236
+ self.attach(module)
237
+ else:
238
+ self.attach_tracker(name)
239
+
240
+ def _tap_disabled(self, via: str) -> bool:
241
+ return f"tap:{via}" in self.disabled
242
+
243
+ # ---- batch probe -------------------------------------------------------------------
244
+
245
+ def _training(self, module: Any) -> bool:
246
+ return bool(getattr(module, "training", False)) and self.torch.is_grad_enabled()
247
+
248
+ def _on_any_forward(self, module: Any, args: Any) -> None:
249
+ """Global pre-hook, used only until the root module is found."""
250
+ try:
251
+ if "batch" in self.disabled or self._root is not None:
252
+ return
253
+ training = self._training(module)
254
+ if training:
255
+ self._saw_forward = True
256
+ if self._barrier:
257
+ self._barrier = False
258
+ if training:
259
+ self._found_root(module, args)
260
+ except Exception as exc: # noqa: BLE001
261
+ self._fail("batch", exc)
262
+
263
+ def _found_root(self, module: Any, args: Any) -> None:
264
+ """The first training-mode module called after a barrier is the outermost one."""
265
+ self._root = module
266
+ if self._global_handle is not None:
267
+ self._global_handle.remove()
268
+ self._global_handle = None
269
+ module.register_forward_pre_hook(self._on_root_forward, with_kwargs=True)
270
+ self._micro += 1 # its tokens are counted in the post-hook, which sees keyword arguments
271
+ try:
272
+ self._params(module)
273
+ except Exception as exc: # noqa: BLE001
274
+ self._fail("params", exc)
275
+
276
+ def _on_first_post(
277
+ self, module: Any, args: Any, kwargs: Any = None, output: Any = None
278
+ ) -> None:
279
+ """Counts the tokens of the very first call of the root module, then removes itself."""
280
+ try:
281
+ if module is not self._root or self._post_handle is None:
282
+ return
283
+ self._post_handle.remove()
284
+ self._post_handle = None
285
+ self._count_tokens(args, kwargs or {})
286
+ except Exception as exc: # noqa: BLE001
287
+ self._fail("batch", exc)
288
+
289
+ def _on_root_forward(self, module: Any, args: Any, kwargs: Any) -> None:
290
+ try:
291
+ if "batch" in self.disabled or not self._training(module):
292
+ return
293
+ self._saw_forward = True
294
+ self._micro += 1
295
+ if self._micro == 1:
296
+ self._maybe_start_count()
297
+ self._count_tokens(args, kwargs)
298
+ except Exception as exc: # noqa: BLE001
299
+ self._fail("batch", exc)
300
+
301
+ # ---- counted FLOPs ----------------------------------------------------------------
302
+
303
+ def _maybe_start_count(self) -> None:
304
+ """Count one whole optimizer step, from its first forward to the step. Never raises.
305
+
306
+ The mode counts operators that execute, from their shapes: no device synchronization, no
307
+ allocation, no change to any value. It uses `FlopCounterMode` and swallows errors in the
308
+ FLOPs formulas, so a formula that cannot handle an operator cannot reach the loop.
309
+ """
310
+ if (
311
+ self._count_state != "idle"
312
+ or "flops" in self.disabled
313
+ or not self._enabled("flops")
314
+ or self._step_index < self._count_after
315
+ or self._eval_start is not None
316
+ ):
317
+ return
318
+ try:
319
+ from torch.utils.flop_counter import FlopCounterMode
320
+
321
+ class Guarded(FlopCounterMode):
322
+ def _count_flops(self, func_packet: Any, out: Any, args: Any, kwargs: Any) -> Any:
323
+ try:
324
+ return super()._count_flops(func_packet, out, args, kwargs)
325
+ except Exception: # noqa: BLE001 - the operator already ran
326
+ return out
327
+
328
+ mode = Guarded(display=False)
329
+ mode.__enter__()
330
+ self._count_mode, self._count_state = mode, "active"
331
+ except Exception as exc: # noqa: BLE001
332
+ self._count_state = "done"
333
+ self._fail("flops", exc)
334
+
335
+ def _cancel_count(self) -> None:
336
+ mode, self._count_mode = self._count_mode, None
337
+ self._count_state = "idle"
338
+ if mode is not None:
339
+ try:
340
+ mode.__exit__(None, None, None)
341
+ except Exception as exc: # noqa: BLE001
342
+ self._fail("flops", exc)
343
+
344
+ def _finish_count(self) -> None:
345
+ """Called at the optimizer step that ends the counted step."""
346
+ mode, self._count_mode = self._count_mode, None
347
+ self._count_state = "done"
348
+ try:
349
+ mode.__exit__(None, None, None)
350
+ total = int(mode.get_total_flops())
351
+ tokens = self._tokens
352
+ if total <= 0 or not tokens:
353
+ print(
354
+ "trainmeter agent: no FLOPs counted for the sampled step (torch.compile, "
355
+ "custom kernels, or unknown token count); the counted convention stays absent",
356
+ file=sys.stderr,
357
+ )
358
+ return
359
+ self._send(
360
+ "fact",
361
+ {
362
+ "key": "counted_flops_per_token",
363
+ "value": total / tokens,
364
+ "source": "measured",
365
+ "note": f"FlopCounterMode, step {self._step_index + 1}, {self._micro} "
366
+ f"micro-batches, {tokens} tokens on this rank",
367
+ },
368
+ )
369
+ except Exception as exc: # noqa: BLE001
370
+ self._fail("flops", exc)
371
+
372
+ def _count_tokens(self, args: Any, kwargs: Any) -> None:
373
+ torch = self.torch
374
+ for value in (*args, *kwargs.values()):
375
+ if (
376
+ torch.is_tensor(value)
377
+ and value.ndim == 2
378
+ and not value.dtype.is_floating_point
379
+ and not value.dtype.is_complex
380
+ and value.dtype != torch.bool
381
+ ):
382
+ if self._tokens is not None:
383
+ self._tokens += int(value.numel())
384
+ if not self._seq_len_sent:
385
+ self._seq_len_sent = True
386
+ self._fact("seq_len", int(value.shape[1]))
387
+ return
388
+ self._tokens = None
389
+
390
+ def _params(self, root: Any) -> None:
391
+ torch = self.torch
392
+ total = sum(p.numel() for p in root.parameters()) # shared weights are counted once
393
+ seen: set[int] = set()
394
+ embedding = 0
395
+ for m in root.modules():
396
+ if isinstance(m, torch.nn.Embedding) and id(m.weight) not in seen:
397
+ seen.add(id(m.weight))
398
+ embedding += m.weight.numel()
399
+ self._fact("n_params", int(total))
400
+ self._fact("n_embedding_params", int(embedding))
401
+ if self._enabled("config"):
402
+ try:
403
+ self._config(root)
404
+ except Exception as exc: # noqa: BLE001
405
+ self._fail("config", exc)
406
+
407
+ def _config(self, root: Any) -> None:
408
+ """Layers and width from a model config, when the family has one. Low confidence."""
409
+ model = root
410
+ while True: # look through DDP, FSDP and torch.compile wrappers
411
+ inner = getattr(model, "module", None) or getattr(model, "_orig_mod", None)
412
+ if inner is None or inner is model:
413
+ break
414
+ model = inner
415
+ config = getattr(model, "config", None)
416
+ if config is None:
417
+ return
418
+ for key, names in (("n_layers", LAYER_ATTRS), ("d_model", WIDTH_ATTRS)):
419
+ for name in names:
420
+ value = getattr(config, name, None)
421
+ if isinstance(value, int) and not isinstance(value, bool) and value > 0:
422
+ self._send(
423
+ "fact",
424
+ {
425
+ "key": key,
426
+ "value": value,
427
+ "source": "detected",
428
+ "note": f"model.config.{name}",
429
+ },
430
+ )
431
+ break
432
+ for name in EXPERT_ATTRS:
433
+ value = getattr(config, name, None)
434
+ if isinstance(value, int) and not isinstance(value, bool) and value > 1:
435
+ self._fact("moe_experts", value)
436
+ break
437
+
438
+ # ---- steps probe -------------------------------------------------------------------
439
+
440
+ def _on_optimizer_step(self, optimizer: Any, args: Any, kwargs: Any) -> None:
441
+ try:
442
+ if "steps" in self.disabled:
443
+ return
444
+ tracking = self._forward_tracking and "batch" not in self.disabled
445
+ if tracking and not self._saw_forward:
446
+ return # optimizer steps back to back count as one step
447
+ self._saw_forward = False
448
+ if self._root is None:
449
+ self._barrier = True
450
+ if self._count_state == "active":
451
+ self._finish_count()
452
+ self._step_index += 1
453
+ if self._step_index == 1:
454
+ self._first_step()
455
+ self._record_step()
456
+ except Exception as exc: # noqa: BLE001
457
+ self._fail("steps", exc)
458
+
459
+ def _first_step(self) -> None:
460
+ torch = self.torch
461
+ hello: dict[str, Any] = {
462
+ "python": sys.version.split()[0],
463
+ "torch": torch.__version__,
464
+ "cuda": getattr(torch.version, "cuda", None),
465
+ "world_size": self.world_size,
466
+ "probes": sorted(self.attached),
467
+ }
468
+ if self.rank is not None:
469
+ hello["local_rank"] = int(os.environ.get("LOCAL_RANK", "0"))
470
+ self._send("hello", hello)
471
+ try:
472
+ if torch.cuda.is_initialized(): # never initializes CUDA itself
473
+ index = torch.cuda.current_device()
474
+ props = torch.cuda.get_device_properties(index)
475
+ device = {"name": props.name, "index": index}
476
+ uuid = getattr(props, "uuid", None)
477
+ if uuid:
478
+ device["uuid"] = str(uuid)
479
+ self._fact("device", device)
480
+ except Exception as exc: # noqa: BLE001
481
+ self._fail("device", exc)
482
+
483
+ def _record_step(self) -> None:
484
+ self._pending_k += 1
485
+ self._pending_micro += self._micro
486
+ if self._tokens is None or self._pending_tokens is None:
487
+ self._pending_tokens = None
488
+ else:
489
+ self._pending_tokens += self._tokens
490
+ self._micro, self._tokens = 0, 0
491
+ if time.monotonic() - self._last_flush >= FLUSH_INTERVAL:
492
+ self.flush()
493
+
494
+ def flush(self) -> None:
495
+ if self._data_n:
496
+ self._send("phase", {"name": "data", "s": self._data_s, "n": self._data_n})
497
+ self._data_s, self._data_n = 0.0, 0
498
+ if not self._pending_k:
499
+ return
500
+ d: dict[str, Any] = {
501
+ "n": self._step_index,
502
+ "k": self._pending_k,
503
+ "micro": self._pending_micro,
504
+ }
505
+ if self._pending_tokens is not None and self._pending_tokens > 0:
506
+ d["tokens"] = self._pending_tokens
507
+ self._send("step", d)
508
+ self._pending_k = self._pending_micro = 0
509
+ self._pending_tokens = 0
510
+ self._last_flush = time.monotonic()
511
+
512
+
513
+ class _Loader:
514
+ """Wraps a module's own loader so the agent can attach right after the import finishes."""
515
+
516
+ def __init__(self, loader: Any, finder: _Finder, name: str) -> None:
517
+ self._loader = loader
518
+ self._finder = finder
519
+ self._name = name
520
+
521
+ def create_module(self, spec: Any) -> Any:
522
+ return self._loader.create_module(spec)
523
+
524
+ def exec_module(self, module: Any) -> None:
525
+ self._loader.exec_module(module)
526
+ self._finder.imported(self._name, module)
527
+
528
+ def __getattr__(self, name: str) -> Any:
529
+ return getattr(self._loader, name)
530
+
531
+
532
+ class _Finder:
533
+ """Meta path finder that runs the agent's attach step once for each module it waits for."""
534
+
535
+ def __init__(self, agent: Agent, names: set[str]) -> None:
536
+ self._agent = agent
537
+ self._waiting = set(names)
538
+
539
+ def find_spec(self, fullname: str, path: Any = None, target: Any = None) -> Any:
540
+ if fullname not in self._waiting:
541
+ return None
542
+ spec = importlib.machinery.PathFinder.find_spec(fullname, path)
543
+ if spec is None or spec.loader is None:
544
+ return None
545
+ spec.loader = _Loader(spec.loader, self, fullname)
546
+ return spec
547
+
548
+ def imported(self, name: str, module: Any) -> None:
549
+ self._waiting.discard(name)
550
+ if not self._waiting:
551
+ sys.meta_path[:] = [f for f in sys.meta_path if f is not self]
552
+ try:
553
+ self._agent.on_import(name, module)
554
+ except Exception as exc: # noqa: BLE001
555
+ self._agent._fail(f"attach:{name}", exc)
556
+
557
+
558
+ def install() -> None:
559
+ """Called from the bootstrap `sitecustomize`. Does nothing outside `tm` or when switched off."""
560
+ if os.environ.get("TRAINMETER_AGENT") == "0":
561
+ return
562
+ sender = get_sender()
563
+ if sender is None:
564
+ return
565
+ agent = Agent(sender)
566
+ waiting: set[str] = set()
567
+ torch = sys.modules.get("torch")
568
+ if torch is not None and hasattr(torch, "optim"):
569
+ agent.attach(torch)
570
+ else:
571
+ waiting.add("torch")
572
+ for name in taps.ADAPTERS:
573
+ if name in sys.modules:
574
+ agent.attach_tracker(name)
575
+ else:
576
+ waiting.add(name)
577
+ if waiting:
578
+ sys.meta_path.insert(0, _Finder(agent, waiting))