Final Performance
Final Performance
import torch
import [Link] as nn
import timm
from torchvision import datasets, transforms
from [Link] import DataLoader
from collections import Counter
# =========================
# SETTINGS
# =========================
BATCH_SIZE = 16
EPOCHS = 30
LR = 5e-5 # safe increase from your working version
# =========================
# TRANSFORMS (SAFE VERSION)
# =========================
train_transform = [Link]([
[Link]((256, 256)),
[Link](224, scale=(0.9, 1.0)),
[Link](p=0.5),
[Link](8),
[Link](
brightness=0.15,
contrast=0.15,
saturation=0.15
),
[Link](),
[Link](
[0.485, 0.456, 0.406],
[0.229, 0.224, 0.225]
)
])
val_transform = [Link]([
[Link]((224, 224)),
[Link](),
[Link](
[0.485, 0.456, 0.406],
[0.229, 0.224, 0.225]
)
])
# =========================
# DATASETS
# =========================
train_data = [Link](
f"{SPLIT_DIR}/train",
transform=train_transform
)
val_data = [Link](
f"{SPLIT_DIR}/val",
transform=val_transform
)
test_data = [Link](
f"{SPLIT_DIR}/test",
transform=val_transform
)
print("Classes:", train_data.class_to_idx)
# =========================
# CLASS WEIGHTS
# =========================
class_counts = Counter(train_data.targets)
class_counts = [Link](list(class_counts.values()), dtype=[Link])
# =========================
# DATALOADERS
# =========================
train_loader = DataLoader(
train_data,
batch_size=BATCH_SIZE,
shuffle=True,
num_workers=2,
pin_memory=True
)
val_loader = DataLoader(
val_data,
batch_size=BATCH_SIZE,
shuffle=False,
num_workers=2,
pin_memory=True
)
test_loader = DataLoader(
test_data,
batch_size=BATCH_SIZE,
shuffle=False,
num_workers=2,
pin_memory=True
)
# =========================
# MODEL (STABLE)
# =========================
model = timm.create_model(
"efficientnet_b0", # back to stable version
pretrained=True,
num_classes=len(train_data.classes),
drop_rate=0.3
).to(device)
# =========================
# LOSS (NO LABEL SMOOTHING)
# =========================
criterion = [Link](weight=class_weights)
# =========================
# OPTIMIZER
# =========================
optimizer = [Link](
[Link](),
lr=LR,
weight_decay=1e-4
)
# =========================
# SCHEDULER
# =========================
scheduler = [Link].lr_scheduler.ReduceLROnPlateau(
optimizer,
mode="min",
factor=0.5,
patience=2
)
# =========================
# MIXED PRECISION (SAFE)
# =========================
scaler = [Link]()
# =========================
# EARLY STOPPING
# =========================
best_val_acc = 0
patience = 5
counter = 0
# =========================
# TRAINING LOOP
# =========================
for epoch in range(EPOCHS):
optimizer.zero_grad()
with [Link]():
outputs = model(images)
loss = criterion(outputs, labels)
[Link](loss).backward()
[Link](optimizer)
[Link]()
train_loss += [Link]()
_, preds = [Link](outputs, 1)
with torch.no_grad():
for images, labels in val_loader:
images, labels = [Link](device), [Link](device)
with [Link]():
outputs = model(images)
loss = criterion(outputs, labels)
val_loss += [Link]()
_, preds = [Link](outputs, 1)
[Link](val_loss)
# =========================
# TESTING
# =========================
print("\n🚀 Evaluating on TEST set...")
model.load_state_dict([Link]("best_model.pth"))
[Link]()
test_loss = 0
test_correct = 0
test_total = 0
with torch.no_grad():
for images, labels in test_loader:
images, labels = [Link](device), [Link](device)
with [Link]():
outputs = model(images)
loss = criterion(outputs, labels)
test_loss += [Link]()
_, preds = [Link](outputs, 1)