File size: 5,057 Bytes
402ab4d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
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)