Funciones de Pérdida en CounterfactualExplanations.jl

Introducción a las funciones de pérdida

En CounterfactualExplanations.jl, el cálculo de las funciones de pérdida y sus gradientes se basa en la funcionalidad ya implementada en Flux. Se han importado todas las funciones de pérdida disponibles en Flux, lo que proporciona una amplia flexibilidad para adaptar el modelo a necesidades específicas. La elección de una función de pérdida adecuada es crucial para obtener explicaciones contrafácticas de alta calidad y relevantes para el problema en cuestión. La herramienta permite experimentar con diferentes funciones de pérdida para optimizar el rendimiento del modelo.

Funciones de pérdida comunes

Para la mayoría de las tareas de clasificación, las funciones :logitbinarycrossentropy (binaria) y :logitcrossentropy suelen ser suficientes. Estas opciones han sido probadas y funcionan de manera nativa. Sin embargo, al usar otras funciones, se recomienda precaución. Algunas funciones basadas en márgenes, como la función de hinge loss, no esperan entradas en el dominio {0, 1}, sino en {-1, 1}. Por lo tanto, es necesario codificar las etiquetas de entrenamiento en consecuencia.

Consideraciones para funciones de distancia

El uso de funciones de pérdida basadas en la distancia, como el error cuadrático medio (MSE), requiere que el cálculo se realice con respecto a probabilidades, no a logits. Actualmente, esta opción no está soportada y, por lo general, se recomienda evitar el uso de funciones basadas en la distancia en el contexto de la clasificación, debido a posibles inconsistencias en el proceso de optimización y la generación de explicaciones válidas. Es fundamental comprender las implicaciones de cada función de pérdida para asegurar la coherencia del modelo.

Aplicación en regresión

Actualmente, CounterfactualExplanations.jl está diseñado principalmente para modelos de clasificación, ya que la mayor parte de la literatura sobre explicaciones contrafácticas se centra en este tipo de problemas. Por defecto, se utilizan funciones de pérdida basadas en márgenes, calculadas con respecto a los logits. Para generar explicaciones contrafácticas en problemas de regresión, los usuarios deben binarizar el problema, por ejemplo, definiendo un umbral y asignando etiquetas 0 o 1 en función de si el valor de la variable dependiente está por debajo o por encima del umbral. Se planea agregar soporte completo para problemas de regresión en futuras versiones.

Fundamentos metodológicos: la función de pérdida ℓ

La función de pérdida, denotada como , juega un papel crucial en la búsqueda contrafáctica, dirigiendo el proceso de optimización hacia soluciones que minimicen la desviación entre la etiqueta objetivo y la predicción del modelo. En la práctica, a menudo se implementa con respecto a los logits a = wTx en lugar de las probabilidades p(y′=1|x′) = σ(a) predichas por el clasificador. El estudio de diferentes elecciones para puede conducir a resultados contrafácticos significativamente diferentes.

Ejemplos de funciones de pérdida: hinge loss y logit binary crossentropy

Entre las opciones comunes para se encuentran la función de hinge loss y la función de pérdida de logit binary crossentropy (o log). La función de hinge loss se define como (1 - a ⋅ t*)+ = max{0, 1 - a ⋅ t*}, donde t* es la etiqueta objetivo en {-1, 1}. Para la función de logit binary crossentropy, la fórmula es - (t ⋅ log(σ(a)) + (1-t) ⋅ log(1-σ(a))). Comprender la derivada de estas funciones es esencial para la optimización.