5-bo‘lim

Multi-head attention

Nima uchun bitta e'tibor yetarli emas, o'lchamni boshlarga bo'lish, view va transpose, chiqish matritsasi Wo va parametrlar soni.

🕑 9 daqiqa o‘qish 📄 671 so‘z 👁 2 marta ko‘rilgan
Ushbu bo‘lim mundarijasi
  1. O'lchamni bo'lish
  2. Shakl o'yini
  3. To'liq amalga oshirish
  4. Har bosh boshqacha qaraydi
  5. Wo nima uchun kerak
  6. Parametrlar soni
  7. Xulosa

Bitta e'tibor bitta savolga javob beradi. Lekin jumlada bir vaqtning o'zida bir nechta bog'liqlik bor: kim nima qildi, qachon, qayerda, qanday.

Yechim oddiy: bir nechta e'tiborni parallel ishlatish.

O'lchamni bo'lish #

Muhim nuqta - boshlar qo'shilganda model kattalashmaydi. Mavjud o'lcham boshlar orasida bo'linadi:

Python
d_model, n_head = 64, 8
d_head = d_model // n_head
print(f"  d_model={d_model}, boshlar={n_head} -> har bosh d_head={d_head}")
print(f"  jami o'lcham o'zgarmaydi: {n_head} x {d_head} = {n_head*d_head}")
Natija
  d_model=64, boshlar=8 -> har bosh d_head=8
  jami o'lcham o'zgarmaydi: 8 x 8 = 64

Har bosh 64 o'lchamli fazoning 8 o'lchamli bo'lagida ishlaydi.

Nega bu foydali

8 ta bosh 8 o'lchamda ishlagani - bitta bosh 64 o'lchamda ishlaganidan yaxshiroq bo'lishi mumkin.

Sabab: bitta e'tiborda bitta softmax bor va u bitta taqsimot beradi. Model «ham egaga, ham zamonga qara» deyishi kerak bo'lsa, bitta taqsimot ikkalasini ham yarim-yorti qoplaydi.

Sakkizta bosh esa sakkizta mustaqil taqsimot beradi - biri egaga, boshqasi zamonga qaray oladi.

Shakl o'yini #

Butun murakkablik shakllarni to'g'ri joylashtirishda:

Python
B, L = 2, 5
x = torch.randn(B, L, d_model)
W = nn.Linear(d_model, d_model, bias=False)

q = W(x)
print("  Linear dan keyin:", tuple(q.shape))

q = q.view(B, L, n_head, d_head).transpose(1, 2)
print("  view + transpose:", tuple(q.shape), " <- (paket, bosh, uzunlik, d_head)")
Natija
  Linear dan keyin: (2, 5, 64)
  view + transpose: (2, 8, 5, 8)  <- (paket, bosh, uzunlik, d_head)

Ikki qadam bajarildi:

AmalNima qildi
.view(B, L, n_head, d_head)64 o'lchamni 8 x 8 ga ajratdi
.transpose(1, 2)Bosh o'lchamini oldinga chiqardi

transpose nima uchun kerak? Shundan keyin oxirgi ikki o'lcham (uzunlik, d_head) bo'lib qoladi - bu aynan 3-bo'limdagi e'tibor kutgan shakl. PyTorch qolgan o'lchamlarni (paket va bosh) avtomatik ravishda paket sifatida qayta ishlaydi.

To'liq amalga oshirish #

Python
class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, n_head):
        super().__init__()
        assert d_model % n_head == 0
        self.n_head, self.d_head = n_head, d_model // n_head
        self.Wq = nn.Linear(d_model, d_model, bias=False)
        self.Wk = nn.Linear(d_model, d_model, bias=False)
        self.Wv = nn.Linear(d_model, d_model, bias=False)
        self.Wo = nn.Linear(d_model, d_model, bias=False)

    def forward(self, x, mask=None):
        B, L, D = x.shape

        def bol(t):
            return t.view(B, L, self.n_head, self.d_head).transpose(1, 2)

        Q, K, V = bol(self.Wq(x)), bol(self.Wk(x)), bol(self.Wv(x))

        ball = Q @ K.transpose(-2, -1) / self.d_head**0.5
        if mask is not None:
            ball = ball.masked_fill(mask == 0, float('-inf'))
        w = F.softmax(ball, dim=-1)

        out = (w @ V).transpose(1, 2).contiguous().view(B, L, D)
        return self.Wo(out), w
Natija
  kirish: (2, 5, 64) -> chiqish: (2, 5, 64)
  e'tibor og'irliklari: (2, 8, 5, 5)  <- har bosh uchun alohida
  parametrlar: 16384

Kirish va chiqish shakli bir xil - shuning uchun bunday bloklarni ustma-ust qo'yish mumkin.

.contiguous() ni tushirib qoldirmang

transpose tenzorni xotirada ko'chirmaydi - u faqat o'lchamlarga boshqacha qaraydi. .view() esa uzluksiz xotira talab qiladi.

Ikkalasini ketma-ket ishlatsangiz, xato chiqadi:

Natija
RuntimeError: view size is not compatible with input tensor's size and
stride (at least one dimension spans across two contiguous subspaces).
Use .reshape(...) instead.

Yechim ikkita: .contiguous().view(...) yoki to'g'ridan-to'g'ri .reshape(...). Xato xabarining o'zi ikkinchisini tavsiya qiladi.

Har bosh boshqacha qaraydi #

Python
for h in range(3):
    print(f"  bosh {h}, 5-so'z og'irliklari:", [round(v, 2) for v in w[0, h, 4].tolist()])
Natija
  bosh 0, 5-so'z og'irliklari: [0.21, 0.26, 0.16, 0.17, 0.2]
  bosh 1, 5-so'z og'irliklari: [0.2, 0.19, 0.19, 0.24, 0.19]
  bosh 2, 5-so'z og'irliklari: [0.17, 0.16, 0.18, 0.13, 0.36]

Bir xil so'z, bir xil jumla - lekin uch xil taqsimot. Ikkinchi bosh 4-so'zga, uchinchisi o'ziga ko'proq qaradi.

Bu model tasodifiy boshlang'ich holatda. O'qitilgan modelda boshlar ancha aniq rollarga ajraladi - ba'zilari sintaksisga, ba'zilari ma'noga qaraydi.

d_model = 64 sakkizta boshga bo'linadi x: (B, L, 64) Wq, Wk, Wv view(B, L, 8, 8) + transpose -> (B, 8, L, 8) bosh 1 8 o'lch. bosh 2 8 o'lch. bosh 3 8 o'lch. bosh 4 8 o'lch. ··· bosh 8 8 o'lch. har biri mustaqil softmax hisoblaydi birlashtirish -> (B, L, 64) Wo - aralashtirish Wo bo'lmasa har bosh o'z bo'lagida yakka qolardi
Bo'lish, mustaqil e'tibor, birlashtirish va Wo bilan aralashtirish

Wo nima uchun kerak #

Boshlar birlashtirilgandan keyin natija yonma-yon yopishtirilgan bo'laklardan iborat: dastlabki 8 o'lcham birinchi boshdan, keyingi 8 tasi ikkinchisidan.

Wo ularni aralashtiradi. Usiz har bosh o'z bo'lagida qolib ketardi va keyingi qatlam ularni birga ishlata olmasdi.

Parametrlar soni #

Bu natija ko'pchilikni hayratda qoldiradi:

Python
for nh in [1, 2, 4, 8, 16]:
    m = MultiHeadAttention(64, nh)
    print(f"  {nh:2d} bosh -> {sum(p.numel() for p in m.parameters()):,} parametr")
Natija
   1 bosh -> 16,384 parametr
   2 bosh -> 16,384 parametr
   4 bosh -> 16,384 parametr
   8 bosh -> 16,384 parametr
  16 bosh -> 16,384 parametr

Bosh soni parametrlar soniga umuman ta'sir qilmaydi. Chunki to'rtala matritsa ham d_model x d_model bo'lib qolaveradi - biz faqat ularning chiqishiga boshqacha qaraymiz.

Demak boshlar soni - bepul sozlama. U hisob-kitobni ham deyarli oshirmaydi.

Modeld_modelBoshlarHar bosh
Asl transformer (base)512864
BERT-base7681264
GPT-2 small7681264

Diqqat qiling: har bosh o'lchami odatda 64 atrofida saqlanadi. Model kattalashganda boshlar soni oshiriladi, har bosh o'lchami emas.

PyTorch da tayyori bor

Amalda nn.MultiheadAttention ishlatiladi - u optimallashtirilgan va tezroq:

Python
mha = nn.MultiheadAttention(d_model, n_head, batch_first=True)
out, w = mha(x, x, x)

batch_first=True ni unutmang - usiz PyTorch (uzunlik, paket, o'lcham) tartibini kutadi va shakl xatosi chiqadi.

Amaliy topshiriq
  1. d_model=128, n_head=8 uchun d_head ni hisoblang.
  2. d_model boshlar soniga bo'linmasa nima bo'lishini sinab ko'ring.
  3. view va transpose dan keyingi shaklni chop eting.
  4. transpose nega kerakligini o'z so'zingiz bilan tushuntiring.
  5. .contiguous() ni olib tashlab, xato xabarini o'qing.
  6. MultiHeadAttention sinfini yozib, kirish va chiqish shakllarini solishtiring.
  7. Uch xil boshning og'irliklarini chop etib, farqini ko'rsating.
  8. Boshlar sonini 1 dan 16 gacha o'zgartirib, parametrlar sonini o'lchang.
  9. Natija nega o'zgarmasligini tushuntiring.
  10. nn.MultiheadAttention bilan o'z yozganingizning chiqish shaklini taqqoslang.

Xulosa #

  • Bitta e'tibor bitta taqsimot beradi - bu bir nechta bog'liqlik uchun kam.
  • Multi-head bir nechta e'tiborni parallel ishlatadi.
  • Boshlar qo'shilganda model kattalashmaydi: d_model ular orasida bo'linadi.
  • view o'lchamni ajratadi, transpose bosh o'lchamini oldinga chiqaradi.
  • transpose dan keyin .view() ishlatish uchun .contiguous() yoki .reshape() kerak.
  • Har bosh mustaqil softmax hisoblaydi va boshqacha taqsimot beradi.
  • Chiqish shakli kirish bilan bir xil - shuning uchun bloklarni ustma-ust qo'yish mumkin.
  • Wo boshlar natijasini aralashtiradi; usiz ular yakka qolardi.
  • Parametrlar soni bosh soniga bog'liq emas - 1 bosh ham, 16 bosh ham 16 384.
  • Amalda har bosh o'lchami 64 atrofida saqlanadi, model kattalashganda boshlar soni oshadi.

Keyingi bo'limda e'tibor atrofiga qolgan qismlarni qo'shib, to'liq transformer blokini yig'amiz.

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.