Lavender825 commited on
Commit
27f47e2
·
1 Parent(s): 306b259

Fix Spaces meta tensor model loading

Browse files
Files changed (1) hide show
  1. src/evaluator.py +34 -2
src/evaluator.py CHANGED
@@ -1,4 +1,4 @@
1
- """Evaluation: per-aspect metrics, aggregated overall comparison.
2
 
3
  Two label encodings live in this project (keep them straight!):
4
  Per-aspect (ACSA): 0=Not_Mentioned, 1=Positive, 2=Negative
@@ -53,7 +53,12 @@ def _move_model_to_device(model, device):
53
  try:
54
  return model.to(device)
55
  except NotImplementedError as exc:
56
- if "meta tensor" not in str(exc):
 
 
 
 
 
57
  raise
58
  logger.warning("Model contains meta tensors; using to_empty() before loading checkpoint weights.")
59
  return model.to_empty(device=device)
@@ -66,6 +71,32 @@ def _load_model_state(model, state_dict, strict=True):
66
  return model.load_state_dict(state_dict, strict=strict)
67
 
68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
69
  def load_meta_acsa(checkpoint_dir: Path = None, meta_encoder: Optional[MetaEncoder] = None,
70
  device=None):
71
  if checkpoint_dir is None:
@@ -94,6 +125,7 @@ def load_meta_acsa(checkpoint_dir: Path = None, meta_encoder: Optional[MetaEncod
94
  model = BertMetaFusionACSAModel(bert_name=bert_name, meta_in_dim=meta_in_dim)
95
  model = _move_model_to_device(model, device)
96
  _load_model_state(model, ckpt["model_state_dict"], strict=False)
 
97
  model.eval()
98
  tokenizer = AutoTokenizer.from_pretrained(checkpoint_dir / "tokenizer")
99
  return model, tokenizer, meta_encoder, device
 
1
+ """Evaluation: per-aspect metrics, aggregated overall comparison.
2
 
3
  Two label encodings live in this project (keep them straight!):
4
  Per-aspect (ACSA): 0=Not_Mentioned, 1=Positive, 2=Negative
 
53
  try:
54
  return model.to(device)
55
  except NotImplementedError as exc:
56
+ if "meta" not in str(exc).lower():
57
+ raise
58
+ logger.warning("Model contains meta tensors; using to_empty() before loading checkpoint weights.")
59
+ return model.to_empty(device=device)
60
+ except RuntimeError as exc:
61
+ if "meta" not in str(exc).lower():
62
  raise
63
  logger.warning("Model contains meta tensors; using to_empty() before loading checkpoint weights.")
64
  return model.to_empty(device=device)
 
71
  return model.load_state_dict(state_dict, strict=strict)
72
 
73
 
74
+
75
+
76
+ def _refresh_runtime_buffers(model, device):
77
+ """Rebuild non-persistent buffers after a Spaces meta-tensor fallback."""
78
+ for module in model.modules():
79
+ if hasattr(module, "source_prior_bias") and hasattr(module, "num_aspects"):
80
+ prior_rows = []
81
+ prior_cfg = getattr(cfg, "CROSS_ATTN_SOURCE_PRIOR", {})
82
+ for aspect in cfg.ASPECTS[: module.num_aspects]:
83
+ row = list(prior_cfg.get(aspect, [1.0 / 3.0] * 3))
84
+ if len(row) < 3:
85
+ row = row + [1.0 / 3.0] * (3 - len(row))
86
+ prior_rows.append(row[:3])
87
+ prior = torch.tensor(prior_rows, dtype=torch.float, device=device)
88
+ prior = prior - prior.mean(dim=-1, keepdim=True)
89
+ module.source_prior_bias = prior
90
+ meta_params = [name for name, param in model.named_parameters() if getattr(param, "is_meta", False)]
91
+ meta_buffers = [name for name, buf in model.named_buffers() if getattr(buf, "is_meta", False)]
92
+ if meta_params or meta_buffers:
93
+ raise RuntimeError(
94
+ "Checkpoint did not materialize model tensors. First meta tensors: "
95
+ + ", ".join((meta_params + meta_buffers)[:20])
96
+ )
97
+ return model
98
+
99
+
100
  def load_meta_acsa(checkpoint_dir: Path = None, meta_encoder: Optional[MetaEncoder] = None,
101
  device=None):
102
  if checkpoint_dir is None:
 
125
  model = BertMetaFusionACSAModel(bert_name=bert_name, meta_in_dim=meta_in_dim)
126
  model = _move_model_to_device(model, device)
127
  _load_model_state(model, ckpt["model_state_dict"], strict=False)
128
+ model = _refresh_runtime_buffers(model, device)
129
  model.eval()
130
  tokenizer = AutoTokenizer.from_pretrained(checkpoint_dir / "tokenizer")
131
  return model, tokenizer, meta_encoder, device