Skip to content

Fix AeroJEPA unbatched mask handling and validate chunked-decode arguments - #1999

Open
harshaa765 wants to merge 5 commits into
NVIDIA:mainfrom
harshaa765:fix/aerojepa-mask-and-chunked-decode
Open

harshaa765 wants to merge 5 commits into
NVIDIA:mainfrom
harshaa765:fix/aerojepa-mask-and-chunked-decode

Conversation

@harshaa765

Copy link
Copy Markdown

PhysicsNeMo Pull Request

Description

Closes #1998.

This PR fixes two groups of input-handling bugs in the experimental AeroJEPA model. The PR has one commit per group.

1. Honor unbatched TokenSet masks. QueryTokenDecoder already drops masked-out tokens. These three helpers now do the same:

  • PrototypeTokenJEPAHead uses context_tokens.mask for unbatched context. Before this fix it built an all-True mask, so masked-out tokens were still used in cross-attention.
  • pad_token_sets copies each set's mask into the packed mask and uses it for the fallback global token. It also checks the mask shape.
  • flatten_valid_token_features applies a rank-2 mask. Before this fix it ignored the mask. The docstring of the example's TokenLatentSIGReg is updated to match.

2. Validate decode_field_chunked arguments.

  • An unknown precision now raises ValueError. Before this fix it silently fell back to fp32. Matching is case-insensitive, the same as the recipe's get_autocast_context. The docstring now describes what the code does.
  • A non-positive chunk_size now raises ValueError up front. Before this fix it failed later with an unrelated torch.cat error.
  • query_sdf is optional, as it already is in decode_field and the decoder. The AeroJEPATrunk.decode_queries annotation and docstring are updated to match.
  • QueryTokenDecoder rejects a non-positive query_chunk_size in its constructor. Before this fix it failed on the first forward call.

Only invalid inputs change behavior: valid calls give the same results as before, and the SuperWing recipe's configs (precision: bf16) keep working.

Tests

  • New tests in test/experimental/models/aerojepa/:
    • Masked unbatched context gives the same prediction as the trimmed context.
    • pad_token_sets keeps per-set masks and averages only valid features into the global token.
    • flatten_valid_token_features applies rank-2 masks and rejects masks with the wrong shape.
    • decode_field_chunked rejects invalid arguments, matches decode_field in fp32 (for both "fp32" and "FP32"), and works without query_sdf.
    • QueryTokenDecoder rejects a non-positive query_chunk_size.
  • 13 of the new tests fail on main and pass with this change.
  • pytest test/experimental/models/aerojepa examples/cfd/external_aerodynamics/aerojepa/tests: 169 passed, 36 skipped. The skipped tests need CUDA. Each commit passes the suite on its own.
  • ruff check/ruff format (v0.12.5, --force-exclude), the docstring-coverage hook, and the doctests in the touched modules all pass.
  • I ran these on CPU only (aarch64, torch 2.14). The autocast tests on CUDA were not run locally.

Checklist

Dependencies

None.

Harshdeep Sharma added 3 commits September 17, 2026 08:51
PrototypeTokenJEPAHead built an all-True context mask for unbatched
inputs, so tokens marked invalid in context_tokens.mask still took part
in cross-attention. The batched path and QueryTokenDecoder already
respect the mask.

pad_token_sets likewise marked every token of each input set as valid
and averaged padding into the synthesised global token, and
flatten_valid_token_features returned rank-2 features unchanged even
when given a mask.

Use the per-set mask in all three places and validate its shape.

Refs NVIDIA#1998

Signed-off-by: Harshdeep Sharma <harsh.sharma52@gmail.com>
AeroJEPA.decode_field_chunked silently decoded in fp32 for any
unrecognised precision string, failed deep inside the decoder with an
unrelated torch.cat error for non-positive chunk sizes, and required
query_sdf even for decoders built with use_sdf=False.

Reject unknown precision values (matching case-insensitively, as the
recipe's autocast helper does) and non-positive chunk sizes up front,
and make query_sdf optional. QueryTokenDecoder now rejects a
non-positive query_chunk_size at construction instead of failing on
the first forward call.

Refs NVIDIA#1998

Signed-off-by: Harshdeep Sharma <harsh.sharma52@gmail.com>
With unbatched masks now honored, an all-False context mask leaves
PrototypeTokenJEPAHead with no context tokens, and the precomputed
cross-attention kNN then fails on the empty point set (the SciPy
backend rejects it). The same failure already existed for batched
inputs where every context token is masked. Skip the cross-neighbor
search when the filtered context is empty; the cross blocks already
pass queries through unchanged in that case.

Also follow the model docstring standards in the touched docstrings:
drop the new Raises section from decode_field_chunked in favour of the
parameter descriptions, put defaults on the parameter lines, and use
:math: notation for shapes in flatten_valid_token_features.

Refs NVIDIA#1998

Signed-off-by: Harshdeep Sharma <harsh.sharma52@gmail.com>
@harshaa765
harshaa765 requested a review from mnabian as a code owner September 17, 2026 15:22
@copy-pr-bot

copy-pr-bot Bot commented Sep 17, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@github-actions

Copy link
Copy Markdown
Contributor

CODEOWNERS review map

Current for commit 2a5c9fd2833a. An approval covers every file listed for that owner; one owner is sufficient for shared files.

⏳ @mnabian — 10 file(s)
  • examples/cfd/external_aerodynamics/aerojepa/src/losses/sigreg.py
  • physicsnemo/experimental/models/aerojepa/aerojepa.py
  • physicsnemo/experimental/models/aerojepa/decoder.py
  • physicsnemo/experimental/models/aerojepa/layers/token_utils.py
  • physicsnemo/experimental/models/aerojepa/predictor.py
  • physicsnemo/experimental/models/aerojepa/trunk.py
  • test/experimental/models/aerojepa/layers/test_token_utils.py
  • test/experimental/models/aerojepa/test_aerojepa.py
  • test/experimental/models/aerojepa/test_decoder.py
  • test/experimental/models/aerojepa/test_predictor.py

No CODEOWNER

  • CHANGELOG.md

Comment /codeowners-info to refresh.

@greptile-apps

greptile-apps Bot commented Sep 17, 2026 •

Copy link
Copy Markdown
Contributor

Retrigger

The PR appears safe to merge; both previous findings are fixed and no new actionable issues remain.

Findings

  1. P1 All-masked context fails ▶
  2. P2 Docstrings violate required format ▶

Summary

AeroJEPA now honors unbatched token masks and validates chunked-decoding arguments.

  • Fully masked predictor context skips cross-neighbor construction and reaches the cross-attention block’s existing empty-context return.
  • Regression coverage checks finite, feature-independent predictions for unbatched and batched fully masked context.
  • The docstring additions identified in the previous review now use the required sections, default declarations, and mathematical shape notation.
  • No new actionable issues were identified.

Reviews (2) · Last reviewed commit: "Handle fully masked predictor context an..."

Comment on lines +198 to +199
if context_tokens.mask is not None:
context_mask = context_tokens.mask.unsqueeze(0)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 All-masked context fails

An all-False unbatched mask now removes every context token, but the predictor still precomputes cross-attention neighbors with empty points and k=0. If SciPy is installed, the default CPU backend passes this to KDTree.query, which rejects k=0. As a result, forward fails before reaching the cross-attention block’s existing empty-context guard. Skip cross-neighbor construction when the filtered context is empty, and add an all-masked regression test.

Comment on lines +342 to +346
Raises
------
ValueError
If ``chunk_size`` is not positive or ``precision`` is not one
of the supported values.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Docstrings violate required format

The new Raises section conflicts with MOD-003d, which requires Parameters and Returns and states that sections outside its listed exceptions “are prohibited.” The newly optional query_sdf also needs default=None on its parameter declaration line under MOD-003h, and the newly documented rank-1 mask shape in flatten_valid_token_features must use :math: notation under MOD-003e. Align these additions with the repository’s documentation requirements before merging.

File Used: CODING_STANDARDS/MODELS_IMPLEMENTATION.md (source)

Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!

@harshaa765

Copy link
Copy Markdown
Author

@greptileai please re-review. a0fef0d addresses both findings: the predictor now skips the cross-neighbor search when the context is fully masked (with a regression test), and the touched docstrings follow MOD-003d/e/h.

Harshdeep Sharma added 2 commits September 17, 2026 18:32
Signed-off-by: Harshdeep Sharma <harsh.sharma52@gmail.com>
Signed-off-by: Harshdeep Sharma <harsh.sharma52@gmail.com>
@harshaa765

Copy link
Copy Markdown
Author

@mnabian can you please review.

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

🐛[BUG]: AeroJEPA ignores unbatched TokenSet masks and decode_field_chunked accepts invalid arguments

1 participant