IlmHamroh
Data Science va sun'iy intellekt/Transformerlar3/12-dars37 daqiqa
Mundarija (21)

24.3-dars: Multi-head attention

24-QISM — TRANSFORMERLAR · 3-dars


1. Kirish va motivatsiya

24.2-darsda self-attention har bir so'zga jumladagi barcha so'zlarga qarash imkonini berdi. Lekin bitta attention bitta taqsimot degani: har pozitsiya uchun bitta og'irliklar qatori, yig'indisi 1. Jumlani o'qiyotgan odam esa bir vaqtning o'zida bir necha narsaga e'tibor beradi: "kesim qaysi so'z?", "bu olmosh kimga ishora qilyapti?", "oldingi so'z nima edi?", "sifat qaysi otga tegishli?". Bularning har biri — alohida bog'lanish, va ular jumlaning turli joylariga qaraydi.

Bitta attention bir nechta joyga qarashi mumkin, lekin og'irliklarni bo'lishishga majbur: 0.5 bu yerga, 0.5 u yerga — natijada qiymatlar aralashadi. Ko'p boshli attention (multi-head attention) boshqacha yo'l tutadi: d_model o'lchamli fazoni h ta kichik bo'lakka bo'ladi va har bo'lakda mustaqil attention hisoblaydi. Har bosh o'z W_q, W_k, W_v proyeksiyasiga ega, demak o'z "savolini" beradi va o'z taqsimotini hosil qiladi. Keyin boshlar natijasi ulanadi va W_o bilan aralashtiriladi.

Eng qiziq jihati: bu bepul. d_model o'zgarmasa, 1 boshli va 8 boshli attention da parametrlar soni ham, skor hisoblash FLOP i ham bir xil. O'zgaradigan narsa — faqat ichki tuzilish.

Real vaziyat. Jamoa Transformer modelini optimallashtirmoqchi bo'ldi: "8 ta bosh — ortiqcha, 1 ta bosh bilan ham parametrlar soni bir xil-ku, soddaroq bo'ladi". Bir seed bilan sinab ko'rishdi — natija deyarli bir xil chiqdi, va o'zgartirish ishlab chiqarishga ketdi. Keyingi qayta o'rgatishlarda model goh yaxshi, goh butunlay yomon natija bera boshladi. 4-misol aynan shu holatni ko'rsatadi: 1 boshli model 4 seeddan faqat 2 tasida vazifani yechdi, 4 va 8 boshli model — 4 tasida ham (4 seedda bu farq hali statistik tasdiqlanmagan — buni ham halol ko'ramiz).

Bu darsda multi-head attention ni noldan yozamiz, nn.MultiheadAttention ga og'irliklarni ko'chirib natija mosligini tekshiramiz, boshlar haqiqatan turli bog'lanishlarni o'rganishini "boshni o'chirish" orqali isbotlaymiz va bosh soni tajribasini halol o'tkazamiz.

Bu darsda:

  • Ko'p boshli attention: bo'lish, parallel attention, ulash, W_o
  • Parametrlar va hisob — bosh soniga bog'liq emas
  • nn.MultiheadAttention: in_proj_weight, maskalar, boshlar og'irligi
  • Boshlar ixtisoslashuvi va boshni o'chirish (ablation)
  • Bosh soni — giperparametr, halol tajriba
  • Cross-attention bir xil modul bilan
  • Tuzoqlar

ℹ Misollar real torch/numpy bilan (Python 3.14, torch 2.14 CPU).


2. Nazariya — chuqur tushuntirish

2.1. Ko'p boshli attention formulasi

text
KIRISH: X - (B, n, d_model);  h - boshlar soni;  d_k = d_model / h

1. PROYEKSIYA (butun d_model uchun bittadan matritsa):
   Q = X W_q,  K = X W_k,  V = X W_v          har biri (B, n, d_model)

2. BOSHLARGA BO'LISH (faqat shaklni o'zgartirish, hisob yo'q):
   (B, n, d_model) -> view (B, n, h, d_k) -> transpose (B, h, n, d_k)
   i-bosh = Q, K, V ning [i*d_k : (i+1)*d_k] ustunlari

3. HAR BOSHDA ALOHIDA ATTENTION:
   W_i = softmax(Q_i K_i^T / sqrt(d_k))       (B, h, n, n)
   O_i = W_i V_i                               (B, h, n, d_k)

4. ULASH VA ARALASHTIRISH:
   transpose (B, n, h, d_k) -> reshape (B, n, d_model)
   chiqish = concat(O_1, ..., O_h) W_o        (B, n, d_model)

MUHIM: masshtab sqrt(d_k), sqrt(d_model) EMAS - har bosh o'z d_k o'lchamida

Nima uchun W_o kerak? Ulashdan keyin har bosh natijasi o'z ustunlarida "yolg'iz" turadi. W_o ularni aralashtiradi: masalan, 0-boshdan kelgan "oldingi so'z" ma'lumoti va 1-boshdan kelgan "birinchi so'z" ma'lumoti keyingi qatlam uchun bitta vektorga birlashadi.

Multi-head = h ta kichik attention parallel: har biri d_k = d_model / h o'lchamda, o'z proyeksiyasi bilan; natijalar ulanib W_o dan o'tadi.

2.2. Parametrlar va hisob

text
PARAMETRLAR (bias bilan):
  W_q, W_k, W_v, W_o - har biri d_model x d_model + d_model
  jami = 4 * d_model^2 + 4 * d_model
  d_model = 32:  4 * 1024 + 128 = 4224     (h = 1, 2, 4, 8, 16 - hammasida)

NIMA UCHUN h GA BOG'LIQ EMAS:
  h ta bosh x (d_model x d_k) = d_model x d_model - bitta matritsa bo'laklari

HISOB (bitta misol, uzunlik n):
  skorlar: h * 2 * n^2 * d_k = 2 * n^2 * d_model   - h ga bog'liq emas
  W V:     h * 2 * n^2 * d_k                        - xuddi shunday
  proyeksiyalar: 4 * 2 * n * d_model^2              - h ga bog'liq emas

NIMA O'ZGARADI:
  og'irliklar matritsalari: h * n^2 son - h marta KO'P xotira
  har boshning "o'lchami" d_k kichrayadi: h = 16, d_model = 32 da d_k = 2
  juda kichik d_k - har bosh juda "tor" savol bera oladi

d_model o'zgarmasa, bosh soni parametr va FLOP ni o'zgartirmaydi — faqat bitta katta attention ni ko'p kichigiga bo'ladi (va og'irliklar xotirasini h marta oshiradi).

2.3. nn.MultiheadAttention ichida

text
PARAMETRLAR:
  in_proj_weight  (3 d, d)  = [W_q; W_k; W_v] ustma-ust
  in_proj_bias    (3 d,)
  out_proj.weight (d, d)    = W_o
  out_proj.bias   (d,)
  (kdim/vdim d dan farq qilsa: q_proj_weight, k_proj_weight, v_proj_weight)

CHAQIRISH:
  O, W = mha(query, key, value,
             key_padding_mask=pad,          # (B, n_k), True = E'TIBORSIZ
             attn_mask=m,                   # (n_q, n_k) yoki (B*h, n_q, n_k)
             need_weights=True,
             average_attn_weights=False)    # W: (B, h, n_q, n_k)
  average_attn_weights=True (sukut) -> W: (B, n_q, n_k) - boshlar O'RTACHASI

SHAKL KONVENSIYASI:
  batch_first=False (SUKUT!) -> (n, B, d)
  batch_first=True           -> (B, n, d)   - bu kursda doim shu

MASKA KONVENSIYASI:
  nn.MultiheadAttention key_padding_mask:  True = e'tiborsiz (PAD)
  F.scaled_dot_product_attention bool:      True = qatnashadi
  -> F.sdpa ga o'tkazishda ~pad

nn.MultiheadAttention — aynan bizning formulamiz; farqi faqat uch proyeksiyaning bitta in_proj_weight da saqlanishi, batch_first sukuti va maska konvensiyasida.

2.4. Boshlar ixtisoslashuvi

text
3-MISOL VAZIFASI - ikki xil bog'lanish:
  har pozitsiya i uchun ikki nishon:
    "oldingi"  = x[i-1]   - NISBIY pozitsiya (har joyda boshqa)
    "birinchi" = x[0]     - MUTLAQ pozitsiya (hamma uchun bir xil)

BITTA BOSH (h = 1):
  ikkala joyga qarashi kerak -> og'irlikni bo'lishadi
  seed 1: 0.32 / 0.64 - ikkalasini ham yechdi (asimmetrik aralashma)
  seed 0 va 2: 0.11 / 0.89 - "birinchi" ga yopishdi, "oldingi" ~0.44

IKKI BOSH (h = 2), 3 seed:
  seed 0, 1: bir bosh i-1 ga 0.97, ikkinchisi 0 ga 1.00 - IXTISOSLASHGAN
  seed 2:    ikkala bosh ham 0 ga qaradi, "oldingi" 0.369 - yechilmadi

BOSHNI O'CHIRISH (ablation) - ROLNI ISBOTLASH:
  bosh chiqishini nolga tenglab, aniqlikni qayta o'lchaymiz
  seed 0: "oldingi" boshi o'chirilsa -> oldingi 0.128, birinchi 1.000
          "birinchi" boshi o'chirilsa -> oldingi 1.000, birinchi 0.176
  (tasodif 1/8 = 0.125)

Og'irliklarga qarash — faqat korrelyatsiya: bosh i-1 ga qaraydi, lekin model undan foydalanadimi? Boshni o'chirish — aralashuv (intervention): agar bosh o'chirilganda aynan bitta vazifa tasodif darajasiga tushsa, ikkinchisi esa o'zgarmasa, bu bosh o'sha vazifa uchun zarur ekani isbotlanadi. Bu kichik mexanistik tahlil — katta modellarni tushunishda ham xuddi shu usul ishlatiladi.

Boshlar turli bog'lanishlarni o'rganishi mumkin, lekin bu kafolat emas — ixtisoslashuvni og'irlik bilan ko'ring, boshni o'chirish bilan isbotlang.

2.5. Bosh soni — giperparametr

text
4-MISOL: d_model = 32, h = 1, 2, 4, 8; 4 seed; 400 qadam
  h:        1        2        4        8
  o'rtacha: 0.7112   0.8510   1.0000   1.0000
  yechgan:  2/4      3/4      4/4      4/4

JUFTLASHGAN FARQ (h = 1 ga nisbatan):
  h = 4:  +0.2888, SE 0.1669 -> 2*SE mezoni bo'yicha SEZILARLI EMAS

NIMA UCHUN "4/4 va 2/4" SEZILARLI EMAS:
  natija ikki xil: seed yo yechadi 1.0-bob, yo qotib qoladi (~0.4)
  bunday "ikki cho'qqili" natijada o'rtacha va SE kam ma'lumot beradi
  4 seed - juda kam; to'g'ri o'lchov: muvaffaqiyat ULUSHI ko'proq seedda

ZAXIRA (redundancy) - bitta bosh o'chirilganda o'rtacha tushish:
  h = 2: 0.688   h = 4: 0.205   h = 8: 0.041
  ko'p boshda vazifa boshlar orasida taqsimlanadi -> barqarorroq

Amaliy qoida: Transformerlarda odatda d_k = 64 atrofida tanlanadi (masalan, d_model = 512, h = 8 yoki d_model = 768, h = 12). Juda ko'p bosh d_k ni juda kichik qiladi — har bosh tor fazoda ishlaydi; juda kam bosh — bir nechta bog'lanishni bitta taqsimotga siqadi.

Bosh sonini "ko'p — yaxshi" yoki "bitta — yetarli" deb emas, tajriba bilan tanlang — va ikki cho'qqili natijalarda ko'p seed ishlating.

2.6. Cross-attention — bir xil modul

text
SELF-ATTENTION:   mha(X, X, X)             - Q, K, V bitta ketma-ketlikdan
CROSS-ATTENTION:  mha(decoder, enc, enc)   - Q decoderdan, K = V encoderdan
  so'rov (B, 5, d), kalit/qiymat (B, 7, d) -> chiqish (B, 5, d)
  og'irliklar (B, h, 5, 7)
  key_padding_mask - ENCODER ning paddingi (kalitlar tomoni)

23.11-DARSDAGI SEQ2SEQ ATTENTION = bitta boshli cross-attention (proyeksiyasiz)
Transformer decoderida: self-attention (kauzal) + cross-attention + FFN

Bitta MultiHeadAttention moduli ham self, ham cross-attention uchun ishlaydi — farq faqat qaysi tensorlar Q, K, V sifatida berilishida.

2.7. Tuzoqlar

Asosiy tuzoqlar: d_model ni h ga bo'linmaydigan qilib tanlash; masshtabni sqrt(d_model) bilan olish (kerak: sqrt(d_k)); boshlarga bo'lishda view(B, h, n, d_k) ni to'g'ridan-to'g'ri qilish (transpose siz — ma'lumot aralashib ketadi); ulashda transpose dan keyin .contiguous() yoki reshape ni unutish; nn.MultiheadAttention da batch_first=False sukutini unutish; key_padding_mask (True = PAD) va F.scaled_dot_product_attention (True = qatnashadi) konvensiyalarini aralashtirish; need_weights=True sukut bo'yicha boshlar o'rtachasini qaytarishini bilmaslik; og'irliklarga qarab bosh rolini "isbotlangan" deb hisoblash (boshni o'chirish kerak); bosh sonini bitta seed bilan tanlash.


3. Tez ma'lumotnoma

python
import torch
import torch.nn as nn
import torch.nn.functional as F

# noldan
Q = W_q(X).view(B, n, h, d_k).transpose(1, 2)            # (B, h, n, d_k)
K = W_k(X).view(B, n, h, d_k).transpose(1, 2)
V = W_v(X).view(B, n, h, d_k).transpose(1, 2)
w = torch.softmax(Q @ K.transpose(-2, -1) / d_k ** 0.5, dim=-1)   # (B, h, n, n)
O = (w @ V).transpose(1, 2).reshape(B, n, h * d_k)
chiqish = W_o(O)

# tayyor modul
mha = nn.MultiheadAttention(d_model, h, batch_first=True)
O, W = mha(X, X, X, key_padding_mask=pad,               # pad: True = PAD
           need_weights=True, average_attn_weights=False)  # W: (B, h, n, n)

# og'irliklarni ko'chirish
mha.in_proj_weight.copy_(torch.cat([W_q.weight, W_k.weight, W_v.weight]))
mha.in_proj_bias.copy_(torch.cat([W_q.bias, W_k.bias, W_v.bias]))
mha.out_proj.weight.copy_(W_o.weight); mha.out_proj.bias.copy_(W_o.bias)

# F.sdpa bilan ko'p bosh
O = F.scaled_dot_product_attention(Q, K, V, attn_mask=~pad[:, None, None, :])

# parametrlar
n_param = 4 * d_model ** 2 + 4 * d_model

# boshni o'chirish (ablation)
O = O * ochiq[None, :, None, None]                        # ochiq: (h,) 0 yoki 1

Multi-head attention xulosasi

d_k = d_model / h; Q, K, V -> (B, h, n, d_k); har boshda attention
ulash (B, n, d_model) -> W_o
parametr 4 d^2 + 4 d va FLOP - h ga bog'liq emas; W xotirasi h marta
nn.MultiheadAttention: in_proj_weight = [W_q; W_k; W_v]; batch_first=True
key_padding_mask: True = PAD; F.sdpa: True = qatnashadi
bosh roli: og'irlik -> gumon, o'chirish -> isbot
bosh soni: tajriba, ko'p seed

4. Batafsil misollar

Misollar real torch/numpy bilan (Python 3.14, torch 2.14 CPU).

Misol 1 — Multi-head attention noldan va parametrlar hisobi

python
"""Ko'p boshli attention noldan: boshlarga bo'lish, birlashtirish va parametrlar hisobi."""

import torch
import torch.nn as nn


class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, h):
        super().__init__()
        assert d_model % h == 0, "d_model boshlar soniga bo'linishi kerak"
        self.h, self.d_k = h, d_model // h
        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def bol(self, x):
        """(B, n, d_model) -> (B, h, n, d_k)"""
        B, n, _ = x.shape
        return x.view(B, n, self.h, self.d_k).transpose(1, 2)

    def forward(self, q_kir, k_kir, v_kir):
        Q, K, V = self.bol(self.W_q(q_kir)), self.bol(self.W_k(k_kir)), self.bol(self.W_v(v_kir))
        skor = Q @ K.transpose(-2, -1) / self.d_k ** 0.5          # (B, h, n_q, n_k)
        w = torch.softmax(skor, dim=-1)
        O = w @ V                                                 # (B, h, n_q, d_k)
        B, _, n_q, _ = O.shape
        O = O.transpose(1, 2).reshape(B, n_q, self.h * self.d_k)  # boshlarni ulash
        return self.W_o(O), w


def main() -> None:
    torch.manual_seed(0)
    d_model, h = 32, 4
    mha = MultiHeadAttention(d_model, h)
    X = torch.randn(2, 6, d_model)

    print("=== 1. Shakllar qadamma-qadam (d_model = 32, h = 4) ===")
    Q = mha.W_q(X)
    print(f"  X {tuple(X.shape)} -> W_q(X) {tuple(Q.shape)}")
    print(f"  view(B, n, h, d_k)  -> {tuple(Q.view(2, 6, h, 8).shape)}")
    print(f"  transpose(1, 2)     -> {tuple(mha.bol(Q).shape)}   (har bosh alohida 'batch')")
    with torch.no_grad():
        O, w = mha(X, X, X)
    print(f"  og'irliklar {tuple(w.shape)} - har boshning O'Z (n x n) matritsasi")
    print(f"  chiqish {tuple(O.shape)} - kirish bilan bir xil shakl")

    print("\n=== 2. Boshlar - Q, K, V ning ustun bo'laklari ===")
    with torch.no_grad():
        Qb = mha.bol(mha.W_q(X))
        qolda = X @ mha.W_q.weight[8:16].T + mha.W_q.bias[8:16]   # 1-bosh ustunlari
    print(f"  1-bosh so'rovi = X @ W_q[8:16].T: maks farq {(Qb[:, 1] - qolda).abs().max().item():.2e}")
    print("  ya'ni h ta kichik (d_k = 8) attention, har biri o'z proyeksiyasi bilan")

    print("\n=== 3. Bitta boshni qo'lda hisoblash va solishtirish ===")
    with torch.no_grad():
        k = 2
        sl = slice(k * 8, (k + 1) * 8)
        q = X @ mha.W_q.weight[sl].T + mha.W_q.bias[sl]
        kk = X @ mha.W_k.weight[sl].T + mha.W_k.bias[sl]
        v = X @ mha.W_v.weight[sl].T + mha.W_v.bias[sl]
        w_k = torch.softmax(q @ kk.transpose(1, 2) / 8 ** 0.5, -1)
    print(f"  {k}-bosh og'irliklari farqi: {(w_k - w[:, k]).abs().max().item():.2e}")
    print(f"  boshlar orasidagi farq (0 va {k}): {(w[:, 0] - w[:, k]).abs().max().item():.3f}"
          " - har bosh boshqacha qaraydi")

    print("\n=== 4. Parametrlar: boshlar soniga bog'liq EMAS ===")
    print(f"  {'h':>3} {'d_k':>5} {'parametr':>9} {'4*d^2 + 4*d':>12} {'skor FLOP':>10} {'W sonlari':>10}")
    n = 128
    for hh in [1, 2, 4, 8, 16]:
        m = MultiHeadAttention(d_model, hh)
        p = sum(t.numel() for t in m.parameters())
        dk = d_model // hh
        flop = hh * 2 * n * n * dk
        w_son = hh * n * n
        print(f"  {hh:>3} {dk:>5} {p:>9} {4 * d_model ** 2 + 4 * d_model:>12} {flop:>10} {w_son:>10}")
    print(f"  n = {n}: skor FLOP = h * 2 n^2 d_k = 2 n^2 d_model - o'zgarmaydi")
    print("  lekin og'irliklar xotirasi h marta oshadi: har boshning o'z (n x n) matritsasi")

    print("\n=== 5. h = 1: oddiy (bitta boshli) self-attention ===")
    torch.manual_seed(1)
    bir = MultiHeadAttention(d_model, 1)
    with torch.no_grad():
        O1, _ = bir(X, X, X)
        q, kk, v = bir.W_q(X), bir.W_k(X), bir.W_v(X)
        oddiy = bir.W_o(torch.softmax(q @ kk.transpose(1, 2) / d_model ** 0.5, -1) @ v)
    print(f"  maks farq: {(O1 - oddiy).abs().max().item():.2e}")
    print("  ⭐ Ko'p boshli = d_model ni h bo'lakka bo'lib, h ta attention parallel + W_o")


if __name__ == "__main__":
    main()

Natijaning muhim qismi:

text
=== 1. Shakllar qadamma-qadam (d_model = 32, h = 4) ===
  X (2, 6, 32) -> W_q(X) (2, 6, 32)
  view(B, n, h, d_k)  -> (2, 6, 4, 8)
  transpose(1, 2)     -> (2, 4, 6, 8)   (har bosh alohida 'batch')
  og'irliklar (2, 4, 6, 6) - har boshning O'Z (n x n) matritsasi
  chiqish (2, 6, 32) - kirish bilan bir xil shakl

=== 2. Boshlar - Q, K, V ning ustun bo'laklari ===
  1-bosh so'rovi = X @ W_q[8:16].T: maks farq 2.38e-07
  ya'ni h ta kichik (d_k = 8) attention, har biri o'z proyeksiyasi bilan

=== 3. Bitta boshni qo'lda hisoblash va solishtirish ===
  2-bosh og'irliklari farqi: 2.98e-08
  boshlar orasidagi farq (0 va 2): 0.248 - har bosh boshqacha qaraydi

=== 4. Parametrlar: boshlar soniga bog'liq EMAS ===
    h   d_k  parametr  4*d^2 + 4*d  skor FLOP  W sonlari
    1    32      4224         4224    1048576      16384
    2    16      4224         4224    1048576      32768
    4     8      4224         4224    1048576      65536
    8     4      4224         4224    1048576     131072
   16     2      4224         4224    1048576     262144
  n = 128: skor FLOP = h * 2 n^2 d_k = 2 n^2 d_model - o'zgarmaydi
  lekin og'irliklar xotirasi h marta oshadi: har boshning o'z (n x n) matritsasi

=== 5. h = 1: oddiy (bitta boshli) self-attention ===
  maks farq: 0.00e+00
  ⭐ Ko'p boshli = d_model ni h bo'lakka bo'lib, h ta attention parallel + W_o

Nima ko'rsatdi:

  • 1-bo'lim: boshlarga bo'lish — faqat shakl o'zgartirish: (2, 6, 32) -> (2, 6, 4, 8) -> (2, 4, 6, 8). transpose(1, 2) dan keyin bosh o'qi batch o'qi yonida turadi, va keyingi barcha matritsa ko'paytmalari har bosh uchun alohida bajariladi. Og'irliklar (2, 4, 6, 6) — har boshning o'z (n x n) matritsasi.
  • 2-bo'lim: 1-boshning so'rovi — W_q ning 8-16 qatorlari bilan hisoblangan kichik proyeksiya (farq 2.38e-07). Ya'ni bitta katta W_q aslida 4 ta mustaqil 8 x 32 matritsaning ustma-ust joylashuvi.
  • 3-bo'lim: 2-boshni butunlay qo'lda hisobladik — og'irliklar ~3e-08 aniqlikda mos. 0- va 2-boshlar og'irliklari orasidagi farq 0.248: o'rgatilmagan bo'lsa ham, har bosh boshqa proyeksiya bilan boshqacha qaraydi.
  • 4-bo'lim: h = 1 dan h = 16 gacha parametrlar 4224 (4 * 32^2 + 4 * 32) va skor FLOP 1 048 576 o'zgarmadi. O'zgargani — og'irliklar matritsalari soni: 16 384 dan 262 144 gacha (h marta).
  • 5-bo'lim: h = 1 bo'lgan multi-head attention — W_o qo'shilgan oddiy self-attention, farq aynan 0.

Misol 2 — nn.MultiheadAttention bilan moslik, maskalar va cross-attention

python
"""nn.MultiheadAttention: og'irliklarni ko'chirib moslikni tekshirish, maskalar va boshlar og'irligi."""

import torch
import torch.nn as nn
import torch.nn.functional as F


class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, h):
        super().__init__()
        self.h, self.d_k = h, d_model // h
        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def bol(self, x):
        B, n, _ = x.shape
        return x.view(B, n, self.h, self.d_k).transpose(1, 2)

    def forward(self, q_kir, k_kir, v_kir, pad=None):
        """pad: (B, n_k) bool, True - PAD (e'tiborsiz), nn.MultiheadAttention kabi."""
        Q, K, V = self.bol(self.W_q(q_kir)), self.bol(self.W_k(k_kir)), self.bol(self.W_v(v_kir))
        skor = Q @ K.transpose(-2, -1) / self.d_k ** 0.5
        if pad is not None:
            skor = skor.masked_fill(pad[:, None, None, :], float("-inf"))
        w = torch.softmax(skor, dim=-1)
        B, _, n_q, _ = Q.shape
        O = (w @ V).transpose(1, 2).reshape(B, n_q, -1)
        return self.W_o(O), w


def main() -> None:
    torch.manual_seed(0)
    d, h = 32, 4
    bizniki = MultiHeadAttention(d, h)
    torch_mha = nn.MultiheadAttention(d, h, batch_first=True)

    print("=== 1. nn.MultiheadAttention ichidagi parametrlar ===")
    for nom, p in torch_mha.named_parameters():
        print(f"  {nom:<20} {tuple(p.shape)}")
    print(f"  jami: {sum(p.numel() for p in torch_mha.parameters())}, "
          f"bizniki: {sum(p.numel() for p in bizniki.parameters())}")
    print("  in_proj_weight = [W_q; W_k; W_v] ustma-ust (3d x d)")

    print("\n=== 2. Og'irliklarni ko'chirish ===")
    with torch.no_grad():
        torch_mha.in_proj_weight.copy_(torch.cat([bizniki.W_q.weight, bizniki.W_k.weight,
                                                  bizniki.W_v.weight]))
        torch_mha.in_proj_bias.copy_(torch.cat([bizniki.W_q.bias, bizniki.W_k.bias,
                                                bizniki.W_v.bias]))
        torch_mha.out_proj.weight.copy_(bizniki.W_o.weight)
        torch_mha.out_proj.bias.copy_(bizniki.W_o.bias)
    X = torch.randn(3, 7, d)
    with torch.no_grad():
        O_b, w_b = bizniki(X, X, X)
        O_t, w_t = torch_mha(X, X, X, need_weights=True, average_attn_weights=False)
        _, w_ort = torch_mha(X, X, X, need_weights=True)
    print(f"  chiqish maks farqi:           {(O_b - O_t).abs().max().item():.2e}")
    print(f"  boshlar og'irligi {tuple(w_t.shape)}: maks farq {(w_b - w_t).abs().max().item():.2e}")
    print(f"  sukut bo'yicha og'irlik {tuple(w_ort.shape)} - boshlar O'RTACHASI: "
          f"farq {(w_b.mean(1) - w_ort).abs().max().item():.2e}")

    print("\n=== 3. key_padding_mask: True = PAD (e'tiborsiz) ===")
    uz = torch.tensor([7, 4, 2])
    pad = torch.arange(7)[None, :] >= uz[:, None]
    with torch.no_grad():
        O_b, _ = bizniki(X, X, X, pad=pad)
        O_t, w_t = torch_mha(X, X, X, key_padding_mask=pad, need_weights=True,
                             average_attn_weights=False)
    for b in range(3):
        print(f"  {b}-jumla (uzunlik {uz[b].item()}): key_padding_mask {pad[b].int().tolist()}")
    print(f"  chiqish maks farqi: {(O_b - O_t).abs().max().item():.2e}")
    print(f"  PAD ustidagi og'irlik massasi: {w_t.masked_select(pad[:, None, None, :]).sum().item():.1f}")
    b = 2
    with torch.no_grad():
        yolgiz, _ = torch_mha(X[b:b + 1, :2], X[b:b + 1, :2], X[b:b + 1, :2])
    print(f"  2-jumla (uzunligi 2) batchda va yolg'iz: farq "
          f"{(O_t[b, :2] - yolgiz[0]).abs().max().item():.2e}")

    print("\n=== 4. F.scaled_dot_product_attention bilan ko'p bosh ===")
    with torch.no_grad():
        Q, K, V = bizniki.bol(bizniki.W_q(X)), bizniki.bol(bizniki.W_k(X)), bizniki.bol(bizniki.W_v(X))
        O_f = F.scaled_dot_product_attention(Q, K, V, attn_mask=~pad[:, None, None, :])
        O_f = bizniki.W_o(O_f.transpose(1, 2).reshape(3, 7, d))
    print(f"  (B, h, n, d_k) tensorlar bilan F.sdpa: farq {(O_f - O_t).abs().max().item():.2e}")
    print("  diqqat: F.sdpa da True = QATNASHADI, shuning uchun ~pad")

    print("\n=== 5. Cross-attention: so'rov va kalit boshqa ketma-ketlikdan ===")
    dekoder = torch.randn(3, 5, d)
    with torch.no_grad():
        O_c, w_c = torch_mha(dekoder, X, X, key_padding_mask=pad, need_weights=True,
                             average_attn_weights=False)
        O_cb, _ = bizniki(dekoder, X, X, pad=pad)
    print(f"  so'rov (3, 5, {d}), kalit/qiymat (3, 7, {d}) -> chiqish {tuple(O_c.shape)}, "
          f"og'irlik {tuple(w_c.shape)}")
    print(f"  bizniki bilan farq: {(O_c - O_cb).abs().max().item():.2e}")
    print("  ⭐ nn.MultiheadAttention = bizning modul; farqi - in_proj_weight va maska konvensiyasi")


if __name__ == "__main__":
    main()

Natijaning muhim qismi:

text
=== 1. nn.MultiheadAttention ichidagi parametrlar ===
  in_proj_weight       (96, 32)
  in_proj_bias         (96,)
  out_proj.weight      (32, 32)
  out_proj.bias        (32,)
  jami: 4224, bizniki: 4224
  in_proj_weight = [W_q; W_k; W_v] ustma-ust (3d x d)

=== 2. Og'irliklarni ko'chirish ===
  chiqish maks farqi:           1.19e-07
  boshlar og'irligi (3, 4, 7, 7): maks farq 8.94e-08
  sukut bo'yicha og'irlik (3, 7, 7) - boshlar O'RTACHASI: farq 2.98e-08

=== 3. key_padding_mask: True = PAD (e'tiborsiz) ===
  0-jumla (uzunlik 7): key_padding_mask [0, 0, 0, 0, 0, 0, 0]
  1-jumla (uzunlik 4): key_padding_mask [0, 0, 0, 0, 1, 1, 1]
  2-jumla (uzunlik 2): key_padding_mask [0, 0, 1, 1, 1, 1, 1]
  chiqish maks farqi: 1.19e-07
  PAD ustidagi og'irlik massasi: 0.0
  2-jumla (uzunligi 2) batchda va yolg'iz: farq 2.38e-07

=== 4. F.scaled_dot_product_attention bilan ko'p bosh ===
  (B, h, n, d_k) tensorlar bilan F.sdpa: farq 1.79e-07
  diqqat: F.sdpa da True = QATNASHADI, shuning uchun ~pad

=== 5. Cross-attention: so'rov va kalit boshqa ketma-ketlikdan ===
  so'rov (3, 5, 32), kalit/qiymat (3, 7, 32) -> chiqish (3, 5, 32), og'irlik (3, 4, 5, 7)
  bizniki bilan farq: 1.19e-07
  ⭐ nn.MultiheadAttention = bizning modul; farqi - in_proj_weight va maska konvensiyasi

Nima ko'rsatdi:

  • 1-bo'lim: nn.MultiheadAttention da to'rtta parametr bor: in_proj_weight (96, 32) — uchta 32 x 32 matritsa ustma-ust, out_proj — bizning W_o. Jami 4224 — bizniki bilan bir xil.
  • 2-bo'lim — asosiy natija: og'irliklarni ko'chirgandan keyin chiqishlar ~1e-7 farq bilan mos, boshlar og'irliklari (3, 4, 7, 7) ham mos. need_weights=True sukut bo'yicha (3, 7, 7) qaytaradi — bu boshlar o'rtachasi (bizning w.mean(1) bilan farq ~3e-08). Boshlarni alohida tahlil qilish uchun average_attn_weights=False shart.
  • 3-bo'lim: key_padding_mask da True = PAD. Natija bizning -inf maskali modulimiz bilan mos, PAD ustidagi og'irlik massasi aynan 0.0, 2 so'zli jumla batchda va yolg'iz hisoblanganda bir xil (2.38e-07).
  • 4-bo'lim: ko'p boshli attention ni F.scaled_dot_product_attention bilan ham yozish mumkin — u (B, h, n, d_k) shaklli tensorlarni qabul qiladi. Lekin maska ~pad: bu funksiyada True = qatnashadi.
  • 5-bo'lim: xuddi shu modul cross-attention uchun: so'rov 5 ta, kalit/qiymat 7 ta — chiqish (3, 5, 32), og'irliklar (3, 4, 5, 7). key_padding_mask kalitlar (encoder) tomonini yopadi.

Misol 3 — Boshlar turli narsaga qaraydi: og'irlik va boshni o'chirish

python
"""Boshlar turli narsaga qaraydi: ikki bog'lanishli vazifa, boshlar og'irligi va boshni o'chirish."""

import torch
import torch.nn as nn

K, N = 8, 12            # 8 xil token, uzunlik 12


def malumot(m, g):
    """Har pozitsiya i >= 1 uchun ikki nishon: oldingi token x[i-1] va birinchi token x[0]."""
    X = torch.randint(0, K, (m, N), generator=g)
    return X, X[:, :-1], X[:, :1].expand(m, N - 1)


class MHA(nn.Module):
    def __init__(self, d, h):
        super().__init__()
        self.h, self.d_k = h, d // h
        self.W_qkv = nn.Linear(d, 3 * d)                 # W_q, W_k, W_v bitta matritsada
        self.W_o = nn.Linear(d, d)

    def forward(self, x, ochiq=None):
        B, n, d = x.shape
        Q, K_, V = self.W_qkv(x).view(B, n, 3, self.h, self.d_k).permute(2, 0, 3, 1, 4)
        w = torch.softmax(Q @ K_.transpose(-2, -1) / self.d_k ** 0.5, -1)   # (B, h, n, n)
        O = w @ V
        if ochiq is not None:                            # boshni "o'chirish" (ablation)
            O = O * ochiq[None, :, None, None]
        return self.W_o(O.transpose(1, 2).reshape(B, n, d)), w


class Model(nn.Module):
    def __init__(self, h, d=32):
        super().__init__()
        self.emb = nn.Embedding(K, d)
        self.poz = nn.Embedding(N, d)                    # o'rgatiladigan pozitsiya
        self.mha = MHA(d, h)
        self.oldingi = nn.Sequential(nn.Linear(d, 64), nn.ReLU(), nn.Linear(64, K))
        self.birinchi = nn.Sequential(nn.Linear(d, 64), nn.ReLU(), nn.Linear(64, K))

    def forward(self, X, ochiq=None):
        h = self.emb(X) + self.poz(torch.arange(X.shape[1]))
        o, w = self.mha(h, ochiq)
        h = (h + o)[:, 1:]
        return self.oldingi(h), self.birinchi(h), w


def orgat(h, seed, qadamlar=400):
    torch.manual_seed(seed)
    model = Model(h)
    opt = torch.optim.Adam(model.parameters(), lr=3e-3)
    g = torch.Generator().manual_seed(seed)
    for _ in range(qadamlar):
        X, ya, yb = malumot(128, g)
        la, lb, _ = model(X)
        loss = (nn.functional.cross_entropy(la.reshape(-1, K), ya.reshape(-1))
                + nn.functional.cross_entropy(lb.reshape(-1, K), yb.reshape(-1)))
        opt.zero_grad()
        loss.backward()
        opt.step()
    return model


def aniqlik(model, test, ochiq=None):
    X, ya, yb = test
    with torch.no_grad():
        la, lb, w = model(X, ochiq)
    return ((la.argmax(-1) == ya).float().mean().item(),
            (lb.argmax(-1) == yb).float().mean().item(), w)


def main() -> None:
    torch.set_num_threads(1)
    test = malumot(1000, torch.Generator().manual_seed(9))
    i = torch.arange(2, N)                               # i = 1 da ikki nishon ustma-ust

    print("=== 1. Vazifa: ikki xil bog'lanish ===")
    X, ya, yb = malumot(1, torch.Generator().manual_seed(1))
    print(f"  x:                {X[0].tolist()}")
    print(f"  nishon 'oldingi':  {['-'] + ya[0].tolist()}")
    print(f"  nishon 'birinchi': {['-'] + yb[0].tolist()}")
    print("  'oldingi' - nisbiy pozitsiya (i-1), 'birinchi' - mutlaq pozitsiya (0)")

    print("\n=== 2. Bitta bosh (h = 1), 3 seed ===")
    for s in range(3):
        a, b, w = aniqlik(orgat(1, s), test)
        print(f"  seed {s}: oldingi {a:.3f}  birinchi {b:.3f}   og'irlik: i-1 ga "
              f"{w[:, 0, i, i - 1].mean():.2f}, 0 ga {w[:, 0, i, 0].mean():.2f}")

    print("\n=== 3. Ikki bosh (h = 2): har bosh qayerga qaraydi va o'chirilsa nima bo'ladi ===")
    print("  jadval: bosh o'rtacha qancha og'irlik beradi (i-1 ga, 0 ga) va shu bosh")
    print("  O'CHIRILGANDA ikki nishondagi aniqlik")
    ixtisos = 0
    for s in range(3):
        model = orgat(2, s)
        a, b, w = aniqlik(model, test)
        print(f"  seed {s}: to'liq model - oldingi {a:.3f}, birinchi {b:.3f}")
        print(f"    {'bosh':>4} {'-> i-1':>7} {'-> 0':>6} {'oldingi':>9} {'birinchi':>9}")
        rollar = []
        for k in range(2):
            ochiq = torch.ones(2)
            ochiq[k] = 0
            a_k, b_k, _ = aniqlik(model, test, ochiq)
            ga_oldingi = w[:, k, i, i - 1].mean().item()
            ga_birinchi = w[:, k, i, 0].mean().item()
            if ga_oldingi > 0.8:
                rollar.append("oldingi")
            else:
                rollar.append("birinchi" if ga_birinchi > 0.8 else "aralash")
            print(f"    {k:>4} {ga_oldingi:>7.2f} {ga_birinchi:>6.2f} {a_k:>9.3f} {b_k:>9.3f}")
        if sorted(rollar) == ["birinchi", "oldingi"]:
            ixtisos += 1
            print(f"    -> seed {s}: boshlar ixtisoslashgan ({rollar[0]} / {rollar[1]})")
        else:
            print(f"    -> seed {s}: ixtisoslashuv yo'q ({rollar[0]} / {rollar[1]})")
    print(f"  ixtisoslashgan seedlar: {ixtisos} / 3")
    print("  ⭐ Har bosh o'z bog'lanishini o'rganishi mumkin - lekin bu kafolatlanmagan")


if __name__ == "__main__":
    main()

Natijaning muhim qismi:

text
=== 1. Vazifa: ikki xil bog'lanish ===
  x:                [5, 3, 4, 0, 7, 1, 3, 5, 7, 0, 0, 1]
  nishon 'oldingi':  ['-', 5, 3, 4, 0, 7, 1, 3, 5, 7, 0, 0]
  nishon 'birinchi': ['-', 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5]
  'oldingi' - nisbiy pozitsiya (i-1), 'birinchi' - mutlaq pozitsiya (0)

=== 2. Bitta bosh (h = 1), 3 seed ===
  seed 0: oldingi 0.445  birinchi 1.000   og'irlik: i-1 ga 0.11, 0 ga 0.89
  seed 1: oldingi 1.000  birinchi 1.000   og'irlik: i-1 ga 0.32, 0 ga 0.64
  seed 2: oldingi 0.434  birinchi 1.000   og'irlik: i-1 ga 0.09, 0 ga 0.88

=== 3. Ikki bosh (h = 2): har bosh qayerga qaraydi va o'chirilsa nima bo'ladi ===
  jadval: bosh o'rtacha qancha og'irlik beradi (i-1 ga, 0 ga) va shu bosh
  O'CHIRILGANDA ikki nishondagi aniqlik
  seed 0: to'liq model - oldingi 1.000, birinchi 1.000
    bosh  -> i-1   -> 0   oldingi  birinchi
       0    0.97   0.01     0.128     1.000
       1    0.00   1.00     1.000     0.176
    -> seed 0: boshlar ixtisoslashgan (oldingi / birinchi)
  seed 1: to'liq model - oldingi 1.000, birinchi 1.000
    bosh  -> i-1   -> 0   oldingi  birinchi
       0    0.00   1.00     1.000     0.170
       1    0.97   0.00     0.132     1.000
    -> seed 1: boshlar ixtisoslashgan (birinchi / oldingi)
  seed 2: to'liq model - oldingi 0.369, birinchi 1.000
    bosh  -> i-1   -> 0   oldingi  birinchi
       0    0.01   0.98     0.362     0.464
       1    0.20   0.80     0.139     0.993
    -> seed 2: ixtisoslashuv yo'q (birinchi / aralash)
  ixtisoslashgan seedlar: 2 / 3
  ⭐ Har bosh o'z bog'lanishini o'rganishi mumkin - lekin bu kafolatlanmagan

Nima ko'rsatdi:

  • 1-bo'lim: har pozitsiyada ikki nishon: "oldingi" (x[i-1], nisbiy pozitsiya) va "birinchi" (x[0], mutlaq pozitsiya). Ikkalasi ham faqat pozitsiyaga bog'liq, shuning uchun o'rgatiladigan pozitsiya embeddingi qo'shilgan (24.4-darsda batafsil).
  • 2-bo'lim — bitta bosh: seed 0 va 2 da bosh 0.89 va 0.88 og'irlikni birinchi tokenga berdi, "oldingi" aniqligi 0.445 va 0.434 da qoldi. Seed 1 da esa bosh og'irlikni 0.32 / 0.64 qilib bo'lishdi va ikkala vazifani ham 1.000 yechdi — bitta bosh ham ikki joyga qaray oladi, asimmetrik aralashma qiymatlarni ajratishga imkon beradi. Lekin bunga faqat 3 seeddan 1 tasida erishildi.
  • 3-bo'lim — ikki bosh: seed 0 va 1 da boshlar aniq ixtisoslashdi: biri i-1 ga 0.97, ikkinchisi 0 ga 1.00. Boshni o'chirish rolni isbotladi: "oldingi" boshi o'chirilganda "oldingi" aniqligi 0.128 / 0.132 ga (tasodif 0.125) tushdi, "birinchi" esa 1.000 qoldi — va aksincha. Seed 2 da esa ikkala bosh ham birinchi tokenga qaradi (0.98 va 0.80): ixtisoslashuv bo'lmadi va "oldingi" vazifasi 0.369 da qoldi. Boshlar soni imkoniyat beradi, lekin optimizatsiya uni har doim ham topmaydi.

Misol 4 — Bosh soni tajribasi: juftlashgan seedlar bilan

python
"""Bosh soni tajribasi: d_model o'zgarmas, h = 1, 2, 4, 8 - juftlashgan seedlar bilan."""

import numpy as np

import torch
import torch.nn as nn

K, N = 8, 12            # 8 xil token, uzunlik 12


def malumot(m, g):
    """Har pozitsiya i >= 1 uchun ikki nishon: oldingi token x[i-1] va birinchi token x[0]."""
    X = torch.randint(0, K, (m, N), generator=g)
    return X, X[:, :-1], X[:, :1].expand(m, N - 1)


class MHA(nn.Module):
    def __init__(self, d, h):
        super().__init__()
        self.h, self.d_k = h, d // h
        self.W_qkv = nn.Linear(d, 3 * d)                 # W_q, W_k, W_v bitta matritsada
        self.W_o = nn.Linear(d, d)

    def forward(self, x, ochiq=None):
        B, n, d = x.shape
        Q, K_, V = self.W_qkv(x).view(B, n, 3, self.h, self.d_k).permute(2, 0, 3, 1, 4)
        w = torch.softmax(Q @ K_.transpose(-2, -1) / self.d_k ** 0.5, -1)   # (B, h, n, n)
        O = w @ V
        if ochiq is not None:                            # boshni "o'chirish" (ablation)
            O = O * ochiq[None, :, None, None]
        return self.W_o(O.transpose(1, 2).reshape(B, n, d)), w


class Model(nn.Module):
    def __init__(self, h, d=32):
        super().__init__()
        self.emb = nn.Embedding(K, d)
        self.poz = nn.Embedding(N, d)                    # o'rgatiladigan pozitsiya
        self.mha = MHA(d, h)
        self.oldingi = nn.Sequential(nn.Linear(d, 64), nn.ReLU(), nn.Linear(64, K))
        self.birinchi = nn.Sequential(nn.Linear(d, 64), nn.ReLU(), nn.Linear(64, K))

    def forward(self, X, ochiq=None):
        h = self.emb(X) + self.poz(torch.arange(X.shape[1]))
        o, w = self.mha(h, ochiq)
        h = (h + o)[:, 1:]
        return self.oldingi(h), self.birinchi(h), w


def orgat(h, seed, qadamlar=400):
    torch.manual_seed(seed)
    model = Model(h)
    opt = torch.optim.Adam(model.parameters(), lr=3e-3)
    g = torch.Generator().manual_seed(seed)
    for _ in range(qadamlar):
        X, ya, yb = malumot(64, g)
        la, lb, _ = model(X)
        loss = (nn.functional.cross_entropy(la.reshape(-1, K), ya.reshape(-1))
                + nn.functional.cross_entropy(lb.reshape(-1, K), yb.reshape(-1)))
        opt.zero_grad()
        loss.backward()
        opt.step()
    return model


def ikkalasi(model, test, ochiq=None):
    """Pozitsiyada IKKALA nishon ham to'g'ri topilgan ulush."""
    X, ya, yb = test
    with torch.no_grad():
        la, lb, _ = model(X, ochiq)
    return ((la.argmax(-1) == ya) & (lb.argmax(-1) == yb)).float().mean().item()


def main() -> None:
    torch.set_num_threads(1)
    test = malumot(1000, torch.Generator().manual_seed(9))
    boshlar = [1, 2, 4, 8]
    seedlar = range(4)

    print("=== 1. Bir xil d_model = 32: parametrlar soni teng ===")
    for h in boshlar:
        m = Model(h)
        print(f"  h = {h}: d_k = {32 // h:>2}, attention parametrlari "
              f"{sum(p.numel() for p in m.mha.parameters())}, jami {sum(p.numel() for p in m.parameters())}")

    print("\n=== 2. Aniqlik (ikkala nishon to'g'ri), 400 qadam, 4 seed ===")
    natija, zaxira = {}, {}
    for h in boshlar:
        acc, tushish = [], []
        for s in seedlar:
            model = orgat(h, s)
            a = ikkalasi(model, test)
            acc.append(a)
            if h > 1:
                for k in range(h):
                    ochiq = torch.ones(h)
                    ochiq[k] = 0
                    tushish.append(a - ikkalasi(model, test, ochiq))
        natija[h] = np.array(acc)
        zaxira[h] = (np.mean(tushish), np.max(tushish)) if tushish else None
        yechdi = int((natija[h] > 0.95).sum())
        print(f"  h = {h}: o'rtacha {natija[h].mean():.4f}  seedlar {np.round(acc, 3).tolist()}"
              f"  yechgan seedlar {yechdi}/4")

    print("\n=== 3. Juftlashgan farq h = 1 ga nisbatan (bir xil seed) ===")
    print(f"  {'h':>3} {'farq':>8} {'SE':>7} {'sezilarli':>10}")
    for h in boshlar[1:]:
        f = natija[h] - natija[1]
        se = f.std(ddof=1) / np.sqrt(len(f))
        print(f"  {h:>3} {f.mean():>+8.4f} {se:>7.4f} {str(abs(f.mean()) > 2 * se):>10}")

    print("\n=== 4. Bitta boshni o'chirganda aniqlik tushishi ===")
    for h in boshlar[1:]:
        o, mx = zaxira[h]
        print(f"  h = {h}: o'rtacha tushish {o:.3f}, eng katta tushish {mx:.3f}")
    print("  boshlar ko'p bo'lsa, bitta bosh kamroq 'yagona mas'ul' bo'ladi")

    print("\n=== 5. Xulosa ===")
    eng_qiymat = max(natija[h].mean() for h in boshlar)
    eng = [h for h in boshlar if natija[h].mean() == eng_qiymat]
    print(f"  eng yuqori o'rtacha ({eng_qiymat:.4f}): h = {eng}")
    ikki_xil = [h for h in boshlar if natija[h].min() < 0.5 and natija[h].max() > 0.95]
    if ikki_xil:
        print(f"  h = {ikki_xil}: natija ikki xil - seed yo 'yechadi' 1.0-bob, yo 'qotib qoladi' (~0.4)")
    sezilarli = [h for h in boshlar[1:]
                 if abs((natija[h] - natija[1]).mean())
                 > 2 * (natija[h] - natija[1]).std(ddof=1) / np.sqrt(len(seedlar))]
    if not sezilarli:
        print("  2*SE mezoni bo'yicha h = 1 dan hech bir farq sezilarli emas: 4 seed kam,"
              " tarqoqlik katta")
    else:
        print(f"  h = 1 dan sezilarli farq: h = {sezilarli}")
    print("  ⭐ Bosh soni - giperparametr: d_model o'zgarmasa, parametr va FLOP bir xil, "
          "natija esa tajribada tekshiriladi")


if __name__ == "__main__":
    main()

Natijaning muhim qismi:

text
=== 1. Bir xil d_model = 32: parametrlar soni teng ===
  h = 1: d_k = 32, attention parametrlari 4224, jami 10128
  h = 2: d_k = 16, attention parametrlari 4224, jami 10128
  h = 4: d_k =  8, attention parametrlari 4224, jami 10128
  h = 8: d_k =  4, attention parametrlari 4224, jami 10128

=== 2. Aniqlik (ikkala nishon to'g'ri), 400 qadam, 4 seed ===
  h = 1: o'rtacha 0.7112  seedlar [0.444, 1.0, 0.401, 1.0]  yechgan seedlar 2/4
  h = 2: o'rtacha 0.8510  seedlar [1.0, 1.0, 0.404, 1.0]  yechgan seedlar 3/4
  h = 4: o'rtacha 1.0000  seedlar [1.0, 1.0, 1.0, 1.0]  yechgan seedlar 4/4
  h = 8: o'rtacha 1.0000  seedlar [1.0, 1.0, 1.0, 1.0]  yechgan seedlar 4/4

=== 3. Juftlashgan farq h = 1 ga nisbatan (bir xil seed) ===
    h     farq      SE  sezilarli
    2  +0.1398  0.1387      False
    4  +0.2888  0.1669      False
    8  +0.2888  0.1669      False

=== 4. Bitta boshni o'chirganda aniqlik tushishi ===
  h = 2: o'rtacha tushish 0.688, eng katta tushish 0.878
  h = 4: o'rtacha tushish 0.205, eng katta tushish 0.786
  h = 8: o'rtacha tushish 0.041, eng katta tushish 0.341
  boshlar ko'p bo'lsa, bitta bosh kamroq 'yagona mas'ul' bo'ladi

=== 5. Xulosa ===
  eng yuqori o'rtacha 1.0000-bob: h = [4, 8]
  h = [1, 2]: natija ikki xil - seed yo 'yechadi' 1.0-bob, yo 'qotib qoladi' (~0.4)
  2*SE mezoni bo'yicha h = 1 dan hech bir farq sezilarli emas: 4 seed kam, tarqoqlik katta
  ⭐ Bosh soni - giperparametr: d_model o'zgarmasa, parametr va FLOP bir xil, natija esa tajribada tekshiriladi

Nima ko'rsatdi:

  • 1-bo'lim: h = 1, 2, 4, 8 da attention parametrlari (4224) va butun model (10128) bir xil — taqqoslash parametrlar bo'yicha adolatli.
  • 2-bo'lim: h = 1 4 seeddan 2 tasida yechdi (0.444, 1.0, 0.401, 1.0), h = 2 — 3 tasida, h = 4 va h = 8 — hammasida (1.0000). Muvaffaqiyatsiz seedlarda aniqlik ~0.4 — bu 3-misoldagi manzaraga mos: model "birinchi" ni yechib, "oldingi" da qotib qolganda aynan shunday qiymat chiqadi (bu yerda nishonlar alohida o'lchanmagan).
  • 3-bo'lim — halol natija: h = 4 ning h = 1 dan ustunligi +0.2888, lekin SE 0.1669 — 2*SE mezoni bo'yicha sezilarli emas. Sabab: natija ikki cho'qqili (1.0 yoki ~0.4), 4 seed esa juda kam; SE ni bitta "omadsiz" seed katta qiladi. "4/4 va 2/4" farqi haqiqiy bo'lishi mumkin, lekin buni tasdiqlash uchun ko'proq seed kerak. Biz buni "isbotlandi" deb yozmaymiz.
  • 4-bo'lim: bitta boshni o'chirish h = 2 da aniqlikni o'rtacha 0.688 ga kamaytiradi, h = 4 da 0.205 ga, h = 8 da atigi 0.041 ga. Ko'p boshli modelda vazifa boshlar orasida taqsimlanadi — bitta bosh yo'qolsa ham model ishlaydi. Ammo h = 8 da ham eng katta tushish 0.341: ba'zi boshlar baribir muhimroq.

5. To'g'ri va noto'g'ri tushunishlar

Noto'g'ri fikr To'g'risi
"Ko'p bosh — ko'p parametr" d_model o'zgarmasa, h = 1..16 da 4224 parametr (1-misol)
"Ko'p bosh — ko'p hisob" Skor FLOP 2 n^2 d_model — h ga bog'liq emas; faqat og'irliklar xotirasi h marta
"Masshtab sqrt(d_model)" Har bosh o'z d_k o'lchamida: sqrt(d_k)
"Bitta bosh ikki joyga qaray olmaydi" Qaray oladi (3-misol, seed 1: 0.32 / 0.64), lekin buni topishi qiyinroq
"Boshlar doim turli narsani o'rganadi" 3-misolda 3 seeddan 2 tasida ixtisoslashdi, seed 2 da ikkala bosh bir joyga qaradi
"Og'irlik bosh rolini isbotlaydi" Og'irlik — gumon; boshni o'chirish — isbot (aniqlik 0.128 ga tushdi)
"need_weights=True har boshni qaytaradi" Sukut bo'yicha boshlar o'rtachasi; kerak: average_attn_weights=False
"4/4 va 2/4 — aniq farq" 4-misolda 2*SE mezoni bo'yicha sezilarli emas; ikki cho'qqili natijada ko'proq seed kerak

6. Keng tarqalgan xatolar va yechimlari

1. Boshlarga noto'g'ri bo'lish

python
Q = W_q(X).view(B, h, n, d_k)                    # tokenlar va boshlar aralashadi  # ⚠️
Q = W_q(X).view(B, n, h, d_k).transpose(1, 2)                                   # ✅

2. Ulashda transpose unutilgan

python
O = (w @ V).reshape(B, n, d_model)               # (B, h, n, d_k) ni noto'g'ri yoyadi # ⚠️
O = (w @ V).transpose(1, 2).reshape(B, n, d_model)                                # ✅

3. Noto'g'ri masshtab

python
skor = Q @ K.transpose(-2, -1) / d_model ** 0.5                                  # ⚠️
skor = Q @ K.transpose(-2, -1) / d_k ** 0.5                                      # ✅

4. batch_first sukuti

python
mha = nn.MultiheadAttention(32, 4); mha(X, X, X)    # X (B, n, d) - n va B almashadi # ⚠️
mha = nn.MultiheadAttention(32, 4, batch_first=True); mha(X, X, X)                 # ✅

5. Maska konvensiyasi

python
mha(X, X, X, key_padding_mask=token_maska)          # True = token -> tokenlar yopiladi # ⚠️
mha(X, X, X, key_padding_mask=~token_maska)         # True = PAD                       # ✅

6. Boshlar o'rtachasini alohida bosh deb o'qish

python
_, W = mha(X, X, X); W[:, 0]                        # W (B, n, n) - 0-bosh EMAS        # ⚠️
_, W = mha(X, X, X, average_attn_weights=False); W[:, 0]   # (B, h, n, n)          # ✅

7. d_model h ga bo'linmaydi

python
nn.MultiheadAttention(30, 4)                        # AssertionError                 # ⚠️
nn.MultiheadAttention(32, 4)                        # d_k = 8                         # ✅

7. Integratsiya — bu bilim qayerda kerak bo'ladi

  • 18-qism (o'tilgan): bir necha seed, juftlashgan farq, SE
  • 21-qism (o'tilgan): view, transpose, reshape, state_dict, parametrlarni ko'chirish
  • 24.1, 24.2-darslar (o'tilgan): scaled dot-product attention, self-attention, niqoblash
  • Keyingi darslar: pozitsion kodlash 24.4-bob, Transformer bloki 24.5-bob va encoder 24.6-bob — nn.MultiheadAttention ularning ichida; decoderda kauzal niqob va cross-attention 24.7-bob; Katta til modellari qismida grouped-query attention va KV-kesh — boshlar sonini xotira uchun optimallashtirish

8. Eng yaxshi amaliyotlar

  1. d_k = d_model / h ni 32-128 oralig'ida saqlang (odatda 64).

  2. O'z implementatsiyangizni nn.MultiheadAttention ga og'irliklarni ko'chirib tekshiring.

  3. nn.MultiheadAttention da doim batch_first=True ni aniq yozing.

  4. Maskani uzatishdan oldin konvensiyani tekshiring: key_padding_mask da True = PAD.

  5. Boshlarni tahlil qilishda average_attn_weights=False ishlating.

  6. Bosh rolini boshni o'chirish (ablation) bilan tasdiqlang, faqat og'irlikka qaramang.

  7. Bosh sonini bir necha seed bilan tanlang; ikki cho'qqili natijada muvaffaqiyat ulushini ham chop eting.

  8. Xotira cheklangan bo'lsa, og'irliklar xotirasi h * n^2 ekanini hisobga oling.


9. Amaliy topshiriq

Vazifa 1: Bashorat qiling

python
1.  # d_model = 64, h = 8: d_k nechaga teng?
2.  # X (4, 10, 64) -> Q boshlarga bo'lingandan keyin shakli?
3.  # og'irliklar shakli?
4.  # d_model = 64 da parametrlar soni (bias bilan)? h = 2 da-chi?
5.  # nn.MultiheadAttention(64, 8).in_proj_weight shakli?
6.  # need_weights=True (sukut) qaytaradigan W shakli?
7.  # skor masshtabi: sqrt(64) mi yoki sqrt(8)?
8.  # h 1 dan 8 ga oshsa skor FLOP va og'irliklar xotirasi qanday o'zgaradi?
9.  # key_padding_mask da True nimani anglatadi?
10. # cross-attention: so'rov (2, 5, 64), kalit (2, 9, 64) - chiqish va W shakli?
11. # boshni o'chirganda faqat bitta vazifa tasodifga tushsa - bu nimani isbotlaydi?
12. # nega 4/4 va 2/4 natija 2*SE mezoni bo'yicha sezilarli bo'lmasligi mumkin?
Javoblar
  1. 8
  2. (4, 8, 10, 8)
  3. (4, 8, 10, 10)
  4. 4 * 64^2 + 4 * 64 = 16640; h = 2 da ham 16640
  5. (192, 64)
  6. (B, n_q, n_k) — boshlar o'rtachasi
  7. sqrt(8) — har bosh o'z d_k o'lchamida
  8. FLOP o'zgarmaydi, og'irliklar xotirasi 8 marta oshadi
  9. Bu pozitsiya PAD — e'tiborsiz qoldiriladi
  10. Chiqish (2, 5, 64), W (2, 8, 5, 9) (average_attn_weights=False bilan)
  11. Shu bosh aynan o'sha vazifa uchun zarur (va boshqa vazifaga keraksiz)
  12. Natija ikki cho'qqili (1.0 yoki ~0.4), seedlar kam — SE katta

Vazifa 2: Xatolarni tuzating

python
1.  Q = self.W_q(X).view(B, self.h, n, self.d_k)

2.  O = (w @ V).reshape(B, n, self.h * self.d_k)

3.  mha = nn.MultiheadAttention(64, 8)
    O, _ = mha(X, X, X)                         # X: (B, n, 64)

4.  token = X_ids != PAD
    O, _ = mha(h, h, h, key_padding_mask=token)

5.  _, W = mha(h, h, h)
    print("0-bosh:", W[:, 0])
Javoblar
python
1.  Q = self.W_q(X).view(B, n, self.h, self.d_k).transpose(1, 2)

2.  O = (w @ V).transpose(1, 2).reshape(B, n, self.h * self.d_k)

3.  mha = nn.MultiheadAttention(64, 8, batch_first=True)
    O, _ = mha(X, X, X)

4.  pad = X_ids == PAD                          # True = PAD
    O, _ = mha(h, h, h, key_padding_mask=pad)

5.  _, W = mha(h, h, h, average_attn_weights=False)   # (B, h, n, n)
    print("0-bosh:", W[:, 0])

Vazifa 3: Noldan multi-head

Modellang:

  1. MultiHeadAttention(d_model, h) — bol va ulash
  2. Bitta boshni qo'lda hisoblab solishtirish
  3. h = 1 da oddiy self-attention bilan moslik
  4. Parametrlar va FLOP jadvali

Vazifa 4: nn.MultiheadAttention bilan moslik

Modellang:

  1. Og'irliklarni in_proj_weight va out_proj ga ko'chirish
  2. Chiqish va boshlar og'irliklari mosligi
  3. key_padding_mask va "yolg'iz" hisoblash
  4. Cross-attention

Vazifa 5: Boshlar ixtisoslashuvi

Modellang:

  1. Ikki nishonli vazifa ("oldingi" va "birinchi")
  2. h = 1 va h = 2, 3 seed
  3. Har boshning i-1 va 0 ga og'irligi
  4. Boshni o'chirish va aniqlik

Vazifa 6: Bosh soni

Modellang:

  1. h = 1, 2, 4, 8 — bir xil d_model
  2. 4+ seed, juftlashgan farq va SE
  3. Muvaffaqiyatli seedlar ulushi
  4. Bitta boshni o'chirishdagi tushish

Vazifa 7: O'ylash

Hamkasbingiz aytdi: "Men 3-misolni takrorladim va seed 0 dagi og'irliklarga qaradim: 0-bosh oldingi tokenga, 1-bosh birinchi tokenga qaraydi. Demak modelimizda boshlar har doim shunday bo'linadi. Keyingi hisobotda yozamiz: '0-bosh — sintaksis boshi, 1-bosh — global kontekst boshi'." Siz nima deysiz?

Javob

Qisqa javob: bitta seeddagi kuzatuv "har doim" degan xulosaga yetmaydi, og'irlikka qarash esa rolni isbotlamaydi. Ikkalasini ham tekshirish kerak.

1. Seedga bog'liqlik. 3-misolda seed 0 da 0-bosh "oldingi", 1-bosh "birinchi" edi; seed 1 da esa teskari — 0-bosh "birinchi", 1-bosh "oldingi". Seed 2 da umuman ixtisoslashuv bo'lmadi (ikkala bosh ham birinchi tokenga qaradi, "oldingi" aniqligi 0.369). Bosh raqami — tasodifiy initsializatsiyaning natijasi, doimiy "nom" emas.

2. Og'irlik — gumon, o'chirish — isbot. Bosh biror joyga qarashi undan foydalanishini anglatmaydi. 3-misolda rolni boshni o'chirish bilan tekshirdik: "oldingi" boshi o'chirilganda "oldingi" aniqligi tasodif darajasiga (0.128) tushdi, "birinchi" esa 1.000 qoldi. Hisobotdagi da'vo shu kabi dalilga tayanishi kerak.

3. Nomlash ehtiyotkorligi. "Sintaksis boshi" — sintetik "oldingi token" vazifasidan ancha katta da'vo. Bizning vazifada faqat "i-1 pozitsiyadagi tokenni ko'chiradi" deyish mumkin.

4. Nima deyish mumkin. "3 seeddan 2 tasida boshlar ixtisoslashdi: bittasi oldingi tokenni, ikkinchisi birinchi tokenni tashiydi; boshni o'chirish bu rollarni tasdiqladi. Qaysi raqamli bosh qaysi rolni olishi seedga bog'liq."

Tavsiya:

python
# 1. Kamida 5 seed; har seedda har bosh uchun og'irlik va o'chirish natijasi
# 2. Rolni bosh raqami bilan emas, xulq-atvori bilan belgilang
# 3. Ixtisoslashgan seedlar ulushini hisobotga yozing
# 4. "Sintaksis" kabi so'zlarni faqat mos vazifada tekshirilgandan keyin ishlating

Hamkasbga javob: "Kuzatuving to'g'ri, lekin faqat seed 0 uchun. Seed 1 da boshlar rollari almashgan, seed 2 da ixtisoslashuv yo'q. Hisobotda bosh raqamini emas, rolni yozamiz va uni boshni o'chirish natijasi bilan tasdiqlaymiz: '3 seeddan 2 tasida bir bosh oldingi tokenni tashiydi — o'chirilganda bu vazifa tasodif darajasiga tushadi'."

Nimani mustahkamlaydi: 2.1, 2.3, 2.4, 2.5-bo'limlar.


Xulosa

Bu darsda multi-head attention ni noldan yozdik, uni nn.MultiheadAttention bilan og'irliklarni ko'chirib solishtirdik va boshlar nimani o'rganishini o'lchadik.

Eng muhim uch fikr:

  1. Multi-head — bir xil narxda ko'p "savol". 1-misolda d_model = 32 da h = 1 dan 16 gacha parametrlar 4224 va skor FLOP 1 048 576 o'zgarmadi — faqat og'irliklar matritsalari soni h marta oshdi. 2-misolda og'irliklari ko'chirilgan nn.MultiheadAttention bizning modul bilan ~1e-7 aniqlikda mos keldi; farqi — in_proj_weight da uch matritsaning ustma-ust saqlanishi, key_padding_mask dagi True = PAD konvensiyasi va need_weights ning sukut bo'yicha boshlar o'rtachasini qaytarishi.

  2. Boshlar ixtisoslashishi mumkin — va buni o'chirish bilan isbotlash kerak. 3-misolda ikki boshli model 3 seeddan 2 tasida bir boshni oldingi tokenga (0.97), ikkinchisini birinchi tokenga (1.00) qaratdi; "oldingi" boshi o'chirilganda shu vazifa tasodif darajasiga (0.128) tushdi, ikkinchisi 1.000 qoldi. Seed 2 da ixtisoslashuv bo'lmadi. Bitta bosh ham ikki joyga qaray oladi (seed 1: 0.32 / 0.64), lekin buni 3 seeddan faqat 1 tasida topdi.

  3. Bosh soni — tajriba bilan tanlanadigan giperparametr. 4-misolda vazifani h = 1 4 seeddan 2 tasida, h = 2 — 3 tasida, h = 4 va h = 8 — 4 tasida ham yechdi. Lekin h = 4 ning ustunligi (+0.2888, SE 0.1669) 2*SE mezoni bo'yicha sezilarli emas: ikki cho'qqili natijada 4 seed kam. Bitta boshni o'chirish h = 2 da aniqlikni o'rtacha 0.688 ga, h = 8 da atigi 0.041 ga kamaytirdi — ko'p boshli model barqarorroq.

Keyingi darsda pozitsion kodlash: 24.2-darsda self-attention tartibni ko'rmasligini isbotladik va 3-4-misollarda o'rgatiladigan pozitsiya embeddingidan foydalandik — endi sinusoidal kodlash, o'rgatiladigan embedding va nisbiy pozitsiyalarni solishtiramiz.

Ulashish:Telegram'da

Izohlar (0)

Izoh yozish uchun kiring.

  • Hozircha izoh yo'q. Birinchi bo'ling!
24.3-dars: Multi-head attention — IlmHamroh