Spaces:
Sleeping
Sleeping
Commit ·
317ce51
1
Parent(s): ce1a6a7
Force real BERT tensor loading in Spaces
Browse files- src/models.py +31 -6
src/models.py
CHANGED
|
@@ -1,4 +1,4 @@
|
|
| 1 |
-
|
| 2 |
|
| 3 |
Three architectures:
|
| 4 |
- BertMetaFusionACSAModel: BERT + Meta Cross-Attention + per-aspect heads (Proposed)
|
|
@@ -47,6 +47,14 @@ logger = logging.getLogger(__name__)
|
|
| 47 |
META_NUM_META_TOKENS = 4
|
| 48 |
|
| 49 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 50 |
# ---------------------------------------------------------------------------
|
| 51 |
# Shared helpers
|
| 52 |
# ---------------------------------------------------------------------------
|
|
@@ -201,6 +209,17 @@ class GatedAspectSemanticCrossAttention(nn.Module):
|
|
| 201 |
self.out_proj = nn.Linear(hidden, hidden)
|
| 202 |
self.gate = nn.Linear(hidden * 2, hidden)
|
| 203 |
nn.init.constant_(self.gate.bias, -0.3)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 204 |
self.dropout = nn.Dropout(dropout)
|
| 205 |
self.norm = nn.LayerNorm(hidden)
|
| 206 |
|
|
@@ -211,7 +230,13 @@ class GatedAspectSemanticCrossAttention(nn.Module):
|
|
| 211 |
q = q.transpose(1, 2)
|
| 212 |
k = self.k_proj(meta_tokens).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
|
| 213 |
v = self.v_proj(meta_tokens).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
|
| 214 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 215 |
ctx = torch.matmul(self.dropout(attn), v)
|
| 216 |
ctx = ctx.transpose(1, 2).contiguous().view(B, self.num_aspects, H)
|
| 217 |
ctx = self.out_proj(ctx)
|
|
@@ -291,7 +316,7 @@ class BertMetaFusionACSAModel(nn.Module):
|
|
| 291 |
class_weights: Optional[torch.Tensor] = None,
|
| 292 |
):
|
| 293 |
super().__init__()
|
| 294 |
-
self.bert =
|
| 295 |
hidden = self.bert.config.hidden_size
|
| 296 |
|
| 297 |
self.num_aspects = num_aspects
|
|
@@ -391,7 +416,7 @@ class GatedAspectSemanticMetaFusionACSAModel(nn.Module):
|
|
| 391 |
expected_dim = cfg.META_TFIDF_DIM + cfg.META_NUM_DIM
|
| 392 |
if meta_in_dim != expected_dim:
|
| 393 |
raise ValueError(f"Semantic meta fusion expects meta_in_dim={expected_dim}, got {meta_in_dim}")
|
| 394 |
-
self.bert =
|
| 395 |
hidden = self.bert.config.hidden_size
|
| 396 |
|
| 397 |
self.num_aspects = num_aspects
|
|
@@ -546,7 +571,7 @@ class BertACSAModel(nn.Module):
|
|
| 546 |
class_weights: Optional[torch.Tensor] = None,
|
| 547 |
):
|
| 548 |
super().__init__()
|
| 549 |
-
self.bert =
|
| 550 |
hidden = self.bert.config.hidden_size
|
| 551 |
self.num_aspects = num_aspects
|
| 552 |
self.num_classes = num_classes
|
|
@@ -581,7 +606,7 @@ class BertOverallModel(nn.Module):
|
|
| 581 |
num_classes: int = cfg.OVERALL_NUM_CLASSES,
|
| 582 |
dropout: float = 0.3):
|
| 583 |
super().__init__()
|
| 584 |
-
self.bert =
|
| 585 |
hidden = self.bert.config.hidden_size
|
| 586 |
self.classifier = nn.Sequential(
|
| 587 |
nn.Dropout(dropout),
|
|
|
|
| 1 |
+
"""Models for ACSA with metadata fusion.
|
| 2 |
|
| 3 |
Three architectures:
|
| 4 |
- BertMetaFusionACSAModel: BERT + Meta Cross-Attention + per-aspect heads (Proposed)
|
|
|
|
| 47 |
META_NUM_META_TOKENS = 4
|
| 48 |
|
| 49 |
|
| 50 |
+
def _load_bert_model(bert_name: str):
|
| 51 |
+
"""Load BERT with real tensors in Spaces instead of lazy meta tensors."""
|
| 52 |
+
try:
|
| 53 |
+
return AutoModel.from_pretrained(bert_name, low_cpu_mem_usage=False)
|
| 54 |
+
except TypeError:
|
| 55 |
+
return AutoModel.from_pretrained(bert_name)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
# ---------------------------------------------------------------------------
|
| 59 |
# Shared helpers
|
| 60 |
# ---------------------------------------------------------------------------
|
|
|
|
| 209 |
self.out_proj = nn.Linear(hidden, hidden)
|
| 210 |
self.gate = nn.Linear(hidden * 2, hidden)
|
| 211 |
nn.init.constant_(self.gate.bias, -0.3)
|
| 212 |
+
prior_rows = []
|
| 213 |
+
prior_cfg = getattr(cfg, "CROSS_ATTN_SOURCE_PRIOR", {})
|
| 214 |
+
for aspect in cfg.ASPECTS[:num_aspects]:
|
| 215 |
+
row = list(prior_cfg.get(aspect, [1.0 / 3.0] * 3))
|
| 216 |
+
if len(row) < 3:
|
| 217 |
+
row = row + [0.0] * (3 - len(row))
|
| 218 |
+
prior_rows.append(row[:3])
|
| 219 |
+
prior = torch.tensor(prior_rows, dtype=torch.float)
|
| 220 |
+
prior = prior - prior.mean(dim=-1, keepdim=True)
|
| 221 |
+
self.register_buffer("source_prior_bias", prior, persistent=False)
|
| 222 |
+
self.source_prior_strength = float(getattr(cfg, "CROSS_ATTN_SOURCE_PRIOR_STRENGTH", 0.0))
|
| 223 |
self.dropout = nn.Dropout(dropout)
|
| 224 |
self.norm = nn.LayerNorm(hidden)
|
| 225 |
|
|
|
|
| 230 |
q = q.transpose(1, 2)
|
| 231 |
k = self.k_proj(meta_tokens).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
|
| 232 |
v = self.v_proj(meta_tokens).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
|
| 233 |
+
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
|
| 234 |
+
if self.source_prior_strength > 0 and T == self.source_prior_bias.shape[-1]:
|
| 235 |
+
prior = self.source_prior_bias[:self.num_aspects, :T].view(
|
| 236 |
+
1, 1, self.num_aspects, T
|
| 237 |
+
)
|
| 238 |
+
scores = scores + self.source_prior_strength * prior
|
| 239 |
+
attn = F.softmax(scores, dim=-1)
|
| 240 |
ctx = torch.matmul(self.dropout(attn), v)
|
| 241 |
ctx = ctx.transpose(1, 2).contiguous().view(B, self.num_aspects, H)
|
| 242 |
ctx = self.out_proj(ctx)
|
|
|
|
| 316 |
class_weights: Optional[torch.Tensor] = None,
|
| 317 |
):
|
| 318 |
super().__init__()
|
| 319 |
+
self.bert = _load_bert_model(bert_name)
|
| 320 |
hidden = self.bert.config.hidden_size
|
| 321 |
|
| 322 |
self.num_aspects = num_aspects
|
|
|
|
| 416 |
expected_dim = cfg.META_TFIDF_DIM + cfg.META_NUM_DIM
|
| 417 |
if meta_in_dim != expected_dim:
|
| 418 |
raise ValueError(f"Semantic meta fusion expects meta_in_dim={expected_dim}, got {meta_in_dim}")
|
| 419 |
+
self.bert = _load_bert_model(bert_name)
|
| 420 |
hidden = self.bert.config.hidden_size
|
| 421 |
|
| 422 |
self.num_aspects = num_aspects
|
|
|
|
| 571 |
class_weights: Optional[torch.Tensor] = None,
|
| 572 |
):
|
| 573 |
super().__init__()
|
| 574 |
+
self.bert = _load_bert_model(bert_name)
|
| 575 |
hidden = self.bert.config.hidden_size
|
| 576 |
self.num_aspects = num_aspects
|
| 577 |
self.num_classes = num_classes
|
|
|
|
| 606 |
num_classes: int = cfg.OVERALL_NUM_CLASSES,
|
| 607 |
dropout: float = 0.3):
|
| 608 |
super().__init__()
|
| 609 |
+
self.bert = _load_bert_model(bert_name)
|
| 610 |
hidden = self.bert.config.hidden_size
|
| 611 |
self.classifier = nn.Sequential(
|
| 612 |
nn.Dropout(dropout),
|