Predictor ciego de calidad de imagen basado en CNN profunda en Python

Introducción

En este tutorial, implementaremos la metodología Deep CNN-Based Blind Image Quality Predictor (DIQA) propuesta por Jongio Kim, Anh-Duc Nguyen y Sanghoon Lee [1]. Además, repasaré los siguientes conceptos de TensorFlow 2.0:

  • Descargar y preparar un conjunto de datos usando un constructor tf.data.Dataset.
  • Definir un pipeline de entrada de TensorFlow para preprocesar los registros del conjunto de datos usando la API tf.data.
  • Crear el modelo CNN usando la API funcional de tf.keras.
  • Definir un bucle de entrenamiento personalizado para el modelo de mapa de error objetivo.
  • Entrenar el modelo de mapa de error objetivo y el modelo de puntuación subjetiva.
  • Usar el modelo de puntuación subjetiva entrenado para hacer predicciones.

Nota: Algunas de las funciones están implementadas en utils.py ya que quedan fuera del alcance de esta guía.

¿Qué es DIQA?

DIQA es una propuesta original que se enfoca en resolver algunos de los retos más importantes de aplicar deep learning a la evaluación de calidad de imagen (IQA). Las ventajas frente a otras metodologías son:

  • El modelo no está limitado a trabajar exclusivamente con imágenes de Estadísticas de Escena Natural (NSS) [1].
  • Previene el sobreajuste al dividir el entrenamiento en dos fases: (1) aprendizaje de características y (2) mapeo de las características aprendidas a puntuaciones subjetivas.

Problema

El costo de generar conjuntos de datos para IQA es alto, ya que requiere supervisión experta. Por lo tanto, los benchmarks fundamentales de IQA están compuestos por solo unos pocos miles de registros. Esto último complica la creación de modelos de deep learning, porque requieren grandes cantidades de muestras de entrenamiento para generalizar.

Por ejemplo, consideremos los conjuntos de datos más usados para entrenar y evaluar métodos de IQA: Live, TID2008, TID2013, CSIQ. Un resumen general de cada conjunto de datos está en la siguiente tabla:

image

La cantidad total de muestras no supera los 4,000 registros en ninguno de ellos.

Conjunto de datos

Los benchmarks de IQA solo contienen una cantidad limitada de registros que podría no ser suficiente para entrenar una CNN. Sin embargo, para el propósito de esta guía, vamos a usar el conjunto de datos Live. Está compuesto por 29 imágenes de referencia y 5 distorsiones distintas con 5 niveles de severidad cada una.

image Fig 1. Un ejemplo de una imagen de referencia en el conjunto de datos Live.

La primera tarea es descargar y preparar el conjunto de datos. He creado un par de constructores de conjuntos de datos de TensorFlow para evaluación de calidad de imagen y los publiqué en el paquete image-quality. Los constructores son una interfaz definida por tensorflow-datasets.

Nota: Este proceso puede tardar varios minutos debido al tamaño del conjunto de datos (700 megabytes).

image

Después de descargar y preparar los datos, convertimos el constructor en un conjunto de datos y lo mezclamos. Nótese que el batch es igual a 1. La razón es que cada imagen tiene una forma distinta. Aumentar el tamaño del batch provocará un error.

image

La salida es un generador; por lo tanto, acceder a las muestras usando el operador de corchetes provoca un error. Hay dos formas de acceder a las imágenes en el generador. La primera es convertir el generador en un iterador y extraer una sola muestra usando la función next.

image

La salida es un diccionario que contiene la representación tensorial de la imagen distorsionada, la imagen de referencia y la puntuación subjetiva (dmos). Otra forma es extraer muestras del generador tomándolas con un bucle for:

image

Metodología

Normalización de imagen

El primer paso de DIQA es preprocesar las imágenes. La imagen se convierte a escala de grises y luego se aplica un filtro paso bajo. El filtro paso bajo se define como:

image

donde la imagen de baja frecuencia es el resultado del siguiente algoritmo:

  1. Difuminar la imagen en escala de grises.
  2. Reducir su escala por un factor de 1/4.
  3. Volver a escalarla al tamaño original.

Las razones principales de esta normalización son (1) el Sistema Visual Humano (HVS) no es sensible a cambios en la banda de baja frecuencia, y (2) las distorsiones de imagen apenas afectan al componente de baja frecuencia de las imágenes.

image

image Fig 2. A la izquierda, la imagen original. A la derecha, la imagen después de aplicar el filtro paso bajo.

Mapa de error objetivo

Para el primer modelo, se usan errores objetivos como proxy para aprovechar el efecto de aumentar los datos. La función de pérdida se define como el error cuadrático medio entre los mapas de error predicho y real.

image

y err(·) puede ser cualquier función de error. Para esta implementación, los autores recomiendan usar

image

con p=0.2. Esto último es para evitar que los valores en el mapa de error sean pequeños o cercanos a cero.

image

image Fig 3. A la izquierda, la imagen original. En el medio, la imagen preprocesada, y finalmente, la representación en imagen del mapa de error.

Mapa de confiabilidad

Según los autores, es probable que el modelo falle al predecir imágenes con regiones homogéneas. Para evitarlo, proponen una función de confiabilidad. El supuesto es que las áreas borrosas tienen menor confiabilidad que las texturizadas. La función de confiabilidad se define como

image

donde α controla la propiedad de saturación del mapa de confiabilidad. La parte positiva de una sigmoide se usa para asignar valores suficientemente grandes a píxeles con baja intensidad.

image

La definición anterior podría afectar directamente la puntuación predicha. Por lo tanto, se usa en su lugar el mapa de confiabilidad promedio.

image

Para la función de Tensorflow, simplemente calculamos el mapa de confiabilidad y lo dividimos entre su media.

image

image Fig 4. A la izquierda, la imagen original, y a la derecha, su mapa de confiabilidad promedio.

Función de pérdida

La función de pérdida se define como el error cuadrático medio del producto entre el mapa de confiabilidad y el mapa de error objetivo. El error es la diferencia entre el mapa de error predicho y el mapa de error real.

image

La función de pérdida requiere multiplicar el error por el mapa de confiabilidad; por lo tanto, no podemos usar la implementación de pérdida por defecto tf.loss.MeanSquareError.

image

Después de crear la pérdida personalizada, necesitamos decirle a TensorFlow cómo diferenciarla. Lo bueno es que podemos aprovechar la diferenciación automática usando tf.GradientTape.

image

Optimizador

Los autores sugirieron usar un optimizador Nadam con una tasa de aprendizaje de 2e-4.

image

Entrenamiento

Modelo de error objetivo

Para la fase de entrenamiento, conviene utilizar los pipelines de entrada de tf.data para producir un código mucho más limpio y legible. El único requisito es crear la función que se aplicará a la entrada.

image

Luego, mapeamos el tf.data.Dataset a la función calculate_error_map.

image

Aplicar la transformación se ejecuta casi de inmediato. La razón es que el procesador aún no realiza ninguna operación sobre los datos; eso ocurre bajo demanda. Este concepto se conoce comúnmente como evaluación perezosa.

Hasta ahora, los siguientes componentes están implementados:

  1. El generador que preprocesa la entrada y calcula el objetivo.
  2. Las funciones de pérdida y gradiente requeridas para el bucle de entrenamiento personalizado.
  3. La función del optimizador.

Lo único que falta es la definición de los modelos.

image Fig 5. La arquitectura para la predicción del mapa de error objetivo. Las flechas roja y azul indican los flujos de la primera y segunda etapa. Fuente: http://bit.ly/2Ldw4PZ

En la imagen anterior, se muestra cómo:

  • La imagen preprocesada entra a la red neuronal convolucional (CNN).
  • Es transformada por 8 convoluciones con la función de activación Relu y padding “same”. Esto se define como f(·).
  • La salida de f(·) es procesada por la última convolución con una función de activación lineal. Esto se define como g(·).

image

Para el bucle de entrenamiento personalizado, es necesario:

  1. Definir una métrica para medir el desempeño del modelo.
  2. Calcular la pérdida y los gradientes.
  3. Usar el optimizador para actualizar los pesos.
  4. Imprimir la precisión.

image

Nota: Sería buena idea usar el coeficiente de correlación de orden de rango de Spearman (SRCC) o el coeficiente de correlación lineal de Pearson (PLCC) como métricas de precisión.

Modelo de puntuación subjetiva

Para crear el modelo de puntuación subjetiva, usemos la salida de f(·) para entrenar un regresor.

image

image

Entrenar un modelo con el método fit de tf.keras.Model espera un conjunto de datos que devuelva dos argumentos. El primero es la entrada y el segundo es el objetivo.

image

Luego, hacemos fit al modelo de puntuación subjetiva.

image

Predicción

Hacer predicciones con el modelo ya entrenado es sencillo. Solo hay que usar el método predict del modelo.

image

Conclusión

En este artículo, aprendimos a utilizar el módulo tf.data para crear pipelines de datos fáciles de leer y eficientes en memoria. Además, implementamos el modelo Deep CNN-Based Blind Image Quality Predictor (DIQA) usando la API funcional de Keras. El modelo se entrenó con un bucle de entrenamiento personalizado que aprovecha la función de diferenciación automática de TensorFlow.

El siguiente paso es encontrar los hiperparámetros que maximicen las métricas de precisión PLCC o SRCC y evaluar el desempeño general del modelo frente a otras metodologías.

Otra idea es usar un conjunto de datos mucho más grande para entrenar el modelo de mapa de error objetivo y observar el desempeño general resultante.

Notebook de Jupyter

Actualización 2020/04/15:* El paquete image-quality y el notebook fueron actualizados para corregir un problema con los conjuntos de datos de TensorFlow LiveIQA y Tid2013. Ahora todo funciona correctamente, échale un vistazo.*

https://github.com/ocampor/image-quality.git

Bibliografía

[1] Kim, J., Nguyen, A. D., & Lee, S. (2019). Deep CNN-Based Blind Image Quality Predictor. IEEE Transactions on Neural Networks and Learning Systems. https://doi.org/10.1109/TNNLS.2018.2829819

← Volver al blog