Fix AeroJEPA unbatched mask handling and validate chunked-decode arguments - #1999
harshaa765 wants to merge 5 commits into
Conversation
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>
CODEOWNERS review mapCurrent for commit ⏳ @mnabian — 10 file(s)
No CODEOWNER
Comment |
|
The PR appears safe to merge; both previous findings are fixed and no new actionable issues remain. FindingsSummaryAeroJEPA now honors unbatched token masks and validates chunked-decoding arguments.
Reviews (2) · Last reviewed commit: "Handle fully masked predictor context an..." |
| if context_tokens.mask is not None: | ||
| context_mask = context_tokens.mask.unsqueeze(0) |
There was a problem hiding this comment.
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.
| Raises | ||
| ------ | ||
| ValueError | ||
| If ``chunk_size`` is not positive or ``precision`` is not one | ||
| of the supported values. |
There was a problem hiding this comment.
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!
|
@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. |
Signed-off-by: Harshdeep Sharma <harsh.sharma52@gmail.com>
Signed-off-by: Harshdeep Sharma <harsh.sharma52@gmail.com>
|
@mnabian can you please review. |
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
TokenSetmasks.QueryTokenDecoderalready drops masked-out tokens. These three helpers now do the same:PrototypeTokenJEPAHeadusescontext_tokens.maskfor unbatched context. Before this fix it built an all-Truemask, so masked-out tokens were still used in cross-attention.pad_token_setscopies 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_featuresapplies a rank-2 mask. Before this fix it ignored the mask. The docstring of the example'sTokenLatentSIGRegis updated to match.2. Validate
decode_field_chunkedarguments.precisionnow raisesValueError. Before this fix it silently fell back to fp32. Matching is case-insensitive, the same as the recipe'sget_autocast_context. The docstring now describes what the code does.chunk_sizenow raisesValueErrorup front. Before this fix it failed later with an unrelatedtorch.caterror.query_sdfis optional, as it already is indecode_fieldand the decoder. TheAeroJEPATrunk.decode_queriesannotation and docstring are updated to match.QueryTokenDecoderrejects a non-positivequery_chunk_sizein its constructor. Before this fix it failed on the firstforwardcall.Only invalid inputs change behavior: valid calls give the same results as before, and the SuperWing recipe's configs (
precision: bf16) keep working.Tests
test/experimental/models/aerojepa/:pad_token_setskeeps per-set masks and averages only valid features into the global token.flatten_valid_token_featuresapplies rank-2 masks and rejects masks with the wrong shape.decode_field_chunkedrejects invalid arguments, matchesdecode_fieldin fp32 (for both"fp32"and"FP32"), and works withoutquery_sdf.QueryTokenDecoderrejects a non-positivequery_chunk_size.mainand 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.Checklist
Dependencies
None.