Lavender825 commited on
Commit
317ce51
·
1 Parent(s): ce1a6a7

Force real BERT tensor loading in Spaces

Browse files
Files changed (1) hide show
  1. src/models.py +31 -6
src/models.py CHANGED
@@ -1,4 +1,4 @@
1
- """Models for ACSA with metadata fusion.
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
- attn = F.softmax(torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim), dim=-1)
 
 
 
 
 
 
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 = AutoModel.from_pretrained(bert_name)
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 = AutoModel.from_pretrained(bert_name)
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 = AutoModel.from_pretrained(bert_name)
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 = AutoModel.from_pretrained(bert_name)
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),