File size: 4,953 Bytes
a79dc56
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
119
120
121
122
123
124
125
126
127
128
129
"""
model.py β€” MARBERT backbone + per-aspect softmax classification heads.

Architecture (v2 β€” per-aspect softmax):
    β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”
    β”‚   MARBERT Encoder      β”‚  (pre-trained, 768-dim hidden states)
    β”‚   (12 layers, 12 heads)β”‚
    β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”¬β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜
               β”‚ [CLS] token representation (768-dim)
    β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β–Όβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”
    β”‚   Dropout (0.2)        β”‚
    β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”¬β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜
               β”‚
    β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β–Όβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”
    β”‚  9 independent Linear heads (768 β†’ 4 each) β”‚
    β”‚  food_head, service_head, price_head, ...   β”‚
    β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”¬β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜
               β”‚
    β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β–Όβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”
    β”‚  Softmax per head      β”‚  4 mutually exclusive classes per aspect
    β”‚  [absent, pos, neg, neu]β”‚
    β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜

Why softmax instead of sigmoid:
- A review cannot be BOTH positive and negative for the same aspect.
- Softmax forces the model to choose exactly one sentiment per aspect.
- The "not_mentioned" class is explicitly learned, eliminating
  threshold-based aspect detection.
"""

import torch
import torch.nn as nn
from transformers import AutoModel
from src.config import MODEL_NAME, NUM_ASPECTS, NUM_SENTIMENT_CLASSES, DROPOUT


class ABSAModel(nn.Module):
    """
    Multi-head classifier built on top of a pre-trained MARBERT.

    Each aspect has its own classification head with 4 outputs:
    [not_mentioned, positive, negative, neutral].

    Parameters
    ----------
    model_name : str
        HuggingFace model identifier (default: UBC-NLP/MARBERT)
    num_aspects : int
        Number of aspect heads (default: 9)
    num_classes : int
        Number of sentiment classes per aspect (default: 4)
    dropout : float
        Dropout rate applied to the [CLS] representation.
    """

    def __init__(self, model_name=MODEL_NAME, num_aspects=NUM_ASPECTS,
                 num_classes=NUM_SENTIMENT_CLASSES, dropout=DROPOUT):
        super().__init__()
        self.encoder = AutoModel.from_pretrained(model_name)
        hidden_size = self.encoder.config.hidden_size  # 768 for MARBERT
        self.num_aspects = num_aspects
        self.num_classes = num_classes

        self.dropout = nn.Dropout(dropout)

        # 9 separate classification heads β€” one per aspect
        self.aspect_heads = nn.ModuleList([
            nn.Linear(hidden_size, num_classes)
            for _ in range(num_aspects)
        ])

    def forward(self, input_ids, attention_mask, labels=None,
                class_weights=None):
        """
        Forward pass.

        Parameters
        ----------
        input_ids : tensor (batch, seq_len)
        attention_mask : tensor (batch, seq_len)
        labels : tensor (batch, 9) of int64, optional
            Per-aspect class indices: 0=absent, 1=pos, 2=neg, 3=neu
        class_weights : tensor (9, 4), optional
            Per-aspect class weights for CrossEntropyLoss

        Returns
        -------
        dict with:
            - "logits": list of 9 tensors, each (batch, 4)
            - "probs": list of 9 tensors, each (batch, 4) β€” softmax probs
            - "loss": scalar (only if labels provided)
        """
        outputs = self.encoder(
            input_ids=input_ids,
            attention_mask=attention_mask,
        )

        # [CLS] token is the first token's hidden state
        cls_output = outputs.last_hidden_state[:, 0, :]
        cls_output = self.dropout(cls_output)

        # Run each aspect head
        all_logits = []
        all_probs = []
        for head in self.aspect_heads:
            logits = head(cls_output)          # (batch, 4)
            probs = torch.softmax(logits, dim=-1)  # (batch, 4)
            all_logits.append(logits)
            all_probs.append(probs)

        result = {"logits": all_logits, "probs": all_probs}

        if labels is not None:
            total_loss = 0.0
            for i, logits in enumerate(all_logits):
                aspect_labels = labels[:, i]  # (batch,) int64

                if class_weights is not None:
                    w = class_weights[i].to(logits.device)
                    loss_fn = nn.CrossEntropyLoss(weight=w)
                else:
                    loss_fn = nn.CrossEntropyLoss()

                total_loss += loss_fn(logits, aspect_labels)

            result["loss"] = total_loss / self.num_aspects

        return result