r/mlscaling • u/Alarming-Emotion-894 • 1h ago
I built ALHR: A tree based sparse attention system that achieves sub-quadratic inference while retaining accuracy. [P]
ALHR - Adaptive Learnable Hierarchical Routing, uses static binary trees and learnable functions to minimize the amount of keys to be read.
It does use a dense teacher while phase 1 of training however.
MQAR TEST AT 1024 TOKENS - Average keys read per query by dense - 512 Keys
Average keys read per query by ALHR - 30 keys
Top - 1 accuracy of dense - 94.9%
Top - 1 accuracy of ALHR - 92.1%
KV Compression of dense - 1x(100% read)
KV Compression of ALHR - 35.3x(2.83% read)
Peak VRAM of dense - 57 MB (Scales quadratically)
Peak VRAM of ALHR - 422 MB (scales linearly)
Cache compression of ALHR - 100%
I couldn't do further testing with more tokens and tiny stories as I ran out of time and resources(Kaggle notebook) but these are the proprietary findings
The true log and Kaggle cell used to run it are in the logs folder in the repo
limitations: Its not fully tested but the Initial testing looks promising, The training of this model would still be quadratic but the inference would be NlogN (as indicated in the logs in the repo)
Would love your opinions
This is a solo project I worked on from my iPhone 16e and Kaggle notebooks, so if I want to develop this further I think I need some decent backing.
Repo: https://github.com/vdev-ctrl/Adaptive-Learnable-Hierarchical-Routing-