Add chunked GNO forward pass - #746
Open
abhs21 wants to merge 5 commits into
Open
Conversation
Process graph messages incrementally to reduce peak memory usage.
Verify chunked execution matches the standard GNO path and preserves gradients.
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.
Adds
chunked_forwardtoGNOBlockandIntegralTransformfor incremental sum and mean aggregation.Feature gathering happens inside each chunk for both unbatched and batched inputs. The earlier implementation gathered the entire unbatched feature array before the loop, which failed when the edge count exceeded
chunk_size.Validation:
A synthetic dense-graph benchmark on the T4 measured:
The workload has 512 input points, 512 queries, 262,144 edges, 16 feature channels, hidden widths [64, 64], nonlinear transform, mean reduction, and no positional embedding. One warmup precedes three measured repeats. Memory is peak
torch.cuda.memory_allocatedabove the pre-call baseline; timing synchronizes CUDA. Backward measures parameter gradients, not coordinate/input gradients.These are measurements for one workload, not general memory or throughput guarantees. Neighbor search still constructs the full graph, and autograd retains activations across chunks during training. Measurements used implementation commit
d9c5060; commit3f7516cadds only CUDA test coverage.Benchmark reproducer
Run in the PR checkout with CUDA PyTorch and the project dependencies installed (
zencfg==0.3.0for this environment).Related to #726.