Skip to content

allscaip: fix radius-graph dtype mismatches under float64 inference - #2184

Open
OwenPriceSkelly wants to merge 1 commit into
facebookresearch:mainfrom
OwenPriceSkelly:allscaip-radius-graph-image-id-dtype
Open

OwenPriceSkelly wants to merge 1 commit into
facebookresearch:mainfrom
OwenPriceSkelly:allscaip-radius-graph-image-id-dtype

Conversation

@OwenPriceSkelly

Copy link
Copy Markdown

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)):

  1. biknn_radius_graph builds the PBC image-offset tensors (image_id) with torch.get_default_dtype() (float32), while the fp64 batch's cell is double:
    File "src/fairchem/core/models/allscaip/utils/allscaip_radius_graph.py", line 529, in build_radius_graph
        src_pos = pos[:, None] + torch.mm(image_id, cell)[None, :]
    RuntimeError: expected mat1 and mat2 to have the same dtype, but got: float != double
    
  2. Once past that, the padded output buffers (padded_disp, src_env, dst_env) are also created at the default dtype and then index-put with double disp/env:
    RuntimeError: Index put requires the source and destination dtypes match, got Float for the destination and Double for the source.
    

Fix: build image_id with data.cell.dtype and 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-omol at fp64 through pretrained_mlip.get_predict_unit + FAIRChemCalculator on 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 only image_id, the second) and pass with this change. test_allscaip_forward.py still passes; ruff check/ruff format clean.

🤖 Generated with Claude Code

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>
@meta-cla

meta-cla Bot commented Sep 1, 2026

Copy link
Copy Markdown

Hi @OwenPriceSkelly!

Thank you for your pull request and welcome to our community.

Action Required

In 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.

Process

In 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 CLA signed. The tagging process may take up to 1 hour after signing. Please give it that time before contacting us about it.

If you have received this in error or have any questions, please contact us at cla@meta.com. Thanks!

@meta-cla

meta-cla Bot commented Sep 1, 2026

Copy link
Copy Markdown

Thank you for signing our Contributor License Agreement. We can now accept your code for this (and any) Meta Open Source project. Thanks!

OwenPriceSkelly added a commit to Garden-AI/rootstock that referenced this pull request Sep 1, 2026
… 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>
@OwenPriceSkelly

Copy link
Copy Markdown
Author

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 test_allscaip_forward.py's config, model.to(torch.float64), double batch) gets past the graph with this patch and then hits:

RuntimeError: mat1 and mat2 must have the same dtype, but got Float and Double

at the energy head's FFN — the backbone hard-casts its node representations to float32 before handing them to heads (AllScAIP.py: "node_reps": neighbor_reps[:, 0].to(torch.float32)), so double head weights mismatch. Since that cast looks deliberate (fp32 heads regardless of base precision?), I've left it out of this PR rather than guess at the intended design — but if InferenceSettings(base_precision_dtype=torch.float64) is meant to be supported for AllScAIP end-to-end, that seam (and the explicit torch.float32 masks in utils/data_preprocess.py) would need a decision from the model owners. Happy to extend this PR or file it separately, whichever you prefer.

OwenPriceSkelly added a commit to Garden-AI/rootstock that referenced this pull request Sep 1, 2026
… 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>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant