kiu098 commited on
Commit
b75531f
·
verified ·
1 Parent(s): be0af1b

Upload Untitled-1.py

Browse files
Files changed (1) hide show
  1. Untitled-1.py +450 -0
Untitled-1.py ADDED
@@ -0,0 +1,450 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # %%
2
+ import pandas
3
+ import torch
4
+ import matplotlib.pyplot as plt
5
+ from torch import nn
6
+
7
+ # %%
8
+ class Patches_to_Embedding(nn.Module):
9
+ def __init__(self):
10
+ super().__init__()
11
+ self.flatten = nn.Flatten(2,-1)
12
+ self.conv = nn.Conv2d(in_channels=3,out_channels=384,kernel_size=8,stride=8)
13
+ self.cls_token = nn.Parameter(torch.randn(1,1,384))
14
+ self.position_embedding = nn.Parameter(torch.randn(1,65,384))
15
+
16
+ def forward(self,x :torch.Tensor):
17
+ batch_size = x.shape[0]
18
+ x = self.conv(x)
19
+ x = self.flatten(x)
20
+ x = x.transpose(1,2)
21
+ cls = self.cls_token.expand(batch_size,-1,-1)
22
+ x = torch.cat((cls,x),dim=1)
23
+ x = x + self.position_embedding
24
+ return x
25
+
26
+ # %%
27
+ class EncoderBlock(nn.Module):
28
+ def __init__(self,dropout):
29
+ super().__init__()
30
+ self.dropout = dropout
31
+ self.Layer_Norm1 = nn.LayerNorm(384)
32
+ self.Layer_Norm2 = nn.LayerNorm(384)
33
+ self.multi_head_attention = nn.MultiheadAttention(384,6,batch_first=True,dropout=dropout)
34
+ self.MLP = nn.Sequential(
35
+ nn.Linear(384,1536),
36
+ nn.GELU(),
37
+ nn.Dropout(self.dropout),
38
+ nn.Linear(1536,384),
39
+ nn.Dropout(self.dropout)
40
+ )
41
+
42
+ def forward(self,x):
43
+ output = self.Layer_Norm1(x)
44
+ output,_ = self.multi_head_attention(output,output,output)
45
+ x = x + output
46
+ output = self.Layer_Norm2(x)
47
+ output = self.MLP(output)
48
+ x = output + x
49
+ return x
50
+
51
+ # %%
52
+ class Encoder(nn.Module):
53
+ def __init__(self,dropout,num_layers):
54
+ super().__init__()
55
+ self.layers = nn.ModuleList([
56
+ EncoderBlock(dropout)
57
+ for _ in range(num_layers)
58
+ ])
59
+ self.norm = nn.LayerNorm(384)
60
+
61
+ def forward(self,x:torch.Tensor) ->torch.Tensor:
62
+ for layer in self.layers:
63
+ x = layer(x)
64
+ x = self.norm(x)
65
+ return x
66
+
67
+ # %%
68
+ class ViT(nn.Module):
69
+ def __init__(self):
70
+ super().__init__()
71
+ self.patch_embedding = Patches_to_Embedding()
72
+ self.encoder = Encoder(0.2,8)
73
+ self.head = nn.Linear(384,39)
74
+
75
+ def forward(self,x):
76
+ x = self.patch_embedding(x)
77
+ x = self.encoder(x)
78
+ cls = x[:,0]
79
+ logits = self.head(cls)
80
+ return logits
81
+
82
+ # %%
83
+ from torchvision import datasets, transforms
84
+ from torch.utils.data import DataLoader, Subset
85
+ import torch
86
+
87
+ DATA_DIR = "/home/ujwal/Documents/Pytorch/GitHub/Vision_Transformer/Plant_leave_diseases_dataset_without_augmentation"
88
+
89
+ train_transform = transforms.Compose([
90
+ transforms.RandomResizedCrop(
91
+ 64,
92
+ scale=(0.8, 1.0),
93
+ ratio=(0.9, 1.1)
94
+ ),
95
+
96
+ transforms.RandomHorizontalFlip(p=0.5),
97
+ transforms.RandomVerticalFlip(p=0.5),
98
+
99
+ transforms.RandomRotation(15),
100
+
101
+ transforms.ColorJitter(
102
+ brightness=0.2,
103
+ contrast=0.2,
104
+ saturation=0.2,
105
+ hue=0.05
106
+ ),
107
+
108
+ transforms.ToTensor(),
109
+
110
+ transforms.Normalize(
111
+ mean=[0.485, 0.456, 0.406],
112
+ std=[0.229, 0.224, 0.225]
113
+ ),
114
+
115
+ transforms.RandomErasing(
116
+ p=0.25,
117
+ scale=(0.02, 0.2),
118
+ ratio=(0.3, 3.3)
119
+ )
120
+ ])
121
+
122
+ test_transform = transforms.Compose([
123
+ transforms.Resize((64, 64)),
124
+
125
+ transforms.ToTensor(),
126
+
127
+ transforms.Normalize(
128
+ mean=[0.485, 0.456, 0.406],
129
+ std=[0.229, 0.224, 0.225]
130
+ )
131
+ ])
132
+
133
+ # %%
134
+ full_dataset = datasets.ImageFolder(DATA_DIR)
135
+
136
+ train_dataset_full = datasets.ImageFolder(
137
+ DATA_DIR,
138
+ transform=train_transform
139
+ )
140
+
141
+ test_dataset_full = datasets.ImageFolder(
142
+ DATA_DIR,
143
+ transform=test_transform
144
+ )
145
+
146
+ generator = torch.Generator().manual_seed(42)
147
+
148
+ train_size = int(0.8 * len(full_dataset))
149
+ test_size = len(full_dataset) - train_size
150
+
151
+ train_subset, test_subset = torch.utils.data.random_split(
152
+ range(len(full_dataset)),
153
+ [train_size, test_size],
154
+ generator=generator
155
+ )
156
+
157
+ train_indices = train_subset.indices
158
+ test_indices = test_subset.indices
159
+
160
+ train_dataset = Subset(
161
+ train_dataset_full,
162
+ train_indices
163
+ )
164
+
165
+ test_dataset = Subset(
166
+ test_dataset_full,
167
+ test_indices
168
+ )
169
+
170
+ print("Train:", len(train_dataset))
171
+ print("Test:", len(test_dataset))
172
+
173
+ # %%
174
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
175
+ device
176
+
177
+ # %%
178
+ import os
179
+
180
+ model = ViT().to(device)
181
+ model = torch.compile(model)
182
+
183
+ checkpoint_path = "checkpoint.pth"
184
+
185
+ if os.path.exists(checkpoint_path):
186
+
187
+ checkpoint = torch.load(
188
+ checkpoint_path,
189
+ map_location=device
190
+ )
191
+
192
+ model.load_state_dict(checkpoint["model"])
193
+
194
+ best_acc = checkpoint["best_acc"]
195
+ start_epoch = checkpoint["epoch"] + 1
196
+
197
+ print(f"Resuming from epoch {start_epoch}")
198
+ print(f"Best validation accuracy: {best_acc:.2f}%")
199
+
200
+ else:
201
+
202
+ best_acc = 0.0
203
+ start_epoch = 0
204
+
205
+ print("No checkpoint found.")
206
+ print("Starting training from scratch.")
207
+
208
+ # %%
209
+ criterion = nn.CrossEntropyLoss(
210
+ label_smoothing=0.1
211
+ )
212
+
213
+ optimizer = torch.optim.AdamW(
214
+ model.parameters(),
215
+ lr=1e-4,
216
+ weight_decay=0.03
217
+ )
218
+
219
+ scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
220
+ optimizer,
221
+ T_max=50
222
+ )
223
+
224
+ # %%
225
+ class AugmentedDataset(torch.utils.data.Dataset):
226
+
227
+ def __init__(self, dataset):
228
+ self.dataset = dataset
229
+
230
+ def __len__(self):
231
+ return len(self.dataset) * 4
232
+
233
+ def __getitem__(self, index):
234
+
235
+ original_index = index // 4
236
+ version = index % 4
237
+
238
+ image, label = self.dataset[original_index]
239
+
240
+ if version == 0:
241
+ return image, label
242
+
243
+ elif version == 1:
244
+ return torch.flip(image, dims=[2]), label
245
+
246
+ elif version == 2:
247
+ return torch.flip(image, dims=[1]), label
248
+
249
+ else:
250
+ return torch.flip(image, dims=[1, 2]), label
251
+
252
+ # %%
253
+ train_dataset = AugmentedDataset(train_dataset)
254
+
255
+ # %%
256
+ print(len(train_dataset))
257
+
258
+ # %%
259
+ from torch.utils.data import DataLoader
260
+
261
+ train_loader = DataLoader(
262
+ train_dataset,
263
+ batch_size=32,
264
+ shuffle=True,
265
+ num_workers=4,
266
+ pin_memory=True,
267
+ persistent_workers=True,
268
+ prefetch_factor=2
269
+ )
270
+
271
+ test_loader = DataLoader(
272
+ test_dataset,
273
+ batch_size=32,
274
+ shuffle=False,
275
+ num_workers=4,
276
+ pin_memory=True,
277
+ persistent_workers=True,
278
+ prefetch_factor=2
279
+ )
280
+
281
+ # %%
282
+ from tqdm.auto import tqdm
283
+
284
+ scaler = torch.amp.GradScaler("cuda")
285
+ total_epochs = 50
286
+
287
+ for epoch in range(start_epoch, total_epochs):
288
+
289
+ model.train()
290
+
291
+ running_loss = 0
292
+ correct = 0
293
+ total = 0
294
+
295
+ progress_bar = tqdm(
296
+ train_loader,
297
+ desc=f"Epoch [{epoch+1}/{total_epochs}]",
298
+ leave=True
299
+ )
300
+
301
+ for images, labels in progress_bar:
302
+
303
+ images = images.to(device, non_blocking=True)
304
+ labels = labels.to(device, non_blocking=True)
305
+
306
+ optimizer.zero_grad()
307
+
308
+ with torch.amp.autocast(device_type="cuda"):
309
+ logits = model(images)
310
+ loss = criterion(logits, labels)
311
+
312
+ scaler.scale(loss).backward()
313
+ scaler.step(optimizer)
314
+ scaler.update()
315
+
316
+ running_loss += loss.item()
317
+
318
+ predictions = logits.argmax(dim=1)
319
+
320
+ correct += (predictions == labels).sum().item()
321
+ total += labels.size(0)
322
+
323
+ progress_bar.set_postfix(
324
+ loss=running_loss / len(progress_bar),
325
+ accuracy=100 * correct / total
326
+ )
327
+
328
+ train_loss = running_loss / len(train_loader)
329
+ train_accuracy = 100 * correct / total
330
+
331
+ model.eval()
332
+
333
+ test_loss = 0
334
+ correct = 0
335
+ total = 0
336
+
337
+ with torch.no_grad():
338
+
339
+ progress_bar = tqdm(
340
+ test_loader,
341
+ desc="Testing",
342
+ leave=False
343
+ )
344
+
345
+ for images, labels in progress_bar:
346
+
347
+ images = images.to(device, non_blocking=True)
348
+ labels = labels.to(device, non_blocking=True)
349
+
350
+ with torch.amp.autocast(device_type="cuda"):
351
+ logits = model(images)
352
+ loss = criterion(logits, labels)
353
+
354
+ test_loss += loss.item()
355
+
356
+ predictions = logits.argmax(dim=1)
357
+
358
+ correct += (predictions == labels).sum().item()
359
+ total += labels.size(0)
360
+
361
+ test_loss /= len(test_loader)
362
+ test_accuracy = 100 * correct / total
363
+
364
+ print(
365
+ f"Epoch {epoch+1}: "
366
+ f"Train Loss={train_loss:.4f}, "
367
+ f"Train Accuracy={train_accuracy:.2f}%, "
368
+ f"Test Loss={test_loss:.4f}, "
369
+ f"Test Accuracy={test_accuracy:.2f}%"
370
+ )
371
+
372
+ scheduler.step()
373
+
374
+ if test_accuracy > best_acc:
375
+ best_acc = test_accuracy
376
+
377
+ torch.save({
378
+ "epoch": epoch,
379
+ "model": model.state_dict(),
380
+ "optimizer": optimizer.state_dict(),
381
+ "scheduler": scheduler.state_dict(),
382
+ "best_acc": best_acc
383
+ }, "checkpoint.pth")
384
+
385
+ # %%
386
+ test_loader = DataLoader(
387
+ test_dataset,
388
+ batch_size=32,
389
+ shuffle=False,
390
+ num_workers=4,
391
+ pin_memory=True
392
+ )
393
+
394
+ # %%
395
+ import random
396
+ import matplotlib.pyplot as plt
397
+
398
+ def predict_test_image(index=None):
399
+
400
+ if index is None:
401
+ index = random.randrange(len(test_dataset))
402
+
403
+ image_tensor, actual_label = test_dataset[index]
404
+
405
+ original_image, _ = full_dataset[test_indices[index]]
406
+
407
+ image_input = image_tensor.unsqueeze(0).to(device)
408
+
409
+ model.eval()
410
+
411
+ with torch.no_grad():
412
+ logits = model(image_input)
413
+
414
+ probabilities = torch.softmax(logits, dim=1)
415
+
416
+ predicted_label = logits.argmax(dim=1).item()
417
+
418
+ confidence = probabilities[0, predicted_label].item()
419
+
420
+ actual_class = full_dataset.classes[actual_label]
421
+ predicted_class = full_dataset.classes[predicted_label]
422
+
423
+ # Show image
424
+ plt.figure(figsize=(6, 6))
425
+ plt.imshow(original_image)
426
+ plt.axis("off")
427
+
428
+ plt.title(
429
+ f"Actual: {actual_class}\n"
430
+ f"Predicted: {predicted_class}\n"
431
+ f"Confidence: {confidence * 100:.2f}%"
432
+ )
433
+
434
+ plt.show()
435
+
436
+ print("Test index:", index)
437
+ print("Actual:", actual_class)
438
+ print("Predicted:", predicted_class)
439
+ print(f"Confidence: {confidence * 100:.2f}%")
440
+
441
+ # %%
442
+ predict_test_image()
443
+
444
+ # %%
445
+ predict_test_image(108)
446
+
447
+ # %%
448
+
449
+
450
+