kingav/bts
0
1import tensorflow as tf
2import segmentation_models_3D as sm
3from tqdm import tqdm
4import numpy as np
5from utils import dice_coef, precision, sensitivity, specificity, sum_scaled_weights, similarity_loss
6
7class Trainer:
8 def __init__(self, model, optimizer, loss_fn, metrics, batch_size, epochs, num_clients, num_comm):
9 self.model = model
10 self.optimizer = optimizer
11 self.loss_fn = loss_fn
12 self.metrics = metrics
13 self.batch_size = batch_size
14 self.epochs = epochs
15 self.num_clients = num_clients
16 self.num_comm = num_comm
17
18 def train(self, clients, final_model, test_data_gen):
19 dice_coee = []
20 sim_loss = []
21 cl_pe_val = []
22 dc_pe_val = []
23
24 for t in range(self.num_comm):
25 print(f'\nCommunication round: {t+1}\n')
26 local_weight_list = list()
27 for epoch in range(self.epochs):
28 print(f'\nEpoch {epoch+1}')
29 client_loss = []
30 for i, client in enumerate(clients):
31 print(f'Training Client {i+1}')
32 model = client['model']
33 optimizer = client['optimizer']
34
35 for batch_idx, (image, target) in enumerate(tqdm(client['original_data'])):
36 mf, _ = final_model['model'](image)
37 with tf.GradientTape() as tape:
38 m, output = model(image)
39 loss1 = self.loss_fn(target, output)
40 dc = dice_coef(target, output)
41 client_loss.append(loss1)
42 dice_coee.append(dc)
43
44 grads = tape.gradient(loss1, model.trainable_weights)
45 optimizer.apply_gradients(zip(grads, model.trainable_variables))
46
47 dc_val_pe = sum(dice_coee) / len(dice_coee)
48 cl_val_pe = sum(sim_loss) / len(sim_loss)
49 print(f'epoch dice coeff mean: {dc_val_pe}')
50 print(f'epoch contrastive loss mean: {cl_val_pe}')
51 dc_pe_val.append(dc_val_pe)
52 cl_pe_val.append(cl_val_pe)
53 dice_coee.clear()
54 sim_loss.clear()
55
56 client['model'] = model
57 local_weight_list.append(client['model'].get_weights())
58
59 average_weight = sum_scaled_weights(local_weight_list, client["length_ratio"])
60 final_model['model'].set_weights(average_weight)
61 local_weight_list.clear()
62
63 self.evaluate(final_model['model'], test_data_gen)
64
65 def evaluate(self, model, test_data_gen):
66 dice = []
67 pre = []
68 batch_loss = []
69 se = []
70 spe = []
71 io = []
72
73 for batch_idx, (images, masks) in enumerate(tqdm(test_data_gen)):
74 _, logits = model(images)
75 loss = self.loss_fn(masks, logits)
76 batch_loss.append(loss)
77 dice.append(dice_coef(masks, logits))
78 pre.append(precision(masks, logits))
79 se.append(sensitivity(masks, logits))
80 spe.append(specificity(masks, logits))
81 io.append(sm.metrics.IOUScore(threshold=0.5)(masks, logits))
82
83 print(f'Test results: Loss: {np.mean(batch_loss)}, Dice Coeff: {np.mean(dice)}, Precision: {np.mean(pre)}, Sensitivity: {np.mean(se)}, Specificity: {np.mean(spe)}, IOU: {np.mean(io)}')