| 57 | |
| 58 | ### Creat model |
| 59 | def get_compiled_model(): |
| 60 | if not args.model_name in ['IRV2', 'ResNet50', 'DenseNet121', 'InceptionV3']: |
| 61 | raise Exception('Pre-trained network not exists. Please choose IRV2/ResNet50/DenseNet121/InceptionV3 instead') |
| 62 | else: |
| 63 | if args.model_name == 'IRV2': |
| 64 | if database == 'RadImageNet': |
| 65 | model_dir ="../RadImageNet_models/RadImageNet-IRV2-notop.h5" |
| 66 | base_model = InceptionResNetV2(weights=model_dir, input_shape=(image_size, image_size, 3), include_top=False,pooling='avg') |
| 67 | else: |
| 68 | base_model = InceptionResNetV2(weights='imagenet', input_shape=(image_size, image_size, 3),include_top=False,pooling='avg') |
| 69 | if args.model_name == 'ResNet50': |
| 70 | if database == 'RadImageNet': |
| 71 | model_dir = "../RadImageNet_models/RadImageNet-ResNet50-notop.h5" |
| 72 | base_model = ResNet50(weights=model_dir, input_shape=(image_size, image_size, 3), include_top=False,pooling='avg') |
| 73 | else: |
| 74 | base_model = ResNet50(weights='imagenet', input_shape=(image_size, image_size, 3), include_top=False,pooling='avg') |
| 75 | if args.model_name == 'DenseNet121': |
| 76 | if database == 'RadImageNet': |
| 77 | model_dir = "../RadImageNet_models/RadImageNet-DenseNet121-notop.h5" |
| 78 | base_model = DenseNet121(weights=model_dir, input_shape=(image_size, image_size, 3), include_top=False,pooling='avg') |
| 79 | else: |
| 80 | base_model = DenseNet121(weights='imagenet', input_shape=(image_size, image_size, 3), include_top=False,pooling='avg') |
| 81 | if args.model_name == 'InceptionV3': |
| 82 | if database == 'RadImageNet': |
| 83 | model_dir = "../RadImageNet_models/RadImageNet-InceptionV3-notop.h5" |
| 84 | base_model = InceptionV3(weights=model_dir, input_shape=(image_size, image_size, 3), include_top=False,pooling='avg') |
| 85 | else: |
| 86 | base_model = InceptionV3(weights='imagenet', input_shape=(image_size, image_size, 3), include_top=False,pooling='avg') |
| 87 | if args.structure == 'freezeall': |
| 88 | for layer in base_model.layers: |
| 89 | layer.trainable = False |
| 90 | if args.structure == 'unfreezeall': |
| 91 | pass |
| 92 | if args.structure == 'unfreezetop10': |
| 93 | for layer in base_model.layers[:-10]: |
| 94 | layer.trainable = False |
| 95 | y = base_model.output |
| 96 | y = Dropout(0.5)(y) |
| 97 | predictions = Dense(2, activation='softmax')(y) |
| 98 | model = Model(inputs=base_model.input, outputs=predictions) |
| 99 | adam = Adam(lr=args.lr) |
| 100 | model.compile(optimizer=adam, loss=BinaryCrossentropy(), metrics=[keras.metrics.AUC(name='auc')]) |
| 101 | return model |
| 102 | |
| 103 | |
| 104 | def run_model(): |