Saltar a contenido

Evaluación de los modelos

Acabamos de terminar de entrenar un modelo, hemos aplicado filtros para mejorar sus detecciones y hemos observado que es capaz de detectar correctamente manzanas, naranjas y peras. Sin embargo, todavía no sabemos qué tan bien funciona realmente el modelo.

En esta lección, aprenderemos a evaluar modelos de detección de objetos. Exploraremos el significado de cada métrica de evaluación y qué revela sobre el rendimiento del modelo. Esta comprensión nos ayudará a evaluar la calidad real de nuestro modelo. Además, nos permitirá comparar varios modelos para identificar el mejor, guiándonos en la selección de los parámetros más adecuados para entrenar modelos futuros.

Al entrenar modelos, como hicimos en la lección anterior, se crea automáticamente una carpeta para cada ejecución en la ruta: runs/detection/created_folder, normalmente llamada " train" seguida de un identificador, a menos que se haya especificado un nombre personalizado en la configuración.

Esta carpeta contiene todos los diagramas y gráficos que comentaremos en esta lección, junto con un archivo results.csv que almacena las métricas de evaluación. Dentro de este directorio, también hay una subcarpeta llamada weights donde se guardan los pesos del modelo.

Ten en cuenta que los gráficos mostrados en esta lección pueden diferir ligeramente de los que aparecen en la carpeta de entrenamiento, ya que son versiones simplificadas destinadas a ayudarte a entender los conceptos con mayor claridad.

Añadir la ruta del módulo de ayuda

Primero, añadimos el módulo de ayuda:

addpath('help-module');

Carga de un modelo

Con este código, estamos cargando el modelo entrenado en la lección anterior.

datasetFolder = "datasets/fruits_3_4998/";    
modelPath = fullfile(pwd, 'runs', 'detect', 'train', 'weights');
configFile  = datasetFolder + "data.yaml";

disp("Loading Model....");
Loading Model....
model = utils.loadModel(modelPath, configFile);
disp("Model loaded!");
Model loaded!

Predicción sobre el conjunto de validación

En machine learning, no basta con que un modelo funcione bien con los datos que ha visto durante el entrenamiento. Si evaluamos el modelo usando esos mismos datos, corremos el riesgo de sobreestimar su rendimiento real. Esto ocurre porque el modelo podría haber "memorizado" las respuestas en lugar de aprender realmente a generalizar.

Para determinar si un modelo ha aprendido de verdad, necesitamos probarlo con datos nuevos que no haya visto antes. Esto se conoce como conjunto de validación. Evaluar el modelo con estos datos no vistos nos da una idea mucho más realista de cómo se comportará en situaciones del mundo real.

En el siguiente código, usamos los datos de validación de nuestro dataset para obtener las predicciones del modelo sobre ese conjunto. En las secciones siguientes, usaremos estas predicciones para calcular las métricas de rendimiento del modelo.

splitFolder = datasetFolder + "val";
classNames =  utils.ReadClassNames(configFile);
disp("Calculating predictions...")
Calculating predictions...
predictions = utils.getPredictions(model, splitFolder, classNames);
disp("Done!")
Done!

calculateIoU

En esta sección, incorporamos la función calculateIoU que implementamos en la lección anterior de posprocesamiento.

function iou = calculateIoU(b1, b2)
    % Calcula la Intersección sobre la Unión (IoU) entre dos bounding boxes.
    %
    % Entradas:
    %   b1 - Vector 1x4 que representa la primera bounding box [x1 y1 x2 y2]
    %   b2 - Vector 1x4 que representa la segunda bounding box [x1 y1 x2 y2]
    %
    % Salida:
    %   iou - Valor escalar entre 0 y 1 que representa la IoU.
    %         IoU = 0 si no hay solapamiento.
    %         IoU = 1 si las cajas son idénticas.
    x1 = max(b1(1), b2(1));
    y1 = max(b1(2), b2(2));
    x2 = min(b1(3), b2(3));
    y2 = min(b1(4), b2(4));
    inter = max(0, x2-x1) * max(0, y2-y1);
    area1 = (b1(3)-b1(1))*(b1(4)-b1(2));
    area2 = (b2(3)-b2(1))*(b2(4)-b2(2));
    union = area1 + area2 - inter;
    if union>0
        iou = inter/union;
    else
        iou = 0;
    end
end

Matriz de confusión

Una matriz de confusión resume con qué frecuencia el modelo predice correcta o incorrectamente cada clase para los objetos detectados.

Cada fila representa la clase real, y cada columna representa la clase predicha.

La diagonal contiene las predicciones correctas (clase predicha = clase real).

Los valores fuera de la diagonal indican confusiones entre clases (por ejemplo, predijo "apple" cuando en realidad era "pear").

Para construir la matriz, solo se consideran las detecciones con suficiente confianza e IoU:

  • Verdadero positivo (TP): confianza > umbral e IoU > umbral.
  • Falso positivo (FP): confianza > umbral e IoU < umbral.
  • Falso/Verdadero negativo (FN/TN): confianza < umbral, lo que significa que el modelo no detectó el objeto.

La matriz ayuda a identificar qué clases se confunden con mayor frecuencia y a evaluar el rendimiento por clase.

Es útil para analizar errores a nivel de clase y evaluar el rendimiento del modelo por clase.

El siguiente código calcula la matriz de confusión para el modelo usando la función que has creado anteriormente:

classNames =  utils.ReadClassNames(configFile);
confidenceThreshold = 0.25;

confMat = utils.computeConfusionMatrix(predictions, classNames, confidenceThreshold, @calculateIoU);

El siguiente código muestra la matriz de confusión:

allClassNames = [classNames(:); {'background'}];
utils.displayConfusionMatrix(confMat, allClassNames);

figure_0.png

Ejercicio 2 - Extraer TP, FP, FN de cada clase

Implementa una función que extraiga los valores de verdaderos positivos, falsos positivos y falsos negativos para cada clase.

function [TP, FP, FN] = extractConfusionMatrixValues(confusionMat)
    % Extrae los valores de verdaderos positivos, falsos positivos y falsos negativos para cada clase
    %
    % Entrada:
    %   confusionMat - Matriz de confusión cuadrada donde las filas representan las clases reales
    %                  y las columnas representan las clases predichas.
    %
    % Salidas:
    %   TP - Vector de verdaderos positivos por clase (predicciones correctas)
    %   FP - Vector de falsos positivos por clase (predicciones incorrectas asignadas a la clase)
    %   FN - Vector de falsos negativos por clase (detecciones omitidas de la clase)

    numClasses = size(confusionMat, 1);
    TP = zeros(numClasses, 1);
    FP = zeros(numClasses, 1);
    FN = zeros(numClasses, 1);


end
function [TP, FP, FN] = extractConfusionMatrixValues(confusionMat)
    % Extrae los valores de verdaderos positivos, falsos positivos y falsos negativos para cada clase
    %
    % Entrada:
    %   confusionMat - Matriz de confusión cuadrada donde las filas representan las clases reales
    %                  y las columnas representan las clases predichas.
    %
    % Salidas:
    %   TP - Vector de verdaderos positivos por clase (predicciones correctas)
    %   FP - Vector de falsos positivos por clase (predicciones incorrectas asignadas a la clase)
    %   FN - Vector de falsos negativos por clase (detecciones omitidas de la clase)

    numClasses = size(confusionMat, 1);
    TP = zeros(numClasses, 1);
    FP = zeros(numClasses, 1);
    FN = zeros(numClasses, 1);
    for i = 1:numClasses
        TP(i) = confusionMat(i, i);
        FP(i) = sum(confusionMat(i, :)) - TP(i);
        FN(i) = sum(confusionMat(:, i)) - TP(i);
    end
end

El siguiente código genera tablas de verdaderos positivos, falsos positivos y falsos negativos para visualizar los valores extraídos para cada clase.

[TP, FP, FN] = extractConfusionMatrixValues(confMat);

tableTP = table(string(allClassNames), TP, 'VariableNames', {'ClassName', 'True Positive'});
tableFP = table(string(allClassNames), FP, 'VariableNames', {'ClassName', 'False Positive'});
tableFN = table(string(allClassNames), FN, 'VariableNames', {'ClassName', 'False Negative'});

disp(tableTP)
     ClassName      True Positive
    ____________    _____________

    "apple"              150     
    "orange"              46     
    "pear"               155     
    "background"           0     
disp(tableFP)
     ClassName      False Positive
    ____________    ______________

    "apple"              221      
    "orange"             137      
    "pear"               439      
    "background"          31      
disp(tableFN)
     ClassName      False Negative
    ____________    ______________

    "apple"                3      
    "orange"              38      
    "pear"                34      
    "background"         753      

Precisión

La precisión mide la exactitud del modelo para una clase determinada. Se define como la proporción de verdaderos positivos (TP) respecto al número total de positivos predichos (TP + FP).

Matemáticamente: \(\textrm{Precisión}=\frac{\textrm{TP}}{\textrm{TP}+\textrm{FP}}\)

Hay diferentes formas de calcular la precisión:

  • Precisión por clase: Un valor para cada clase.
  • Precisión global: Precisión general en todas las clases.
  • Precisión balanceada (promediada macro): Promedio de la precisión por clase (excluyendo el background).

El rango de la precisión va de 0 a 1:

  • Precisión = 1 significa que todas las predicciones positivas fueron correctas.
  • Precisión = 0 significa que ninguna de las predicciones positivas fue correcta.

Ejercicio 3 - Calcular la precisión del modelo

En este ejercicio, debes implementar una función de MATLAB para calcular la precisión para cada clase y la precisión global (precisión promediada macro).

Recuerda

  • (/) es división matricial.
  • (./) es división elemento a elemento
function precision = precisionFunc(predictions, classNames, confidenceThreshold, calculateIoU)   
 % Calcula la precisión por clase y la precisión promediada macro a partir de las predicciones.
    % Salida:
    %   precision            - Vector fila de valores de precisión:
    %                          * Un valor por clase (en el mismo orden que classNames).
    %                          * Un valor final adicional que representa la precisión
    %                            promediada macro (media de las precisiones por clase, excluyendo el background).
    confMat = utils.computeConfusionMatrix(predictions, classNames, confidenceThreshold, calculateIoU);
    [TP, FP, FN] = extractConfusionMatrixValues(confMat);


    % Gestiona los NaNs si alguna clase tiene 0 TP+FP
    precisionMacro = mean(precisionWithoutBG(~isnan(precisionWithoutBG)));
    precision = [precisionWithoutBG; precisionMacro];

end

Solución

function precision = precisionFunc(predictions, classNames, confidenceThreshold, calculateIoU)   
 % Calcula la precisión por clase y la precisión promediada macro a partir de las predicciones.
    % Salida:
    %   precision            - Vector fila de valores de precisión:
    %                          * Un valor por clase (en el mismo orden que classNames).
    %                          * Un valor final adicional que representa la precisión
    %                            promediada macro (media de las precisiones por clase, excluyendo el background).
    confMat = utils.computeConfusionMatrix(predictions, classNames, confidenceThreshold, calculateIoU);
    [TP, FP, FN] = extractConfusionMatrixValues(confMat);

    denom = (TP + FP);

    precision = TP ./ denom;
    totalTP = sum(TP);
    total = totalTP + sum(FP);

    if total == 0
        total = NaN;
    end

    precisionWithoutBG = precision(1:end-1);

    % Gestiona los NaNs si alguna clase tiene 0 TP+FP
    precisionMacro = mean(precisionWithoutBG(~isnan(precisionWithoutBG)));
    precision = [precisionWithoutBG; precisionMacro];

end

Tabla de precisión

confidenceThreshold = 0.5;
precision = precisionFunc(predictions, classNames, confidenceThreshold, @calculateIoU);
metricsClassNames = string([classNames(:); {'Overall'}]);
precisionTable = table(metricsClassNames, precision, 'VariableNames', {'ClassName', 'Precision'});
disp(precisionTable)
    ClassName    Precision
    _________    _________

    "apple"        0.4698 
    "orange"      0.24444 
    "pear"        0.23604 
    "Overall"     0.31676 

Recall

El recall mide la sensibilidad del modelo para una clase determinada. Se define como la proporción de verdaderos positivos (TP) respecto al número total de positivos reales (TP + FN).

Matemáticamente: \(\textrm{Recall}=\frac{\textrm{TP}}{\textrm{TP}+\textrm{FN}}\)

Al igual que con la precisión, hay diferentes formas de calcular el recall:

  • Recall por clase: Un valor para cada clase.
  • Recall global: Recall general en todas las clases.
  • Recall balanceado (promediado macro): Promedio del recall por clase (excluyendo el background).

El rango del recall va de 0 a 1:

  • Recall = 1 significa que el modelo encontró todas las muestras positivas reales.
  • Recall = 0 significa que el modelo omitió todos los positivos reales.

Ejercicio 4 - Calcular el recall del modelo

En este ejercicio, debes implementar una función de MATLAB para calcular el recall para cada clase y el recall global (recall promediado macro).

Recuerda

  • (/) es división matricial.
  • (./) es división elemento a elemento
function recall = recallFunc(predictions, classNames, confidenceThreshold, calculateIoU)
    % Calcula el recall por clase y el recall promediado macro a partir de las predicciones.
    % Salida:
    %   recall - Vector fila de recall por clase, con una entrada adicional
    %            al final que representa el recall promediado macro.

    confMat = utils.computeConfusionMatrix(predictions, classNames, confidenceThreshold, calculateIoU);
    [TP, FP, FN] = extractConfusionMatrixValues(confMat);

    % Gestiona los NaNs si alguna clase tiene 0 TP+FN
    recallMacro = mean(recallWithoutBG(~isnan(recallWithoutBG)));
    recall = [recallWithoutBG; recallMacro];
end
function recall = recallFunc(predictions, classNames, confidenceThreshold, calculateIoU)
    % Calcula el recall por clase y el recall promediado macro a partir de las predicciones.
    % Salida:
    %   recall - Vector fila de recall por clase, con una entrada adicional
    %            al final que representa el recall promediado macro.

    confMat = utils.computeConfusionMatrix(predictions, classNames, confidenceThreshold, calculateIoU);
    [TP, FP, FN] = extractConfusionMatrixValues(confMat);

    denom = (TP + FN);
    recall = TP ./ denom;

    totalTP = sum(TP);
    total = totalTP + sum(FN);

    if total == 0
        total = NaN;
    end

    recallWithoutBG = recall(1:end-1);

    % Gestiona los NaNs si alguna clase tiene 0 TP+FN
    recallMacro = mean(recallWithoutBG(~isnan(recallWithoutBG)));
    recall = [recallWithoutBG; recallMacro];
end

Tabla de recall

confidenceThreshold = 0.5;
recall = recallFunc(predictions, classNames, confidenceThreshold, @calculateIoU);
metricsClassNames = string([classNames(:); {'Overall'}]);
recallTable = table(metricsClassNames, recall, 'VariableNames', {'ClassName', 'Recall'});
disp(recallTable)
    ClassName    Recall 
    _________    _______

    "apple"      0.91503
    "orange"     0.52381
    "pear"       0.69312
    "Overall"    0.71065

Ejercicio 5 - Precisión/Recall

En este ejercicio, debes cambiar el valor de confidenceThreshold y responder a las siguientes preguntas.

confidenceThreshold = 0.5;
precision = precisionFunc(predictions, classNames, confidenceThreshold, @calculateIoU);
recall = recallFunc(predictions, classNames, confidenceThreshold, @calculateIoU);
metricsClassNames = string([classNames(:); {'Overall'}]);
metricsTable = table(metricsClassNames, precision(:), recall(:), ...
    'VariableNames', {'ClassName', 'Precision', 'Recall'});
disp(metricsTable)
    ClassName    Precision    Recall 
    _________    _________    _______

    "apple"        0.4698     0.91503
    "orange"      0.24444     0.52381
    "pear"        0.23604     0.69312
    "Overall"     0.31676     0.71065

Pregunta:

¿Qué observaste al cambiar el valor de confidenceThreshold? ¿Por qué ocurre esto?

Respuesta:

Al cambiar el valor de confidenceThreshold, observé un compromiso entre precisión y recall. Concretamente:

  • Aumentar el umbral de confianza produjo una precisión más alta pero un recall más bajo. Esto se debe a que el modelo se vuelve más selectivo y solo realiza predicciones cuando tiene más confianza. Como resultado, reduce los falsos positivos, pero también omite algunos verdaderos positivos.
  • Disminuir el umbral de confianza aumentó el recall, pero redujo la precisión, ya que el modelo realiza más predicciones, capturando más verdaderos positivos, pero también aumentando el número de falsos positivos.

Este comportamiento refleja el típico compromiso precisión-recall, donde mejorar una métrica a menudo conlleva una disminución de la otra. Encontrar el equilibrio adecuado depende de los requisitos específicos de la aplicación.

F1-Score

Como observamos en el ejercicio anterior, existe un compromiso entre precisión y recall. Para resumir ambas métricas en un único valor, usamos el F1-score, que equilibra las dos.

El F1-score es la media armónica de la precisión y el recall, y se define como:

$$ {\textrm{f1}}_{\textrm{score}} =\frac{2\cdot \textrm{presicion}\cdot \textrm{recall}}{\textrm{precision}+\textrm{recall}} $$

Al igual que la precisión y el recall, el F1-score puede calcularse de varias formas:

  • F1-score por clase: Un valor para cada clase.
  • F1-score global: Rendimiento general en todas las clases.
  • F1-score balanceado (promediado macro): Promedio de los F1-score por clase (excluyendo la clase background, si existe).

El rango del F1-score va de 0 a 1:

  • F1-score = 1 significa precisión y recall perfectos.
  • F1-score = 0 significa que la precisión o el recall es cero, y que el modelo falla de alguna forma clave.

Un buen F1-score indica un buen equilibrio entre evitar falsos positivos (alta precisión) y no omitir positivos (alto recall).

Ejercicio 6 - Calcular el F1-Score del modelo

En este ejercicio, debes implementar una función de MATLAB para calcular:

  • El F1-score para cada clase, y
  • El F1-score global usando promediado macro.

Para ello, reutilizarás las funciones precisionFunc y recallFunc que definiste anteriormente.

function f1 = f1ScoreFunc(predictions, classNames, confidenceThreshold, calculateIoU)
    % Calcula el F1-score por clase y promediado macro usando recall y precisión.
    % Salida:
    %   f1 - Vector fila de F1-score por clase, con una entrada adicional
    %        al final que representa el F1-score promediado macro.

    % Obtiene la precisión y el recall
    recall = recallFunc(predictions, classNames, confidenceThreshold, calculateIoU);
    precision = precisionFunc(predictions, classNames, confidenceThreshold, calculateIoU);



    % Concatena el F1 por clase y el macro
    f1 = [f1PerClass; f1Macro];
end
function f1 = f1ScoreFunc(predictions, classNames, confidenceThreshold, calculateIoU)
    % Calcula el F1-score por clase y promediado macro usando recall y precisión.
    % Salida:
    %   f1 - Vector fila de F1-score por clase, con una entrada adicional
    %        al final que representa el F1-score promediado macro.

    % Obtiene la precisión y el recall
    recall = recallFunc(predictions, classNames, confidenceThreshold, calculateIoU);
    precision = precisionFunc(predictions, classNames, confidenceThreshold, calculateIoU);

    % Elimina los valores macro (último elemento)
    recallPerClass = recall(1:end-1);
    precisionPerClass = precision(1:end-1);

    % Calcula el F1-score por clase
    f1PerClass = 2 * (precisionPerClass .* recallPerClass) ./ ...
                 (precisionPerClass + recallPerClass);

    % Gestiona los NaNs cuando precision + recall = 0
    f1PerClass(isnan(f1PerClass)) = NaN;

    % F1 promediado macro (excluyendo background y NaNs)
    f1Macro = mean(f1PerClass(~isnan(f1PerClass)));

    % Concatena el F1 por clase y el macro
    f1 = [f1PerClass; f1Macro];
end

Tabla de F1-Score

confidenceThreshold = 0.53;
precision = precisionFunc(predictions, classNames, confidenceThreshold, @calculateIoU);
recall = recallFunc(predictions, classNames, confidenceThreshold, @calculateIoU);
f1 = f1ScoreFunc(predictions, classNames, confidenceThreshold, @calculateIoU);

metricsClassNames = string([classNames(:); {'Overall'}]);
metricsTable = table(metricsClassNames, precision(:), recall(:), f1(:), ...
    'VariableNames', {'ClassName', 'Precision', 'Recall', 'F1-Score'});
disp(metricsTable)
    ClassName    Precision    Recall     F1-Score
    _________    _________    _______    ________

    "apple"       0.48929     0.89542    0.63279 
    "orange"      0.24852         0.5    0.33202 
    "pear"        0.25097     0.68254      0.367 
    "Overall"     0.32959     0.69265    0.44394 

Curvas

En esta sección, presentamos las siguientes curvas: precisión vs. confianza, recall vs. confianza y F1-score vs. confianza. Estos gráficos proporcionan una visualización más clara del compromiso entre precisión y recall comentado en el Ejercicio 5, e ilustran cómo el F1-score equilibra estas dos métricas.

Curva de precisión / confianza

El siguiente código genera la curva de precisión vs confianza:

utils.plotMetric('Precision', predictions, classNames, @precisionFunc, @calculateIoU); 
Computing plot...
     1    21

figure_1.png

Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.
Warning: Graphics timeout occurred. To share details of this issue with MathWorks technical support, please include that this is an unresponsive graphics client with your service request.

La imagen siguiente muestra ejemplos de cómo puede verse la curva de precisión vs confianza en escenarios buenos, regulares y malos.

image_0.png

Curva de recall / confianza

El siguiente código genera la curva de recall vs confianza:

utils.plotMetric('Recall', predictions, classNames, @recallFunc, @calculateIoU);
Computing plot...
     1    21

figure_2.png

La imagen siguiente muestra ejemplos de cómo puede verse la curva de recall vs confianza en escenarios buenos, regulares y malos.

image_1.png

Curva de F1-Score / confianza

El siguiente código genera la curva de f1-score vs confianza:

utils.plotMetric('F1-score', predictions, classNames, @f1ScoreFunc, @calculateIoU);  
Computing plot...
     1    21

figure_3.png

La imagen siguiente muestra ejemplos de cómo puede verse la curva de f1-score vs confianza en escenarios buenos, regulares y malos.

image_2.png

Curva de precisión / recall y área bajo la curva

En esta sección, usaremos las funciones integradas de MATLAB para calcular y visualizar la curva precisión-recall.

Aunque podríamos implementarlo usando nuestras propias funciones, esta es una gran oportunidad para familiarizarnos con las utilidades de evaluación de detección de objetos de MATLAB.

Cargar el dataset

Empezamos cargando los datos de validación, que incluyen rutas de imágenes y bounding boxes:

data = load("fruitsValidationData.mat");
validationData = data.validationData;
imds = imageDatastore(validationData.imageFilename);
blds = boxLabelDatastore(validationData(:,2:end));
Loading Model....
Model loaded!

Ejecutar la inferencia

Usa el modelo para generar predicciones sobre las imágenes de validación:

results = detect(model,imds,Threshold=0.01);

Evaluar las predicciones del modelo

metrics = evaluateObjectDetection(results, blds);

metrics = evaluateObjectDetection(results,blds);

Calcular y visualizar la curva precisión-recall

La imagen siguiente muestra ejemplos de cómo puede verse la curva de precisión vs recall en escenarios buenos, regulares y malos.

image_3.png

Extraemos los vectores de recall, precisión y puntuación:

[recall,precision,scores] = precisionRecall(metrics);

Ahora trazamos la curva precisión-recall para una clase específica (por ejemplo, la primera):

figure
plot(recall{3},precision{3})
grid on
title("Precision vs Recall");
xlabel("Recall");
ylabel("Precision");

figure_4.png

Resumen

En este punto, hemos cubierto lo siguiente:

  • Qué es la IoU (Intersección sobre la Unión).
  • Cómo funciona la matriz de confusión en la detección de objetos.
  • Cómo calcular precisión, recall y F1-score.
  • Cómo trazar precisión, recall y F1-score vs. confianza.
  • Cómo calcular e interpretar curvas de precisión vs. recall.

Sin embargo, las métricas que hemos calculado hasta ahora no capturan completamente qué tan bien están localizadas las bounding boxes. Para este propósito, en detección de objetos se utiliza una métrica más completa.

Mean Average Precision (mAP)

La mean Average Precision (mAP) es una de las métricas más comunes en detección de objetos para resumir tanto la precisión como la calidad de localización.

La mAP se calcula de la siguiente manera:

  1. Partiendo de la matriz de confusión, se calculan la precisión y el recall del modelo para cada clase.
  2. Al variar el umbral de decisión, se genera la curva precisión-recall correspondiente para cada clase, junto con su área bajo la curva (AUC) asociada. El área bajo cada curva precisión-recall se conoce como Average Precision (AP).
  3. Finalmente, la mean Average Precision (mAP) se calcula promediando los valores AP de todas las clases.

mAP@50 significa que la AP se calcula con un umbral de IoU de 0.50.

mAP@95 se refiere a la AP calculada con un umbral de IoU de 0.95 o, más comúnmente, al promedio de AP calculado en varios umbrales de IoU desde 0.50 hasta 0.95 en incrementos de 0.05.

El rango de la mAP va de 0 a 1:

  • mAP = 1 significa que el modelo tiene precisión y localización perfectas en todas las clases.
  • mAP = 0 significa que el modelo falla completamente al detectar objetos correctamente.

En el siguiente fragmento de código, mostramos la Average Precision para cada clase con un umbral de IoU de 0.50:

ap = averagePrecision(metrics);
disp("ap")
ap
disp(ap)
    0.2110
    0.3773
    0.2517
disp("mAP") 
mAP
disp(sum(ap)/3)
    0.2800

Podemos mostrar métricas resumidas para todo el dataset y para clases individuales usando:

[summaryDataset,summaryClass] = summarize(metrics);
disp(summaryDataset)
    NumObjects    mAPOverlapAvg    mAP0.5 
    __________    _____________    _______

       426           0.27999       0.27999
disp(summaryClass)
              NumObjects    APOverlapAvg     AP0.5 
              __________    ____________    _______

    apple        153          0.21101       0.21101
    orange        84          0.37729       0.37729
    pear         189          0.25167       0.25167