PROJECT 15 / 18NEURAL ARCHITECTURESPYTHON / PYTORCH

Architecture research prototype

GRAFT-Net.

Explore attention, learned topology, and expert routing together.

3mechanism hypotheses
B × N × Npairwise topology shape
Top-kper-token expert selection
01 / IDEA02 / SYSTEM03 / PLAYGROUND04 / DECISIONS05 / SOURCE
01 / THE IDEA

A closer look.

An experimental PyTorch network that combines future-state query prediction, a learned token graph, and top-k expert routing. Its configuration and ablation paths make each mechanism separately inspectable.

Attention, graph aggregation, and expert selection offer different ways to mix information. GRAFT-Net puts these mechanisms inside one residual block so their interactions can be studied with explicit ablations.

01

Future-state query path

The attention module predicts a latent state with an MLP and uses it for query projection, while keys and values come from the current state. A flag restores the ordinary query path.

02

Learned relational structure

Pairwise MLP scores produce a soft adjacency for diagnostics and a hard top-k mask for message aggregation. Sigmoid gates control the graph contribution.

03

Inspectable expert routing

Per-token utility scores choose top-k experts. The output includes selected indices and load fractions, with a dense FFN bypass for ablation.

04

Ablation-oriented training

The repository contains separate mechanisms, task heads, auxiliary losses, baselines, and ablation configuration, supporting experiments without hiding the paths being compared.

02 / UNDER THE SURFACE

Inside a GRAFT block

Trace tensor shapes through future queries, a learned graph, gated fusion, and experts. The training-target limitation is shown explicitly.

DRAG TO PAN · SELECT A NODE · + / − TO ZOOM

Read the architecture as text
  1. Positioned states — The backbone accepts embedded states, adds a learned position embedding, applies input dropout, and runs the configured stack of GraftBlock layers.
  2. Pre-normalization — The first normalization feeds attention. Later normalizations feed topology and experts after their preceding residual updates.
  3. Future predictor — An MLP predicts a same-shaped future_state. When predictive attention is disabled, q_src is simply the current state.
  4. Predictive attention — Queries use predicted states; keys and values use current states. Scaled dot-product scores are masked, softmaxed, and multiplied by values before output projection.
  5. Attention residual — Attention output is added to x. The topology module then sees a second normalized view of that updated stream.
  6. Pairwise edge scores — MLPEdgeScorer concatenates token-pair features and scores each pair with an MLP. It materializes a dense pairwise representation before sparsification.
  7. Top-k adjacency — topk_adjacency keeps k entries per row of the edge-score matrix. The code clamps k to the final dimension and scatters selected indices into a boolean mask.
  8. Message aggregation — GraphMessagePassing projects token states, averages selected neighbor messages using batch matrix multiplication, and applies an output projection.
  9. Gated topology fusion — The topology module gates its own graph state. The block then uses another sigmoid gate to combine attention output and graph state before adding the topology delta.
  10. Utility predictor — A learned utility_predictor produces one score per token and expert. This is a learned routing signal; the name alone does not establish that gradient-derived targets are supplied.
  11. Top-k routing — topk_route selects the highest k expert scores for each token and softmax-normalizes only those selected scores.
  12. Expert mixture — For every expert selected anywhere in the batch, the implementation evaluates expert(x) on the full tensor, then weights its contribution per token. Routing sparsity therefore does not imply a measured compute speedup.
  13. Task outputs — Task modules attach heads to the backbone and return task loss plus mechanism diagnostics. Their routing_targets currently reuse routing_scores.
  14. Combined objective — compute_total_loss combines task, future prediction, routing KL, topology, and balance terms with configured coefficients. The routing KL detaches its supplied target.
  15. Optimization loop — The training loop forwards model outputs directly into the combined loss, backpropagates, clips gradient norms, and steps the optimizer. This inspected path does not replace routing targets with independent gradient utilities.
03 / INTERACTIVE STUDY

Route a token through the block

Adjust expert count, selected experts, and token-graph degree to inspect two different top-k operations on illustrative scores.

CHANGE THE INPUTS

Illustrative routing only. Top-k is clamped to available experts. This is neither a trained model nor a speed or quality benchmark.

ILLUSTRATIVE MODELLIVE

04 / ENGINEERING CHOICES

Why it works this way.

01

Keep mechanisms switchable

Predictive attention, latent topology, and routed experts each have a bypass. This creates explicit experimental comparisons rather than assuming every mechanism helps.

02

Expose dense intermediate work

Topology constructs dense pair features before choosing neighbors. Expert routing selects sparse contributions, but selected experts still process full input tensors in this implementation.

03

Separate design from evidence

The routing loss accepts gradient-utility targets, but the inspected task modules supply the routing scores themselves. The architecture should be presented as an experiment, not a demonstrated gradient-supervised advantage.

05 / OPEN THE SOURCE

Trace it back.

Implementation details, examples, and project documentation.

Scope & limitations

  • The inspected training path reuses routing scores as targets; independent gradient-utility supervision is not established.
  • No accuracy, efficiency, or generalization improvement is asserted. The interactive graph uses illustrative values rather than a trained checkpoint.

Architecture and descriptions reflect the linked repository snapshot. The playground explains a mechanism; it does not execute the repository or report measured performance.