Tworzenie i trenowanie sieci CNN w PyTorch
nn.Conv2d, nn.MaxPool2d, nn.Linear, pętla trenowania, optymalizator, funkcja straty, śledzenie dokładności.
Tworzenie i trenowanie sieci CNN w PyTorch to bezpłatna lekcja Learn AI with Python na CoddyKit. To lekcja 3 z 4. Możesz przeczytać całą lekcję poniżej za darmo — a potem ćwiczyć ją interaktywnie w przeglądarce z wbudowanym edytorem kodu i tutorem AI dostępnym 24/7. To część ścieżki edukacyjnej Learn AI with Python, a Twój postęp synchronizuje się między webem a aplikacją CoddyKit. Kurs Learn AI with Python zawiera 4 lekcji w sumie.
Dlaczego CNN do obrazów?
Konwolucyjne sieci neuronowe uczą się wzorców przestrzennych: krawędzi, tekstur, a następnie kształtów i obiektów. Operacje konwolucji współdzielą wagi w obrębie obrazu, dzięki czemu sieci CNN są wydajne i uwzględniają przesunięcia. Stanowią podstawę widzenia komputerowego.
import torch
import torch.nn as nnKlasa bazowa nn.Module
Modele dziedziczą po nn.Module. Warstwy definiuje się w __init__, a przepływ danych w forward. PyTorch automatycznie śledzi parametry i gradienty.
class CNN(nn.Module):
def __init__(self):
super().__init__()Warstwy konwolucyjne
nn.Conv2d(in_channels, out_channels, kernel_size) przesuwa wyuczalne filtry po obrazie, aby utworzyć mapy cech. Pierwsza warstwa konwolucyjna przyjmuje 3 kanały (RGB) i generuje więcej kanałów, które wychwytują różne cechy.
self.conv1 = nn.Conv2d(3, 16, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1)Warstwy poolingowe
nn.MaxPool2d zmniejsza mapy cech, zachowując maksimum w każdym oknie. Zmniejsza to rozmiar przestrzenny, ogranicza liczbę obliczeń i zapewnia pewną odporność na przesunięcia.
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)Warstwy w pełni połączone
Po wykonaniu konwolucji należy spłaszczyć cechy i przekazać je przez warstwy nn.Linear, aby uzyskać wyniki dla klas. Końcowa warstwa Linear zwraca jedną wartość dla każdej klasy.
self.fc1 = nn.Linear(32 * 8 * 8, 128)
self.fc2 = nn.Linear(128, 10) # 10 classesMetoda forward
Metoda forward definiuje przepływ danych: od conv przez ReLU do pool, powtarzając ten schemat, a następnie spłaszczając dane i przekazując je do warstw liniowych. F.relu dodaje nieliniowość, która pozwala sieci uczyć się złożonych wzorców.
import torch.nn.functional as F
def forward(self, x):
x = self.pool(F.relu(self.conv1(x)))
x = self.pool(F.relu(self.conv2(x)))
x = x.view(x.size(0), -1) # flatten
x = F.relu(self.fc1(x))
return self.fc2(x)Funkcja straty i optymalizator
Do klasyfikacji należy użyć CrossEntropyLoss. Optymalizator taki jak Adam aktualizuje wagi za pomocą gradientów. Należy przekazać mu parametry modelu i współczynnik uczenia.
model = CNN()
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)Pętla treningowa: zerowanie gradientów
Każdy krok należy rozpocząć od wyczyszczenia starych gradientów za pomocą optimizer.zero_grad(). Pominięcie tego kroku powoduje akumulowanie gradientów między partiami i prowadzi do nieprawidłowych aktualizacji.
for images, labels in loader:
optimizer.zero_grad()Przejście w przód i funkcja straty
Należy przekazać partię przez model, aby uzyskać predykcje, a następnie obliczyć funkcję straty, porównując predykcje z prawdziwymi etykietami. Strata to pojedyncza liczba określająca, jak bardzo model się myli.
outputs = model(images)
loss = criterion(outputs, labels)Propagacja wsteczna i aktualizacja
loss.backward() oblicza gradienty za pomocą autograd, a optimizer.step() nieznacznie modyfikuje wagi, aby zmniejszyć stratę. Razem stanowią jeden krok uczenia.
loss.backward()
optimizer.step()Śledzenie dokładności
Postęp należy monitorować, zliczając poprawne predykcje. Należy zastosować argmax do wyników, aby uzyskać przewidywane klasy, porównać je z etykietami i podzielić liczbę trafień przez ich łączną liczbę.
preds = outputs.argmax(dim=1)
correct = (preds == labels).sum().item()
acc = correct / labels.size(0)
print("batch acc:", acc)Szybkie sprawdzenie
Proszę sprawdzić znajomość pętli treningowej.
Podsumowanie: budowanie i trenowanie sieci CNN
Zbudowano sieć CNN przez utworzenie podklasy nn.Module z warstwami nn.Conv2d, nn.MaxPool2d i nn.Linear oraz metodą forward. Wytrenowano ją za pomocą pętli obejmującej: zero_grad, przejście w przód, loss.backward(), optimizer.step(), a dokładność śledzono za pomocą argmax.
Często zadawane pytania
Czy lekcja „Tworzenie i trenowanie sieci CNN w PyTorch” jest bezpłatna?
Tak — pełny tekst „Tworzenie i trenowanie sieci CNN w PyTorch” jest dostępny za darmo tutaj w sieci. Aby ćwiczyć ją interaktywnie (wbudowany edytor kodu i tutor AI dostępny 24/7) i odblokować resztę kursu Learn AI with Python, przejdź na CoddyKit PRO. Kurs Learn AI with Python zawiera 4 lekcji w sumie.
Co nauczysz się w „Tworzenie i trenowanie sieci CNN w PyTorch”?
nn.Conv2d, nn.MaxPool2d, nn.Linear, pętla trenowania, optymalizator, funkcja straty, śledzenie dokładności. Ćwiczysz Learn AI with Python z praktycznym kodem, który uruchamiasz bezpośrednio w przeglądarce, a tutor AI dostępny 24/7 odpowiada na Twoje pytania podczas pracy nad lekcją.
Czy potrzebuję doświadczenia, aby zacząć Learn AI with Python?
Nie wymagamy żadnego doświadczenia. Learn AI with Python w CoddyKit jest strukturyzowany dla początkujących i zaawansowanych użytkowników, więc możesz zacząć tutaj lub od początku i uczyć się w swoim tempie. To lekcja 3 z 4.
Ile czasu zajmuje lekcja „Tworzenie i trenowanie sieci CNN w PyTorch”?
Większość lekcji CoddyKit trwa około 5–10 minut. Każda lekcja to mały, interaktywny krok, dzięki czemu robisz systematyczne postępy i zawsze wracasz dokładnie do tego samego miejsca — na webie i w aplikacji.
Czy mogę pisać i uruchamiać kod w tej lekcji Learn AI with Python?
Tak. Każda lekcja Learn AI with Python zawiera wbudowany edytor kodu, więc piszesz i uruchamiasz prawdziwy kod bezpośrednio w przeglądarce i od razu otrzymujesz sprzężenie zwrotne od AI — bez konfiguracji na komputerze.
Wszystkie lekcje w tym kursie
- Tensory PyTorch i Autograd
- Własne zbiory danych i moduły DataLoader
- Tworzenie i trenowanie sieci CNN w PyTorch
- Wykrywanie obiektów za pomocą YOLOv8