Skip to content

Force float32 precision on CPU/MPS instead of bf16-mixed - #258

Open
fnachon wants to merge 9 commits into
HannesStark:mainfrom
fnachon:fix/force-fp32-precision-cpu-mps
Open

Force float32 precision on CPU/MPS instead of bf16-mixed#258
fnachon wants to merge 9 commits into
HannesStark:mainfrom
fnachon:fix/force-fp32-precision-cpu-mps

Conversation

@fnachon

@fnachon fnachon commented Jul 10, 2026

Copy link
Copy Markdown

design.yaml, fold.yaml, and affinity.yaml all hardcode trainer.precision: bf16-mixed, and nothing in the CLI overrides it based on accelerator. On CPUs with AVX-512 BF16 support, PyTorch Lightning's bf16-mixed silently runs ops in actual bfloat16 rather than falling back to float32, which produces structures with wrong bond lengths and atom clashes instead of erroring. MPS has the same reliability problem.

boltz hit and fixed the identical bug: jwohlwend/boltz#653. This ports that fix: in Predict.run(), force precision=32 whenever trainer.accelerator is explicitly "cpu" or "mps" (e.g. via --config <step> trainer.accelerator=mps), regardless of what the step config requested.

Given the existing MPS work in #145, this seemed worth prioritizing separately since it's a silent correctness bug, not just a crash.

fnachon and others added 9 commits January 10, 2026 15:54
Changes made to run without errors on the Mac MPS device: torch.autocast, number of devices and workers to use on M1-5 chips, workaround for CUDA-specific code, handling of float64 incompatibilities for MPS.
Replace hardcoded torch.autocast("cuda") with device-agnostic
device_type=tensor.device.type in confidence_utils, inverse_fold,
and writer modules introduced in the upstream merge.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
Python pickle does not preserve RDKit atom-level SetProp values. When
PyTorch DataLoader spawns worker processes (default num_workers=1 on
macOS), self.canonicals is pickled and all atom 'name' properties are
lost, causing KeyError in process_atom_features.

Fix: load all required molecules directly from the moldir zip inside
each get_sample() / get_feat() call instead of using the pickled
self.canonicals. The moldir zip handle is cached per-process by
_get_zipfile(), so there is no repeated I/O overhead.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
…ders

- Disable pin_memory on MPS (unsupported, causes UserWarning)
- Enable persistent_workers when num_workers > 0 (avoids repeated
  worker init overhead and the PL suggestion warning)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
design.yaml, fold.yaml, and affinity.yaml all hardcode
trainer.precision: bf16-mixed with no accelerator-conditional
override anywhere in the CLI. On CPUs with AVX-512 BF16 support,
Lightning's bf16-mixed silently runs ops in bfloat16 rather than
falling back, producing structures with wrong bond lengths and atom
clashes instead of an error. MPS has the same reliability problem.
boltz hit and fixed the identical bug (jwohlwend/boltz#653); this
port forces precision=32 whenever trainer.accelerator is explicitly
"cpu" or "mps".

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant