در معماری مخلوط متخصص‌ها، یک لایه به جای یک شبکه‌ی کامل، چند شبکه‌ی کوچک دارد و برای هر توکن فقط دو تای آن‌ها را حساب می‌کند. در Mixtral 8x7B این یعنی ۴۶٫۷ میلیارد پارامتر روی دیسک، اما فقط ۱۲٫۹ میلیارد پارامتر در هر توکن [1]. در این پست از config.json خود مدل شروع می‌کنیم، مسیر هشت‌خطی روتر را در کد ترنسفورمرها می‌خوانیم، و با یک اسکریپت کوچک نشان می‌دهیم چرا توزیع کار بین متخصص‌ها نامتوازن می‌شود.

روتر: هشت خطی که تصمیم می‌گیرند کدام متخصص کار کند

پیش از هر چیز باید بدانیم اعداد از کجا می‌آیند. یک مدل MoE توپولوژی خود را در config.json منتشر می‌کند و همان فایل تنها سند رسمی است.

# تعداد متخصص‌ها و تعداد انتخاب به‌ازای هر توکن را از خود مدل بخوان
$ curl -s https://huggingface.co/mistralai/Mixtral-8x7B-Instruct-v0.1/raw/main/config.json | grep -E "num_local_experts|num_experts_per_tok"
  "num_experts_per_tok": 2,
  "num_local_experts": 8,

پس Mixtral هشت متخصص دارد و برای هر توکن دو تای آن‌ها را انتخاب می‌کند [2]. یعنی ۲۵ درصد عرضه‌ی متخصص‌ها در هر لحظه محاسبه می‌شود و ۷۵ درصد بی‌کار می‌ماند.

روتر نه یک لایه‌ی عمیق، که یک ضرب خطی ساده است [4]. کد ترنسفورمرها برای Mixtral تمام مسیر را در هشت خط نشان می‌دهد.

# پیاده‌سازی واقعی روتر در کتابخانه‌ی ترنسفورمرها
$ sed -n "103,110p" modeling_mixtral.py
    def forward(self, hidden_states):
        hidden_states = hidden_states.reshape(-1, self.hidden_dim)
        router_logits = F.linear(hidden_states, self.weight)  # (seq_len, num_experts)
        router_probs = torch.nn.functional.softmax(router_logits.float(), dim=-1)
        router_top_value, router_indices = torch.topk(router_probs, self.top_k, dim=-1)
        router_top_value /= router_top_value.sum(dim=-1, keepdim=True)
        router_scores = router_top_value
        return router_logits, router_scores, router_indices

چهار گام در این هشت خط اتفاق می‌افتد و ترتیبشان مهم است: بردار حالت پنهان ضرب خطی می‌شود، نرم‌سازی نرمی اعمال می‌شود، topk دو بزرگ‌ترین را برمی‌دارد، و در گام آخر همان دو وزن دوباره بر مجموعشان تقسیم می‌شوند.

گام آخر همان چیزی است که بیشتر پیاده‌سازی‌های خانگی از قلم می‌اندازند [4]. پیش از نرمال‌سازی، مجموع وزن دو متخصص برنده در اجرای ما ۰٫۴۸۳۳ بود؛ یعنی نیمی از جرم احتمالی دور ریخته می‌شد.

بعد از تقسیم بر مجموع، همان دو وزن به ۰٫۶۴۵۸ و ۰٫۳۵۴۲ رسیدند. جمعشان دقیقاً ۱ است.

نکته‌ی دوم از همین کد بیرون می‌آید و تفاوت دو نام است. روتر router_logits برمی‌گرداند و آن‌ها با output_router_logits قابل استخراج‌اند؛ این‌ها ورودی خام‌اند و برای دیدن اینکه چه اتفاقی افتاده به کار می‌آیند. آنچه به لایه‌ی بعد می‌رود router_scores است، یعنی همان دو وزن نرمال‌شده.

چند پارامتر واقعا در هر توکن محاسبه می‌شود

حالا می‌توان ادعای اصلی را راستی‌آزمایی کرد. ادعای رسمی میسترال، منتشرشده در ۱۱ دسامبر ۲۰۲۳، این است که مدل ۴۶٫۷ میلیارد پارامتر دارد و در هر توکن ۱۲٫۹ میلیارد پارامتر را استفاده می‌کند [1]. به‌جای باور کردن، هر دو عدد را از config.json درمی‌آوریم.

قاعده‌ی محاسبه ساده است: هر متخصص سه ماتریس دارد [2]، یعنی دروازه، بالا و پایین، پس پارامتر یک متخصص برابر 3 × d_model × d_ff است. تعداد لایه‌های MoE را در این ضرب می‌کنیم، بعد یک بار با E و یک بار با k.

بخشکل مدلدر هر توکناز کجا آمد
متخصص‌ها۴۵٫۱۰ میلیارد۱۱٫۲۷ میلیارد۳۲ لایه × ۸ متخصص × ۳ × ۴۰۹۶ × ۱۴۳۳۶
توجه۱٫۳۴ میلیارد۱٫۳۴ میلیاردهمیشه اجرا؛ k و v به‌خاطر GQA کوچک‌ترند
جاسازی۰٫۲۶ میلیارد۰٫۲۶ میلیاردورودی و خروجی؛ tie_word_embeddings خاموش
روتر۰٫۰۰۱ میلیارد۰٫۰۰۱ میلیارد۳۲ لایه × ۴۰۹۶ × ۸؛ عملا بی‌هزینه
مجموع۴۶٫۷۰ میلیارد۱۲٫۸۸ میلیاردکسر فعال ۲۷٫۶ درصد

محاسبه‌ی خودم به ۴۶٫۷۰ و ۱۲٫۸۸ می‌رسد، در برابر ۴۶٫۷ و ۱۲٫۹ رسمی [1][2]. اختلاف در رقم دوم ۰٫۰۲ میلیارد است که از گرد کردن دو رقم اعشاری می‌آید، نه از اختلاف روش. این تطبیق یعنی فرمول بالا همان چیزی است که فروشنده می‌شمارد.

یک نکته که این جدول پنهان می‌کند، به زبان اجرا برمی‌گردد. در bfloat16 هر پارامتر دو بایت است، پس خواندن کل وزن‌های هر توکن ۹۳٫۴ گیگابایت می‌شد و بخش فعال ۲۵٫۸ گیگابایت. اگر k را زیاد کنید این عدد بالا می‌رود و اگر کم کنید پایین می‌آید. پس مزیت MoE یک رایگان نیست.

توزیع کار بین متخصص‌ها را خودتان ببینید

تئوری تا اینجا بود. حالا همان روتر را بدون torch و بدون numpy اجرا می‌کنیم تا ببینیم چه اتفاقی برای توزیع کار می‌افتد.

# -*- coding: utf-8 -*-
# روتر Mixtral را با کتابخانه‌ی استاندارد بازنویسی می‌کنیم
import random

E, K = 8, 2          # num_local_experts, num_experts_per_tok
random.seed(7)

def softmax(xs):
    m = max(xs)
    es = [pow(2.718281828459045, x - m) for x in xs]
    s = sum(es)
    return [e / s for e in es]

def route(logits):
    p = softmax(logits)
    # مرزهای مساوی با اندیس کوچک‌تر شکسته می‌شوند، مثل torch.topk
    idx = sorted(range(len(p)), key=lambda i: (-p[i], i))[:K]
    kept = [p[i] for i in idx]
    # خط تقسیم بر مجموع: وزن‌های بازمانده دوباره به ۱ نرمال می‌شوند
    return idx, [w / sum(kept) for w in kept], p

logits = [random.gauss(0, 1) for _ in range(E)]
idx, norm, p = route(logits)

for i in range(E):
    mark = "  -> fires" if i in idx else ""
    print(f"  expert {i}  logit {logits[i]:+.4f}  p {p[i]:.4f}{mark}")
print(f"  weights after renorm {[round(w, 4) for w in norm]} sum={sum(norm):.4f}")
print(f"  experts computed: {K} of {E} = {K / E:.0%}")

خروجی واقعی همین اجراست و با همان بذر تصادفی دوباره همان اعداد را می‌دهد.

$ python3 route.py
  expert 0  logit -0.2559  p 0.0795
  expert 1  logit +0.5114  p 0.1712  -> fires
  expert 6  logit +1.1119  p 0.3121  -> fires
  weights after renorm [0.6458, 0.3542] sum=1.0000
  experts computed: 2 of 8 = 25%

دو متخصص برنده‌اند و شش متخصص دیگر برای این توکن اصلا اجرا نمی‌شوند. حالا همان روتر را روی ۵۱۲ توکن اجرا کنیم تا ببینیم آیا بار بین هشت متخصص منصفانه پخش می‌شود.

# آیا یک متخصص کل batch را می‌بلعد؟ سهم هر متخصص از ۵۱۲ توکن
TOKENS = 512
counts = [0] * E
for _ in range(TOKENS):
    idx, _, _ = route([random.gauss(0, 1) for _ in range(E)])
    for i in idx:
        counts[i] += 1

ideal = TOKENS * K / E
print("  " + "  ".join(f"e{i}:{c}" for i, c in enumerate(counts)))
fracs = [c / (TOKENS * K) for c in counts]
cv = (sum((f - 1 / E) ** 2 for f in fracs) / E) ** 0.5 * E
print(f"  ideal {ideal:.0f} per expert, cv {cv:.4f}, "
      f"busiest/quietest {max(counts) / min(counts):.2f}")

نتیجه‌ی واقعی: e0:142 e1:109 e2:130 e3:123 e4:135 e5:136 e6:118 e7:131 در برابر سهم ایده‌آل ۱۲۸. ضریب تغییرات ۰٫۰۷۸ و نسبت شلوغ‌ترین به ساکت‌ترین ۱٫۳۰ است.

یعنی در یک اجرای سالم، بار تقریبا یکنواخت پخش می‌شود. اما این نتیجه‌ی یک بذر تصادفی است، نه تضمین. روتر یک لایه‌ی خطی آموختنی است و اگر به هم بریزد، همین جدول می‌تواند به شکل e1:900 e2:3 دربیاید. مقاله‌ی اصلی MoE همین را مسئله‌ی بارگذاری می‌داند و راه‌حلش یک جریمه‌ی کمکی برای تعادل است [8].

مدل‌های تازه چگونه تنگ‌تر کرده‌اند

نسبت k/E در مدل‌های بعدی به‌شدت کوچک شده [5][6][7]. جدول زیر را عدد به عدد از config.json هر مدل خوانده‌ام و برای مدل‌هایی که چند لایه‌ی متراکم در ابتدا دارند، تعداد لایه‌های MoE را کسر کرده‌ام.

مدلمتخصصkتراکمپارامتر کلفعال در هر توکن
Mixtral 8x7B۸۲۲۵٫۰٪۴۶٫۷۰ میلیارد۱۲٫۸۸ میلیارد
Qwen3-30B-A3B۱۲۸۸۶٫۲۵٪۳۰٫۰۸ میلیارد۲٫۹۰ میلیارد
DeepSeek-V3۲۵۶۸۳٫۱۲٪۶۶۸٫۴۱ میلیارد۳۴٫۹۴ میلیارد
Kimi K2۳۸۴۸۲٫۰۸٪۱۰۲۹٫۷۴ میلیارد۳۶٫۱۹ میلیارد

ستون آخر این جدول یک نکته را پنهان می‌کند. Mixtral با ۴۶٫۷ میلیارد پارامتر، ۱۲٫۹ میلیارد فعال دارد. DeepSeek-V3 با ۶۶۸ میلیارد پارامتر، فقط ۳۴٫۹ میلیارد فعال دارد. پس تراکم کمتر شده، اما هزینه‌ی هر توکن بالا رفته است.

تفاوت دوم در فیلدهای کنترلی است که فقط مدل‌های تازه دارند [5][6][7]. Qwen3 و DeepSeek و Kimi هر سه norm_topk_prob را روشن دارند که همان نرمال‌سازی گام آخر است. DeepSeek و Kimi علاوه بر آن n_shared_experts دارند؛ یعنی یک متخصص که همیشه اجرا می‌شود و به‌عنوان مسیر پشتیبان کار می‌کند. Mixtral چنین چیزی ندارد و فقط به روتر امتیاز می‌دهد.

سرانجام، اگر دارید مدلی را سرو می‌کنید، همین config.json نخستین جایی است که باید نگاه کنید. با دانستن E و k می‌توانید بار حافظه را پیش از دانلود وزن‌ها تخمین بزنید، و اگر moe_layer_freq داشت بدانید که همه‌ی لایه‌ها MoE نیستند. همین دو عدد تفاوت بین یک اجرای روان و یک اجرایی که در bf16 حافظه‌ی بیشتری از کارت گرافیک می‌خواهد را رقم می‌زنند.

برای زمینه‌ی دقیق‌تر، توجه چگونه کار می‌کند لایه‌ی متراکمی است که در هر توکن کامل اجرا می‌شود و هیچ‌وقت اسپارس نمی‌شود، و تفاوت RMSNorm و LayerNorm همان نرمال‌سازی پیش از روتر را پوشش می‌دهد.

منابع

  1. اعلان رسمی Mixtral، ۱۱ دسامبر ۲۰۲۳ — اعداد ۴۶٫۷B و ۱۲٫۹B که محاسبه‌ی این پست بازتولید می‌کند
  2. پیکربندی Mixtral 8x7B روی Hugging Face — منبع همه‌ی اعداد جدول‌ها
  3. پیاده‌سازی روتر در کتابخانه‌ی ترنسفورمرها، خطوط ۱۰۳ تا ۱۱۰
  4. پیکربندی Qwen3-30B-A3B — ۱۲۸ متخصص و norm_topk_prob
  5. پیکربندی DeepSeek-V3 — ۲۵۶ متخصص، متخصص مشترک و first_k_dense_replace
  6. پیکربندی Kimi K2 — ۳۸۴ متخصص و routed_scaling_factor
  7. مقاله‌ی اصلی MoE، ۲۰۱۷ — تعریف لایه‌ی کم‌گسسته و مسئله‌ی بارگذاری
  8. مقاله‌ی MegaBlocks، ۲۰۲۲ — چرا محاسبه‌ی تنک به کرنل اختصاصی نیاز دارد
  9. مرجع torch.topk — رفتار تساوی و ترتیب خروجی
  10. مستند مدل Mixtral در ترنسفورمرها
  11. اعلان نسخه‌ی Instruct Mixtral — همان روتور در یک مدل گفتگومحور