Introduction | Environment | Experiments with Real-world Data | Experiments with Synthetic Data | Contact Us
We explore feature attribution for black-box human sensing models based on radio signals by proposing a Matryoshka-like saliency method, MatryMask. This repository provides the PyTorch implementation of MatryMask as well as eight attribution baselines on real-world and synthetic datasets, including the experiments of quantitative analysis and ablation study.
- Ubuntu 20.04
- Python 3.9.12
- PyTorch 2.0.1
- Captum 0.7.0
- Time Interpret 0.3.0
We suggest creating the environment by installing environment.yaml with Anaconda.
conda env create -f environment.yaml
conda activate ExplainWe evaluate the performance of attribution methods for three human sensing tasks with three black-box models on three real-world datasets.
For human identification, we implement a CNN-1D model (cnn_1d) on the MM-Fi dataset (mmfi).
For human localization, we implement an ABLSTM model (ablstm) on the ARIL dataset (aril).
For human activity recognition, we implement a THAT model (that) on the Office dataset (office).
Run run_model.py to train black-box human sensing models with the following arguments:
--model: select the model to train (default:cnn_1d; options:cnn_1d,ablstm,that)--data: select the dataset (default:mmfi; options:mmfi,aril,office)--n_fold: define the number of folds for cross-validation (default:5)
The following example trains an ABLSTM model on the ARIL dataset with 5-fold cross-validation.
python run_model.py --model ablstm --data arilThe trained model weights are saved in ./result/model.
After training models, we can run evaluate_model.py to evaluate their performance for human sensing.
The following example evaluates the performance of ABLSTM via 5-fold cross-validation.
python evaluate_model.py --model ablstm --data arilFor the feature attribution of trained black-box models, we implement MatryMask (matrymask) in comparison with Lime (lime), Retain (retain), Integral Gradient (ig), GradSHAP (gs), Feature Occlusion (fo), Augmented Feature Occlusion (afo), DynaMask (dynamask), and ExtrMask (extrmask).
The default hyperparameters of different attribution methods are in preset.py.
Specifically, the settings of MatryMask are defined by preset["explainer"]["matrymask"] in the preset.py file.
preset["explainer"]["matrymask"]["perturbation"]: define the perturbations (default:frequency_removal; options:frequency_removal,average,moving_average,past_moving_average,gaussian_blur)preset["explainer"]["matrymask"]["error"]: define the error term (default:cross_entropy; options:cross_entorpy,mean_squared_error)preset["explainer"]["matrymask"]["reg"]: define the regularization term (default:vecsort; options:vecsort,l1)preset["explainer"]["matrymask"]["mode"]: define the preservation/deletion mode (default:preserve; options:preserve,delete)preset["explainer"]["matrymask"]["num_mask"]: define the number of masks (default:6)preset["explainer"]["matrymask"]["ratio_keep_range"]: define the range of a^n (default:(0.2, 0.8))preset["explainer"]["matrymask"]["ratio_frequency_remove"]: define the removal ratio for frequency-removal perturbations (default:0.5)
Run run_explainer.py for feature attribution with the following arguments:
--explainer: select the attribution method (default:matrymask; options:lime,retain,ig,gs,fo,afo,dynamask,extrmask,matrymask)--model: select the model to perform feature attribution--data: select the dataset--n_fold: define the number of folds for cross-validation
The following example uses MatryMask to highlight important features for the ABLSTM model on the ARIL dataset with 5-fold cross-validation.
python run_explainer.py --model ablstm --data aril --explainer matrymaskThe learned attribution masks are saved in ./result/saliency.
Run evaluate_explainer.py to evaluate the attribution masks with the following arguments:
--ratio: define the ratio to perturb/preserve most important features (default:0.2)--explainer: select the attribution method--model: select the model--data: select the dataset--n_fold: define the number of folds for cross-validation
The following example evaluates the attribution masks from MatryMask for explaining the ABLSTM model on the ARIL dataset via 5-fold cross validation with the perturbed/preserved ratio of 0.2.
python evaluate_explainer.py --model ablstm --data aril --explainer matrymask --ratio 0.2The evaluation results are saved in ./result.
Following previous works, we use a synthetic dataset based on a hidden Markov model (HMM) to evaluate the performance of MatryMask for the attribution of time-series models.
Run run_hmm.py to train black-box models, highlight important features, and evaluate the attribution results of black-box models.
--run: define the running (options:train,explain,evaluate)--explainer: select the attribution method (only useful if--runisexplainorevaluate)--n_fold: define the number of folds for cross validation (default:5)
By default, we implement a single-layer Recurrent Neural Network (RNN) with Gated Recurrent Units (GRUs) to perform time-series classification.
The hyperparameters of the RNN are defined by preset in run_hmm.py.
The following example trains the RNN model with 5-fold cross-validation.
python run_hmm.py --run trainThe trained model weights are saved in ./result/hmm/model.
We evaluate different attribution methods for the RNN model.
The default hyperparameters of different attribution methods are defined by preset in run_hmm.py.
The following example uses MatryMask to highlight important features for the RNN model with 5-fold cross-validation.
python run_hmm.py --run explain --explainer matrymaskThe learned attribution masks are saved in ./result/hmm/saliency.
Finally, the attribution masks are evaluated.
The following example evaluates the attribution masks from MatryMask via 5-fold cross validation.
python run_hmm.py --run evaluate --explainer matrymaskThe evaluation results are saved in ./result/hmm.
Please feel free to contact us if you have any questions about MatryMask.
- Shuokang Huang: s.huang21@imperial.ac.uk
- Professor Julie A. McCann: j.mccann@imperial.ac.uk