3-bo‘lim

E'tibor mexanizmi - Q, K va V

Skalyar ko'paytma orqali o'xshashlik, so'rov-kalit-qiymat uchligi, ballar matritsasi va softmax bilan e'tibor og'irliklari.

🕑 12 daqiqa o‘qish 📄 669 so‘z 👁 2 marta ko‘rilgan
Ushbu bo‘lim mundarijasi
  1. Boshlang'ich nuqta: vektorlar
  2. O'xshashlikni qanday o'lchaymiz
  3. Uch rol: so'rov, kalit, qiymat
  4. 1-qadam: ballar
  5. 2-qadam: ildizga bo'lish
  6. 3-qadam: softmax va yig'ish
  7. Hammasi bitta funksiyada
  8. Xulosa

Endi transformerning yuragiga kiramiz. E'tibor mexanizmi murakkab ko'rinadi, lekin u aslida uchta oddiy qadamdan iborat - va biz ularni birma-bir quramiz.

Boshlang'ich nuqta: vektorlar #

2-bo'limda matn token id lariga aylandi. Endi har id vektorga aylanadi - bu PyTorch darsligining 17-bo'limidagi nn.Embedding:

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

torch.manual_seed(42)

emb = nn.Embedding(20, 4)
ids = torch.tensor([[3, 7, 1]])
v = emb(ids)
print("id lar:", ids.tolist(), "-> shakl:", tuple(v.shape))
Natija
id lar: [[3, 7, 1]] -> shakl: (1, 3, 4)

Uch token, har biri 4 o'lchamli vektor.

O'xshashlikni qanday o'lchaymiz #

E'tiborning butun g'oyasi «qaysi so'z menga muhim» degan savolga javob berishdir. Buning uchun o'xshashlik o'lchovi kerak, va u juda oddiy: skalyar ko'paytma.

Python
a = torch.tensor([1.0, 0.0, 1.0])
b = torch.tensor([1.0, 0.0, 0.9])
c = torch.tensor([0.0, 1.0, 0.0])

print("a . b =", round(torch.dot(a, b).item(), 2), "(o'xshash)")
print("a . c =", round(torch.dot(a, c).item(), 2), "(o'xshash emas)")
Natija
a . b = 1.9 (o'xshash)
a . c = 0.0 (o'xshash emas)

Yo'nalishi yaqin vektorlarning ko'paytmasi katta, perpendikulyarlarniki nol. Butun e'tibor mexanizmi shu bitta amalga tayanadi.

Uch rol: so'rov, kalit, qiymat #

Har bir so'z vektori uchta turli rolda ishlatiladi. Buning uchun u uchta alohida Linear qatlamdan o'tkaziladi:

BelgiNomiRoli
Q (Query)So'rov«Men nimani qidiryapman?»
K (Key)Kalit«Menda nima bor?»
V (Value)Qiymat«Meni tanlasang, nima beraman?»
Python
d_model = 4
x = torch.randn(1, 3, d_model)

Wq = nn.Linear(d_model, d_model, bias=False)
Wk = nn.Linear(d_model, d_model, bias=False)
Wv = nn.Linear(d_model, d_model, bias=False)

Q, K, V = Wq(x), Wk(x), Wv(x)
print("x:", tuple(x.shape), "-> Q:", tuple(Q.shape), "K:", tuple(K.shape), "V:", tuple(V.shape))
Natija
x: (1, 3, 4) -> Q: (1, 3, 4) K: (1, 3, 4) V: (1, 3, 4)
Kutubxona tashbehi

Kutubxonaga kirdingiz va bitta savolingiz bor - bu so'rov (Q).

Javonlardagi kitoblarning muqovasida sarlavha yozilgan - bular kalitlar (K). Siz savolingizni har bir sarlavha bilan solishtirasiz.

Kitobning ichidagi matn esa qiymat (V). Sarlavha mos kelgan kitoblardan ko'proq o'qiysiz, mos kelmaganlaridan - kamroq.

Muhimi: sarlavha va ichki matn boshqa-boshqa narsa. Shuning uchun K va V alohida qatlamlardan chiqadi.

1-qadam: ballar #

Har so'rov har kalit bilan solishtiriladi:

Python
ball = Q @ K.transpose(-2, -1)
Natija
=== Xom ballar: Q @ K.T ===
   -0.07   -0.20   -0.20
    0.39    1.64    0.97
   -0.18   -0.50   -0.48

Bu 3x3 matritsa: ball[i][j] - i-so'zning j-so'zga qanchalik qiziqishi.

Ikkinchi qatorga qarang: 0.39, 1.64, 0.97. Ikkinchi so'z ikkinchi so'zning o'ziga eng ko'p (1.64) qiziqyapti.

2-qadam: ildizga bo'lish #

Formulada √d_k ga bo'lish bor. Nega?

Python
for d in [4, 64, 512]:
    q = torch.randn(2000, d)
    k = torch.randn(2000, d)
    raw = (q * k).sum(-1)
    print(f"  d_k={d:4d}: xom ballar std = {raw.std():7.2f} | ildizga bo'lingach = {(raw/d**0.5).std():5.2f}")
Natija
  d_k=   4: xom ballar std =    2.02 | ildizga bo'lingach =  1.01
  d_k=  64: xom ballar std =    8.11 | ildizga bo'lingach =  1.01
  d_k= 512: xom ballar std =   22.90 | ildizga bo'lingach =  1.01

Vektor uzunlashgan sari ballar tarqalishi o'sadi: 2.02 -> 8.11 -> 22.90. Ildizga bo'lgandan keyin esa u har doim 1 atrofida qoladi.

Nima uchun bu muhim? Chunki keyingi qadamda softmax keladi:

Python
for katta in [1.0, 10.0, 50.0]:
    s = torch.tensor([katta, 0.0, 0.0])
    print(f"  ballar {s.tolist()} -> {F.softmax(s, dim=-1).round(decimals=4).tolist()}")
Natija
  ballar [1.0, 0.0, 0.0] -> [0.5761, 0.2119, 0.2119]
  ballar [10.0, 0.0, 0.0] -> [0.9999, 0.0, 0.0]
  ballar [50.0, 0.0, 0.0] -> [1.0, 0.0, 0.0]
Ildizga bo'lishni tashlab ketsangiz

Ballar katta bo'lsa, softmax bittasiga 1.0, qolganiga 0.0 beradi.

Bu ikki tomonlama yomon. Birinchidan, model faqat bitta so'zga qaraydi va qolganini butunlay e'tiborsiz qoldiradi. Ikkinchidan - bu jiddiyroq - softmax ning hosilasi nolga aylanadi va gradient oqmay qoladi. Model o'qishni to'xtatadi.

Maqolada bu bitta jumla bilan aytilgan, lekin √d_k aynan shuning uchun turibdi.

Natija
=== Ildizga bo'lingach ===
   -0.04   -0.10   -0.10
    0.19    0.82    0.49
   -0.09   -0.25   -0.24

3-qadam: softmax va yig'ish #

Python
w = F.softmax(ball / d_model**0.5, dim=-1)
Natija
=== softmax dan keyin (e'tibor og'irliklari) ===
    0.35    0.33    0.33
    0.24    0.44    0.32
    0.37    0.31    0.32
  har qator yig'indisi: [1.0, 1.0, 1.0]

Har qator yig'indisi aynan 1.0 - bu taqsimot. Ikkinchi so'z e'tiborining 44% ini ikkinchi so'zga, 24% ini birinchiga, 32% ini uchinchiga berdi.

Oxirgi amal - shu og'irliklar bilan qiymatlarni yig'ish:

Python
out = w @ V
Natija
  shakl: (1, 3, 4)
    0.32    0.10    0.34    0.47
    0.42    0.16    0.45    0.62
    0.31    0.09    0.33    0.45
Attention(Q, K, V) = softmax(Q Kᵀ / √d_k) V Q so'rov K kalit V qiymat 1. Q Kᵀ ballar 2. ÷ √d_k barqarorlash 3. softmax yig'indi = 1 og'irliklar 4. og'irlik × V vaznli yig'indi = chiqish yangi vektorlar Butun mexanizm - ikkita matritsa ko'paytmasi va bitta softmax
To'rt qadam: ballar, masshtablash, softmax, vaznli yig'indi

Hammasi bitta funksiyada #

Python
def attention(Q, K, V, mask=None):
    d_k = Q.size(-1)
    ball = Q @ K.transpose(-2, -1) / d_k**0.5
    if mask is not None:
        ball = ball.masked_fill(mask == 0, float('-inf'))
    w = F.softmax(ball, dim=-1)
    return w @ V, w

Olti qator. Sohani o'zgartirgan mexanizm shundan iborat.

Bu yerda o'qitiladigan nima

attention funksiyasining o'zida birorta ham parametr yo'q - u faqat matritsa amallari.

O'qitiladigan narsa - Wq, Wk va Wv qatlamlari. Model aynan shu uch matritsani sozlash orqali «nimani qidirish» va «nimani taklif qilish»ni o'rganadi.

Amaliy topshiriq
  1. Ikki vektorning skalyar ko'paytmasini hisoblab, o'xshashlik bilan bog'liqligini ko'rsating.
  2. Perpendikulyar vektorlar uchun natija nega nol ekanini tushuntiring.
  3. Wq, Wk, Wv qatlamlarini yaratib, Q, K va V shakllarini chop eting.
  4. Q @ K.T ni hisoblab, matritsa o'lchamini izohlang.
  5. d_k ni 4, 64 va 512 qilib, ballar standart og'ishini o'lchang.
  6. Ildizga bo'lgandan keyin og'ish nega o'zgarmasligini tushuntiring.
  7. [1, 0, 0] va [50, 0, 0] ballarga softmax qo'llab, farqni yozing.
  8. Ildizga bo'lishni olib tashlab, og'irliklar qanday o'zganini kuzating.
  9. attention funksiyasini yozib, og'irliklar yig'indisi 1 ekanini tekshiring.
  10. Q, K va V uchun bitta umumiy matritsa ishlatilsa nima yo'qolishini muhokama qiling.

Xulosa #

  • E'tibor skalyar ko'paytma orqali o'xshashlikni o'lchaydi.
  • Har so'z uchta rolda ishlatiladi: so'rov (Q), kalit (K) va qiymat (V).
  • Ular uchta alohida Linear qatlamdan chiqadi - kalit va qiymat boshqa-boshqa narsa.
  • Q @ K.T ballar matritsasini beradi: kim kimga qanchalik qiziqishi.
  • Ballar √d_k ga bo'linadi, chunki vektor uzunlashgan sari ular tarqalib ketadi.
  • Bizning o'lchovda tarqalish 2.02 -> 8.11 -> 22.90 bo'ldi, bo'lingach esa doim 1.01.
  • Bu shart: katta ballarda softmax bittasiga 1.0 beradi va gradient to'xtaydi.
  • softmax ballarni taqsimotga aylantiradi - har qator yig'indisi aynan 1.
  • Oxirgi qadam - og'irliklar bilan V ni yig'ish.
  • Mexanizmning o'zida parametr yo'q; o'qitiladigani - Wq, Wk, Wv.

Keyingi bo'limda self-attention, niqob va kvadratik xarajat masalasini ko'ramiz.

Xatolik topdingizmi?

Imlo xatosi, ishlamaydigan kod yoki noto‘g‘ri ma‘lumotni ko‘rsangiz - bizga xabar bering. Har bir xabar administrator tomonidan ko‘rib chiqiladi.