Detecting Metal GPU Faults in PyTorch, When PyTorch Tells You Nothing
Collective Library edition. This is the complete technical report. Private filesystem paths, internal run identifiers, campaign-control notes, and repository navigation were removed. Technical claims, code, measurements, evidence labels, citations, corrections, and falsification criteria are preserved.
Written 2026-08-01 on torch 2.11.0 / macOS 26.5.2 / M2 Ultra 64 GB. Every claim marked [MEASURED] was produced by running the thing on that stack. Claims marked [REASONED] follow from measured facts but were not themselves executed. Nothing here is recalled from training data.
The short version: a Metal command-buffer abort does not raise a Python exception, and there is no error state you can poll. Training continues on uninitialised GPU memory, the tokenizer silently swallows the resulting garbage, and you publish a finite, plausible-looking WER computed from nothing. The only observable signal is the token ids themselves, and you have to look before anything decodes them.
1. What the failure actually looks like
The driver writes to stderr, not to Python:
Execution of the command buffer was aborted due to an error during execution.
Caused GPU Timeout Error (kIOGPUCommandBufferCallbackErrorTimeout)
... Impacting Interactivity (kIOGPUCommandBufferCallbackErrorImpactingInteractivity)
After that, tensors read back from the device contain uninitialised memory. Its
byte pattern is 0x01 repeated, which surfaces as recognisable constants:
| Width | Value you will see |
|---|---|
| int32 | 16843009 (0x01010101) |
| int64 | 72340172838076673 (0x0101010101010101) |
[MEASURED] Seeing either of those numbers in an index, a token id, or an "out of bounds" message is close to proof of a killed command buffer. They are worth memorising.
A related tell: processes stuck in state ?E (ps -eo pid,stat,comm | awk '$2 ~ /E/')
— unkillable, RSS 0, wedged in the kernel exit path. [MEASURED] These hold a
poisoned Metal/IOGPU context and every subsequent run on that machine inherits
it, which is why a fault can look like "my code broke" when the previous run is
the culprit. [MEASURED] They can clear on their own; verify with a probe
rather than assuming a reboot is required.
2. Three detection strategies, two of which do not work
✗ Polling torch.mps for an error state — not implementable
[MEASURED] dir(torch.mps) in torch 2.11 exposes no command-buffer error
accessor and no failed-status query. And:
strings libtorch_cpu.dylib | grep -i "command buffer"
[MEASURED] returns only Objective-C selector names — zero error strings. The message you see on stderr is emitted by the driver, below PyTorch entirely. There is nothing in the Python process to read. Do not design around a hook that does not exist.
✗ Non-finite guards — real, but they cannot reach this
Guarding loss and gradients for NaN/Inf is worth doing and catches a different failure. It cannot catch this one: generation produces integer token ids, and there is no such thing as a non-finite integer. A corrupted decode passes every finiteness check by construction.
[MEASURED] On the run this guide comes from, a separate non-finite-loss guard
did fire on a different day (grad_norm spiked to 58,463 against neighbours of
13–45, then the raw loss went non-finite). Both guards are needed; neither
substitutes for the other.
✓ Range-check token ids before anything decodes them
This is the one that works. Every id must lie in [0, len(tokenizer)).
3. Why the tokenizer hides the corruption
This is the part that makes the bug dangerous rather than merely annoying.
[MEASURED] against a real WhisperProcessor on openai/whisper-tiny:
| Input id | WhisperTokenizerFast.batch_decode |
|---|---|
16843009 |
decodes silently to '.' |
99999 |
dropped silently |
4294967295 |
dropped silently |
2**31 |
dropped silently |
72340172838076673 |
OverflowError |
-1 |
OverflowError |
[MEASURED] The slow WhisperTokenizer drops even the int64 value silently.
So a faulted batch does not crash. It returns a short, plausible string. That
string scores a finite WER — typically near 100 — and any downstream check
asking "is this a real, finite evaluation metric?" answers yes. The number is
then written to trainer_state.json, to an eval result, or to a receipt, and it
is indistinguishable from a genuinely bad model.
A corrupt number that looks valid is worse than a crash. Optimise your guards for "can this produce a plausible wrong answer?", not for "can this throw?"
4. Getting the bound right — the part that can make the fix worse than the bug
Use len(tokenizer), never tokenizer.vocab_size.
[MEASURED] Whisper's language, task, and timestamp tokens are added tokens
that live above vocab_size. Checking against vocab_size would reject
legitimate output on every single healthy run — a guard that breaks working
systems is strictly worse than the fault it prevents.
[MEASURED] across cached checkpoints:
| Checkpoint | len(tokenizer) == config.vocab_size |
|---|---|
| whisper-tiny / base | 51865 |
.en variants |
51864 |
| large-v3 | 51866 |
[REASONED] Because the logit tensor has exactly len(tokenizer) columns, a
healthy model cannot emit an id this check rejects. The guard therefore has no
false-positive mode as long as nothing calls resize_token_embeddings — worth
grepping for in your own tree before relying on it. If vocabulary surgery is ever
on the table, vocabulary and tokenizer surgery
covers what resizing does to config.vocab_size and why this guard's premise
has to be re-derived rather than merely re-run.
Two more practical notes:
- Pass the tokenizer, not the processor. The bound is
len(), and a processor has no length. - [MEASURED] Ragged input will crash the guard.
torch.as_tensorraises on ragged nested lists, which pipelines produce whenever padding happens per-batch and trailing tokens are trimmed per-row. Flatten first. This bug bit us during implementation and would have crashed every healthy pseudo-label run.
5. Where to put the check, and the trap waiting there
Put it immediately before every batch_decode of device-generated ids.
The trap: a guard inside a broad except becomes a silent no-op.
[MEASURED] In this repo, the teacher pseudo-labelling path wrapped its body
in except Exception and downgraded non-systemic failures to a per-item
decode_error, classifying "systemic" by matching strings in the message. A
newly added guard raised an exception whose message matched none of those
markers — so the raise would have been absorbed into a per-sample error, and the
run would have kept generating on faulted hardware while dropping every result
one at a time, looking exactly like ordinary audio attrition.
The guard would have fired, and the system would have reported it as normal.
Rules that follow:
- Classify hardware faults by exception type, never by message text.
- After adding a guard, trace every
exceptbetween it and the top of the run, and confirm the raise actually escapes. - Prefer a distinct exception class for "the machine miscomputed" versus "this result is not good enough". They demand different remediation, and every artifact produced after a hardware fault is suspect, not just the one that tripped.
6. Keeping coverage from rotting
Individual guards decay as new code arrives. What holds is an AST ledger: a test that enumerates every decode site in the package and fails when one appears that is neither guarded nor carries a written exemption.
Exemptions must be decided per site by reading the code. The one that matters
here: tensors that are dataloader labels still carrying the HuggingFace -100
padding sentinel are exempt — -100 is out of range by design, so guarding
them false-alarms on every healthy run, and CPU-resident labels never passed
through the compute that faults. Where -100 has already been replaced with
pad_token_id and the labels were gathered back off the device, the guard
does apply.
[MEASURED] The inventory in this repo found 13 decode sites. Anyone guessing would have said three or four.
7. Verifying the machine before trusting any number
Cheap, and it settles arguments. [MEASURED] this exact probe returned
MPS_OK loss=0.09864 finite=True at exit 0 after a fault cleared on its own:
import torch, torch.nn as nn
d = torch.device("mps")
m = nn.Sequential(nn.Linear(512, 512), nn.LayerNorm(512), nn.GELU(), nn.Linear(512, 512)).to(d)
opt = torch.optim.AdamW(m.parameters(), lr=1e-4, fused=True)
for _ in range(60):
loss = (m(torch.randn(32, 512, device=d)) ** 2).mean()
loss.backward(); opt.step(); opt.zero_grad(set_to_none=True)
torch.mps.synchronize()
print("MPS_OK", loss.item(), torch.isfinite(loss).item())
It deliberately exercises layer_norm backward and fused AdamW, the two
paths implicated in the hangs and NaNs on this stack.
Then, before trusting any number a run produced:
grep -icE "command buffer|kIOGPUCommandBufferCallbackError|Impacting Interactivity|GPU Timeout" \
run.log train.stderr train.stdout
Cross-check log-line ordering. If the fault line precedes the evaluation that
produced your metric, the metric is poisoned — regardless of how reasonable it
looks. [MEASURED] We discarded an eval_wer of 99.98 on exactly that
ordering evidence.
8. Operational notes that cost us real time
- [MEASURED] A long run must not be a child of its supervisor. A WindowServer
watchdog termination tore down the GUI login session and SIGTERM'd a training
run.
nohupblocks SIGHUP, not SIGTERM. Useos.setsid()(macOS ships nosetsid(1)) pluscaffeinate -dimsu. - [MEASURED] "Everything vanished" is not proof of a reboot. Check
last rebootagainstuptime. A login-session teardown is indistinguishable from the outside, and the crash report's ownuptimefield will confirm it. - [MEASURED] Peak driver allocation is not the tell. Ours died at 1.25 GiB and ran clean at 40.82 GiB. Do not diagnose Metal faults by memory pressure.
9. Checklist
- Range-check every device-generated token tensor before it is decoded.
- Bound on
len(tokenizer), notvocab_size. - Pass the tokenizer, not the processor.
- Flatten ragged inputs before
torch.as_tensor. - Exempt
-100-bearing dataloader labels; guard post-replacement device labels. - Raise a distinct exception type for hardware faults; classify by type.
- Trace every
exceptbetween the guard and the top of the run. - Inventory decode sites with an AST ledger so new ones cannot arrive silently.
- Grep logs for fault markers and check line ordering before quoting a metric.
- Detach long runs with
os.setsid+caffeinate.
Related
pytorch-mps.md— general MPS production notes.../whisper/whisper-training-modes-field-guide.md— per-mode settings and failure modes.../python/architectural-gates-that-do-not-rot.md— the ledger pattern in §6, and the "self-defeating gate" class §5 is an instance of.../../notes/2026-08-01-supervisor-teardown-killed-the-arm.md— the incident forensics.