-
Notifications
You must be signed in to change notification settings - Fork 5
Expand file tree
/
Copy pathtrain_unet_segmentation.py
More file actions
executable file
·64 lines (54 loc) · 2.63 KB
/
Copy pathtrain_unet_segmentation.py
File metadata and controls
executable file
·64 lines (54 loc) · 2.63 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
import torch
import datetime
from learner.UnetSegmentationLearner import UnetSegmentationLearner
from common.model.Unet3D import Unet3D
from common import data, util, metrics
def train():
args = util.get_args_unet_training()
# Params / Config
batchsize = 6 # 17 training, 6 validation
learning_rate = 1e-3
momentums_cae = (0.99, 0.999)
criterion = metrics.BatchDiceLoss([1.0]) # nn.BCELoss()
path_saved_model = args.unetpath
channels = args.channels
pad = args.padding
cuda = True
# Unet model
unet = Unet3D(channels)
if cuda:
unet = unet.cuda()
# Model params
params = [p for p in unet.parameters() if p.requires_grad]
print('# optimizing params', sum([p.nelement() * p.requires_grad for p in params]),
'/ total: unet', sum([p.nelement() for p in unet.parameters()]))
# Optimizer with scheduler
optimizer = torch.optim.Adam(params, lr=learning_rate, weight_decay=1e-5, betas=momentums_cae)
if args.lrsteps:
scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, args.lrsteps)
else:
scheduler = None
# Data
train_transform = [data.ResamplePlaneXY(args.xyresample),
data.HemisphericFlipFixedToCaseId(split_id=args.hemisflipid),
data.PadImages(pad[0], pad[1], pad[2], pad_value=0),
data.RandomPatch(104, 104, 68, pad[0], pad[1], pad[2]),
data.ToTensor()]
valid_transform = [data.ResamplePlaneXY(args.xyresample),
data.HemisphericFlipFixedToCaseId(split_id=args.hemisflipid),
data.PadImages(pad[0], pad[1], pad[2], pad_value=0),
data.RandomPatch(104, 104, 68, pad[0], pad[1], pad[2]),
data.ToTensor()]
ds_train, ds_valid = data.get_stroke_shape_training_data(train_transform, valid_transform, args.fold,
args.validsetsize, batchsize=batchsize)
print('Size training set:', len(ds_train.sampler.indices), 'samples | Size validation set:', len(ds_valid.sampler.indices),
'samples | Capacity batch:', batchsize, 'samples')
print('# training batches:', len(ds_train), '| # validation batches:', len(ds_valid))
# Training
learner = UnetSegmentationLearner(ds_train, ds_valid, unet, path_saved_model, optimizer, scheduler, criterion,
path_previous_base=args.inbasepath, path_outputs_base=args.outbasepath)
learner.run_training()
if __name__ == '__main__':
print(datetime.datetime.now())
train()
print(datetime.datetime.now())