Skip to content

Latest commit

 

History

21 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Visual-Steering-For-Deep-Neural-Networks

Abstract

This repository contains the source code to visually steer the CNN models into alligning with the human knowledge fed into them. This method is provided for integrating human knowledge with the convolutional neural network to improve its performance, reduce the biases that arise, and leverage human experience. [paper]

Citation

If you find this repository useful, please cite the following reference.

@InProceedings{10.1007/978-3-031-43247-7_4,
author="Elbeialy, Hossam and Ghaith, Arwa and ElBaradei, Mai H. and Alkhodary, Mohamed and Eissa, Alaa M. and Farag, Esraa A. and Fdawey, Omnia and Emad, Noran and Elshenawy, Ahmed and Shabaan, Mariam and Gamal, Marwa",
editor="Hassanien, AboulElla and Rizk, Rawya Y. and Pamucar, Dragan and Darwish, Ashraf and Chang, Kuo-Chi",
title="Visual Steering for Deep Neural Networks Using Explainable Artificial Intelligence",
booktitle="Proceedings of the 9th International Conference on Advanced Intelligent Systems and Informatics 2023",
year="2023",
publisher="Springer Nature Switzerland",
address="Cham",
pages="43--52",
abstract="Despite their remarkable revolution in artificial intelligence and promising performance, deep neural networks are susceptible to biases due to imbalanced training data or the way that training data is represented as features for the model; thus, they need a high level of transparency in critical fields, especially in the medical sector. Several explainability methods have been developed, aiming at enhancing the transparency of the model and increasing its trustworthiness. One of these methods is Grad-CAM, which can highlight the significant parts of an input that most influence the model's prediction and detect biases in the dataset. In this paper, a method is provided for integrating human knowledge with the deep neural network to improve its performance, reduce the biases that arise, and leverage human experience. This can be done via the attention maps generated by Grad-CAM, which are then edited by domain experts. The proposed method enables the edited attention maps to contribute to the training of the model and updating its parameters, leading to better generalisation and fewer inductive biases. This approach is applied to several image classification datasets: Imagenette2, BUSI, and skin cancer ISIC. Experimental results show that the generated attention maps highly correspond to the edited ones while having high classification performance.",
isbn="978-3-031-43247-7"
}

Methodology

The proposed approach is divided into 3 stages: explaining the model, manually editing the attention maps, and embedding human knowledge using a specific loss. First, Grad-CAM was used to generate attention maps describing the importance values of each image region according to the model to make a specific decision. These attention maps were then annotated by domain experts to represent what they believed were the image regions important to make that decision. Finally, the annotated attention maps were used to steer the model into aligning its behaviour with the human behaviour hence, embedding the human knowledge into the learning process of the model. methodology diagram

Enviroment

to install the required packages run the following command:

pip install -r requirements.txt

Execution

to run the code, you need to run the following command:

python Training.py [arguments]

Model-related Arguments

Argument Description Default Value
-m, --model Model to train, check parameters.py for the list of supported models
--pretrained Use pretrained model's weights from torchvision True
--pretrained_weights_path Path to pretrained model's weights
--model_path Path to save the trained model (full model not the model weights)
--gamma Gamma value of attention loss 10

Training-related Arguments

Argument Description Default Value
-b, --batch_size Batch size for training 32
-e, --epochs Number of epochs for training 100
--lr Learning rate for training 0.001
-s, --seed Random seed for reproducibility 42
-l, --loss Loss function to use, check parameters.py for the list of supported loss functions
-o, --optimizer Optimizer to use, check parameters.py for the list of supported optimizers
--class_weights Use class weights to balance the loss function

Data-related Arguments

Argument Description Default Value
--image_folder Path to the folder containing the images imagenette2/Images/Train
--masks_folder Path to the folder containing the masks imagenette2/Masks/Train
--image_size Size of the images (224,224)
--test_image_folder Path to the folder containing the test images imagenette2All/val
-n, --num_classes Number of classes in the dataset 10

Optimizer-related Arguments

Argument Description Default Value
--momentum Momentum value of optimizer for SGD and RMSprop 0.9
--weight_decay Weight decay value of optimizer 5e-4

Other Arguments

Argument Description Default Value
--out_path Path to save the model results
--gpu_id Id of the GPU to use 0

Datasets

We applied this approach on 3 datasets:

Models

We applied this approach on several models from the ResNet family:

  • ResNet50
  • ResNet101
  • ResNet152

Comparison of Accuracy, AUC, and Qualitative assessments of attention maps for Imagenette2 Dataset while using the annotation of 3948 images.

Approach Model Accuracy AUC Qualitative Measure
Without ResNet50 99.49% 0.9999 3.4084
ResNet101 97.73% 0.9995 3.5392
ResNet152 98.22% 0.9997 3.2401
ABN ResNet50 97.91% 0.9994 2.3048
ResNet101 96.46% 0.9982 2.7847
ResNet152 98.32% 0.9997 2.4483
Proposed Approach ResNet50 97.58% 0.9995 1.8623
ResNet101 98.11% 0.9996 3.1043
ResNet152 97.89% 0.9991 2.4673

Comparison of F2 score and Qualitative assessments of attention maps for BUSI dataset while using the annotation of the whole dataset.

Approach Model F2 Score Qualitative Measure
Without ResNet50 0.8298 1.6049
ResNet101 0.7736 1.8655
ResNet152 0.7860 1.3899
ABN ResNet50 0.8238 2.0944
ResNet101 0.8066 1.9425
ResNet152 0.7720 2.4787
Proposed Approach ResNet50 0.8599 1.3292
ResNet101 0.8429 1.4367
ResNet152 0.8225 1.3881

Comparison of F2 score and Qualitative assessments of attention maps for ISIC dataset while using the annotation of the whole dataset.

Approach Model F2 Score Qualitative Measure
Without ResNet50 0.7061 2.9413
ResNet101 0.5715 3.5000
ResNet152 0.7156 2.9188
ABN ResNet50 0.7207 2.2736
ResNet101 0.6137 2.5259
ResNet152 0.6961 2.2733
Proposed Approach ResNet50 0.7405 1.6235
ResNet101 0.7014 2.8837
ResNet152 0.7126 2.8722

About

No description, website, or topics provided.

Resources

Stars

5 stars

Watchers

1 watching

Forks

Releases

Packages

Used by

Contributors

Languages