20-bo‘lim

Yakuniy loyiha

Barcha bo'limlarni bitta quvurda birlashtirish - BatchNorm, Dropout, AdamW, kengaytirish, lr rejalashtiruvchi va erta to'xtatish.

🕑 10 daqiqa o‘qish 📄 858 so‘z 👁 0 marta ko‘rilgan
Ushbu bo‘lim mundarijasi
  1. Reja
  2. Sozlamalar
  3. Ma'lumot: uchga bo'lish
  4. Model
  5. Optimizator va rejalashtiruvchi
  6. O'quv sikli
  7. Natija
  8. Test - bir marta
  9. Sinflar bo'yicha
  10. Nimani o'zgartirdik
  11. Xulosa

Oxirgi bo'lim. Endi 19 ta bo'limda o'rganilgan hamma narsani bitta to'liq loyihada birlashtiramiz va 14-bo'limdagi 89.34% natijani yaxshilashga harakat qilamiz.

Reja #

Bo'limNimani qo'shamiz
8To'liq o'quv sikli, train() / eval()
9Device-agnostik kod
10Nazorat nuqtasi va eng yaxshi modelni saqlash
12Dropout, weight decay, erta to'xtatish
13, 14Chuqurroq CNN
15Normallash va kengaytirish
18Sinflar bo'yicha tahlil

Sozlamalar #

Barcha o'zgarmas qiymatlarni bir joyga yig'amiz - bu keyin nazorat fayliga ham tushadi:

Python
import torch, time, os
import torch.nn as nn
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms

QURILMA = torch.device("cuda" if torch.cuda.is_available() else "cpu")
ORTACHA, OGISH = 0.2860, 0.3530
SINFLAR = ['Futbolka', 'Shim', 'Sviter', 'Libos', 'Palto',
           'Sandal', "Ko'ylak", 'Krossovka', 'Sumka', 'Botinka']
torch.manual_seed(42)
print("qurilma:", QURILMA)

Ma'lumot: uchga bo'lish #

15-bo'limdagi qoida: kengaytirish faqat o'quv to'plamiga. Shuning uchun bir xil ma'lumotni ikki xil transform bilan ochamiz:

Python
oquv_t = transforms.Compose([
    transforms.RandomHorizontalFlip(),
    transforms.RandomAffine(degrees=8, translate=(0.08, 0.08)),
    transforms.ToTensor(), transforms.Normalize((ORTACHA,), (OGISH,))])

baho_t = transforms.Compose([
    transforms.ToTensor(), transforms.Normalize((ORTACHA,), (OGISH,))])

toliq   = datasets.FashionMNIST("./data", train=True, transform=oquv_t)
toliq_v = datasets.FashionMNIST("./data", train=True, transform=baho_t)
test    = datasets.FashionMNIST("./data", train=False, transform=baho_t)

g = torch.Generator().manual_seed(42)
o_idx, v_idx = random_split(range(len(toliq)), [54000, 6000], generator=g)
oquv  = torch.utils.data.Subset(toliq,   list(o_idx))
valid = torch.utils.data.Subset(toliq_v, list(v_idx))
Natija
qurilma: cpu
o'quv: 54000 | validatsiya: 6000 | test: 10000
Test to'plamiga o'qitish davomida tegmang

Bu yerda uchta to'plam bor: o'quv, validatsiya va test.

Sozlamalarni (davrlar soni, lr, model o'lchami) faqat validatsiya natijasiga qarab tanlang. Test to'plami oxirida, bir marta ishlatiladi.

Agar test natijasiga qarab sozlama tanlasangiz, test ham amalda validatsiyaga aylanadi va yakuniy baho haqiqatdan yuqori chiqadi.

Model #

13-bo'limdagi CNN ga ikkita yangilik qo'shamiz: har blokda ikkita konvolyutsiya va BatchNorm2d.

Python
def blok(kir, chiq):
    return nn.Sequential(
        nn.Conv2d(kir, chiq, 3, padding=1), nn.BatchNorm2d(chiq), nn.ReLU(),
        nn.Conv2d(chiq, chiq, 3, padding=1), nn.BatchNorm2d(chiq), nn.ReLU(),
        nn.MaxPool2d(2))

class YaxshiCNN(nn.Module):
    def __init__(self, sinf=10):
        super().__init__()
        self.xususiyat = nn.Sequential(blok(1, 32), blok(32, 64))
        self.tasnif = nn.Sequential(
            nn.Flatten(), nn.Dropout(0.3),
            nn.Linear(64 * 7 * 7, 128), nn.ReLU(),
            nn.Dropout(0.3), nn.Linear(128, sinf))

    def forward(self, x):
        return self.tasnif(self.xususiyat(x))
Natija
parametrlar: 468,202

14-bo'limdagi model 20 490 parametrga ega edi - bu 23 barobar kattaroq.

BatchNorm2d nima qiladi

U har paketdagi qiymatlarni normallaydi - 15-bo'limdagi Normalize ning tarmoq ichidagi varianti.

Foydasi: o'qitish barqarorroq kechadi, kattaroq lr ishlatish mumkin va u o'zi ham yengil regularizatsiya beradi. Dropout singari, u ham train() va eval() rejimlarida boshqacha ishlaydi - shuning uchun ularni to'g'ri chaqirish yana bir bor muhim bo'ladi.

Optimizator va rejalashtiruvchi #

Python
kriteriya = nn.CrossEntropyLoss()
opt = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)
DAVRLAR = 15
rejalash = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=DAVRLAR)

CosineAnnealingLR lr ni asta-sekin kamaytiradi: boshida katta qadamlar bilan tez o'rganadi, oxirida esa kichik qadamlar bilan aniq sozlanadi.

O'quv sikli #

Python
@torch.no_grad()
def baholash(dl):
    model.eval(); j = t = 0; yo = 0.0
    for xb, yb in dl:
        xb, yb = xb.to(QURILMA), yb.to(QURILMA)
        o = model(xb); yo += kriteriya(o, yb).item() * len(yb)
        t += (o.argmax(1) == yb).sum().item(); j += len(yb)
    return yo / j, t / j

eng_yaxshi, eng_davr, sabr, kutildi = float('inf'), 0, 4, 0

for davr in range(1, DAVRLAR + 1):
    model.train(); s = 0.0; n = 0
    for xb, yb in o_dl:
        xb, yb = xb.to(QURILMA), yb.to(QURILMA)
        opt.zero_grad()
        l = kriteriya(model(xb), yb)
        l.backward(); opt.step()
        s += l.item() * len(yb); n += len(yb)

    vy, va = baholash(v_dl)
    rejalash.step()

    if vy < eng_yaxshi:
        eng_yaxshi, eng_davr, kutildi = vy, davr, 0
        torch.save({"model": model.state_dict(), "davr": davr,
                    "ortacha": ORTACHA, "ogish": OGISH,
                    "sinflar": SINFLAR}, "yakuniy.pt")
    else:
        kutildi += 1
        if kutildi >= sabr:
            print(f"erta to'xtatildi ({davr}-davr)")
            break

Nazorat fayliga model bilan birga ORTACHA, OGISH va SINFLAR ham yozildi - 19-bo'limdagi eng ko'p uchraydigan xatoning oldini olish uchun.

Natija #

Natija
Davr | O'quv yo'q. | Valid yo'q. | Valid aniq. |   lr    | Vaqt
-----|------------|-------------|-------------|---------|------
   1 |     0.6049 |      0.4099 |      85.08% | 0.001000 |   72s
   2 |     0.4232 |      0.3783 |      87.03% | 0.000989 |   71s
   3 |     0.3756 |      0.2892 |      89.30% | 0.000957 |   69s
   4 |     0.3485 |      0.2648 |      90.08% | 0.000905 |   67s
   5 |     0.3261 |      0.2708 |      89.73% | 0.000835 |   69s
   6 |     0.3104 |      0.2546 |      90.68% | 0.000750 |   67s
   7 |     0.2953 |      0.2275 |      91.78% | 0.000655 |   63s
   8 |     0.2806 |      0.2205 |      91.57% | 0.000552 |   68s
   9 |     0.2705 |      0.2231 |      91.80% | 0.000448 |   71s
  10 |     0.2594 |      0.2079 |      92.17% | 0.000345 |   76s
  11 |     0.2472 |      0.2044 |      92.48% | 0.000250 |   70s
  12 |     0.2441 |      0.1970 |      92.47% | 0.000165 |   70s
  13 |     0.2364 |      0.1932 |      92.65% | 0.000095 |   65s
  14 |     0.2347 |      0.1932 |      92.80% | 0.000043 |   68s
  15 |     0.2308 |      0.1921 |      92.95% | 0.000011 |   67s

eng yaxshi davr: 15  (valid yo'qotish 0.1921)
Yakuniy modelning o'qish egri chizig'i 0.16 0.32 0.48 0.63 1 3 6 9 12 15 davr eng yaxshi (15-davr) o'quv yo'qotishi validatsiya yo'qotishi Ikkala chiziq ham birga tushdi - kengaytirish va regularizatsiya ishladi
Validatsiya yo'qotishi barqaror tushdi - ortiqcha moslashish belgisi yo'q

Test - bir marta #

Python
nz = torch.load("yakuniy.pt", weights_only=False)
model.load_state_dict(nz["model"])
ty, ta = baholash(t_dl)
print(f"TEST: yo'qotish={ty:.4f}  aniqlik={ta:.2%}")
Natija
TEST: yo'qotish=0.2023  aniqlik=92.72%

14-bo'limdagi oddiy CNN 89.34% bergan edi. Yangi quvur 92.72% beradi - ya'ni +3.38 foiz yaxshilanish.

Erta to'xtatish bu safar ishga tushmadi

Jadvalga qarang: validatsiya yo'qotishi oxirgi, 15-davrgacha tushishda davom etdi (0.1932 -> 0.1921). Shuning uchun sabr=4 sharti bajarilmadi va sikl to'liq aylandi.

Bu yomon xabar emas, balki ma'lumot: model hali ortiqcha moslashmagan, demak davrlar sonini oshirish natijani yana yaxshilashi mumkin. 11-bo'limdagi holat bunga teskari edi - u yerda validatsiya 3-davrdayoq yomonlasha boshlagan.

Erta to'xtatish - xavfsizlik kamari: kerak bo'lmasa ishlamaydi, kerak bo'lganda esa sizni qutqaradi.

Xato qilingan rasmlar ulushi 10.66% dan 7.28% ga tushdi: bu xatolarning taxminan 32% i yo'qoldi.

Foiz oshgani sayin har qadam qiyinlashadi

89% dan 92% ga ko'tarilish 60% dan 70% ga ko'tarilishdan ancha qiyin. Oxirgi foizlar odatda eng ko'p mehnat talab qiladi - chunki qolgan xatolar aynan 14-bo'limda ko'rgan o'xshash sinflar orasida bo'ladi.

Sinflar bo'yicha #

18-bo'limdagi qoida: umumiy songa ishonmang.

Natija
Sinflar bo'yicha:
  Futbolka    88.5%
  Shim        98.7%
  Sviter      91.7%
  Libos       92.8%
  Palto       90.5%
  Sandal      97.8%
  Ko'ylak     73.9%
  Krossovka   98.0%
  Sumka       99.0%
  Botinka     96.3%

Eng qiyin sinf yana Ko'ylak (73.9%), eng osoni Sumka (99.0%). Farq hamon katta - 25.1 foiz.

Ya'ni model umumiy jihatdan yaxshilandi, lekin muammoning tabiati o'zgarmadi: ustki kiyimlar 28x28 kulrang rasmda haqiqatan ham bir-biriga o'xshaydi. Buni hal qilish uchun boshqa ma'lumot kerak - kattaroq yoki rangli rasmlar.

Nimani o'zgartirdik #

O'zgarishManba
Ikki konvolyutsiyali blok13-bo'lim
BatchNorm2dshu bo'lim
Dropout(0.3)12-bo'lim
AdamW + weight_decay6 va 12-bo'limlar
Kengaytirish15-bo'lim
CosineAnnealingLRshu bo'lim
Erta to'xtatish12-bo'lim
Nazorat nuqtasi10-bo'lim
Keyin nima qilish kerak

Bu darslik PyTorch ning asoslarini qamrab oldi. Keyingi yo'nalishlar:

  • Rasm: segmentatsiya (U-Net), obyekt aniqlash (YOLO, Faster R-CNN).
  • Matn: transformers kutubxonasi, tayyor modelni sozlash.
  • Vositalar: PyTorch Lightning (qoliplarni kamaytiradi), Weights & Biases (tajribalarni kuzatish).
  • Tezlik: torch.compile, mixed precision (torch.autocast), ko'p GPU.

Eng foydalisi esa - o'zingiz qiziqqan masalada kichik loyiha qilish. Tayyor to'plamda emas, o'zingiz yig'gan ma'lumotda.

Yakuniy topshiriq
  1. Yuqoridagi quvurni to'liq ishga tushiring va test aniqligini yozing.
  2. Natijangizni 14-bo'limdagi 89.34% bilan solishtiring.
  3. BatchNorm2d ni olib tashlab, farqni o'lchang.
  4. Dropout ni 0.3 dan 0.5 ga oshirib, natijani kuzating.
  5. CosineAnnealingLR ni o'chirib, doimiy lr bilan sinang.
  6. Uchinchi konvolyutsion blok qo'shib, parametrlar soni va aniqlikni yozing.
  7. Kengaytirishni o'chirib, validatsiya yo'qotishi qanday o'zganini ko'ring.
  8. Sinflar bo'yicha aniqlikni chiqarib, eng yomon sinfni toping.
  9. Nazorat faylini yuklab, 19-bo'limdagi bashorat funksiyasi bilan ishlating.
  10. O'z ma'lumotingizda (yoki boshqa torchvision to'plamida) shu quvurni takrorlang.

Xulosa #

  • To'liq quvur 92.72% test aniqligini berdi - 14-bo'limdagi 89.34% ga nisbatan +3.38 foiz.
  • Ma'lumot uchga bo'linadi: o'quv, validatsiya va test.
  • Sozlamalar validatsiyada tanlanadi, test esa oxirida bir marta ishlatiladi.
  • Kengaytirish faqat o'quv to'plamiga qo'llanadi - shuning uchun ma'lumot ikki xil transform bilan ochiladi.
  • BatchNorm2d tarmoq ichida normallaydi: o'qitish barqarorroq kechadi.
  • U ham Dropout kabi train() va eval() rejimlarida boshqacha ishlaydi.
  • AdamW weight decay ni to'g'ri qo'llaydi - katta modellarda standart tanlov.
  • CosineAnnealingLR lr ni asta kamaytiradi: avval tez, oxirida aniq o'rganish.
  • Erta to'xtatish eng yaxshi modelni saqlaydi va keraksiz davrlarni kesadi.
  • Nazorat fayliga model bilan birga Normalize qiymatlari va sinf nomlari yoziladi.
  • Umumiy aniqlik oshsa ham, sinflar orasidagi farq saqlanib qoladi.
  • Eng qiyin sinf hamon Ko'ylak (73.9%) - muammo modelda emas, ma'lumotda.

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.