JeeKay/brain-tumor-segmentation
0
1#!/usr/bin/env python2# coding: utf-83 4# In[ ]:5 6 7import tensorflow as tf8from tensorflow.keras import layers, Model, Input9from tensorflow.keras import backend as K10 11# ------------------ U-Net Blocks ------------------12def conv_block(input_tensor, num_filters):13 x = layers.Conv2D(num_filters, 3, padding='same')(input_tensor)14 x = layers.BatchNormalization()(x)15 x = layers.Activation('relu')(x)16 17 x = layers.Conv2D(num_filters, 3, padding='same')(x)18 x = layers.BatchNormalization()(x)19 x = layers.Activation('relu')(x)20 return x21 22def encoder_block(input_tensor, num_filters):23 x = conv_block(input_tensor, num_filters)24 p = layers.MaxPooling2D(2)(x)25 return x, p26 27def decoder_block(input_tensor, skip_tensor, num_filters):28 x = layers.Conv2DTranspose(num_filters, 2, strides=2, padding='same')(input_tensor)29 x = layers.Concatenate()([x, skip_tensor])30 x = conv_block(x, num_filters)31 return x32 33# ------------------ Multi-Output U-Net ------------------34def build_unet_multioutput(input_shape=(240, 240, 4)):35 inputs = Input(shape=input_shape)36 37 # Encoder38 s1, p1 = encoder_block(inputs, 64)39 s2, p2 = encoder_block(p1, 128)40 s3, p3 = encoder_block(p2, 256)41 s4, p4 = encoder_block(p3, 512)42 43 # Bottleneck + Dropout44 b = conv_block(p4, 1024)45 b = layers.Dropout(0.5)(b)46 47 # Decoder48 d1 = decoder_block(b, s4, 512)49 d2 = decoder_block(d1, s3, 256)50 d3 = decoder_block(d2, s2, 128)51 d4 = decoder_block(d3, s1, 64)52 53 # Multi-task Heads54 wt_out = layers.Conv2D(1, 1, activation='sigmoid', name='wt_head')(d4) # Whole Tumor55 tc_out = layers.Conv2D(1, 1, activation='sigmoid', name='tc_head')(d4) # Tumor Core56 et_out = layers.Conv2D(1, 1, activation='sigmoid', name='et_head')(d4) # Enhancing Tumor57 58 model = Model(inputs=[inputs], outputs=[wt_out, tc_out, et_out], name="U-Net-MultiOutput")59 return model60 61# ------------------ Loss and Metrics ------------------62def dice_coefficient(y_true, y_pred, smooth=1e-6):63 y_true_f = K.flatten(y_true)64 y_pred_f = K.flatten(y_pred)65 intersection = K.sum(y_true_f * y_pred_f)66 return (2. * intersection + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth)67 68def focal_tversky_loss(alpha=0.5, beta=0.7, gamma=1.33):69 def loss(y_true, y_pred):70 y_true = tf.cast(y_true, tf.float32)71 y_pred = tf.clip_by_value(y_pred, 1e-7, 1.0 - 1e-7)72 73 tp = tf.reduce_sum(y_true * y_pred, axis=[1, 2, 3])74 fp = tf.reduce_sum((1 - y_true) * y_pred, axis=[1, 2, 3])75 fn = tf.reduce_sum(y_true * (1 - y_pred), axis=[1, 2, 3])76 77 tversky = (tp + 1e-7) / (tp + alpha * fp + beta * fn + 1e-7)78 return tf.reduce_mean(tf.pow((1 - tversky), gamma))79 return loss80import tensorflow as tf81from tensorflow.keras import layers, Model, Input82from tensorflow.keras import backend as K83 84# ------------------ U-Net Blocks ------------------85def conv_block(input_tensor, num_filters):86 x = layers.Conv2D(num_filters, 3, padding='same')(input_tensor)87 x = layers.BatchNormalization()(x)88 x = layers.Activation('relu')(x)89 90 x = layers.Conv2D(num_filters, 3, padding='same')(x)91 x = layers.BatchNormalization()(x)92 x = layers.Activation('relu')(x)93 return x94 95def encoder_block(input_tensor, num_filters):96 x = conv_block(input_tensor, num_filters)97 p = layers.MaxPooling2D(2)(x)98 return x, p99 100def decoder_block(input_tensor, skip_tensor, num_filters):101 x = layers.Conv2DTranspose(num_filters, 2, strides=2, padding='same')(input_tensor)102 x = layers.Concatenate()([x, skip_tensor])103 x = conv_block(x, num_filters)104 return x105 106# ------------------ Multi-Output U-Net ------------------107def build_unet_multioutput(input_shape=(240, 240, 4)):108 inputs = Input(shape=input_shape)109 110 # Encoder111 s1, p1 = encoder_block(inputs, 64)112 s2, p2 = encoder_block(p1, 128)113 s3, p3 = encoder_block(p2, 256)114 s4, p4 = encoder_block(p3, 512)115 116 # Bottleneck + Dropout117 b = conv_block(p4, 1024)118 b = layers.Dropout(0.5)(b)119 120 # Decoder121 d1 = decoder_block(b, s4, 512)122 d2 = decoder_block(d1, s3, 256)123 d3 = decoder_block(d2, s2, 128)124 d4 = decoder_block(d3, s1, 64)125 126 # Multi-task Heads127 wt_out = layers.Conv2D(1, 1, activation='sigmoid', name='wt_head')(d4) # Whole Tumor128 tc_out = layers.Conv2D(1, 1, activation='sigmoid', name='tc_head')(d4) # Tumor Core129 et_out = layers.Conv2D(1, 1, activation='sigmoid', name='et_head')(d4) # Enhancing Tumor130 131 model = Model(inputs=[inputs], outputs=[wt_out, tc_out, et_out], name="U-Net-MultiOutput")132 return model133 134 