8-bo‘lim

To'liq o'quv sikli

train() va eval() rejimlari, davr tushunchasi, validatsiya bilan baholash va o'qish jarayonini kuzatish.

🕑 7 daqiqa o‘qish 📄 635 so‘z 👁 0 marta ko‘rilgan
Ushbu bo‘lim mundarijasi
  1. Masala
  2. Ma'lumotni tayyorlash
  3. Model, yo'qotish, optimizator
  4. Baholash funksiyasi
  5. O'quv sikli
  6. Sikl ichida nima bo'ladi
  7. Raqamlarni qanday o'qish kerak
  8. Xulosa

Bizda model, yo'qotish, optimizator va DataLoader bor. Endi ularni bitta ishlaydigan quvurga birlashtiramiz. Bu bo'limdagi kod - keyingi barcha bo'limlar uchun asos.

Masala #

Nuqta markazdan 1 birlik radius ichidami yoki tashqarisidami - shuni aniqlaydigan model quramiz. Bu masalani chiziqli model yecholmaydi, shuning uchun u yashirin qatlamlarni sinash uchun qulay:

Python
import torch
import torch.nn as nn
from torch.utils.data import TensorDataset, DataLoader, random_split

torch.manual_seed(42)

X = torch.randn(600, 2)
y = ((X[:, 0] ** 2 + X[:, 1] ** 2) < 1.0).long()
print("sinf taqsimoti:", torch.bincount(y).tolist())
Natija
sinf taqsimoti: [361, 239]

361 ta nuqta tashqarida, 239 tasi ichkarida. Taqsimot bir xil emas, lekin o'ta nomutanosib ham emas - bu yaxshi.

Ma'lumotni tayyorlash #

Python
tds = TensorDataset(X, y)
oquv, valid = random_split(tds, [500, 100], generator=torch.Generator().manual_seed(42))

oquv_dl = DataLoader(oquv, batch_size=32, shuffle=True)
valid_dl = DataLoader(valid, batch_size=32)

Diqqat: valid_dl da shuffle yo'q - baholashda tartibning ahamiyati yo'q.

Model, yo'qotish, optimizator #

Python
model = nn.Sequential(
    nn.Linear(2, 16), nn.ReLU(),
    nn.Linear(16, 16), nn.ReLU(),
    nn.Linear(16, 2),
)

kriteriya = nn.CrossEntropyLoss()
opt = torch.optim.Adam(model.parameters(), lr=0.01)

Oxirgi qatlam 2 ta chiqish beradi - har sinf uchun bitta logit. Softmax yo'q, chunki uni CrossEntropyLoss o'zi bajaradi.

Baholash funksiyasi #

O'qish natijasini ko'rish uchun alohida funksiya yozamiz:

Python
def baholash(dl):
    model.eval()
    jami, togri, yoqotish = 0, 0, 0.0
    with torch.no_grad():
        for xb, yb in dl:
            chiqish = model(xb)
            yoqotish += kriteriya(chiqish, yb).item() * len(yb)
            togri += (chiqish.argmax(dim=1) == yb).sum().item()
            jami += len(yb)
    return yoqotish / jami, togri / jami

Uchta muhim nuqta bor:

  • model.eval() - modelni baholash rejimiga o'tkazadi.
  • torch.no_grad() - hosila hisoblanmaydi, tezroq va kam xotira.
  • Yo'qotish len(yb) ga ko'paytirilib qo'shiladi - oxirgi paket kichik bo'lgani uchun oddiy o'rtacha noto'g'ri chiqadi.
train() va eval() ni almashtirishni unutmang

Dropout va BatchNorm qatlamlari ikki rejimda boshqacha ishlaydi (11 va 12-bo'limlarda ko'ramiz).

eval() ni unutsangiz, validatsiya natijasi tasodifiy bo'lib sakraydi. train() ga qaytishni unutsangiz, model o'qishni to'xtatadi - lekin xato chiqmaydi. Shuning uchun ikkalasini ham sikl ichiga yozib qo'ying.

O'quv sikli #

Python
print("Davr | O'quv yo'q. | Valid. yo'q. | Valid. aniqlik")
print("-----|------------|--------------|---------------")

for davr in range(1, 21):
    model.train()
    jami_yoqotish, jami = 0.0, 0

    for xb, yb in oquv_dl:
        opt.zero_grad()
        chiqish = model(xb)
        yoqotish = kriteriya(chiqish, yb)
        yoqotish.backward()
        opt.step()

        jami_yoqotish += yoqotish.item() * len(yb)
        jami += len(yb)

    oquv_y = jami_yoqotish / jami
    val_y, val_a = baholash(valid_dl)

    if davr <= 5 or davr % 5 == 0:
        print(f"{davr:4d} | {oquv_y:10.4f} | {val_y:12.4f} | {val_a:13.2%}")
Natija
Davr | O'quv yo'q. | Valid. yo'q. | Valid. aniqlik
-----|------------|--------------|---------------
   1 |     0.6037 |       0.5337 |        61.00%
   2 |     0.4671 |       0.3828 |        88.00%
   3 |     0.2980 |       0.2109 |        97.00%
   4 |     0.1767 |       0.1300 |        96.00%
   5 |     0.1132 |       0.0886 |        98.00%
  10 |     0.0564 |       0.0635 |        97.00%
  15 |     0.0450 |       0.0294 |        99.00%
  20 |     0.0405 |       0.0411 |        99.00%

Model 61% dan 99% gacha ko'tarildi. Ikkala yo'qotish ham birga tushdi - bu sog'lom o'qish belgisi.

Sikl ichida nima bo'ladi #

Bitta davr (epoch) ichida model.train() har bir paket uchun: zero_grad() tozalash model(xb) bashorat backward() gradient step() yangilash model.eval() + torch.no_grad() validatsiya to'plami bo'ylab yurib, yo'qotish va aniqlikni o'lchaydi parametrlar bu yerda o'zgarmaydi - faqat o'lchov olinadi keyingi davr shu tartibda qaytadan boshlanadi
Har davrda avval o'quv paketlari, so'ng validatsiya o'tkaziladi

Raqamlarni qanday o'qish kerak #

Ko'rinishMa'nosiNima qilish
Ikkala yo'qotish tushmoqdaSog'lom o'qishDavom eting
O'quv tushdi, validatsiya o'smoqdaOrtiqcha moslashish11-bo'limga qarang
Ikkalasi ham turib qoldiModel kuchsiz yoki lr kichikModel kattalashtiring yoki lr oshiring
Yo'qotish sakraydi yoki nanlr katta yoki zero_grad() yo'qlr ni 10 barobar kichraytiring
Aniqlik sinf ulushiga tengModel hech nima o'rganmagansuper().__init__() va optimizatorni tekshiring

Oxirgi qatorga alohida e'tibor bering: bizning ma'lumotda 361/600 = 60.2% nuqta tashqarida. Agar model 60% atrofida qotib qolsa, demak u hamma narsani bitta sinfga tashlayapti - o'rgangani yo'q.

Avval kichik namunada sinang

Yangi quvurni 20-30 ta namunada ishga tushiring va modelni ataylab ortiqcha moslashtiring. Agar model shuncha kichik ma'lumotni ham yodlay olmasa, demak kodda xato bor - katta ma'lumotda vaqt yo'qotishdan oldin shuni tekshiring.

.item() nima uchun kerak

yoqotish - bu tenzor, u hisoblash grafiga bog'langan. Uni to'g'ridan- to'g'ri ro'yxatga yig'sangiz, butun graf xotirada qoladi va xotira asta-sekin to'ladi. .item() esa oddiy Python sonini qaytaradi.

Amaliy topshiriq
  1. Yuqoridagi kodni to'liq ishga tushiring va 20 davr natijasini oling.
  2. Sinf taqsimotini bincount bilan chop eting.
  3. Yashirin qatlamlarni olib tashlab, faqat nn.Linear(2, 2) qoldiring va aniqlikni solishtiring.
  4. Yashirin qatlamni 16 dan 64 ga oshirib, natijani yozing.
  5. lr ni 0.001 va 0.5 qilib, ikkala holatni taqqoslang.
  6. model.eval() ni o'chirib, validatsiya natijasi o'zgarganini kuzating.
  7. opt.zero_grad() ni o'chirib, yo'qotish qanday sakraganini ko'ring, keyin qaytaring.
  8. Davrlar sonini 100 ga oshirib, validatsiya yo'qotishi qachon to'xtaganini toping.
  9. batch_size ni 8 va 128 qilib, bir davr qancha davom etishini solishtiring.
  10. Har davrdagi yo'qotishni ro'yxatga yig'ib, oxirida eng yaxshi davrni toping.

Xulosa #

  • O'quv sikli ikki qismdan iborat: o'qitish va baholash.
  • O'qitishdan oldin model.train(), baholashdan oldin model.eval() chaqiriladi.
  • Baholash har doim torch.no_grad() ichida bajariladi.
  • Bitta davr - butun o'quv to'plamining bir marta to'liq aylanishi.
  • Har paketda o'sha uchlik: zero_grad() -> backward() -> step().
  • O'rtacha yo'qotishni hisoblashda paket hajmiga ko'paytiring - oxirgi paket kichik bo'lishi mumkin.
  • Yo'qotishni .item() bilan yig'ing, aks holda graf xotirada to'planadi.
  • Ikkala yo'qotish birga tushsa - o'qish sog'lom.
  • Validatsiya yo'qotishi o'sa boshlasa - ortiqcha moslashish boshlangan.
  • Aniqlik eng katta sinf ulushi atrofida qotib qolsa, model hech nima o'rganmagan.

Keyingi bo'limda shu quvurni GPU ga ko'chiramiz va tezlikni o'lchaymiz.

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.