MissingManual

Metal / data integrity · Collective Library

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_tensor raises 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:

  1. Classify hardware faults by exception type, never by message text.
  2. After adding a guard, trace every except between it and the top of the run, and confirm the raise actually escapes.
  3. 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. nohup blocks SIGHUP, not SIGTERM. Use os.setsid() (macOS ships no setsid(1)) plus caffeinate -dimsu.
  • [MEASURED] "Everything vanished" is not proof of a reboot. Check last reboot against uptime. A login-session teardown is indistinguishable from the outside, and the crash report's own uptime field 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

  1. Range-check every device-generated token tensor before it is decoded.
  2. Bound on len(tokenizer), not vocab_size.
  3. Pass the tokenizer, not the processor.
  4. Flatten ragged inputs before torch.as_tensor.
  5. Exempt -100-bearing dataloader labels; guard post-replacement device labels.
  6. Raise a distinct exception type for hardware faults; classify by type.
  7. Trace every except between the guard and the top of the run.
  8. Inventory decode sites with an AST ledger so new ones cannot arrive silently.
  9. Grep logs for fault markers and check line ordering before quoting a metric.
  10. 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.