Use native MPS SVD when available; document MPS usage in README - #261
Open
fnachon wants to merge 9 commits into
Open
Use native MPS SVD when available; document MPS usage in README#261fnachon wants to merge 9 commits into
fnachon wants to merge 9 commits into
Conversation
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>
torch.linalg.svd on MPS silently round-trips through CPU on stable PyTorch (with a UserWarning), and boltzgen's weighted_rigid_align already did that round-trip manually and unconditionally. PyTorch nightlies after the 2.8 branch cut ship a native MPS kernel instead. Add a cached runtime capability check (torch.backends.mps.is_available() plus a canary SVD call with warnings escalated to exceptions) so the manual CPU round-trip is skipped whenever the installed PyTorch already supports it natively, saving a GPU<->CPU sync every diffusion step. Verified against both an installed stable release (torch 2.13.0, correctly falls back) and a nightly build (torch 2.14.0.dev20260710, runs natively on-device, matches CPU output to float32 precision). Also documents how to actually run BoltzGen on MPS in the README — every step config hardcodes trainer.accelerator: gpu, and there was no existing guidance on overriding it per step, or on the nightly torch tip for this SVD path. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Follow-up to #145.
weighted_rigid_alignunconditionally round-trips the covariance matrix through CPU fortorch.linalg.svdon MPS. That's necessary on stable PyTorch, which silently does the same round-trip internally (with aUserWarning) since it has no native MPS kernel for this op. PyTorch nightlies after the 2.8 branch cut do have a native kernel.Adds a cached runtime check (
torch.backends.mps.is_available()+ a canary SVD call with warnings escalated to exceptions so the fallback is detectable) and skips the manual round-trip whenever the installed PyTorch already runs it natively — saving a GPU↔CPU sync on every diffusion step for anyone on nightly, no behavior change on stable.Verified against both:
mps:0, output matches the CPU path to float32 precision (max abs diff ~9.5e-7).Also adds a README section on actually running BoltzGen on MPS — every step config hardcodes
trainer.accelerator: gpuand there was no documented way to override it per step, nor any mention of the nightly-torch performance tip.