-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
101 lines (78 loc) · 2.6 KB
/
Copy pathtrain.py
File metadata and controls
101 lines (78 loc) · 2.6 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
import torch
import torch.nn as nn
import torch.optim as optim
from tqdm import tqdm
import torchvision.transforms as Transforms
from loss import bce_dice_loss
from metric import iou
from dataset import create_io_pairs, Brain_Segmentation_Dataset
from model import Unet
from utils import load_checkpoint, save_checkpoint, get_loaders
# Config params
LEARNING_RATE = 1e-4
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
BATCH_SIZE = 6
NUM_EPOCHS = 20
NUM_WORKERS = 2
IMAGE_HEIGHT = 256
IMAGE_WIDTH = 256
LOAD_MODEL = False
DATA_DIR = "D:/U-Net/unet/data/kaggle_3m"
def train(loader, val_loader, model, optimizer, loss_fn, metric):
loop = tqdm(loader)
model.train()
loss, score = run_model(loop, model, optimizer, loss_fn, metric)
model.eval()
loop = tqdm(val_loader)
val_loss, val_score = run_model(loop, model, optimizer, loss_fn, metric)
return loss, score, val_loss, val_score
def run_model(loop, model, optimizer, loss_fn, metric):
losses = []
scores = []
for batch_idx, (data, target) in enumerate(loop):
data = data.to(device=DEVICE)
target = target.float().to(device=DEVICE)
predictions = model(data)
loss = loss_fn(predictions, target)
losses.append(loss)
optimizer.zero_grad()
loss.backward()
optimizer.step()
score = metric(predictions.detach(), target)
scores.append(score)
loss = sum(losses) / len(losses)
scores = sum(scores) / len(scores)
return loss, score
def main():
transforms = Transforms.Compose(
[
Transforms.Normalize(
mean=[0.0, 0.0, 0.0],
std=[1.0, 1.0, 1.0]
)
],
)
model = Unet(in_channels=3, out_channels=1).to(DEVICE)
print('Model Initialized')
loss_fn = bce_dice_loss
optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)
metric = iou
train_loader, val_loader = get_loaders(
DATA_DIR,
BATCH_SIZE,
transforms
)
if LOAD_MODEL:
load_checkpoint(torch.load("my_checkpoint.pth.tar"), model)
print('Starting training')
for epoch in range(NUM_EPOCHS):
loss, score, val_loss, val_score = train(train_loader, val_loader, model, optimizer, loss_fn, metric)
# save model
checkpoint = {
"state_dict": model.state_dict(),
"optimizer":optimizer.state_dict(),
}
save_checkpoint(checkpoint)
print(f'Epoch: {epoch}, Loss: {loss}, score: {score}, Val_loss: {val_loss}, Val_Score: {val_score}')
if __name__ == "__main__":
main()