CoolFace
Apppublic

kingav/bts

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
train.py83 linesDownload Raw Back to root
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)}')