4-bo‘lim

Self-attention, niqob va kvadratik xarajat

Q, K va V ning bir manbadan olinishi, kelajakni yopuvchi niqob, -inf ning ma'nosi va uzunlik oshgandagi xarajat.

🕑 9 daqiqa o‘qish 📄 707 so‘z 👁 2 marta ko‘rilgan
Ushbu bo‘lim mundarijasi
  1. Self-attention
  2. Muammo: model kelajakni ko'rmasligi kerak
  3. Ta'sirini ko'ramiz
  4. Nega -inf, nol emas
  5. Kvadratik xarajat
  6. Amalda o'lchaymiz
  7. Xulosa

3-bo'limda mexanizmni qurdik. Endi undagi eng muhim tafsilotga o'tamiz: Q, K va V qayerdan kelishi.

Self-attention #

Agar uchalasi ham bir xil manbadan olinsa, bu self-attention (o'ziga e'tibor) deyiladi:

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

torch.manual_seed(0)

d = 8
x = torch.randn(1, 5, d)

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

Q, K, V = Wq(x), Wk(x), Wv(x)   # uchalasi ham x dan

Shuning uchun «o'ziga»: jumla o'z-o'ziga qaraydi. Har so'z o'sha jumladagi boshqa so'zlardan qaysi biri o'ziga muhimligini hal qiladi.

TurQ qayerdanK va V qayerdanQayerda ishlatiladi
Self-attentionBir xil matndanO'sha matndanGPT, BERT - asosiy qism
Cross-attentionChiqish matnidanKirish matnidanTarjima: nemischa so'z inglizchaga qaraydi

Bu darslikda asosan self-attention bilan ishlaymiz.

Muammo: model kelajakni ko'rmasligi kerak #

Matn yozuvchi model (GPT kabi) keyingi so'zni bashorat qilib o'qiydi. «Men bozorga ___» jumlasida u «bordim» ni topishi kerak.

Lekin self-attention hamma so'zga qaraydi - shu jumladan javobning o'ziga ham. Model javobni ko'rib turib «bashorat qilsa», u hech nima o'rganmaydi.

Yechim - niqob (mask):

Python
n = 5
mask = torch.tril(torch.ones(n, n))
Natija
   1.00   0.00   0.00   0.00   0.00
   1.00   1.00   0.00   0.00   0.00
   1.00   1.00   1.00   0.00   0.00
   1.00   1.00   1.00   1.00   0.00
   1.00   1.00   1.00   1.00   1.00

torch.tril quyi uchburchak matritsa yasaydi. Birinchi so'z faqat o'ziga, ikkinchisi o'ziga va birinchisiga qaraydi va hokazo.

Ta'sirini ko'ramiz #

Python
ball = Q @ K.transpose(-2, -1) / d**0.5
w1 = F.softmax(ball, dim=-1)
w2 = F.softmax(ball.masked_fill(mask == 0, float('-inf')), dim=-1)
Natija
  niqobsiz (har so'z hammaga qaraydi):
   0.25   0.12   0.23   0.18   0.22
   0.17   0.20   0.15   0.20   0.28
   0.18   0.20   0.20   0.19   0.23
   0.18   0.25   0.17   0.21   0.19
   0.14   0.16   0.17   0.16   0.36
  niqob bilan (faqat orqaga qaraydi):
   1.00   0.00   0.00   0.00   0.00
   0.45   0.55   0.00   0.00   0.00
   0.32   0.34   0.34   0.00   0.00
   0.22   0.31   0.21   0.26   0.00
   0.14   0.16   0.17   0.16   0.36

Ikkinchi jadvalning yuqori o'ng qismi butunlay nol. Birinchi so'z faqat o'ziga qaraydi - shuning uchun uning og'irligi 1.00.

Oxirgi qator ikkala holatda ham bir xil: oxirgi so'zdan keyin kelajak yo'q.

Har so'z qaysi so'zlarga qaray oladi Niqobsiz (BERT uslubi) Niqob bilan (GPT uslubi) kimga → 1 2 3 4 5 1 2 3 4 5 1 2 3 4 5 1 2 3 4 5 butun matnni ko'radi - tushunish uchun faqat orqaga qaraydi - yozish uchun
Niqob modelni javobni ko'rishdan to'sadi; bu GPT va BERT o'rtasidagi asosiy farq

Nega -inf, nol emas #

Niqobda 0 emas, -inf qo'yiladi. Sababi softmax:

Python
print("  oxirgisini 0 qilsak ->",
      F.softmax(torch.tensor([2.0, 1.0, 0.0]), dim=-1).round(decimals=3).tolist())
print("  oxirgisini -inf qilsak ->",
      F.softmax(torch.tensor([2.0, 1.0, float('-inf')]), dim=-1).round(decimals=3).tolist())
Natija
  oxirgisini 0 qilsak -> [0.665, 0.245, 0.09]  (hali ham ulush oladi)
  oxirgisini -inf qilsak -> [0.731, 0.269, 0.0]  (aniq nol)

Ball nolga tenglashtirilsa, u hamon 9% ulush oldi - chunki exp(0) = 1. Faqat -inf da exp(-inf) = 0 bo'ladi va token butunlay o'chadi.

Niqobni softmax dan keyin qo'llamang

Ba'zan niqobni og'irliklarga ko'paytirish orqali qo'llashga urinishadi: w = w * mask.

Bu noto'g'ri. Softmax allaqachon hisoblab bo'lgan va yig'indi 1 ga teng edi; ko'paytirgandan keyin yig'indi 1 dan kichik bo'lib qoladi.

Niqob har doim softmax dan oldin, ballarga -inf sifatida qo'llanadi.

Kvadratik xarajat #

1-bo'limda aytilgan kamchilikni endi hisoblaymiz. Har so'z har so'zga qarasa, juftliklar soni uzunlikning kvadrati bo'ladi:

Python
for L in [10, 100, 512, 2048, 8192]:
    juft = L * L
    print(f"  {L:8d} {juft:12,d} {juft*4/1048576:20.1f}")
Natija
   uzunlik   juftliklar    xotira (MB, fp32)
        10          100                  0.0
       100       10,000                  0.0
       512      262,144                  1.0
      2048    4,194,304                 16.0
      8192   67,108,864                256.0

Uzunlik 16 barobar oshdi (512 dan 8192 ga), xotira esa 256 barobar.

Va bu faqat bitta e'tibor qatlamining bitta boshi uchun. 12 qatlamli, 12 boshli modelda buni 144 ga ko'paytiring.

Amalda o'lchaymiz #

Python
import time

for L in [128, 256, 512, 1024]:
    q = torch.randn(1, L, 64); k = torch.randn(1, L, 64); v = torch.randn(1, L, 64)
    for _ in range(2):                      # isitish (PyTorch 9-bo'lim)
        F.softmax(q @ k.transpose(-2, -1) / 8, dim=-1) @ v
    t = time.perf_counter()
    for _ in range(5):
        F.softmax(q @ k.transpose(-2, -1) / 8, dim=-1) @ v
    print(f"  uzunlik {L:5d}: {(time.perf_counter()-t)/5*1000:7.2f} ms")
Natija
  uzunlik   128:    0.06 ms
  uzunlik   256:    0.11 ms
  uzunlik   512:    1.48 ms
  uzunlik  1024:    4.88 ms
Kichik uzunlikda kvadratik o'sish ko'rinmaydi

128 dan 256 ga o'tganda vaqt atigi 1.8 barobar oshdi - kvadratik bo'lsa 4 barobar bo'lishi kerak edi.

Sabab: bu o'lchamlarda qo'shimcha xarajatlar (funksiya chaqiruvi, xotira ajratish) asosiy hisobdan katta. 512 dan 1024 ga o'tganda esa o'sish 3.3 barobar - haqiqiy nisbatga yaqinlashdi.

Bu o'lchashning umumiy saboqi: kichik namunada tendensiya ko'rinmaydi. Xulosa chiqarishdan oldin masshtabni yetarlicha kattalashtiring.

UzunlikNima qilish mumkin
512 gachaOddiy e'tibor yetarli
2048 gachaFlashAttention kabi tezroq amalga oshirish
Undan uzunMatnni bo'laklash yoki boshqa arxitektura

Aynan shuning uchun til modellarida «kontekst oynasi» degan chegara bor - va uni oshirish qimmatga tushadi.

O'zbek tili uchun bu ikki barobar muhim

2-bo'limda ko'rdik: o'zbekcha matn 3 barobar ko'p token oladi.

Xarajat esa token sonining kvadrati bilan o'sadi. Demak bir xil ma'noli matn uchun e'tibor hisobi taxminan 9 barobar qimmatroq tushadi.

Amaliy topshiriq
  1. Self-attention va cross-attention farqini o'z so'zingiz bilan yozing.
  2. torch.tril bilan 6x6 niqob yasab, chop eting.
  3. Niqobsiz va niqob bilan e'tibor og'irliklarini solishtiring.
  4. Birinchi qatordagi og'irlik nega aynan 1.00 ekanini tushuntiring.
  5. Oxirgi qator nega ikkala holatda bir xil ekanini yozing.
  6. Niqobga 0 va -inf qo'yib, softmax natijasini solishtiring.
  7. exp(0) = 1 ekani nega muammo tug'dirishini tushuntiring.
  8. Uzunlik 256, 1024 va 4096 uchun juftliklar sonini hisoblang.
  9. Turli uzunliklarda e'tibor vaqtini o'lchab, jadval tuzing.
  10. O'zbekcha matn uchun xarajat nega 9 barobar oshishini hisoblab ko'rsating.

Xulosa #

  • Self-attention da Q, K va V bir xil manbadan olinadi.
  • Cross-attention da so'rov bir matndan, kalit va qiymat boshqasidan keladi.
  • Matn yozuvchi model kelajakni ko'rmasligi kerak - buning uchun niqob qo'yiladi.
  • torch.tril quyi uchburchak niqob yasaydi.
  • Niqobli modelda birinchi so'zning og'irligi aynan 1.00 bo'ladi.
  • Niqob 0 emas, -inf bo'lishi shart: exp(0) = 1 va token ulush olib qoladi.
  • Niqob har doim softmax dan oldin qo'llanadi.
  • E'tibor xarajati uzunlikning kvadrati bilan o'sadi.
  • 512 dan 8192 ga o'tganda xotira 256 barobar oshadi.
  • Kichik uzunlikda bu o'sish ko'rinmaydi - qo'shimcha xarajatlar hisobni yashiradi.

Keyingi bo'limda bir nechta e'tibor boshini birlashtiramiz - multi-head attention.

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.