aimen-ameer's picture
Upload folder using huggingface_hub
402ab4d verified
Raw
History Blame Contribute Delete
5.06 kB
import torch
import torch.nn as nn
from transformers import ViTModel, BertModel
# -------------------------------------------------------------------
# 1. Configuration
# -------------------------------------------------------------------
class Config:
"""Central configuration for model hyperparameters"""
def __init__(self):
self.VIT_MODEL_NAME = 'google/vit-base-patch16-224-in21k'
self.BERT_MODEL_NAME = 'bert-base-uncased'
self.IMAGE_SIZE = (224, 224)
self.BERT_MAX_LENGTH = 128
self.VIT_EMBED_DIM = 768
self.BERT_EMBED_DIM = 768
self.FUSION_NUM_HEADS = 8
self.FUSION_NUM_LAYERS = 2
self.FUSION_FF_DIM = 2048
self.DROPOUT_RATE = 0.3
self.NUM_CLASSES = 6
self.CLASS_NAMES = ['True', 'Satire', 'Misleading', 'Imposter',
'False_Connection', 'Manipulated']
self.DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# -------------------------------------------------------------------
# 2. Fusion Layer
# -------------------------------------------------------------------
class CrossAttentionFusion(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.img_to_text_attn = nn.MultiheadAttention(
embed_dim=config.VIT_EMBED_DIM,
num_heads=config.FUSION_NUM_HEADS,
dropout=config.DROPOUT_RATE,
batch_first=True
)
self.text_to_img_attn = nn.MultiheadAttention(
embed_dim=config.BERT_EMBED_DIM,
num_heads=config.FUSION_NUM_HEADS,
dropout=config.DROPOUT_RATE,
batch_first=True
)
self.norm_img1 = nn.LayerNorm(config.VIT_EMBED_DIM, eps=1e-6)
self.norm_img2 = nn.LayerNorm(config.VIT_EMBED_DIM, eps=1e-6)
self.norm_text1 = nn.LayerNorm(config.BERT_EMBED_DIM, eps=1e-6)
self.norm_text2 = nn.LayerNorm(config.BERT_EMBED_DIM, eps=1e-6)
self.ffn_img = nn.Sequential(
nn.Linear(config.VIT_EMBED_DIM, config.FUSION_FF_DIM),
nn.GELU(), nn.Dropout(config.DROPOUT_RATE),
nn.Linear(config.FUSION_FF_DIM, config.VIT_EMBED_DIM),
)
self.ffn_text = nn.Sequential(
nn.Linear(config.BERT_EMBED_DIM, config.FUSION_FF_DIM),
nn.GELU(), nn.Dropout(config.DROPOUT_RATE),
nn.Linear(config.FUSION_FF_DIM, config.BERT_EMBED_DIM),
)
self.dropout_img = nn.Dropout(config.DROPOUT_RATE)
self.dropout_text = nn.Dropout(config.DROPOUT_RATE)
def forward(self, image_features, text_features):
# Image attends to text
img_cross_attn, _ = self.img_to_text_attn(
query=image_features, key=text_features, value=text_features, need_weights=False
)
image_features = self.norm_img1(image_features + self.dropout_img(img_cross_attn))
image_features = self.norm_img2(image_features + self.ffn_img(image_features))
# Text attends to image
text_cross_attn, _ = self.text_to_img_attn(
query=text_features, key=image_features, value=image_features, need_weights=False
)
text_features = self.norm_text1(text_features + self.dropout_text(text_cross_attn))
text_features = self.norm_text2(text_features + self.ffn_text(text_features))
return image_features, text_features
# -------------------------------------------------------------------
# 3. Main Model
# -------------------------------------------------------------------
class MultimodalFakeNewsDetector(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.vision_encoder = ViTModel.from_pretrained(config.VIT_MODEL_NAME)
self.text_encoder = BertModel.from_pretrained(config.BERT_MODEL_NAME)
self.fusion_layers = nn.ModuleList([
CrossAttentionFusion(config) for _ in range(config.FUSION_NUM_LAYERS)
])
self.classifier = nn.Sequential(
nn.Linear(config.VIT_EMBED_DIM + config.BERT_EMBED_DIM, 1024),
nn.GELU(), nn.Dropout(config.DROPOUT_RATE), nn.LayerNorm(1024, eps=1e-6),
nn.Linear(1024, 512),
nn.GELU(), nn.Dropout(config.DROPOUT_RATE), nn.LayerNorm(512, eps=1e-6),
nn.Linear(512, config.NUM_CLASSES)
)
def forward(self, pixel_values, input_ids, attention_mask):
vit_outputs = self.vision_encoder(pixel_values=pixel_values)
image_features = vit_outputs.last_hidden_state
bert_outputs = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask)
text_features = bert_outputs.last_hidden_state
for fusion_layer in self.fusion_layers:
image_features, text_features = fusion_layer(image_features, text_features)
img_repr = image_features.mean(dim=1)
text_repr = text_features.mean(dim=1)
combined = torch.cat([img_repr, text_repr], dim=1)
return self.classifier(combined)