Lavender825 commited on
Commit
cfc9571
·
1 Parent(s): a32b4ae

Fix meta tensor checkpoint materialization

Browse files
Files changed (1) hide show
  1. src/evaluator.py +64 -5
src/evaluator.py CHANGED
@@ -50,6 +50,34 @@ def _load_ckpt(path: Path, device):
50
  return torch.load(path, map_location="cpu", weights_only=False)
51
 
52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
53
  def _move_model_to_device(model, device):
54
  """Move safely when Transformers creates meta tensors in Spaces."""
55
  try:
@@ -61,12 +89,43 @@ def _move_model_to_device(model, device):
61
  return model.to_empty(device=device)
62
 
63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64
  def _load_model_state(model, state_dict, strict=True):
65
  """Load checkpoint weights and report mismatches instead of hiding them."""
 
 
 
66
  try:
67
  result = model.load_state_dict(state_dict, strict=strict, assign=True)
68
  except TypeError:
69
  result = model.load_state_dict(state_dict, strict=strict)
 
70
  missing = list(getattr(result, "missing_keys", []))
71
  unexpected = list(getattr(result, "unexpected_keys", []))
72
  if missing:
@@ -130,12 +189,12 @@ def load_meta_acsa(checkpoint_dir: Path = None, meta_encoder: Optional[MetaEncod
130
  f"Re-fit the encoder on the same train split used for training."
131
  )
132
  if architecture == "gated_aspect_semantic_meta_acsa":
133
- model = GatedAspectSemanticMetaFusionACSAModel(
134
  bert_name=bert_name,
135
  meta_in_dim=meta_in_dim,
136
- )
137
  else:
138
- model = BertMetaFusionACSAModel(bert_name=bert_name, meta_in_dim=meta_in_dim)
139
  _load_model_state(model, ckpt["model_state_dict"], strict=False)
140
  model = _fix_meta_buffers_only(model, device)
141
  model = model.to(device)
@@ -152,7 +211,7 @@ def load_acsa(checkpoint_dir: Path = None, device=None):
152
  device = get_device()
153
  ckpt = _load_ckpt(checkpoint_dir / "best.pt", device)
154
  bert_name = ckpt.get("config", {}).get("bert_name", cfg.BERT_MODEL_NAME)
155
- model = BertACSAModel(bert_name=bert_name)
156
  _load_model_state(model, ckpt["model_state_dict"])
157
  model = _fix_meta_buffers_only(model, device)
158
  model = model.to(device)
@@ -169,7 +228,7 @@ def load_bert_overall(checkpoint_dir: Path = None, device=None):
169
  device = get_device()
170
  ckpt = _load_ckpt(checkpoint_dir / "best.pt", device)
171
  bert_name = ckpt.get("config", {}).get("bert_name", cfg.BERT_MODEL_NAME)
172
- model = BertOverallModel(bert_name=bert_name)
173
  _load_model_state(model, ckpt["model_state_dict"])
174
  model = _fix_meta_buffers_only(model, device)
175
  model = model.to(device)
 
50
  return torch.load(path, map_location="cpu", weights_only=False)
51
 
52
 
53
+ def _construct_model_on_cpu(factory):
54
+ """Build modules with real CPU tensors even if the runtime default device is meta."""
55
+ old_device = None
56
+ can_restore = False
57
+ if hasattr(torch, "get_default_device") and hasattr(torch, "set_default_device"):
58
+ try:
59
+ old_device = torch.get_default_device()
60
+ torch.set_default_device("cpu")
61
+ can_restore = True
62
+ except Exception:
63
+ old_device = None
64
+ can_restore = False
65
+ try:
66
+ try:
67
+ with torch.device("cpu"):
68
+ return factory()
69
+ except Exception:
70
+ return factory()
71
+ finally:
72
+ # Do not restore a meta default device inside Spaces; that would make
73
+ # later tensors empty again. Restoring normal devices is safe.
74
+ if can_restore and old_device is not None and str(old_device) != "meta":
75
+ try:
76
+ torch.set_default_device(old_device)
77
+ except Exception:
78
+ pass
79
+
80
+
81
  def _move_model_to_device(model, device):
82
  """Move safely when Transformers creates meta tensors in Spaces."""
83
  try:
 
89
  return model.to_empty(device=device)
90
 
91
 
92
+ def _has_meta_tensors(model) -> bool:
93
+ return any(p.device.type == "meta" for p in model.parameters()) or any(
94
+ b is not None and b.device.type == "meta" for b in model.buffers()
95
+ )
96
+
97
+
98
+ def _reset_known_bert_buffers(model, device="cpu"):
99
+ """Reset non-persistent BERT embedding buffers after to_empty()."""
100
+ bert = getattr(model, "bert", None)
101
+ emb = getattr(bert, "embeddings", None)
102
+ if emb is None:
103
+ return model
104
+ device = torch.device(device)
105
+ for name, buffer in list(emb._buffers.items()):
106
+ if buffer is None or name not in {"position_ids", "token_type_ids"}:
107
+ continue
108
+ shape = tuple(buffer.shape)
109
+ if not shape:
110
+ continue
111
+ if name == "position_ids":
112
+ values = torch.arange(shape[-1], dtype=torch.long, device=device)
113
+ emb._buffers[name] = values.view(1, -1).expand(shape).clone()
114
+ elif name == "token_type_ids":
115
+ emb._buffers[name] = torch.zeros(shape, dtype=torch.long, device=device)
116
+ return model
117
+
118
+
119
  def _load_model_state(model, state_dict, strict=True):
120
  """Load checkpoint weights and report mismatches instead of hiding them."""
121
+ if _has_meta_tensors(model):
122
+ logger.warning("Model was built with meta tensors; materializing empty CPU tensors before loading checkpoint.")
123
+ model.to_empty(device="cpu")
124
  try:
125
  result = model.load_state_dict(state_dict, strict=strict, assign=True)
126
  except TypeError:
127
  result = model.load_state_dict(state_dict, strict=strict)
128
+ _reset_known_bert_buffers(model, device="cpu")
129
  missing = list(getattr(result, "missing_keys", []))
130
  unexpected = list(getattr(result, "unexpected_keys", []))
131
  if missing:
 
189
  f"Re-fit the encoder on the same train split used for training."
190
  )
191
  if architecture == "gated_aspect_semantic_meta_acsa":
192
+ model = _construct_model_on_cpu(lambda: GatedAspectSemanticMetaFusionACSAModel(
193
  bert_name=bert_name,
194
  meta_in_dim=meta_in_dim,
195
+ ))
196
  else:
197
+ model = _construct_model_on_cpu(lambda: BertMetaFusionACSAModel(bert_name=bert_name, meta_in_dim=meta_in_dim))
198
  _load_model_state(model, ckpt["model_state_dict"], strict=False)
199
  model = _fix_meta_buffers_only(model, device)
200
  model = model.to(device)
 
211
  device = get_device()
212
  ckpt = _load_ckpt(checkpoint_dir / "best.pt", device)
213
  bert_name = ckpt.get("config", {}).get("bert_name", cfg.BERT_MODEL_NAME)
214
+ model = _construct_model_on_cpu(lambda: BertACSAModel(bert_name=bert_name))
215
  _load_model_state(model, ckpt["model_state_dict"])
216
  model = _fix_meta_buffers_only(model, device)
217
  model = model.to(device)
 
228
  device = get_device()
229
  ckpt = _load_ckpt(checkpoint_dir / "best.pt", device)
230
  bert_name = ckpt.get("config", {}).get("bert_name", cfg.BERT_MODEL_NAME)
231
+ model = _construct_model_on_cpu(lambda: BertOverallModel(bert_name=bert_name))
232
  _load_model_state(model, ckpt["model_state_dict"])
233
  model = _fix_meta_buffers_only(model, device)
234
  model = model.to(device)