# %% import pandas import torch import matplotlib.pyplot as plt from torch import nn # %% class Patches_to_Embedding(nn.Module): def __init__(self): super().__init__() self.flatten = nn.Flatten(2,-1) self.conv = nn.Conv2d(in_channels=3,out_channels=384,kernel_size=8,stride=8) self.cls_token = nn.Parameter(torch.randn(1,1,384)) self.position_embedding = nn.Parameter(torch.randn(1,65,384)) def forward(self,x :torch.Tensor): batch_size = x.shape[0] x = self.conv(x) x = self.flatten(x) x = x.transpose(1,2) cls = self.cls_token.expand(batch_size,-1,-1) x = torch.cat((cls,x),dim=1) x = x + self.position_embedding return x # %% class EncoderBlock(nn.Module): def __init__(self,dropout): super().__init__() self.dropout = dropout self.Layer_Norm1 = nn.LayerNorm(384) self.Layer_Norm2 = nn.LayerNorm(384) self.multi_head_attention = nn.MultiheadAttention(384,6,batch_first=True,dropout=dropout) self.MLP = nn.Sequential( nn.Linear(384,1536), nn.GELU(), nn.Dropout(self.dropout), nn.Linear(1536,384), nn.Dropout(self.dropout) ) def forward(self,x): output = self.Layer_Norm1(x) output,_ = self.multi_head_attention(output,output,output) x = x + output output = self.Layer_Norm2(x) output = self.MLP(output) x = output + x return x # %% class Encoder(nn.Module): def __init__(self,dropout,num_layers): super().__init__() self.layers = nn.ModuleList([ EncoderBlock(dropout) for _ in range(num_layers) ]) self.norm = nn.LayerNorm(384) def forward(self,x:torch.Tensor) ->torch.Tensor: for layer in self.layers: x = layer(x) x = self.norm(x) return x # %% class ViT(nn.Module): def __init__(self): super().__init__() self.patch_embedding = Patches_to_Embedding() self.encoder = Encoder(0.2,8) self.head = nn.Linear(384,39) def forward(self,x): x = self.patch_embedding(x) x = self.encoder(x) cls = x[:,0] logits = self.head(cls) return logits # %% from torchvision import datasets, transforms from torch.utils.data import DataLoader, Subset import torch DATA_DIR = "/home/ujwal/Documents/Pytorch/GitHub/Vision_Transformer/Plant_leave_diseases_dataset_without_augmentation" train_transform = transforms.Compose([ transforms.RandomResizedCrop( 64, scale=(0.8, 1.0), ratio=(0.9, 1.1) ), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomVerticalFlip(p=0.5), transforms.RandomRotation(15), transforms.ColorJitter( brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05 ), transforms.ToTensor(), transforms.Normalize( mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225] ), transforms.RandomErasing( p=0.25, scale=(0.02, 0.2), ratio=(0.3, 3.3) ) ]) test_transform = transforms.Compose([ transforms.Resize((64, 64)), transforms.ToTensor(), transforms.Normalize( mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225] ) ]) # %% full_dataset = datasets.ImageFolder(DATA_DIR) train_dataset_full = datasets.ImageFolder( DATA_DIR, transform=train_transform ) test_dataset_full = datasets.ImageFolder( DATA_DIR, transform=test_transform ) generator = torch.Generator().manual_seed(42) train_size = int(0.8 * len(full_dataset)) test_size = len(full_dataset) - train_size train_subset, test_subset = torch.utils.data.random_split( range(len(full_dataset)), [train_size, test_size], generator=generator ) train_indices = train_subset.indices test_indices = test_subset.indices train_dataset = Subset( train_dataset_full, train_indices ) test_dataset = Subset( test_dataset_full, test_indices ) print("Train:", len(train_dataset)) print("Test:", len(test_dataset)) # %% device = torch.device("cuda" if torch.cuda.is_available() else "cpu") device # %% import os model = ViT().to(device) model = torch.compile(model) checkpoint_path = "checkpoint.pth" if os.path.exists(checkpoint_path): checkpoint = torch.load( checkpoint_path, map_location=device ) model.load_state_dict(checkpoint["model"]) best_acc = checkpoint["best_acc"] start_epoch = checkpoint["epoch"] + 1 print(f"Resuming from epoch {start_epoch}") print(f"Best validation accuracy: {best_acc:.2f}%") else: best_acc = 0.0 start_epoch = 0 print("No checkpoint found.") print("Starting training from scratch.") # %% criterion = nn.CrossEntropyLoss( label_smoothing=0.1 ) optimizer = torch.optim.AdamW( model.parameters(), lr=1e-4, weight_decay=0.03 ) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=50 ) # %% class AugmentedDataset(torch.utils.data.Dataset): def __init__(self, dataset): self.dataset = dataset def __len__(self): return len(self.dataset) * 4 def __getitem__(self, index): original_index = index // 4 version = index % 4 image, label = self.dataset[original_index] if version == 0: return image, label elif version == 1: return torch.flip(image, dims=[2]), label elif version == 2: return torch.flip(image, dims=[1]), label else: return torch.flip(image, dims=[1, 2]), label # %% train_dataset = AugmentedDataset(train_dataset) # %% print(len(train_dataset)) # %% from torch.utils.data import DataLoader train_loader = DataLoader( train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True, persistent_workers=True, prefetch_factor=2 ) test_loader = DataLoader( test_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True, persistent_workers=True, prefetch_factor=2 ) # %% from tqdm.auto import tqdm scaler = torch.amp.GradScaler("cuda") total_epochs = 50 for epoch in range(start_epoch, total_epochs): model.train() running_loss = 0 correct = 0 total = 0 progress_bar = tqdm( train_loader, desc=f"Epoch [{epoch+1}/{total_epochs}]", leave=True ) for images, labels in progress_bar: images = images.to(device, non_blocking=True) labels = labels.to(device, non_blocking=True) optimizer.zero_grad() with torch.amp.autocast(device_type="cuda"): logits = model(images) loss = criterion(logits, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() running_loss += loss.item() predictions = logits.argmax(dim=1) correct += (predictions == labels).sum().item() total += labels.size(0) progress_bar.set_postfix( loss=running_loss / len(progress_bar), accuracy=100 * correct / total ) train_loss = running_loss / len(train_loader) train_accuracy = 100 * correct / total model.eval() test_loss = 0 correct = 0 total = 0 with torch.no_grad(): progress_bar = tqdm( test_loader, desc="Testing", leave=False ) for images, labels in progress_bar: images = images.to(device, non_blocking=True) labels = labels.to(device, non_blocking=True) with torch.amp.autocast(device_type="cuda"): logits = model(images) loss = criterion(logits, labels) test_loss += loss.item() predictions = logits.argmax(dim=1) correct += (predictions == labels).sum().item() total += labels.size(0) test_loss /= len(test_loader) test_accuracy = 100 * correct / total print( f"Epoch {epoch+1}: " f"Train Loss={train_loss:.4f}, " f"Train Accuracy={train_accuracy:.2f}%, " f"Test Loss={test_loss:.4f}, " f"Test Accuracy={test_accuracy:.2f}%" ) scheduler.step() if test_accuracy > best_acc: best_acc = test_accuracy torch.save({ "epoch": epoch, "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "best_acc": best_acc }, "checkpoint.pth") # %% test_loader = DataLoader( test_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True ) # %% import random import matplotlib.pyplot as plt def predict_test_image(index=None): if index is None: index = random.randrange(len(test_dataset)) image_tensor, actual_label = test_dataset[index] original_image, _ = full_dataset[test_indices[index]] image_input = image_tensor.unsqueeze(0).to(device) model.eval() with torch.no_grad(): logits = model(image_input) probabilities = torch.softmax(logits, dim=1) predicted_label = logits.argmax(dim=1).item() confidence = probabilities[0, predicted_label].item() actual_class = full_dataset.classes[actual_label] predicted_class = full_dataset.classes[predicted_label] # Show image plt.figure(figsize=(6, 6)) plt.imshow(original_image) plt.axis("off") plt.title( f"Actual: {actual_class}\n" f"Predicted: {predicted_class}\n" f"Confidence: {confidence * 100:.2f}%" ) plt.show() print("Test index:", index) print("Actual:", actual_class) print("Predicted:", predicted_class) print(f"Confidence: {confidence * 100:.2f}%") # %% predict_test_image() # %% predict_test_image(108) # %%