跳到正文
r/MachineLearning· /u/Alarming-Emotion-894·· 2 小时前AI 评分43

ALHR:基于树的稀疏注意力系统实现亚二次方推理并保持精度

I built ALHR: A tree based sparse attention system that achieves sub-quadratic inference while retaining accuracy. [P]

AI 导读

开发者构建了 ALHR(Adaptive Learnable Hierarchical Routing),用静态二叉树和可学习函数减少注意力需读取的 key 数量,在 1024 token 的 MQAR 测试中每查询平均只读 30 个 key,而稠密模型需读 512 个,KV 压缩达 35.3 倍(仅读取 2.83%)。

正文

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%

The true log and Kaggle cell used to run it are in the logs folder in the repo

limitations: Full scale tests are still not completed, 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

ALHR Repository: https://github.com/vdev-ctrl/Adaptive-Learnable-Hierarchical-Routing-

submitted by /u/Alarming-Emotion-894
[link] [留言]

来源:r/MachineLearning · reddit.com