MSA؛ کرنلهای اتنشن تُنُک مینیمکس برای GPUهای SM100
خلاصهٔ کاملتر
پروژهی MSA (با اسم پکیج fmha_sm100) که تیم مینیمکس منتشر کرده، کرنلهای FlashAttention متراکم و اتنشن تُنُکِ top-k رو مخصوص معماری SM100 انویدیا فراهم میکنه. منظور از اتنشن تُنُک اینه که بهجای اینکه هر توکن به همهی توکنهای دیگه نگاه کنه، فقط مهمترین بلوکها انتخاب میشن؛ این کار باری که روی GPU میفته رو کم میکنه و برای مدلهای با متن طولانی بهصرفهتره. کل پروژه با لایسنس MIT عرضه شده.
طبق توضیحات، پکیج از دو پشتهی جداگانه ساخته شده که زیر یه پکیج پایتون کنار هم میشینن. پشتهی csrc JIT نسخهی متراکمِ FMHA بهعلاوهی ایندکسرِ sparse_topk_select رو میده که موقع اجرا از قالبهای Jinja کامپایل میشه. پشتهی CuTe-DSL هم اتنشن تُنُکِ کامل (شامل مسیر forward و decode صفحهبندیشدهی FP8) رو با فرمتهای BF16، FP8، NVFP4 و FP4 پوشش میده و موقع اجرا با cute.compile ساخته میشه. یه فایل پل (bridge) هم هست که API متراکم رو به تابع تُنُک وصل میکنه.
برای اجرا چند پیشنیاز لازمه: یه GPU از نوع SM100، نصبِ CUDA Toolkit با nvcc تو مسیر، پایتون ۳.۱۰ به بالا و لینوکسِ x86_64. مستندات هشدار میده که اولین اجرا چون sparse_topk_select رو JIT کامپایل میکنه، ممکنه از ۳۰ ثانیه تا چند دقیقه طول بکشه؛ این طبیعیه و هنگ نکرده، چون اجراهای بعدی از کشِ JIT استفاده میکنن و چند ثانیهای تموم میشن.
سادهترین راه شروع، استفاده از کتابخونهی kernels هاگینگفیسه که خودش کرنل رو میگیره و آماده میکنه:
from kernels import get_kernel
kernel_module = get_kernel("MiniMaxAI/msa", version=0)
sparse_atten_func = kernel_module.sparse_atten_func
sparse_atten_func(...)جریان اصلیِ کار سه مرحلهست: اول یه پاسِ ارزون با fmha_sm100_plan و fmha_sm100 بیشترین امتیاز هر بلوک رو حساب میکنه، بعد sparse_topk_select از روی این امتیازها مهمترین بلوکهای KV رو انتخاب میکنه، و آخرش اتنشن فقط روی همون بلوکهای انتخابشده اجرا میشه:
kv_block_indexes = sparse_topk_select(
max_score.contiguous(), topk, num_valid_pages=num_pages,
)
out, _ = fmha_sm100(
q, k_pages, v_pages, sparse_plan,
kv_indices=kv_indices,
kv_block_indexes=kv_block_indexes,
)پروژه یه مجموعهی کامل از تستها (smoke، integration و regression) و یه بنچمارک به اسم bench_sparse_attention_ops.py داره که حالتهای prefill و decode رو تو دقتهای مختلف میسنجه. بخشی از کدها هم از پروژههای شناختهشدهای مثل CUTLASS انویدیا، FlashInfer و TensorRT-LLM گرفته یا مشتق شده که هرکدوم لایسنس خودشون رو نگه داشتن.
نکات کلیدی:
- کرنلهای FlashAttention متراکم و اتنشن تُنُکِ top-k مخصوص GPUهای SM100 انویدیا، با لایسنس MIT
- دو پشته: csrc JIT (متراکم + ایندکسر) و CuTe-DSL (تُنُک کامل با FP8/NVFP4/FP4)
- اولین اجرا بهخاطر JIT چند دقیقه طول میکشه، بعدش از کش سریع میشه
- میتونی مستقیم نصبش کنی یا از کتابخونهی kernels هاگینگفیس بگیریش




