Weiter zum Inhalt

Bild-Superauflösung mit Multi-Decoder-Framework: Tutorial

In diesem Tutorial setzt du ein Paper zur medizinischen Bildgebung mit Deep Learning in Python und Keras um.
Aktualisiert 18. Sept. 2026  · 15 Min. lesen

Mit KI erkunden

ChatGPTClaudePerplexity

Im Trainingsskript arbeitest du an einem Superauflösungsproblem für Bilder mit einer neuartigen Deep-Learning-Architektur. Die Aufgabe ist eine nichtlineare Abbildung von einer niederfeldigen 3-Tesla-Gehirn-MRT-Aufnahme zu einer hochfeldigen 7-Tesla-Gehirn-MRT-Aufnahme. Im Testskript verwendest du die in Teil 1 trainierten Gewichte und sagst auf ungesehenen Daten voraus. Außerdem lernst du, wie du die 2D-Bilder am Ende als kombiniertes Volumen speicherst.

Das Tutorial ist in zwei Teile gegliedert: Der erste Teil führt dich durch den Trainingsprozess, der zweite Teil durch den Testprozess.

Hinweis: Wenn dich das Paper interessiert, findest du den Artikel hier.

Trainingsskript

Kurz gesagt, behandelst du in Teil 1 des Tutorials die folgenden Themen:

  • Du beginnst mit dem Import der benötigten Module, um dein Deep-Learning-Modell zu trainieren,
  • dann bekommst du einen Überblick über den 3T- und 7T-MRT-Datensatz,
  • anschließend definierst du die Initialisierer und lädst den 3T- und 7T-Datensatz; beim Laden werden die Bilder on the fly skaliert,
  • als Nächstes preprocessst du die geladenen Daten: Konvertiere die Train- und Test-Listen in NumPy-Matrizen, setze den Typ auf float32, skaliere mit der Min-Max-Strategie, forme die Arrays um und teile die Daten schließlich in 80 % Training und 20 % Validierung,
  • dann erstellst du die 1-Encoder-3-Decoder-Architektur mit Merge-Verbindungen und Multi-Decodern,
  • als Nächstes definierst du die Loss-Funktion, erstellst drei Modelle und kompilierst sie,
  • zum Schluss trainierst du dein Merge- und Multi-Decoder-Modell, testest es auf den Validierungsdaten und berechnest die quantitativen Ergebnisse.

Abhängigkeiten der Python-Module

Bevor du loslegst, stelle sicher, dass du exakt die folgenden Modulversionen installiert hast:

  • Keras==2.0.4
  • tensorflow==1.8.0
  • scipy==0.19.0
  • numpy==1.14.5
  • Pillow==4.1.1
  • nibabel==2.1.0
  • scikit_learn==0.18.1

Hinweis: Das Modell wird auf einem System mit Nvidia 1080 Ti GPU, Xeon e5 GeForce Prozessor und 32 GB RAM trainiert. Wenn du Jupyter Notebook verwendest, füge drei Zeilen hinzu, um die CUDA-Gerätereihe und sichtbaren CUDA-Geräte über das Modul os zu setzen.

Im folgenden Code setzt du Umgebungsvariablen im Notebook via os.environ. Es ist sinnvoll, das vor der Initialisierung von Keras zu tun, um das Keras-Backend TensorFlow auf die erste GPU zu begrenzen. Wenn deine Maschine die GPU auf 0 hat, verwende 0 statt 1. Das kannst du z. B. mit nvidia-smi im Terminal prüfen.

import os
os.environ["CUDA_DEVICE_ORDER"]="PCI_BUS_ID"
os.environ["CUDA_VISIBLE_DEVICES"]="0" #model will be trained on GPU

Module importieren

Zuerst importierst du alle benötigten Module wie tensorflow, numpy und vor allem keras sowie die benötigten Funktionen bzw. Layer wie Input, Conv2D, MaxPooling2D usw., da du all diese Frameworks zum Trainieren des Modells brauchst!
Zum Einlesen von NIfTI-Bildern importierst du außerdem das Modul nibabel.

import os
import numpy as np
import math
import tensorflow as tf
import nibabel as nib
import numpy as np
from keras.layers import Input,Dense,merge,Reshape,Conv2D,MaxPooling2D,UpSampling2D
from keras.layers.normalization import BatchNormalization
from keras.models import Model,Sequential
from keras.callbacks import ModelCheckpoint
from keras.optimizers import RMSprop
from keras import backend as K
import scipy.misc
from sklearn.utils import shuffle
from sklearn.cross_validation import train_test_split
import matplotlib.pyplot as plt
from keras.models import model_from_json
Using TensorFlow backend.
/usr/local/lib/python3.5/dist-packages/sklearn/cross_validation.py:41: DeprecationWarning: This module was deprecated in version 0.18 in favor of the model_selection module into which all the refactored classes and functions are moved. Also, note that the interface of the new CV iterators is different from that of this module. This module will be removed in 0.20.
  "This module will be removed in 0.20.", DeprecationWarning)

Den 3T-/7T-Gehirn-MRT-Datensatz verstehen

Der 3T- und 7T-Gehirn-MRT-Datensatz besteht aus 3D-Volumina; jedes Volumen enthält insgesamt 207 Scheiben/Bilder von MRT-Aufnahmen unterschiedlicher Schnitte des Gehirns. Jede Scheibe hat die Abmessungen 173 x 173. Die Bilder sind ein-kanalige Graustufenbilder. Insgesamt gibt es 39 Probanden, pro Person je einen MRT-Scan. Das Bildformat ist nicht jpeg, png etc., sondern NIfTI. Später siehst du, wie man NIfTI-Bilder einliest.

Der Datensatz enthält MR-Bilder der T1-Modaliät; T1-Sequenzen eignen sich traditionell gut zur Beurteilung anatomischer Strukturen. Der Datensatz, mit dem du heute arbeitest, enthält 3T- und 7T-Gehirn-MRTs.

Der Datensatz ist öffentlich und kann unter dieser Quelle heruntergeladen werden.

28 Probanden werden zum Training verwendet, die verbleibenden 11 zur Evaluierung.

Initialisierer definieren

Definieren wir zunächst die Datenabmessungen. Die Bildgröße wird beim Einlesen von 173x173 auf 176x176 vergrößert. Hier legst du auch das Datenverzeichnis, die Batchgröße fürs Training, die Kanalzahl, eine Input()-Schicht, Train- und Testmatrizen als Listen fest und lädst für das Rescaling die Textdatei mit den Minimal- und Maximalwerten des MRT-Datensatzes.

Hinweis: Für das Rescaling kannst du auch Minimum und Maximum deines Datensatzes berechnen, da du die Datei maxANDmin.txt vermutlich nicht hast.

x,y = 173,173
full_z = 207
resizeTo=176
batch_size = 32
inChannel = outChannel = 1
input_shape=(x,y,inChannel)
input_img = Input(shape = (resizeTo, resizeTo, inChannel))
inp = "ground3T"
out = "ground7T"     
train_matrix = []
test_matrix = []
min_max = np.loadtxt('maxANDmin.txt')

Daten laden

Als Nächstes lädst du die MRT-Daten mit nibabel und skalierst die Bilder von 173 x 173 auf 176 x 176, indem du in x- und y-Richtung mit Nullen auffüllst.

Beachte: Beim Laden eines NIfTI-Volumens lädt Nibabel das Bildarray nicht sofort. Es wartet, bis du die Daten explizit anforderst. Üblicherweise geschieht das über die Methode get_data().

Da du 2D-Slices statt 3D willst, nutzt du die zuvor initialisierten Listen train und test: Jedes gelesene Volumen wird über alle 207 Scheiben iteriert, und jede Scheibe wird einzeln an die Liste angehängt.

folder = os.listdir(inp)

Laden wir zuerst die 3T-Daten.

for f in folder:
    temp = np.zeros([resizeTo,full_z,resizeTo])
    a = nib.load(inp + f)
    a = a.get_data()
    temp[3:,:,3:] = a
    a = temp
    for j in range(full_z):
        train_matrix.append(a[:,j,:])

Dann lädst du die 7T-Daten. Du verwendest dieselbe Variable folder, da Anzahl der 3T- und 7T-Volumina gleich ist.

for f in folder:
    temp = np.zeros([resizeTo,full_z,resizeTo])
    b = nib.load(out + f)
    b = b.get_data()
    temp[3:,:,3:] = b
    b = temp
    for j in range(full_z):
        test_matrix.append(b[:,j,:])

Datenvorverarbeitung

Da Train- und Testmatrizen als Listen vorliegen, konvertierst du sie mit NumPy in Arrays.

Anschließend setzt du den type auf float32 und skalierst sowohl input als auch ground truth.

train_matrix = np.asarray(train_matrix)
train_matrix = train_matrix.astype('float32')
m = min_max[0]
mi = min_max[1]
train_matrix = (train_matrix - mi) / (m - mi)

test_matrix = np.asarray(test_matrix)
test_matrix = test_matrix.astype('float32')
test_matrix = (test_matrix - mi) / (m - mi)

Lass uns kurz die Shapes von train_matrix (3T) und test_matrix (7T) prüfen. Es sollten 28 x 207 = 5796 Bilder mit je 176 x 176 sein.

train_matrix.shape
(5796, 176, 176)
test_matrix.shape
(5796, 176, 176)

Als Nächstes erstellst du zwei neue Variablen augmented_images (3T/Input) und Haugmented_images (7T/Ground Truth) mit dem Shape der Train- bzw. Testmatrix. Das ergibt eine 4D-Matrix: erste Dimension Bildanzahl, zweite und dritte Bildgröße, letzte die Anzahl Kanäle (hier 1).

augmented_images=np.zeros(shape=[(train_matrix.shape[0]),(train_matrix.shape[1]),(train_matrix.shape[2]),(1)])
Haugmented_images=np.zeros(shape=[(train_matrix.shape[0]),(train_matrix.shape[1]),(train_matrix.shape[2]),(1)])

Dann iterierst du über alle Bilder, formst Train- und Testmatrix jeweils zu 176 x 176 und fügst sie den Arrays augmented_images (3T/Input) bzw. Haugmented_images (7T/Ground Truth) hinzu.

for i in range(train_matrix.shape[0]):
    augmented_images[i,:,:,0] = train_matrix[i,:,:].reshape(resizeTo,resizeTo)
    Haugmented_images[i,:,:,0] = test_matrix[i,:,:].reshape(resizeTo,resizeTo)

Nun teilst du die Daten, damit dein Modell gut generalisiert: 80 % zum Trainieren, 20 % zur Validierung. Das senkt die Overfitting-Gefahr, da auf nicht gesehenen Daten validiert wird.

Zum sauberen Split verwendest du das zuvor importierte train_test_split aus scikit-learn:

data,Label = shuffle(augmented_images,Haugmented_images, random_state=2)
X_train, X_test, y_train, y_test = train_test_split(data, Label, test_size=0.2, random_state=2)
X_test = np.array(X_test)
y_test = np.array(y_test)
X_test = X_test.astype('float32')
y_test = y_test.astype('float32')
X_train = np.array(X_train)
y_train = np.array(y_train)
X_train = X_train.astype('float32')
y_train = y_train.astype('float32')

Das Modell: 1 Encoder, 3 Decoder!

model
Abbildung: Selektives Autoencoder-Backpropagation
Bild aus diesem Paper.

Merge-Verbindungen

Als Nächstes definierst du die vorgeschlagene Architektur mit Blöcken aufeinanderfolgender Filterlayer, gefolgt von Max-Pooling im Encoder, wie in der nächsten Zelle gezeigt. Um die Originalgröße am Ausgang wiederherzustellen, wird in jedem Decoder-Block ein Upsampling-Layer verwendet. Beim Upsampling können Artefakte entstehen, da Details im heruntergesampelten Decoder-Input fehlen. Daher konkateniert man den Decoder-Input mit der hochskalierten Version aus dem Encoder, um die Art der hochskalierten Details für eine bessere Rekonstruktion bereitzustellen. Merge-Verbindungen bringen messbare PSNR-Verbesserungen (Größenordnung 5 dB). Dieses Setting ist inspiriert von diesem Paper.

Multi-Decoder

Der Ansatz nutzt einen einzelnen Encoder und mehrere Decoder bei ein-kanaligern Eingaben. In jedem Block des Encoders und in allen drei Decodern kommen drei Convolution-Layer zum Einsatz, gefolgt von Batch Normalization für numerische Stabilität.

Im Encoder hat der erste Convolution-Block 32 Filter, danach verdoppelt sich die Filteranzahl nach jedem Block. In allen Decodern hat der erste Block 256 Filter, danach halbiert sich die Filteranzahl nach jedem Block. Du verwendest überall 3x3-Filter.

Als Aktivierung dient ReLU in allen Layern außer dem letzten. Da die Daten auf 0 bis 1 normalisiert sind, wird in der letzten Schicht Sigmoid genutzt.

Lokale Bilddetails auf verschiedenen Skalen sind entscheidend für die Rekonstruktion. Die Architektur berücksichtigt Bilder auf mehreren Skalen via hierarchischem Downsampling (MaxPooling) und Upsampling (jeweils Faktor 2) in Encoder und Decodern. Die kodierte Repräsentation nach drei Downsamplings überführt die hochdimensionale Eingabe in einen latenten Raum.

Encoder

def encoder(input_img):
    conv1 = Conv2D(32, (3, 3), activation='relu', padding='same')(input_img)
    conv1 = BatchNormalization()(conv1)
    conv1 = Conv2D(32, (3,3), activation='relu', padding='same')(conv1)
    conv1 = BatchNormalization()(conv1)
    conv1 = Conv2D(32, (3,3), activation='relu', padding='same')(conv1)
    conv1 = BatchNormalization()(conv1)
    pool1 = MaxPooling2D(pool_size=(2, 2))(conv1)
    conv2 = Conv2D(64, (3, 3), activation='relu', padding='same')(pool1)
    conv2 = BatchNormalization()(conv2)
    conv2 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv2)
    conv2 = BatchNormalization()(conv2)
    conv2 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv2)
    conv2 = BatchNormalization()(conv2)
    pool2 = MaxPooling2D(pool_size=(2, 2))(conv2)
    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same')(pool2)
    conv3 = BatchNormalization()(conv3)
    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv3)
    conv3 = BatchNormalization()(conv3)
    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv3)
    conv3 = BatchNormalization()(conv3)
    pool3 = MaxPooling2D(pool_size=(2, 2))(conv3)
    conv4 = Conv2D(256, (3, 3), activation='relu', padding='same')(pool3)
    conv4 = BatchNormalization()(conv4)
    conv4 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv4)
    conv4 = BatchNormalization()(conv4)
    conv4 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv4)
    conv4 = BatchNormalization()(conv4)
    conv5 = Conv2D(512, (3, 3), activation='relu', padding='same')(conv4)
    conv5 = BatchNormalization()(conv5)
    conv5 = Conv2D(512, (3, 3), activation='relu', padding='same')(conv5)
    conv5 = BatchNormalization()(conv5)
    conv5 = Conv2D(512, (3, 3), activation='sigmoid', padding='same')(conv5)
    conv5 = BatchNormalization()(conv5)
    return conv5,conv4,conv3,conv2,conv1

Decoder 1

def decoder_1(conv5,conv4,conv3,conv2,conv1):
    up6 = merge([conv5, conv4], mode='concat', concat_axis=3)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(up6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    up7 = UpSampling2D((2,2))(conv6)
    up7 = merge([up7, conv3], mode='concat', concat_axis=3)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(up7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    up8 = UpSampling2D((2,2))(conv7)
    up8 = merge([up8, conv2], mode='concat', concat_axis=3)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(up8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    up9 = UpSampling2D((2,2))(conv8)
    up9 = merge([up9, conv1], mode='concat', concat_axis=3)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(up9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    decoded_1 = Conv2D(1, (3, 3), activation='sigmoid', padding='same')(conv9)
    return decoded_1

Decoder 2

def decoder_2(conv5,conv4,conv3,conv2,conv1):
    up6 = merge([conv5, conv4], mode='concat', concat_axis=3)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(up6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    up7 = UpSampling2D((2,2))(conv6)
    up7 = merge([up7, conv3], mode='concat', concat_axis=3)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(up7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    up8 = UpSampling2D((2,2))(conv7)
    up8 = merge([up8, conv2], mode='concat', concat_axis=3)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(up8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    up9 = UpSampling2D((2,2))(conv8)
    up9 = merge([up9, conv1], mode='concat', concat_axis=3)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(up9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    decoded_2 = Conv2D(1, (3, 3), activation='sigmoid', padding='same')(conv9)
    return decoded_2

Decoder 3

def decoder_3(conv5,conv4,conv3,conv2,conv1):
    up6 = merge([conv5, conv4], mode='concat', concat_axis=3)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(up6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    up7 = UpSampling2D((2,2))(conv6)
    up7 = merge([up7, conv3], mode='concat', concat_axis=3)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(up7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    up8 = UpSampling2D((2,2))(conv7)
    up8 = merge([up8, conv2], mode='concat', concat_axis=3)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(up8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    up9 = UpSampling2D((2,2))(conv8)
    up9 = merge([up9, conv1], mode='concat', concat_axis=3)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(up9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    decoded_3 = Conv2D(1, (3, 3), activation='sigmoid', padding='same')(conv9)
    return decoded_3

In den obigen vier Zellen hast du vier Funktionen definiert: eine für den Encoder und drei für die Decoder. Da es Funktionen sind, könntest du decoder() einmal definieren und dreimal aufrufen. Zur besseren Nachvollziehbarkeit und um Zufallseffekte der drei Decoder zu vermeiden, definieren wir sie dreifach.

Loss-Funktion

Als Nächstes verwendest du den Mean-Square-Error, wobei Werte (Pixel) in y_t und y_p, die Null sind, ausgeschlossen werden.

def root_mean_sq_GxGy(y_t, y_p):
    a1=1
    where = tf.not_equal(y_t, 0)
    a_t=tf.boolean_mask(y_t,where,name='boolean_mask')
    a_p=tf.boolean_mask(y_p,where,name='boolean_mask')
    return a1*(K.sqrt(K.mean((K.square(a_t-a_p)))))

Modelldefinition und Kompilierung

Zuerst rufst du die encoder-Funktion mit dem Input auf. Da du Merge-Verbindungen nutzt, gibt die encoder-Funktion die Ausgaben von fünf convolution-Layern zurück, die du dann jeweils mit allen drei Decodern zusammenführst.

conv5,conv4,conv3,conv2,conv1 = encoder(input_img)
autoencoder_1 = Model(input_img, decoder_1(conv5,conv4,conv3,conv2,conv1))
autoencoder_1.compile(loss=root_mean_sq_GxGy, optimizer = RMSprop())

autoencoder_2 = Model(input_img, decoder_2(conv5,conv4,conv3,conv2,conv1))
autoencoder_2.compile(loss=root_mean_sq_GxGy, optimizer = RMSprop())

autoencoder_3 = Model(input_img, decoder_3(conv5,conv4,conv3,conv2,conv1))
autoencoder_3.compile(loss=root_mean_sq_GxGy, optimizer = RMSprop())
/usr/local/lib/python3.5/dist-packages/ipykernel_launcher.py:2: UserWarning: The `merge` function is deprecated and will be removed after 08/2017. Use instead layers from `keras.layers.merge`, e.g. `add`, `concatenate`, etc.

/usr/local/lib/python3.5/dist-packages/keras/legacy/layers.py:460: UserWarning: The `Merge` layer is deprecated and will be removed after 08/2017. Use instead layers from `keras.layers.merge`, e.g. `add`, `concatenate`, etc.
  name=name)
/usr/local/lib/python3.5/dist-packages/ipykernel_launcher.py:10: UserWarning: The `merge` function is deprecated and will be removed after 08/2017. Use instead layers from `keras.layers.merge`, e.g. `add`, `concatenate`, etc.
  # Remove the CWD from sys.path while we load stuff.
/usr/local/lib/python3.5/dist-packages/ipykernel_launcher.py:18: UserWarning: The `merge` function is deprecated and will be removed after 08/2017. Use instead layers from `keras.layers.merge`, e.g. `add`, `concatenate`, etc.
/usr/local/lib/python3.5/dist-packages/ipykernel_launcher.py:26: UserWarning: The `merge` function is deprecated and will be removed after 08/2017. Use instead layers from `keras.layers.merge`, e.g. `add`, `concatenate`, etc.


WARNING:tensorflow:From /usr/local/lib/python3.5/dist-packages/keras/backend/tensorflow_backend.py:1257: calling reduce_mean (from tensorflow.python.ops.math_ops) with keep_dims is deprecated and will be removed in a future version.
Instructions for updating:
keep_dims is deprecated, use keepdims instead

Modell trainieren

Die Gewichte speicherst du nur, wenn sich das Peak Signal-to-Noise Ratio auf den Validierungsdaten verbessert. Dazu definierst du eine Liste psnr_gray_channel und fügst als Dummy-Wert 1 hinzu, da der PSNR selbst in frühen Trainingsphasen deutlich größer sein wird.

Die Learning Rate initialisierst du mit 1e-3 und reduzierst sie alle 20 Epochen um 10 %.

psnr_gray_channel = []
psnr_gray_channel.append(1)
learning_rate = 0.001
j=0

Hinweis: Der nächste Codeblock sollte in einer Zelle laufen, um den gesamten Trainingsprozess nachvollziehbar zu machen. Hier ist er zur Erklärung in kleinere Teile zerlegt.

Das Modell wird 500 Epochen trainiert. Die anfängliche Learning Rate ist 1e-3. Nach jeder Epoche speicherst du PSNR- und MSE-Werte zwischen 7T und der Vorhersage. Mit K.set_value passt du die Learning Rate aller drei Modelle alle 20 Epochen um 10 % nach unten an.

for jj in range(500):
    myfile_valid_psnr_7T = open('../1_encoder_3_decoders_complete_slices_single_channel/validation_psnr7T_1encoder_3decoders.txt', 'a')
    myfile_valid_mse_7T = open('../1_encoder_3_decoders_complete_slices_single_channel/validation_mse7T_1encoder_3decoders.txt', 'a')

    K.set_value(autoencoder_1.optimizer.lr, learning_rate)
    K.set_value(autoencoder_2.optimizer.lr, learning_rate)
    K.set_value(autoencoder_3.optimizer.lr, learning_rate)

Als Nächstes mischst du die 3T-Input- und 7T-Ground-Truth-Bilder, um Overfitting zu vermeiden. Dann berechnest du die Anzahl Batches anhand der zuvor definierten batch_size und iterierst über num_batches.

train_X,train_Y = shuffle(X_train,y_train)
print ("Epoch is: %d\n" % j)
print ("Number of batches: %d\n" % int(len(train_X)/batch_size))
num_batches = int(len(train_X)/batch_size)
for batch in range(num_batches):

Neben den PSNR-Werten speicherst du auch die Verluste aller drei Autoencoder sowie der drei Decoder separat.

myfile_ae1_loss = open('../1_encoder_3_decoders_complete_slices_single_channel/ae1_train_loss_1encoder_3decoders.txt', 'a')
myfile_ae2_loss = open('../1_encoder_3_decoders_complete_slices_single_channel/ae2_train_loss_1encoder_3decoders.txt', 'a')
myfile_ae3_loss = open('../1_encoder_3_decoders_complete_slices_single_channel/ae3_train_loss_1encoder_3decoders.txt', 'a')
myfile_dec1_loss = open('../1_encoder_3_decoders_complete_slices_single_channel/dec1_train_loss_1encoder_3decoders.txt', 'a')
myfile_dec2_loss = open('../1_encoder_3_decoders_complete_slices_single_channel/dec2_train_loss_1encoder_3decoders.txt', 'a')
myfile_dec3_loss = open('../1_encoder_3_decoders_complete_slices_single_channel/dec3_train_loss_1encoder_3decoders.txt', 'a')

Damit das Modell pro Batch die nächsten 32 (batch_size) Samples sieht, erledigt die folgende Zelle genau das.

batch_train_X = train_X[batch*batch_size:min((batch+1)*batch_size,len(train_X)),:]
batch_train_Y = train_Y[batch*batch_size:min((batch+1)*batch_size,len(train_Y)),:]

Für die Minimum-Autoencoder-Strategie testest du nach jeder Epoche auf den Trainingsdaten mit test_on_batch von Keras. Das liefert drei Verluste, die du ausgibst.

loss_1 = autoencoder_1.test_on_batch(batch_train_X,batch_train_Y)
loss_2 = autoencoder_2.test_on_batch(batch_train_X,batch_train_Y)
loss_3 = autoencoder_3.test_on_batch(batch_train_X,batch_train_Y)
print ('epoch_num: %d batch_num: %d Test_loss_1: %f\n' % (j,batch,loss_1))
print ('epoch_num: %d batch_num: %d Test_loss_2: %f\n' % (j,batch,loss_2))
print ('epoch_num: %d batch_num: %d Test_loss_3: %f\n' % (j,batch,loss_3))

Es gibt sechs mögliche Fälle im Netzwerk:
loss_1: Verlust Autoencoder 1
loss_2: Verlust Autoencoder 2
loss_3: Verlust Autoencoder 3

    • loss_1 kann größer als loss_2 und loss_3 sein. Dann trainierst du nur Autoencoder 1. Bei Autoencoder 2 und 3 setzt du den Encoder-Teil auf False und trainierst nur deren Decoder. Alle Verluste schreibst du in die Textdateien.
    • loss_2 kann größer als loss_1 und loss_3 sein. Dann trainierst du nur Autoencoder 2. Bei Autoencoder 1 und 3 setzt du den Encoder-Teil auf False und trainierst nur deren Decoder. Alle Verluste werden protokolliert.
    • loss_3 kann größer als loss_1 und loss_2 sein. Dann trainierst du nur Autoencoder 3. Bei Autoencoder 1 und 2 setzt du den Encoder-Teil auf False und trainierst nur deren Decoder. Abschließend schreibst du die Verluste in die Dateien.
    • loss_1 kann gleich loss_2 sein. Dann trainierst du entweder Autoencoder 1 oder 2. Autoencoder 3 sowie den nicht gewählten Autoencoder setzt du im Encoder-Teil auf False und trainierst nur deren Decoder. Danach schreibst du alle Verluste in die Dateien.
    • loss_2 kann gleich loss_3 sein. Dann trainierst du entweder Autoencoder 2 oder 3. Autoencoder 1 sowie den nicht gewählten Autoencoder setzt du im Encoder-Teil auf False und trainierst nur deren Decoder. Abschließend protokollierst du die Verluste.
    • loss_3 kann gleich loss_1 sein. Dann trainierst du entweder Autoencoder 3 oder 1. Autoencoder 2 sowie den nicht gewählten Autoencoder setzt du im Encoder-Teil auf False und trainierst nur deren Decoder. Danach werden alle Verluste geschrieben.
model
Abbildung: Architektur des Modells
Bild aus diesem Paper.
if loss_1 < loss_2 and loss_1 < loss_3:
    train_1 = autoencoder_1.train_on_batch(batch_train_X,batch_train_Y)
    myfile_ae1_loss.write("%f \n" % (train_1))
    print ('epoch_num: %d batch_num: %d AE_Train_loss_1: %f\n' % (j,batch,train_1))
    #myfile.write('epoch_num: %d batch_num: %d AE_Train_loss_1: %f\n' % (j,batch,train_1))
    for layer in autoencoder_2.layers[:34]:
        layer.trainable = False
    for layer in autoencoder_3.layers[:34]:
        layer.trainable = False
    #autoencoder_2.summary()
    #autoencoder_3.summary()
    train_2 = autoencoder_2.train_on_batch(batch_train_X,batch_train_Y)
    train_3 = autoencoder_3.train_on_batch(batch_train_X,batch_train_Y)
    myfile_dec2_loss.write("%f \n" % (train_2))
    myfile_dec3_loss.write("%f \n" % (train_3))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_2: %f\n' % (j,batch,train_2))
    print ('epoch_num: %d batch_num: %d Decoder_loss_3: %f\n' % (j,batch,train_3))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_2: %f\n' % (j,batch,train_2))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_loss_3: %f\n' % (j,batch,train_3))
elif loss_2 < loss_1 and loss_2 < loss_3:
    train_2 = autoencoder_2.train_on_batch(batch_train_X,batch_train_Y)
    myfile_ae2_loss.write("%f \n" % (train_2))
    print ('epoch_num: %d batch_num: %d AE_Train_loss_2: %f\n' % (j,batch,train_2))
    #myfile.write('epoch_num: %d batch_num: %d AE_Train_loss_2: %f\n' % (j,batch,train_2))
    for layer in autoencoder_1.layers[:34]:
        layer.trainable = False
    for layer in autoencoder_3.layers[:34]:
        layer.trainable = False
    #autoencoder_1.summary()
    #autoencoder_3.summary()
    train_1 = autoencoder_1.train_on_batch(batch_train_X,batch_train_Y)
    train_3 = autoencoder_3.train_on_batch(batch_train_X,batch_train_Y)
    myfile_dec1_loss.write("%f \n" % (train_1))
    myfile_dec3_loss.write("%f \n" % (train_3))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_1: %f\n' % (j,batch,train_1))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_3: %f\n' % (j,batch,train_3))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_1: %f\n' % (j,batch,train_1))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_3: %f\n' % (j,batch,train_3))
elif loss_3 < loss_1 and loss_3 < loss_2:
    train_3 = autoencoder_3.train_on_batch(batch_train_X,batch_train_Y)
    myfile_ae3_loss.write("%f \n" % (train_3))
    print ('epoch_num: %d batch_num: %d AE_Train_loss_3: %f\n' % (j,batch,train_3))
    #myfile.write('epoch_num: %d batch_num: %d AE_Train_loss_3: %f\n' % (j,batch,train_3))
    for layer in autoencoder_1.layers[:34]:
        layer.trainable = False
    for layer in autoencoder_2.layers[:34]:
        layer.trainable = False
    #autoencoder_1.summary()
    #autoencoder_2.summary()
    train_1 = autoencoder_1.train_on_batch(batch_train_X,batch_train_Y)
    train_2 = autoencoder_2.train_on_batch(batch_train_X,batch_train_Y)
    myfile_dec1_loss.write("%f \n" %(train_1))
    myfile_dec2_loss.write("%f \n" % (train_2))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_1: %f\n' % (j,batch,train_1))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_2: %f\n' % (j,batch,train_2))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_1: %f\n' % (j,batch,train_1))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_2: %f\n' % (j,batch,train_2))
elif loss_1 == loss_2:
    train_1 = autoencoder_1.train_on_batch(batch_train_X,batch_train_Y)
    myfile_ae1_loss.write("%f \n" % (train_1))
    print ('epoch_num: %d batch_num: %d AE_Train_loss_1_equal_state: %f\n' % (j,batch,train_1))
    #myfile.write('epoch_num: %d batch_num: %d AE_Train_loss_1_equal_state: %f\n' % (j,batch,train_1))
    for layer in autoencoder_3.layers[:34]:
        layer.trainable = False
    for layer in autoencoder_2.layers[:34]:
        layer.trainable = False
    #autoencoder_2.summary()
    #autoencoder_3.summary()
    train_2 = autoencoder_2.train_on_batch(batch_train_X,batch_train_Y)
    train_3 = autoencoder_3.train_on_batch(batch_train_X,batch_train_Y)
    myfile_dec2_loss.write("%f \n" % (train_2))
    myfile_dec3_loss.write("%f \n" % (train_3))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_2_equal_state: %f\n' % (j,batch,train_2))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_3_equal_state: %f\n' % (j,batch,train_3))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_2_equal_state: %f\n' % (j,batch,train_2))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_3_equal_state: %f\n' % (j,batch,train_3))
elif loss_2 == loss_3:
    train_2 = autoencoder_2.train_on_batch(batch_train_X,batch_train_Y)
    myfile_ae2_loss.write("%f \n" % (train_2))
    print ('epoch_num: %d batch_num: %d AE_Train_loss_2_equal_state: %f\n' % (j,batch,train_2))
    #myfile.write('epoch_num: %d batch_num: %d AE_Train_loss_2_equal_state: %f\n' % (j,batch,train_2))
    for layer in autoencoder_1.layers[:34]:
        layer.trainable = False
    for layer in autoencoder_3.layers[:34]:
        layer.trainable = False

    #autoencoder_2.summary()
    #autoencoder_3.summary()
    train_1 = autoencoder_1.train_on_batch(batch_train_X,batch_train_Y)
    train_3 = autoencoder_3.train_on_batch(batch_train_X,batch_train_Y)
    myfile_dec1_loss.write("%f \n" % (train_1))
    myfile_dec3_loss.write("%f \n" % (train_3))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_1_equal_state: %f\n' % (j,batch,train_1))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_3_equal_state: %f\n' % (j,batch,train_3))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_1_equal_state: %f\n' % (j,batch,train_1))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_3_equal_state: %f\n' % (j,batch,train_3))
elif loss_3 == loss_1:
    train_1 = autoencoder_1.train_on_batch(batch_train_X,batch_train_Y)
    myfile_ae1_loss.write("%f \n" % (train_1))
    print ('epoch_num: %d batch_num: %d AE_Train_loss_1_equal_state: %f\n' % (j,batch,train_1))
    #myfile.write('epoch_num: %d batch_num: %d AE_Train_loss_1_equal_state: %f\n' % (j,batch,train_1))
    for layer in autoencoder_2.layers[:34]:
        layer.trainable = False
    for layer in autoencoder_3.layers[:34]:
        layer.trainable = False
    #autoencoder_2.summary()
    #autoencoder_3.summary()
    train_2 = autoencoder_2.train_on_batch(batch_train_X,batch_train_Y)
    train_3 = autoencoder_3.train_on_batch(batch_train_X,batch_train_Y)
    myfile_dec2_loss.write("%f \n" % (train_2))
    myfile_dec3_loss.write("%f \n" % (train_3))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_2_equal_state: %f\n' % (j,batch,train_2))
    print ('epoch_num: %d batch_num: %d Decoder_Train_loss_3_equal_state: %f\n' % (j,batch,train_3))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_2_equal_state: %f\n' % (j,batch,train_2))
    #myfile.write('epoch_num: %d batch_num: %d Decoder_Train_loss_3_equal_state: %f\n' % (j,batch,train_3))


    myfile_ae1_loss.close()
    myfile_ae2_loss.close()
    myfile_ae3_loss.close()
    myfile_dec1_loss.close()
    myfile_dec2_loss.close()
    myfile_dec3_loss.close()

Wichtig: Da du oben in manchen Fällen Layer auf False gesetzt hast, musst du diese wieder auf True setzen, damit alle Autoencoder-Layer für test_on_batch verwendet werden und nicht dauerhaft deaktiviert bleiben.

for layer in autoencoder_1.layers[:34]:
            layer.trainable = True
for layer in autoencoder_2.layers[:34]:
            layer.trainable = True
for layer in autoencoder_3.layers[:34]:
            layer.trainable = True

Die Gewichte speicherst du auf drei Arten; zwei davon sind unten: einmal alle 100 Epochen, außerdem nach jeder Epoche (die Dateien werden dabei überschrieben).

if jj % 100 ==0:
            autoencoder_1.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE1_" + str(jj)+".h5")
            autoencoder_2.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE2_" + str(jj)+".h5")
            autoencoder_3.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE3_" + str(jj)+".h5")


autoencoder_1.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE1.h5")
autoencoder_2.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE2.h5")
autoencoder_3.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE3.h5")

Validierungsdaten testen

Du mischst die Validierungsdaten und verwendest dann dieselben sechs Bedingungen wie oben. Trifft eine Bedingung zu, testest du dein Modell mit dem entsprechenden Autoencoder.

Dann berechnest du die Metriken MSE und PSNR zwischen decoded_imgs (Vorhersage) und Ground Truth. Zum Schluss speicherst du sie in Textdateien.

X_test,y_test = shuffle(X_test,y_test)
if loss_1 < loss_2 and loss_1 < loss_3:
    decoded_imgs = autoencoder_1.predict(X_test)
    mse_7T=  np.mean((y_test[:,:,:,0] - decoded_imgs[:,:,:,0]) ** 2)
    check_7T = math.sqrt(mse_7T)
    psnr_7T = 20 * math.log10( 1.0 / check_7T)


    myfile_valid_psnr_7T.write("%f \n" % (psnr_7T))
    myfile_valid_mse_7T.write("%f \n" % (mse_7T))

    #print (check)
elif loss_2 < loss_1 and loss_2 < loss_3:
    decoded_imgs = autoencoder_2.predict(X_test)
    mse_7T=  np.mean((y_test[:,:,:,0] - decoded_imgs[:,:,:,0]) ** 2)
    check_7T = math.sqrt(mse_7T)
    psnr_7T = 20 * math.log10( 1.0 / check_7T)

    myfile_valid_psnr_7T.write("%f \n" % (psnr_7T))
    myfile_valid_mse_7T.write("%f \n" % (mse_7T))

    #print (check)
elif loss_3 < loss_2 and loss_3 < loss_1:
    decoded_imgs = autoencoder_3.predict(X_test)
    mse_7T=  np.mean((y_test[:,:,:,0] - decoded_imgs[:,:,:,0]) ** 2)
    check_7T = math.sqrt(mse_7T)
    psnr_7T = 20 * math.log10( 1.0 / check_7T)

    myfile_valid_psnr_7T.write("%f \n" % (psnr_7T))
    myfile_valid_mse_7T.write("%f \n" % (mse_7T))
    #print (check)

elif loss_1 == loss_2:
    decoded_imgs = autoencoder_1.predict(X_test)
    mse_7T=  np.mean((y_test[:,:,:,0] - decoded_imgs[:,:,:,0]) ** 2)
    check_7T = math.sqrt(mse_7T)
    psnr_7T = 20 * math.log10( 1.0 / check_7T)

    myfile_valid_psnr_7T.write("%f \n" % (psnr_7T))
    myfile_valid_mse_7T.write("%f \n" % (mse_7T))


elif loss_2 == loss_3:
    decoded_imgs = autoencoder_2.predict(X_test)
    mse_7T=  np.mean((y_test[:,:,:,0] - decoded_imgs[:,:,:,0]) ** 2)
    check_7T = math.sqrt(mse_7T)
    psnr_7T = 20 * math.log10( 1.0 / check_7T)

    myfile_valid_psnr_7T.write("%f \n" % (psnr_7T))
    myfile_valid_mse_7T.write("%f \n" % (mse_7T))


elif loss_3 == loss_1:
    decoded_imgs = autoencoder_3.predict(X_test)
    mse_7T=  np.mean((y_test[:,:,:,0] - decoded_imgs[:,:,:,0]) ** 2)
    check_7T = math.sqrt(mse_7T)
    psnr_7T = 20 * math.log10( 1.0 / check_7T)

    myfile_valid_psnr_7T.write("%f \n" % (psnr_7T))
    myfile_valid_mse_7T.write("%f \n" % (mse_7T))

Hier speicherst du die Gewichte nur, wenn der PSNR zwischen der Vorhersage (7T) und Ground Truth (7T) größer ist als alle zuvor in psnr_gray_channel gespeicherten Werte.

if max(psnr_gray_channel) < psnr_7T:
            autoencoder_1.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE1_" + str(jj)+".h5")
            autoencoder_2.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE2_" + str(jj)+".h5")
            autoencoder_3.save_weights("../Model/CROSSVAL1/CROSSVAL1_AE3_" + str(jj)+".h5")

    psnr_gray_channel.append(psnr_7T)

Input, Ground Truth und Vorhersage speichern: Quantitative Ergebnisse

Du definierst eine NumPy-Matrix temp mit 176 x 528, da drei Bilder mit je 176 x 176 in einer Zeile gespeichert werden. Du speicherst eines der Validierungsbilder (3T, 7T und Vorhersage) und multiplizierst die Matrix mit 255, da die Bilder zwischen 0 und 1 skaliert wurden.

Mit scipy speicherst du anschließend das Bild.

temp = np.zeros([resizeTo,resizeTo*3])
temp[:resizeTo,:resizeTo] = X_test[0,:,:,0]
temp[:resizeTo,resizeTo:resizeTo*2] = y_test[0,:,:,0]
temp[:resizeTo,2*resizeTo:] = decoded_imgs[0,:,:,0]
temp = temp*255
scipy.misc.imsave('../Results/1_encoder_3_decoders_complete_slices_single_channel/' + str(j) + '.jpg', temp)
j +=1

Zum Schluss schließt du die PSNR- und MSE-Dateien.

myfile_valid_psnr_7T.close()
myfile_valid_mse_7T.close()

Und am Ende reduzierst du die Learning Rate alle 20 Epochen um 10 % ihres aktuellen Werts.

if jj % 20 ==0:
        learning_rate = learning_rate - learning_rate * 0.10

Testskript

In Teil 2 des Tutorials behandelst du Folgendes:

  • Du startest mit dem Import der benötigten Module für das Training deines Deep-Learning-Modells,
  • dann folgt ein Überblick über den 3T- und 7T-MRT-Datensatz,
  • anschließend definierst du die Initialisierer und lädst die 3T- und 7T-Testdaten; beim Laden werden die Bilder on the fly skaliert,
  • als Nächstes preprocessst du die Daten: Konvertiere die Listen in NumPy-Matrizen, setze den Typ auf float32, skaliere mit Min-Max, forme Arrays um und teile schließlich in 80 % Training und 20 % Validierung,
  • dann erstellst du die 1-Encoder-3-Decoder-Architektur mit Merge-Verbindungen und Multi-Decodern,
  • als Nächstes definierst du Loss-Funktion, drei Modelle und lädst die trainierten Gewichte,
  • zum Schluss sagst du mit deinem Merge- und Multi-Decoder-Modell auf ungesehenen Daten voraus und speicherst quantitative und qualitative Ergebnisse. Außerdem lernst du, wie du 2D-Bilder mit nibabel als kombiniertes Volumen speicherst.

Abhängigkeiten der Python-Module

Bevor du loslegst, stelle sicher, dass du exakt die folgenden Modulversionen installiert hast:

  • Keras==2.0.4
  • tensorflow==1.8.0
  • scipy==0.19.0
  • numpy==1.14.5
  • Pillow==4.1.1
  • nibabel==2.1.0
  • scikit_learn==0.18.1

Hinweis: Das Modell wird auf einem System mit Nvidia 1080 Ti GPU, Xeon e5 GeForce Prozessor und 32 GB RAM trainiert. Wenn du Jupyter Notebook verwendest, füge drei Zeilen hinzu, um die CUDA-Gerätereihe und sichtbaren CUDA-Geräte über das Modul os zu setzen.

Im folgenden Code setzt du Umgebungsvariablen im Notebook via os.environ. Es ist sinnvoll, das vor der Initialisierung von Keras zu tun, um das Keras-Backend TensorFlow auf die erste GPU zu begrenzen. Wenn deine Maschine die GPU auf 0 hat, verwende 0 statt 1. Das kannst du z. B. mit nvidia-smi prüfen.

import os
os.environ["CUDA_DEVICE_ORDER"]="PCI_BUS_ID"
os.environ["CUDA_VISIBLE_DEVICES"]="0" #model will be trained on GPU

Module importieren

Zuerst importierst du alle benötigten Module wie tensorflow, numpy und vor allem keras sowie Funktionen/Layer wie Input, Conv2D, MaxPooling2D usw., da du all diese Frameworks zum Trainieren brauchst!
Zum Einlesen von NIfTI-Bildern importierst du außerdem nibabel.

import os
from keras.layers import Input,Dense,Flatten,Dropout,merge,Reshape,Conv2D,MaxPooling2D,UpSampling2D
from keras.layers.normalization import BatchNormalization
from keras.models import Model,Sequential
from keras.callbacks import ModelCheckpoint
from keras.optimizers import Adadelta, RMSprop,SGD,Adam
from keras import regularizers
from keras import backend as K
import numpy as np
import scipy.misc
import numpy.random as rng
from sklearn.utils import shuffle
import nibabel as nib
from sklearn.cross_validation import train_test_split
import math

Den 3T-/7T-Gehirn-MRT-Datensatz verstehen

Der 3T- und 7T-Gehirn-MRT-Datensatz besteht aus 3D-Volumina; jedes Volumen enthält insgesamt 207 Scheiben/Bilder von MRT-Aufnahmen unterschiedlicher Hirnschnitte. Jede Scheibe hat die Abmessungen 173 x 173. Die Bilder sind ein-kanalige Graustufenbilder. Insgesamt gibt es 39 Probanden, je ein MRT-Scan pro Patient. Das Bildformat ist NIfTI (nicht jpeg, png etc.). Später siehst du, wie NIfTI-Bilder gelesen werden.

Der Datensatz enthält MR-Bilder der T1-Modaliät, die sich traditionell gut zur Beurteilung anatomischer Strukturen eignen. In diesem Tutorial arbeitest du mit 3T- und 7T-Gehirn-MRTs.

Der Datensatz ist öffentlich unter dieser Quelle verfügbar.

Initialisierer definieren

Definieren wir zunächst die Datenabmessungen. Die Bildgröße wird beim Einlesen von 173x173 auf 176x176 vergrößert. Hier definieren wir außerdem Datenverzeichnis, Batchgröße fürs Training, Kanalanzahl, eine Input()-Schicht, Train- und Testmatrizen als Listen und laden für das Rescaling die Textdatei mit Minimal- und Maximalwerten des MRT-Datensatzes.

x,y = 173,173
full_z = 207
resizeTo=176
inChannel = outChannel = 1
input_shape=(x,y,inChannel)
input_img = Input(shape = (resizeTo, resizeTo, inChannel))
train_matrix = []
test_matrix = []
ff = os.listdir("../test_crossval1")
save = "../Result_nii_crossval1"
folder_ground = os.listdir("../test_g_crossval1")
ToPredict_images=[]
predict_matrix=[]
ground_images=[]
ground_matrix=[]
min_max=np.loadtxt('../maxANDmin.txt')

Testvolumina laden

Als Nächstes laden wir die MRT-Daten mit nibabel und skalieren die Bilder von 173 x 173 auf 176 x 176 durch Padding mit Nullen in x- und y-Richtung.

Beachte: Beim Laden eines NIfTI-Volumens lädt Nibabel das Array erst, wenn du die Daten anforderst (get_data()).

Da du 2D-Slices statt 3D brauchst, verwenden wir die zuvor initialisierten Listen train und test: Für jedes gelesene Volumen iterieren wir über alle 207 Scheiben und fügen sie der Liste hinzu.

for f in ff:
    temp = np.zeros([resizeTo,full_z,resizeTo])
    a = nib.load("../test_crossval1" + f)
    affine = a.affine
    a = a.get_data()
    temp[3:,:,3:] = a
    a = temp
    for j in range(full_z):
        predict_matrix.append(a[:,j,:])
for f in ff:
    temp = np.zeros([resizeTo,full_z,resizeTo])
    a = nib.load("../test_g_crossval1" + f)
    affine = a.affine
    a = a.get_data()
    temp[3:,:,3:] = a
    a = temp
    for j in range(full_z):
        ground_matrix.append(a[:,j,:])

Datenvorverarbeitung

Da 3T- und 7T-Testmatrizen Listen sind, konvertierst du sie mit NumPy in Arrays.

Dann setzt du den type auf float32 und skalierst sowohl input als auch ground truth.

ToPredict_images = np.asarray(predict_matrix)
ToPredict_images = ToPredict_images.astype('float32')
mx = min_max[0]
mn = min_max[1]
ToPredict_images[:,:,:,0] = (ToPredict_images[:,:,:,0] - mn ) / (mx - mn)
ground_images = np.asarray(ground_matrix)
ground_images = ground_images.astype('float32')
ground_images[:,:,:,0] = (ground_images[:,:,:,0] - mn ) / (mx - mn)

Als Nächstes erstellst du zwei neue Variablen ToPredict_images (3T-Test/Input) und ground_images (7T-Test/Ground Truth) mit dem Shape der Train- und Testmatrix. Das ergibt eine 4D-Matrix: erste Dimension Anzahl Bilder, zweite und dritte Bildgröße, letzte die Kanalzahl (hier 1).

ToPredict_images=np.zeros(shape=[(ToPredict_images.shape[0]),(ToPredict_images.shape[1]),(ToPredict_images.shape[2]),(1)])
ground_images=np.zeros(shape=[(ground_images.shape[0]),(ground_images.shape[1]),(ground_images.shape[2]),(1)])

Dann iterierst du über alle Bilder, formst die Matrizen jeweils zu 176 x 176 und schreibst sie in ToPredict_images (3T Test/Input) bzw. ground_images (7T Test/Ground Truth).

for i in range(ToPredict_images.shape[0]):
    ToPredict_images[i,:,:,0] = ToPredict_images[i,:,:].reshape(resizeTo,resizeTo)
for i in range(ground_images.shape[0]):
    ground_images[i,:,:,0] = ground_images[i,:,:].reshape(resizeTo,resizeTo)

Das Modell!

model
Abbildung: Architektur des Modells
Bild aus diesem Paper.

Encoder

def encoder(input_img):
    conv1 = Conv2D(32, (3, 3), activation='relu', padding='same')(input_img)
    conv1 = BatchNormalization()(conv1)
    conv1 = Conv2D(32, (3,3), activation='relu', padding='same')(conv1)
    conv1 = BatchNormalization()(conv1)
    conv1 = Conv2D(32, (3,3), activation='relu', padding='same')(conv1)
    conv1 = BatchNormalization()(conv1)
    pool1 = MaxPooling2D(pool_size=(2, 2))(conv1)
    conv2 = Conv2D(64, (3, 3), activation='relu', padding='same')(pool1)
    conv2 = BatchNormalization()(conv2)
    conv2 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv2)
    conv2 = BatchNormalization()(conv2)
    conv2 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv2)
    conv2 = BatchNormalization()(conv2)
    pool2 = MaxPooling2D(pool_size=(2, 2))(conv2)
    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same')(pool2)
    conv3 = BatchNormalization()(conv3)
    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv3)
    conv3 = BatchNormalization()(conv3)
    conv3 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv3)
    conv3 = BatchNormalization()(conv3)
    pool3 = MaxPooling2D(pool_size=(2, 2))(conv3)
    conv4 = Conv2D(256, (3, 3), activation='relu', padding='same')(pool3)
    conv4 = BatchNormalization()(conv4)
    conv4 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv4)
    conv4 = BatchNormalization()(conv4)
    conv4 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv4)
    conv4 = BatchNormalization()(conv4)
    conv5 = Conv2D(512, (3, 3), activation='relu', padding='same')(conv4)
    conv5 = BatchNormalization()(conv5)
    conv5 = Conv2D(512, (3, 3), activation='relu', padding='same')(conv5)
    conv5 = BatchNormalization()(conv5)
    conv5 = Conv2D(512, (3, 3), activation='sigmoid', padding='same')(conv5)
    conv5 = BatchNormalization()(conv5)
    return conv5,conv4,conv3,conv2,conv1

Decoder

def decoder(conv5,conv4,conv3,conv2,conv1):
    up6 = merge([conv5, conv4], mode='concat', concat_axis=3)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(up6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    conv6 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv6)
    conv6 = BatchNormalization()(conv6)
    up7 = UpSampling2D((2,2))(conv6)
    up7 = merge([up7, conv3], mode='concat', concat_axis=3)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(up7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    conv7 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv7)
    conv7 = BatchNormalization()(conv7)
    up8 = UpSampling2D((2,2))(conv7)
    up8 = merge([up8, conv2], mode='concat', concat_axis=3)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(up8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    conv8 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv8)
    conv8 = BatchNormalization()(conv8)
    up9 = UpSampling2D((2,2))(conv8)
    up9 = merge([up9, conv1], mode='concat', concat_axis=3)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(up9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    conv9 = Conv2D(32, (3, 3), activation='relu', padding='same')(conv9)
    conv9 = BatchNormalization()(conv9)
    decoded = Conv2D(1, (3, 3), activation='sigmoid', padding='same')(conv9)
    return decoded

Loss-Funktion

Als Nächstes verwendest du den Mean-Square-Error und schließt dabei Pixelwerte aus y_t und y_p aus, die gleich Null sind.

def root_mean_sq_GxGy(y_t, y_p):
    a1=1
    zero = tf.constant(0, dtype=tf.float32)
    where = tf.not_equal(y_t, zero)
    a_t=tf.boolean_mask(y_t,where,name='boolean_mask')
    a_p=tf.boolean_mask(y_p,where,name='boolean_mask')
    return a1*(K.sqrt(K.mean((K.square(a_t-a_p)))))

Modelldefinition und Laden der Gewichte in allen drei Autoencodern

Zuerst rufen wir die encoder-Funktion mit dem Input auf. Da du Merge-Verbindungen nutzt, gibt die encoder-Funktion die Ausgaben von fünf convolution-Layern zurück, die du dann mit allen drei Decodern zusammenführst.

conv5,conv4,conv3,conv2,conv1 = encoder(input_img)

Jetzt erstellen wir drei Modelle und laden die trainierten Gewichte.

autoencoder_1 = Model(input_img, decoder(conv5,conv4,conv3,conv2,conv1))

autoencoder_1.load_weights("../Model/CROSSVAL1/CROSSVAL1_AE1.h5")
autoencoder_2 = Model(input_img, decoder(conv5,conv4,conv3,conv2,conv1))

autoencoder_2.load_weights("../Model/CROSSVAL1/CROSSVAL1_AE2.h5")
autoencoder_3 = Model(input_img, decoder(conv5,conv4,conv3,conv2,conv1))

autoencoder_3.load_weights("../Model/CROSSVAL1/CROSSVAL1_AE3.h5")

Vorhersagen auf Testvolumina: Quantitative und qualitative Ergebnisse

Initialisiere zwei NumPy-Arrays mit je 11 x 3 x 1. Die erste Dimension sind die 11 Testvolumina. Die zweite Dimension repräsentiert MSE und PSNR für: Vorhersage vs. 7T-Ground-Truth, 3T-Input vs. 7T-Ground-Truth, 3T-Input vs. Vorhersage. Die dritte Dimension steht für die Kanalzahl des Modelleingangs.

mse= np.zeros([11,3,1])
psnr= np.zeros([11,3,1])
i=0 #for iterating over the slices of all the 11 volumes

Im nächsten Schritt iterierst du über alle 11 Volumina. Pro Volumen sagst du mit allen drei Autoencodern voraus und bildest anschließend den Mittelwert der Vorhersagen.

Innerhalb eines Volumens iterierst du über die Kanäle und berechnest PSNR und MSE für die drei genannten Fälle.

Mit nibabel speicherst du anschließend die Vorhersage, das Input (3T) und die Ground Truth (7T) als .nii-Dateien: Jedes der 11 Volumina hat 207 Scheiben.

Zum Schluss speicherst du die PSNR-Matrix mit NumPy in einer Textdatei.

Wie im Paper beschrieben, reduziert das Mitteln der Vorhersagen Rauscheffekte und erhält lokale Merkmale in den rekonstruierten Bildern, wodurch sich der PSNR gegenüber einzelnen Decoder-Ausgaben verbessert.

for j in range(11):
    decoded_imgs_1 = autoencoder_1.predict(ToPredict_images[i:i+207,:,:,:])
    decoded_imgs_2 = autoencoder_2.predict(ToPredict_images[i:i+207,:,:,:])
    decoded_imgs_3 = autoencoder_3.predict(ToPredict_images[i:i+207,:,:,:])
    decoded_imgs = np.mean( np.array([ decoded_imgs_1, decoded_imgs_2,decoded_imgs_3 ]), axis=0 )
    for channel in range(1):
        mse[j,0,channel]=  np.mean((ground_images[i:i+207,:,:,channel] - decoded_imgs[:,:,:,channel]) ** 2)
        psnr[j,0,channel] = 20 * math.log10( 1.0 / math.sqrt(mse[j,0,channel]))
        mse[j,1,channel]=  np.mean((ground_images[i:i+207,:,:,channel] - ToPredict_images[i:i+207,:,:,channel])** 2)
        psnr[j,1,channel] = 20 * math.log10( 1.0 / math.sqrt(mse[j,1,channel]))
        mse[j,2,channel] =  np.mean((ToPredict_images[i:i+207,:,:,channel] - decoded_imgs[:,:,:,channel]) ** 2)
        checklt = math.sqrt(mse[j,2,channel])
        psnr[j,2,channel] = 20 * math.log10( 1.0 / math.sqrt(mse[j,2,channel]))
    obj = nib.Nifti1Image(decoded_imgs, affine)
    string =str(j)+'_crossval1.nii'
    nib.save(obj, save + string)
    obj = nib.Nifti1Image(ground_images[i:i+207,:,:,:], affine)
    string =str(j)+'_ground_images_crossval1.nii'
    nib.save(obj, save + string)
    obj = nib.Nifti1Image(ToPredict_images[i:i+207,:,:,:], affine)
    string =str(j)+'_ToPredict_images_crossval1.nii'
    nib.save(obj, save + string)
    i=i+207


np.savetxt('psnr_all_slices.txt',psnr[:,:,0])

Wenn du mehr über Python lernen möchtest, schau dir den DataCamp-Kurs Introduction to Data Visualization with Matplotlib an.

Außerdem empfehlenswert: Keras Tutorial: Deep Learning in Python.

Themen
Python
Maschinelles Lernen
Deep Learning

Mehr über Python und Deep Learning

Kurs

Introduction to Deep Learning with Keras

4 Std.
46.3K
Learn to start developing deep learning models with Keras.
Details anzeigenRight Arrow
Kurs Starten
Mehr anzeigenRight Arrow