MamaPearl's picture
Update train.py
00538cb verified
Raw
History Blame
4.39 kB
from datasets import load_dataset
import torchvision.transforms as T
from torch.utils.data import Dataset, DataLoader
from PIL import Image
from tqdm.auto import tqdm
dataset = load_dataset("uoft-cs/cifar10")
train_transform = T.Compose([
T.RandomCrop(32, padding=4),
T.RandomHorizontalFlip(p=0.5),
T.AutoAugment(policy=T.AutoAugmentPolicy.CIFAR10),
T.ToTensor(),
T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
])
test_transform = T.Compose([
T.ToTensor(),
T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
])
class CIFAR10Wrapper(Dataset):
def __init__(self, hf_ds, transform):
self.ds = hf_ds
self.transform = transform
def __len__(self):
return len(self.ds)
def __getitem__(self, i):
ex = self.ds[i]
img = ex["img"]
label = ex["label"]
if not isinstance(img, Image.Image):
img = Image.fromarray(img)
x = self.transform(img)
return { "pixel_values": x, "labels": label }
train_ds = CIFAR10Wrapper(dataset["train"], train_transform)
test_ds = CIFAR10Wrapper(dataset["test"], test_transform)
train_loader = DataLoader(
train_ds,
batch_size=128,
shuffle=True,
num_workers=0,
pin_memory=True
)
test_loader = DataLoader(
test_ds,
batch_size=256,
shuffle=False,
num_workers=0,
pin_memory=True
)
batch = next(iter(train_loader))
print(batch["pixel_values"].shape)
print(batch['labels'].shape)
from torch.optim.lr_scheduler import SequentialLR, LinearLR, CosineAnnealingLR
cfg = NulaConfig(
block_channels=(128, 256, 512),
classifier_hidden_dim=512,
use_se=True
)
DEVICE = "cuda" if pt.cuda.is_available() else "cpu"
model = NulaForImageClassification(cfg).to(DEVICE)
optimizer = pt.optim.AdamW(
model.parameters(),
lr=1e-3,
weight_decay=0.01
)
warmup = LinearLR(optimizer, start_factor=0.1, end_factor=1.0, total_iters=5)
cosine = CosineAnnealingLR(optimizer, T_max=45)
scheduler = SequentialLR(optimizer, schedulers=[warmup, cosine], milestones=[5])
def train_one_epoch(model, loader, optimizer, device, grad_clip=1.0):
model.train()
total_loss = 0.0
total_correct = 0
total_examples = 0
pbar = tqdm(loader, desc="training...", leave=False)
for batch in pbar:
x = batch["pixel_values"].to(device, non_blocking=True)
y = batch["labels"].to(device, non_blocking=True)
optimizer.zero_grad(set_to_none=True)
out = model(pixel_values=x, labels=y)
loss = out.loss
logits = out.logits
preds = logits.argmax(dim=1)
loss.backward()
pt.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)
optimizer.step()
total_loss += loss.item() * y.size(0)
total_correct += (preds == y).sum().item()
total_examples += y.size(0)
pbar.set_postfix(loss=f"{loss.item():.4f}", acc=f"{100 * total_correct / total_examples:.2f}%")
return total_loss / total_examples, total_correct / total_examples
@pt.no_grad()
def evaluate(model, loader, device):
model.eval()
total_loss = 0.0
total_correct = 0
total_examples = 0
for batch in loader:
x = batch["pixel_values"].to(device, non_blocking=True)
y = batch["labels"].to(device, non_blocking=True)
out = model(pixel_values=x, labels=y)
loss = out.loss
logits = out.logits
preds = logits.argmax(dim=1)
total_loss += loss.item() * y.size(0)
total_correct += (preds == y).sum().item()
total_examples += y.size(0)
return total_loss / total_examples, total_correct / total_examples
num_epochs = 50
best_val_acc = 0.0
for epoch in range(1, num_epochs + 1):
train_loss, train_acc = train_one_epoch(model, train_loader, optimizer, DEVICE)
val_loss, val_acc = evaluate(model, test_loader, DEVICE)
scheduler.step()
best_val_acc = max(best_val_acc, val_acc)
current_lr = optimizer.param_groups[0]["lr"]
print(f"{'='}{'-'*60}{'='}")
print(f"Epoch [{epoch}/{num_epochs}]")
print(f"learning_rate : {current_lr:.6f}")
print(f"train_loss : {train_loss:.6f}")
print(f"train_acc : {train_acc * 100:.2f}%")
print(f"val_loss : {val_loss:4f}")
print(f"val_acc : {val_acc * 100:.2f}")
print(f"best_val_acc : {100 * best_val_acc:.2f}%")