244 lines
8.8 KiB
Python
244 lines
8.8 KiB
Python
"""Code for FineNet in paper "Robust Minutiae Extractor: Integrating Deep Networks and Fingerprint Domain Knowledge" at ICB 2018
|
|
https://arxiv.org/pdf/1712.09401.pdf
|
|
|
|
If you use whole or partial function in this code, please cite paper:
|
|
|
|
@inproceedings{Nguyen_MinutiaeNet,
|
|
author = {Dinh-Luan Nguyen and Kai Cao and Anil K. Jain},
|
|
title = {Robust Minutiae Extractor: Integrating Deep Networks and Fingerprint Domain Knowledge},
|
|
booktitle = {The 11th International Conference on Biometrics, 2018},
|
|
year = {2018},
|
|
}
|
|
"""
|
|
|
|
from __future__ import absolute_import
|
|
from __future__ import division
|
|
|
|
from keras.models import Model
|
|
from keras.layers import Activation, AveragePooling2D, BatchNormalization, Concatenate, Conv2D, Dense, GlobalAveragePooling2D
|
|
from keras.layers import Input, Lambda, MaxPooling2D
|
|
from keras.applications.imagenet_utils import _obtain_input_shape
|
|
from keras import backend as K
|
|
import matplotlib.pyplot as plt
|
|
import numpy as np
|
|
import itertools
|
|
|
|
def preprocess_input(x):
|
|
"""Preprocesses a numpy array encoding a batch of images.
|
|
|
|
"""
|
|
return keras.applications.imagenet_utils.preprocess_input(x, mode='tf')
|
|
|
|
|
|
def conv2d_bn(x,
|
|
filters,
|
|
kernel_size,
|
|
strides=1,
|
|
padding='same',
|
|
activation='relu',
|
|
use_bias=False,
|
|
name=None):
|
|
"""Utility function to apply conv + BN.
|
|
|
|
"""
|
|
x = Conv2D(filters,
|
|
kernel_size,
|
|
strides=strides,
|
|
padding=padding,
|
|
use_bias=use_bias,
|
|
name=name)(x)
|
|
if not use_bias:
|
|
bn_axis = 1 if K.image_data_format() == 'channels_first' else 3
|
|
bn_name = None if name is None else name + '_bn'
|
|
x = BatchNormalization(axis=bn_axis, scale=False, name=bn_name)(x)
|
|
if activation is not None:
|
|
ac_name = None if name is None else name + '_ac'
|
|
x = Activation(activation, name=ac_name)(x)
|
|
return x
|
|
|
|
|
|
def inception_resnet_block(x, scale, block_type, block_idx, activation='relu'):
|
|
"""Inception-ResNet block.
|
|
|
|
"""
|
|
if block_type == 'block35':
|
|
branch_0 = conv2d_bn(x, 32, 1)
|
|
branch_1 = conv2d_bn(x, 32, 1)
|
|
branch_1 = conv2d_bn(branch_1, 32, 3)
|
|
branch_2 = conv2d_bn(x, 32, 1)
|
|
branch_2 = conv2d_bn(branch_2, 48, 3)
|
|
branch_2 = conv2d_bn(branch_2, 64, 3)
|
|
branches = [branch_0, branch_1, branch_2]
|
|
elif block_type == 'block17':
|
|
branch_0 = conv2d_bn(x, 192, 1)
|
|
branch_1 = conv2d_bn(x, 128, 1)
|
|
branch_1 = conv2d_bn(branch_1, 160, [1, 7])
|
|
branch_1 = conv2d_bn(branch_1, 192, [7, 1])
|
|
branches = [branch_0, branch_1]
|
|
elif block_type == 'block8':
|
|
branch_0 = conv2d_bn(x, 192, 1)
|
|
branch_1 = conv2d_bn(x, 192, 1)
|
|
branch_1 = conv2d_bn(branch_1, 224, [1, 3])
|
|
branch_1 = conv2d_bn(branch_1, 256, [3, 1])
|
|
branches = [branch_0, branch_1]
|
|
else:
|
|
raise ValueError('Unknown Inception-ResNet block type. '
|
|
'Expects "block35", "block17" or "block8", '
|
|
'but got: ' + str(block_type))
|
|
|
|
block_name = block_type + '_' + str(block_idx)
|
|
channel_axis = 1 if K.image_data_format() == 'channels_first' else 3
|
|
mixed = Concatenate(axis=channel_axis, name=block_name + '_mixed')(branches)
|
|
up = conv2d_bn(mixed,
|
|
K.int_shape(x)[channel_axis],
|
|
1,
|
|
activation=None,
|
|
use_bias=True,
|
|
name=block_name + '_conv')
|
|
|
|
x = Lambda(lambda inputs, scale: inputs[0] + inputs[1] * scale,
|
|
output_shape=K.int_shape(x)[1:],
|
|
arguments={'scale': scale},
|
|
name=block_name)([x, up])
|
|
if activation is not None:
|
|
x = Activation(activation, name=block_name + '_ac')(x)
|
|
return x
|
|
|
|
def FineNetmodel(num_classes = 2, pretrained_path = None, input_shape = None):
|
|
"""Create FineNet architecture.
|
|
|
|
"""
|
|
# Determine proper input shape
|
|
input_shape = _obtain_input_shape(
|
|
input_shape,
|
|
default_size=299,
|
|
min_size=139,
|
|
data_format=K.image_data_format(),
|
|
require_flatten=False,
|
|
weights=pretrained_path)
|
|
|
|
|
|
img_input = Input(shape=input_shape)
|
|
|
|
# Stem block: 35 x 35 x 192
|
|
x = conv2d_bn(img_input, 32, 3, strides=2, padding='valid')
|
|
x = conv2d_bn(x, 32, 3, padding='valid')
|
|
x = conv2d_bn(x, 64, 3)
|
|
x = MaxPooling2D(3, strides=2)(x)
|
|
x = conv2d_bn(x, 80, 1, padding='valid')
|
|
x = conv2d_bn(x, 192, 3, padding='valid')
|
|
x = MaxPooling2D(3, strides=2)(x)
|
|
|
|
# Mixed 5b (Inception-A block): 35 x 35 x 320
|
|
branch_0 = conv2d_bn(x, 96, 1)
|
|
branch_1 = conv2d_bn(x, 48, 1)
|
|
branch_1 = conv2d_bn(branch_1, 64, 5)
|
|
branch_2 = conv2d_bn(x, 64, 1)
|
|
branch_2 = conv2d_bn(branch_2, 96, 3)
|
|
branch_2 = conv2d_bn(branch_2, 96, 3)
|
|
branch_pool = AveragePooling2D(3, strides=1, padding='same')(x)
|
|
branch_pool = conv2d_bn(branch_pool, 64, 1)
|
|
branches = [branch_0, branch_1, branch_2, branch_pool]
|
|
channel_axis = 1 if K.image_data_format() == 'channels_first' else 3
|
|
x = Concatenate(axis=channel_axis, name='mixed_5b')(branches)
|
|
|
|
# 10x block35 (Inception-ResNet-A block): 35 x 35 x 320
|
|
for block_idx in range(1, 11):
|
|
x = inception_resnet_block(x,
|
|
scale=0.17,
|
|
block_type='block35',
|
|
block_idx=block_idx)
|
|
|
|
# Mixed 6a (Reduction-A block): 17 x 17 x 1088
|
|
branch_0 = conv2d_bn(x, 384, 3, strides=2, padding='valid')
|
|
branch_1 = conv2d_bn(x, 256, 1)
|
|
branch_1 = conv2d_bn(branch_1, 256, 3)
|
|
branch_1 = conv2d_bn(branch_1, 384, 3, strides=2, padding='valid')
|
|
branch_pool = MaxPooling2D(3, strides=2, padding='valid')(x)
|
|
branches = [branch_0, branch_1, branch_pool]
|
|
x = Concatenate(axis=channel_axis, name='mixed_6a')(branches)
|
|
|
|
# 20x block17 (Inception-ResNet-B block): 17 x 17 x 1088
|
|
for block_idx in range(1, 21):
|
|
x = inception_resnet_block(x,
|
|
scale=0.1,
|
|
block_type='block17',
|
|
block_idx=block_idx)
|
|
|
|
# Mixed 7a (Reduction-B block): 8 x 8 x 2080
|
|
branch_0 = conv2d_bn(x, 256, 1)
|
|
branch_0 = conv2d_bn(branch_0, 384, 3, strides=2, padding='valid')
|
|
branch_1 = conv2d_bn(x, 256, 1)
|
|
branch_1 = conv2d_bn(branch_1, 288, 3, strides=2, padding='valid')
|
|
branch_2 = conv2d_bn(x, 256, 1)
|
|
branch_2 = conv2d_bn(branch_2, 288, 3)
|
|
branch_2 = conv2d_bn(branch_2, 320, 3, strides=2, padding='valid')
|
|
branch_pool = MaxPooling2D(3, strides=2, padding='valid')(x)
|
|
branches = [branch_0, branch_1, branch_2, branch_pool]
|
|
x = Concatenate(axis=channel_axis, name='mixed_7a')(branches)
|
|
|
|
# 10x block8 (Inception-ResNet-C block): 8 x 8 x 2080
|
|
for block_idx in range(1, 10):
|
|
x = inception_resnet_block(x,
|
|
scale=0.2,
|
|
block_type='block8',
|
|
block_idx=block_idx)
|
|
x = inception_resnet_block(x,
|
|
scale=1.,
|
|
activation=None,
|
|
block_type='block8',
|
|
block_idx=10)
|
|
|
|
# Final convolution block: 8 x 8 x 1536
|
|
x = conv2d_bn(x, 1536, 1, name='conv_7b')
|
|
|
|
# Classification block
|
|
x = GlobalAveragePooling2D(name='avg_pool')(x)
|
|
x = Dense(num_classes, activation='softmax', name='predictions')(x)
|
|
|
|
|
|
inputs = img_input
|
|
|
|
# Create model
|
|
model = Model(inputs, x, name='FineNet')
|
|
|
|
# Load weights
|
|
if pretrained_path != None:
|
|
print 'Loading FineNet weights from %s'%(pretrained_path)
|
|
model.load_weights(pretrained_path)
|
|
|
|
return model
|
|
|
|
def plot_confusion_matrix(cm, classes,
|
|
normalize=False,
|
|
title='Confusion matrix',
|
|
cmap=plt.cm.Blues):
|
|
"""
|
|
This function prints and plots the confusion matrix.
|
|
Normalization can be applied by setting `normalize=True`.
|
|
"""
|
|
plt.imshow(cm, interpolation='nearest', cmap=cmap)
|
|
plt.title(title)
|
|
plt.colorbar()
|
|
tick_marks = np.arange(len(classes))
|
|
plt.xticks(tick_marks, classes, rotation=45)
|
|
plt.yticks(tick_marks, classes)
|
|
|
|
if normalize:
|
|
cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
|
|
print("Normalized confusion matrix")
|
|
else:
|
|
print('Confusion matrix, without normalization')
|
|
|
|
print(cm)
|
|
|
|
thresh = cm.max() / 2.
|
|
for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):
|
|
plt.text(j, i, cm[i, j],
|
|
horizontalalignment="center",
|
|
color="white" if cm[i, j] > thresh else "black")
|
|
|
|
plt.tight_layout()
|
|
plt.ylabel('True label')
|
|
plt.xlabel('Predicted label')
|
|
plt.show() |