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.
- trainmeter/__init__.py +17 -0
- trainmeter/__main__.py +5 -0
- trainmeter/agent/__init__.py +1 -0
- trainmeter/agent/bootstrap/sitecustomize.py +28 -0
- trainmeter/agent/core.py +578 -0
- trainmeter/agent/sender.py +68 -0
- trainmeter/agent/taps.py +144 -0
- trainmeter/agent/wrap.py +23 -0
- trainmeter/catalog.py +100 -0
- trainmeter/cli.py +382 -0
- trainmeter/commands.py +99 -0
- trainmeter/config.py +65 -0
- trainmeter/doctor.py +118 -0
- trainmeter/emit.py +64 -0
- trainmeter/engine.py +621 -0
- trainmeter/export/__init__.py +0 -0
- trainmeter/export/files.py +24 -0
- trainmeter/export/wandb.py +87 -0
- trainmeter/facts.py +50 -0
- trainmeter/flops.py +52 -0
- trainmeter/metrics.py +68 -0
- trainmeter/passport.py +327 -0
- trainmeter/peaks.py +92 -0
- trainmeter/records.py +97 -0
- trainmeter/replay.py +110 -0
- trainmeter/report.py +47 -0
- trainmeter/sources/__init__.py +1 -0
- trainmeter/sources/gpu.py +458 -0
- trainmeter/sources/host.py +135 -0
- trainmeter/supervisor/__init__.py +1 -0
- trainmeter/supervisor/ingest.py +86 -0
- trainmeter/supervisor/launcher.py +90 -0
- trainmeter/supervisor/live.py +133 -0
- trainmeter/supervisor/store.py +118 -0
- trainmeter/timeline.py +102 -0
- trainmeter/viewer.py +273 -0
- trainmeter/web/__init__.py +0 -0
- trainmeter/web/server.py +192 -0
- trainmeter/web/static/app.js +571 -0
- trainmeter/web/static/index.html +43 -0
- trainmeter/web/static/style.css +157 -0
- trainmeter-0.0.2.dist-info/METADATA +109 -0
- trainmeter-0.0.2.dist-info/RECORD +46 -0
- trainmeter-0.0.2.dist-info/WHEEL +4 -0
- trainmeter-0.0.2.dist-info/entry_points.txt +3 -0
- 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 @@
|
|
|
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
|
trainmeter/agent/core.py
ADDED
|
@@ -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))
|