allscaip: fix radius-graph dtype mismatches under float64 inference - #2184
OwenPriceSkelly wants to merge 1 commit into
Conversation
biknn_radius_graph built its PBC image-offset tensors (image_id) with
torch.get_default_dtype(), and the padded output buffers (padded_disp,
src_env, dst_env) likewise. With
InferenceSettings(base_precision_dtype=torch.float64) the batch is double
while the default dtype is still float32, so inference crashed first at
torch.mm(image_id, cell) in build_radius_graph ("expected mat1 and mat2
to have the same dtype, but got: float != double") and, once past that,
at the padded_disp index_put ("Index put requires the source and
destination dtypes match").
Build image_id with the cell dtype and the padded buffers with the
dtypes of the tensors scattered into them. No behavior change at the
default float32.
Found running allscaip-md-{conserving,direct}-all-omol at float64 on
ALCF Aurora (torch 2.13 xpu), but the bug is device-independent — the
new test reproduces both crashes on CPU.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
|
Hi @OwenPriceSkelly! Thank you for your pull request and welcome to our community. Action RequiredIn order to merge any pull request (code, docs, etc.), we require contributors to sign our Contributor License Agreement, and we don't seem to have one on file for you. ProcessIn order for us to review and merge your suggested changes, please sign at https://code.facebook.com/cla. If you are contributing on behalf of someone else (eg your employer), the individual CLA may not be sufficient and your employer may need to sign the corporate CLA. Once the CLA is signed, our tooling will perform checks and validations. Afterwards, the pull request will be tagged with If you have received this in error or have any questions, please contact us at cla@meta.com. Thanks! |
|
Thank you for signing our Contributor License Agreement. We can now accept your code for this (and any) Meta Open Source project. Thanks! |
… inference Reverts the FP64 forcing (and #221's image_id patch rationale) after CPU probing against fairchem main showed AllScAIP cannot run at base_precision_dtype=float64 at all, in three layers: 1. the radius graph builds image_id at the torch default dtype -> torch.mm(image_id, cell) dtype crash (the Aurora verify failure); 2. past that, the padded disp/envelope buffers are also default-dtype -> index_put dtype crash; 3. past both (fixed upstream in facebookresearch/fairchem#2184), the backbone hard-casts its node representations to float32 before the output heads (AllScAIP.py), which then mismatch the doubled head weights -- an explicit cast no default-dtype workaround can reach. So FP64 AllScAIP needs upstream design work, not an env shim. Run the fairchem default float32 instead -- the precision the NVIDIA deployments verify at. If XPU float32 proves numerically inadequate (the reason uma.py forces FP64), verification will say so and allscaip-on-Aurora is blocked on upstream FP64 support. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
|
Scope note from further testing: this PR fixes the radius-graph construction path (both crashes above reproduce and pass in the added test), but it is not the whole float64 story for AllScAIP. Running a full tiny-model forward at fp64 on CPU (mirroring at the energy head's FFN — the backbone hard-casts its node representations to float32 before handing them to heads ( |
… inference (#224) Reverts the FP64 forcing (and #221's image_id patch rationale) after CPU probing against fairchem main showed AllScAIP cannot run at base_precision_dtype=float64 at all, in three layers: 1. the radius graph builds image_id at the torch default dtype -> torch.mm(image_id, cell) dtype crash (the Aurora verify failure); 2. past that, the padded disp/envelope buffers are also default-dtype -> index_put dtype crash; 3. past both (fixed upstream in facebookresearch/fairchem#2184), the backbone hard-casts its node representations to float32 before the output heads (AllScAIP.py), which then mismatch the doubled head weights -- an explicit cast no default-dtype workaround can reach. So FP64 AllScAIP needs upstream design work, not an env shim. Run the fairchem default float32 instead -- the precision the NVIDIA deployments verify at. If XPU float32 proves numerically inadequate (the reason uma.py forces FP64), verification will say so and allscaip-on-Aurora is blocked on upstream FP64 support. Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
What does this PR do?
Fixes two dtype crashes in AllScAIP's radius-graph construction when running float64 inference (
InferenceSettings(base_precision_dtype=torch.float64)):biknn_radius_graphbuilds the PBC image-offset tensors (image_id) withtorch.get_default_dtype()(float32), while the fp64 batch'scellis double:padded_disp,src_env,dst_env) are also created at the default dtype and then index-put with doubledisp/env:Fix: build
image_idwithdata.cell.dtypeand the padded buffers with the dtypes of the tensors scattered into them. No behavior change at the default float32 (the new dtype expressions resolve to float32 there).Found running
allscaip-md-{conserving,direct}-all-omolat fp64 throughpretrained_mlip.get_predict_unit+FAIRChemCalculatoron ALCF Aurora (torch 2.13, xpu), but the bug is device-independent — the added test reproduces both crashes on CPU at upstream main.Test
tests/core/models/allscaip/test_radius_graph_dtype.py— parametrized over float32/float64 × PBC/non-PBC; the float64 cases fail on main (first crash above; after fixing onlyimage_id, the second) and pass with this change.test_allscaip_forward.pystill passes;ruff check/ruff formatclean.🤖 Generated with Claude Code