Downloads · 30 days
0
Armaan2340/CF_training_modal
CF_training_modal is a machine learning model from Armaan2340. Use it for the machine learning task on the model card, and read the license before you ship it in a product. The card lists the license as apache-2.0.
[](https://github.com/BLANK-2340/A-Unified-Approach-for-Multimodal-Emotion-Recognition-Using-Counterfactual-Learning.git) [](https://python.org) [](https://pytorch.org)
Downloads · 30 days
0
Access
Public
Updated Dec 8, 2025
Repo size
9.5 GB
Likes
0
Public
Click a slice to open those files.
.pth9.3 GB · 98%
From the Hugging Face model README
This repository contains the official implementation of "A Unified Approach for Multimodal Emotion Recognition Using Counterfactual Learning", a novel framework designed to enhance multimodal emotion recognition through advanced sequential modeling and structured counterfactual learning.
The framework is evaluated on two benchmark datasets for conversational emotion recognition:
data/
├── MELD/
│ ├── train/
│ ├── dev/
│ └── test/
└── IEMOCAP/
├── train/
├── dev/
└── test/
git clone https://github.com/BLANK-2340/A-Unified-Approach-for-Multimodal-Emotion-Recognition-Using-Counterfactual-Learning.git
cd A-Unified-Approach-for-Multimodal-Emotion-Recognition-Using-Counterfactual-Learning
# Create virtual environment
python -m venv venv
source venv/bin/activate # On Windows: venv\Scripts\activate
# Install required packages
pip install -r requirements.txt
# Install additional dependencies for audio processing
pip install librosa soundfile
# Install pre-trained models
python -c "from transformers import RobertaTokenizer, Wav2Vec2Processor; RobertaTokenizer.from_pretrained('roberta-base'); Wav2Vec2Processor.from_pretrained('facebook/wav2vec2-base-960h')"
# Set environment variables
export CUDA_VISIBLE_DEVICES=0
export TOKENIZERS_PARALLELISM=false
# Train the model with default settings
python train.py --dataset MELD --batch_size 16 --epochs 40
# Evaluate trained model
python evaluate.py --model_path checkpoints/best_model.pth --dataset MELD
# Generate predictions
python predict.py --input_file data/sample.csv --output_file results/predictions.csv
# Custom training with specific parameters
python train.py \
--dataset MELD \
--batch_size 16 \
--learning_rate 5e-5 \
--weight_decay 0.01 \
--phase1_epochs 15 \
--phase2_epochs 15 \
--phase3_epochs 10 \
--save_dir checkpoints/
from model import MultimodalEmotionRecognizer
from utils import load_data, preprocess
# Load pre-trained model
model = MultimodalEmotionRecognizer.from_pretrained('checkpoints/best_model.pth')
# Process input data
text, audio, video = preprocess(input_data)
# Predict emotion
prediction = model.predict(text, audio, video)
print(f"Predicted emotion: {prediction}")
You can download our pre-trained model checkpoint from the following Google Drive link:
This checkpoint can be used for inference as demonstrated in the Inference section.
Our approach integrates advanced sequential modeling with a novel three-phase counterfactual learning strategy to address key challenges in multimodal emotion recognition.
Architecture Overview: The framework processes synchronized audio, visual, and textual data through three main components:
Modality-Specific Feature Extraction: Each modality (text, audio, video) is processed through dedicated pipelines:
Progressive BiLSTM Processing: Each modality's features undergo hierarchical temporal modeling with adaptive gating mechanisms that progressively refine representations across multiple layers.
Cross-Modal Fusion: All-pairs bidirectional cross-attention mechanism facilitates rich information exchange between modalities, followed by standardization and MLP-based classification.
Counterfactual Feature Generation (CFG): Generates minimally perturbed feature versions to enhance model robustness through the three-phase training strategy.
Cross-Modal Attention (CMA) enables bidirectional information exchange across all modality pairs. The mechanism computes attention weights between query features from one modality and key-value pairs from another:
Q_{m_i} = W_Q^{m_1} F_{m_i}, \quad K_{m_j} = W_K^{m_2} F_{m_j}, \quad V_{m_j} = W_V^{m_2} F_{m_j}
A_{m_i \rightarrow m_j} = \text{softmax}\left(\frac{Q_{m_i} K_{m_j}^T}{\sqrt{256}}\right) V_{m_j}
Process Flow:
The training strategy progressively incorporates counterfactual reasoning to enhance model robustness:
Phase 1: Base Model Training
\mathcal{L}_{Phase1} = \mathcal{L}_{Focal}(\hat{y}_{P1}, y_{True})
Phase 2: Counterfactual Alignment
\mathcal{L}_{Phase2} = \mathcal{L}_{Focal}(\hat{y}_{P2}, y_{True}) + \lambda_{align}(e) \cdot \mathcal{L}_{align}(H_{org}, H_{cf}, y_{True})
Phase 3: Intention-Guided Refinement
\mathcal{L}_{Phase3} = \mathcal{L}_{Focal} + \lambda_{align} \cdot \mathcal{L}_{align} + \lambda_{intent}(e) \cdot \mathcal{L}_{intent}
Purpose: Models the relationship between original and counterfactual features to understand prediction rationale.
Architecture:
H_{diff} = H_{cf} - H_{org}
\tilde{H}_{diff} = [H_{diff}; H_{org}] \in \mathbb{R}^{B \times 512}
\hat{y}^{int} = \text{MLP}(\tilde{H}_{diff})
Intention Loss: Confidence-weighted cross-entropy with label smoothing, focusing on samples where the base model is already confident (confidence > 0.7).
Progressive BiLSTM with Adaptive Gating:
H_{lstm}^{(i)} = \text{BiLSTM}^{(i)}(F^{(i-1)}) \in \mathbb{R}^{B \times SeqLen \times 2h}
H_{attn}^{(i)} = \text{MultiHeadAttn}(H_{lstm}^{(i)}, H_{lstm}^{(i)}, H_{lstm}^{(i)}) \in \mathbb{R}^{B \times SeqLen \times 2h}
G^{(i)} = \sigma(\text{Linear}(\text{Concat}(T^{(i)}, F^{(i-1)})))
F^{(i)} = G^{(i)} \odot T^{(i)} + (1 - G^{(i)}) \odot F^{(i-1)}
Alignment Loss (InfoNCE-based):
S_{ij} = \frac{\langle \hat{F}_i, \hat{F}_{cf,j} \rangle}{\tau}
\mathcal{L}_{align} = -\text{mean}_i \left( \log \frac{\exp(\sum_{j \in Pos(i)} S_{ij})}{\exp(\sum_{j \in Pos(i)} S_{ij}) + \exp(\sum_{k \in Neg(i)} [\max(S_{ik} + m, 0)])} \right)
Where τ = 0.1 (temperature), m = 0.5 (margin), and Pos(i)/Neg(i) represent positive/negative sample sets.
Our framework achieves state-of-the-art results on both benchmark datasets:
| Dataset | Method | Accuracy (%) | Weighted F1 (%) |
|---|---|---|---|
| MELD | MMGCN | 60.42 | 58.65 |
| DER-GCN | 66.80 | 66.10 | |
| ELR-GNN | 68.70 | 69.90 | |
| AMuSE | 73.28 | 71.32 | |
| Ours | 94.99 | 94.87 | |
| IEMOCAP | MMGCN | 67.40 | 66.22 |
| DER-GCN | 69.70 | 69.40 | |
| ELR-GNN | 70.60 | 70.90 | |
| AMuSE | 74.49 | 73.91 | |
| Ours | 87.11 | 86.68 |
| MELD Dataset | IEMOCAP Dataset |
|---|---|
![]() | ![]() |
![]() | ![]() |
Learning Curve Analysis:
ROC Analysis:
Confusion Matrix Insights:
Cross-Attention t-SNE Analysis:
Progressive BiLSTM t-SNE Analysis:
| Emotion | Accuracy (%) | Precision (%) | Recall (%) | F1-Score (%) |
|---|---|---|---|---|
| Anger | 98.85 | 93.82 | 98.41 | 96.06 |
| Disgust | 99.83 | 98.85 | 100.00 | 99.42 |
| Fear | 99.83 | 98.85 | 100.00 | 99.42 |
| Joy | 97.42 | 90.04 | 92.14 | 91.08 |
| Neutral | 96.04 | 93.05 | 78.13 | 84.94 |
| Sadness | 99.32 | 95.63 | 99.79 | 97.66 |
| Surprise | 98.70 | 94.49 | 96.50 | 95.48 |
| Emotion | Accuracy (%) | Precision (%) | Recall (%) | F1-Score (%) |
|---|---|---|---|---|
| Anger | 97.23 | 82.42 | 98.68 | 89.82 |
| Excitement | 97.88 | 89.87 | 93.42 | 91.61 |
| Fear | 100.00 | 100.00 | 100.00 | 100.00 |
| Frustration | 92.33 | 70.27 | 67.53 | 68.87 |
| Happy | 98.37 | 95.89 | 90.91 | 93.33 |
| Neutral | 92.17 | 74.14 | 56.58 | 64.18 |
| Sad | 96.74 | 85.19 | 89.61 | 87.34 |
| Surprised | 99.51 | 96.25 | 100.00 | 98.08 |
Performance Insights:
├── Counterfactual_Training_Run/ # Main training scripts and implementations
│ ├── train.py # Training pipeline
│ ├── model.py # Model architecture
│ └── losses.py # Loss functions
├── Video_vector/ # Video feature extraction modules
│ ├── frame_extraction.py # Entropy-based frame selection
│ ├── resnet_features.py # ResNet-50 feature extraction
│ └── video_lstm.py # BiLSTM for video sequences
├── Audio_vector/ # Audio feature processing modules
│ ├── wav2vec_features.py # Wav2Vec2 voice embeddings
│ ├── mfcc_extraction.py # MFCC coefficient extraction
│ └── spectral_features.py # Spectral feature computation
├── Text_vector/ # Text feature processing modules
│ ├── roberta_features.py # RoBERTa text embeddings
│ └── text_preprocessing.py # Text preprocessing utilities
├── images/ # Visualization results and figures
│ ├── Architecture diagrams/ # Model architecture illustrations
│ ├── ROC curves/ # ROC analysis across phases
│ ├── Confusion matrices/ # Classification performance matrices
│ ├── t-SNE visualizations/ # Feature space visualizations
│ └── Learning curves/ # Training progress plots
├── model/ # Core model architecture components
│ ├── progressive_bilstm.py # Progressive BiLSTM implementation
│ ├── cross_attention.py # Cross-modal attention mechanism
│ ├── counterfactual_generator.py # CFG module
│ └── intention_predictor.py # Intention prediction module
├── utils/ # Utility functions and helpers
│ ├── data_loader.py # Dataset loading and preprocessing
│ ├── metrics.py # Evaluation metrics
│ └── visualization.py # Plotting and visualization tools
├── config/ # Configuration files
│ ├── meld_config.yaml # MELD dataset configuration
│ └── iemocap_config.yaml # IEMOCAP dataset configuration
├── requirements.txt # Python dependencies
└── README.md # Project documentation
Novel Progressive Architecture: Modality-specific Progressive BiLSTM with internal self-attention and adaptive gating mechanisms for hierarchical temporal modeling that captures complex intra-modal dynamics
Three-Phase Counterfactual Learning: Structured training strategy incorporating:
Comprehensive Cross-Modal Fusion: All-pairs bidirectional attention mechanism enabling rich information exchange between text, audio, and visual modalities through six distinct attention pathways
State-of-the-Art Performance: Achieving significant improvements over existing methods:
Robust Feature Learning: Enhanced discriminability and generalization through counterfactual reasoning, demonstrated via comprehensive t-SNE visualizations showing progressive cluster formation
Armaan Singh
Department of Electronics and Communication Engineering
Delhi Technological University, Delhi, India
Chhavi Dhiman
Department of Electronics and Communication Engineering
Delhi Technological University, Delhi, India
This project is licensed under the MIT License - see the LICENSE file for details.
⭐ Star this repository if you find it helpful!
</div>