Mundarija (21)
- 1. Kirish va motivatsiya
- 2. Nazariya — chuqur tushuntirish
- 2.1. Ko'p boshli attention formulasi
- 2.2. Parametrlar va hisob
- 2.3. nn.MultiheadAttention ichida
- 2.4. Boshlar ixtisoslashuvi
- 2.5. Bosh soni — giperparametr
- 2.6. Cross-attention — bir xil modul
- 2.7. Tuzoqlar
- 3. Tez ma'lumotnoma
- 4. Batafsil misollar
- Misol 1 — Multi-head attention noldan va parametrlar hisobi
- Misol 2 — nn.MultiheadAttention bilan moslik, maskalar va cross-attention
- Misol 3 — Boshlar turli narsaga qaraydi: og'irlik va boshni o'chirish
- Misol 4 — Bosh soni tajribasi: juftlashgan seedlar bilan
- 5. To'g'ri va noto'g'ri tushunishlar
- 6. Keng tarqalgan xatolar va yechimlari
- 7. Integratsiya — bu bilim qayerda kerak bo'ladi
- 8. Eng yaxshi amaliyotlar
- 9. Amaliy topshiriq
- Xulosa
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
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'lchamidaNima 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
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
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
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
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 -> barqarorroqAmaliy 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
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
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 1Multi-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 seed4. Batafsil misollar
Misollar real torch/numpy bilan (Python 3.14, torch 2.14 CPU).
Misol 1 — Multi-head attention noldan va parametrlar hisobi
"""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:
=== 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_oNima 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_qning 8-16 qatorlari bilan hisoblangan kichik proyeksiya (farq2.38e-07). Ya'ni bitta kattaW_qaslida 4 ta mustaqil8 x 32matritsaning ustma-ust joylashuvi. - 3-bo'lim: 2-boshni butunlay qo'lda hisobladik — og'irliklar
~3e-08aniqlikda mos. 0- va 2-boshlar og'irliklari orasidagi farq0.248: o'rgatilmagan bo'lsa ham, har bosh boshqa proyeksiya bilan boshqacha qaraydi. - 4-bo'lim:
h = 1danh = 16gacha parametrlar4224(4 * 32^2 + 4 * 32) va skor FLOP1 048 576o'zgarmadi. O'zgargani — og'irliklar matritsalari soni:16 384dan262 144gacha (hmarta). - 5-bo'lim:
h = 1bo'lgan multi-head attention —W_oqo'shilgan oddiy self-attention, farq aynan0.
Misol 2 — nn.MultiheadAttention bilan moslik, maskalar va cross-attention
"""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:
=== 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 konvensiyasiNima ko'rsatdi:
- 1-bo'lim:
nn.MultiheadAttentionda to'rtta parametr bor:in_proj_weight (96, 32)— uchta32 x 32matritsa ustma-ust,out_proj— bizningW_o. Jami4224— bizniki bilan bir xil. - 2-bo'lim — asosiy natija: og'irliklarni ko'chirgandan keyin chiqishlar
~1e-7farq bilan mos, boshlar og'irliklari(3, 4, 7, 7)ham mos.need_weights=Truesukut bo'yicha(3, 7, 7)qaytaradi — bu boshlar o'rtachasi (bizningw.mean(1)bilan farq~3e-08). Boshlarni alohida tahlil qilish uchunaverage_attn_weights=Falseshart. - 3-bo'lim:
key_padding_maskda True = PAD. Natija bizning-infmaskali modulimiz bilan mos, PAD ustidagi og'irlik massasi aynan0.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_attentionbilan 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_maskkalitlar (encoder) tomonini yopadi.
Misol 3 — Boshlar turli narsaga qaraydi: og'irlik va boshni o'chirish
"""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:
=== 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 kafolatlanmaganNima 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.89va0.88og'irlikni birinchi tokenga berdi, "oldingi" aniqligi0.445va0.434da qoldi. Seed 1 da esa bosh og'irlikni0.32 / 0.64qilib bo'lishdi va ikkala vazifani ham1.000yechdi — 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-1ga0.97, ikkinchisi0ga1.00. Boshni o'chirish rolni isbotladi: "oldingi" boshi o'chirilganda "oldingi" aniqligi0.128/0.132ga (tasodif0.125) tushdi, "birinchi" esa1.000qoldi — va aksincha. Seed 2 da esa ikkala bosh ham birinchi tokenga qaradi (0.98va0.80): ixtisoslashuv bo'lmadi va "oldingi" vazifasi0.369da qoldi. Boshlar soni imkoniyat beradi, lekin optimizatsiya uni har doim ham topmaydi.
Misol 4 — Bosh soni tajribasi: juftlashgan seedlar bilan
"""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:
=== 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 tekshiriladiNima ko'rsatdi:
- 1-bo'lim:
h = 1, 2, 4, 8da attention parametrlari (4224) va butun model (10128) bir xil — taqqoslash parametrlar bo'yicha adolatli. - 2-bo'lim:
h = 14 seeddan 2 tasida yechdi (0.444,1.0,0.401,1.0),h = 2— 3 tasida,h = 4vah = 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 = 4ningh = 1dan ustunligi+0.2888, lekin SE0.1669—2*SEmezoni 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 = 2da aniqlikni o'rtacha0.688ga kamaytiradi,h = 4da0.205ga,h = 8da atigi0.041ga. Ko'p boshli modelda vazifa boshlar orasida taqsimlanadi — bitta bosh yo'qolsa ham model ishlaydi. Ammoh = 8da ham eng katta tushish0.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
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
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
skor = Q @ K.transpose(-2, -1) / d_model ** 0.5 # ⚠️
skor = Q @ K.transpose(-2, -1) / d_k ** 0.5 # ✅4. batch_first sukuti
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
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
_, 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
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.MultiheadAttentionularning 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
d_k = d_model / hni 32-128 oralig'ida saqlang (odatda 64).O'z implementatsiyangizni
nn.MultiheadAttentionga og'irliklarni ko'chirib tekshiring.nn.MultiheadAttentionda doimbatch_first=Trueni aniq yozing.Maskani uzatishdan oldin konvensiyani tekshiring:
key_padding_maskda True = PAD.Boshlarni tahlil qilishda
average_attn_weights=Falseishlating.Bosh rolini boshni o'chirish (ablation) bilan tasdiqlang, faqat og'irlikka qaramang.
Bosh sonini bir necha seed bilan tanlang; ikki cho'qqili natijada muvaffaqiyat ulushini ham chop eting.
Xotira cheklangan bo'lsa, og'irliklar xotirasi
h * n^2ekanini hisobga oling.
9. Amaliy topshiriq
Vazifa 1: Bashorat qiling
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
- 8
(4, 8, 10, 8)(4, 8, 10, 10)4 * 64^2 + 4 * 64 = 16640;h = 2da ham 16640(192, 64)(B, n_q, n_k)— boshlar o'rtachasisqrt(8)— har bosh o'zd_ko'lchamida- FLOP o'zgarmaydi, og'irliklar xotirasi 8 marta oshadi
- Bu pozitsiya PAD — e'tiborsiz qoldiriladi
- Chiqish
(2, 5, 64), W(2, 8, 5, 9)(average_attn_weights=Falsebilan) - Shu bosh aynan o'sha vazifa uchun zarur (va boshqa vazifaga keraksiz)
- Natija ikki cho'qqili (1.0 yoki ~0.4), seedlar kam — SE katta
Vazifa 2: Xatolarni tuzating
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
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:
MultiHeadAttention(d_model, h)—bolva ulash- Bitta boshni qo'lda hisoblab solishtirish
h = 1da oddiy self-attention bilan moslik- Parametrlar va FLOP jadvali
Vazifa 4: nn.MultiheadAttention bilan moslik
Modellang:
- Og'irliklarni
in_proj_weightvaout_projga ko'chirish - Chiqish va boshlar og'irliklari mosligi
key_padding_maskva "yolg'iz" hisoblash- Cross-attention
Vazifa 5: Boshlar ixtisoslashuvi
Modellang:
- Ikki nishonli vazifa ("oldingi" va "birinchi")
h = 1vah = 2, 3 seed- Har boshning
i-1va0ga og'irligi - Boshni o'chirish va aniqlik
Vazifa 6: Bosh soni
Modellang:
h = 1, 2, 4, 8— bir xild_model- 4+ seed, juftlashgan farq va SE
- Muvaffaqiyatli seedlar ulushi
- 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:
# 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 ishlatingHamkasbga 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:
Multi-head — bir xil narxda ko'p "savol". 1-misolda
d_model = 32dah = 1dan16gacha parametrlar4224va skor FLOP1 048 576o'zgarmadi — faqat og'irliklar matritsalari sonihmarta oshdi. 2-misolda og'irliklari ko'chirilgannn.MultiheadAttentionbizning modul bilan~1e-7aniqlikda mos keldi; farqi —in_proj_weightda uch matritsaning ustma-ust saqlanishi,key_padding_maskdagi True = PAD konvensiyasi vaneed_weightsning sukut bo'yicha boshlar o'rtachasini qaytarishi.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, ikkinchisi1.000qoldi. 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.Bosh soni — tajriba bilan tanlanadigan giperparametr. 4-misolda vazifani
h = 14 seeddan 2 tasida,h = 2— 3 tasida,h = 4vah = 8— 4 tasida ham yechdi. Lekinh = 4ning ustunligi (+0.2888, SE0.1669)2*SEmezoni bo'yicha sezilarli emas: ikki cho'qqili natijada 4 seed kam. Bitta boshni o'chirishh = 2da aniqlikni o'rtacha0.688ga,h = 8da atigi0.041ga 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.
Izohlar (0)
Izoh yozish uchun kiring.
- Hozircha izoh yo'q. Birinchi bo'ling!