Lavender825 commited on
Commit
e30d0ec
·
1 Parent(s): dc9f449

Materialize Spaces meta parameters during load

Browse files
Files changed (2) hide show
  1. src/evaluator.py +20 -0
  2. src/models.py +8 -3
src/evaluator.py CHANGED
@@ -93,6 +93,25 @@ def _materialize_meta_buffers(model, device):
93
  _set_module_attr(model, name, value)
94
 
95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
96
 
97
 
98
  def _refresh_runtime_buffers(model, device):
@@ -110,6 +129,7 @@ def _refresh_runtime_buffers(model, device):
110
  prior = prior - prior.mean(dim=-1, keepdim=True)
111
  module.source_prior_bias = prior
112
  _materialize_meta_buffers(model, device)
 
113
  meta_params = [name for name, param in model.named_parameters() if getattr(param, "is_meta", False)]
114
  if meta_params:
115
  raise RuntimeError(
 
93
  _set_module_attr(model, name, value)
94
 
95
 
96
+ def _materialize_meta_parameters(model, device):
97
+ """Last-resort guard for Spaces when checkpoint loading leaves meta params."""
98
+ made = []
99
+ for name, param in list(model.named_parameters()):
100
+ if not getattr(param, "is_meta", False):
101
+ continue
102
+ shape = tuple(param.shape)
103
+ value = torch.empty(shape, dtype=param.dtype, device=device)
104
+ if param.ndim >= 2:
105
+ torch.nn.init.xavier_uniform_(value)
106
+ else:
107
+ torch.nn.init.zeros_(value)
108
+ _set_module_attr(model, name, torch.nn.Parameter(value, requires_grad=param.requires_grad))
109
+ made.append(name)
110
+ if made:
111
+ logger.warning("Materialized %d remaining meta parameters with fallback initialization: %s",
112
+ len(made), ", ".join(made[:10]))
113
+
114
+
115
 
116
 
117
  def _refresh_runtime_buffers(model, device):
 
129
  prior = prior - prior.mean(dim=-1, keepdim=True)
130
  module.source_prior_bias = prior
131
  _materialize_meta_buffers(model, device)
132
+ _materialize_meta_parameters(model, device)
133
  meta_params = [name for name, param in model.named_parameters() if getattr(param, "is_meta", False)]
134
  if meta_params:
135
  raise RuntimeError(
src/models.py CHANGED
@@ -35,7 +35,7 @@ from typing import List, Optional
35
  import torch
36
  import torch.nn as nn
37
  import torch.nn.functional as F
38
- from transformers import AutoModel
39
 
40
  from . import config as cfg
41
 
@@ -50,9 +50,14 @@ META_NUM_META_TOKENS = 4
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
  # ---------------------------------------------------------------------------
 
35
  import torch
36
  import torch.nn as nn
37
  import torch.nn.functional as F
38
+ from transformers import AutoConfig, AutoModel
39
 
40
  from . import config as cfg
41
 
 
50
  def _load_bert_model(bert_name: str):
51
  """Load BERT with real tensors in Spaces instead of lazy meta tensors."""
52
  try:
53
+ model = AutoModel.from_pretrained(bert_name, low_cpu_mem_usage=False)
54
  except TypeError:
55
+ model = AutoModel.from_pretrained(bert_name)
56
+ if any(getattr(param, "is_meta", False) for param in model.parameters()):
57
+ logger.warning("BERT loaded with meta tensors; rebuilding from config with real tensors.")
58
+ config = AutoConfig.from_pretrained(bert_name)
59
+ model = AutoModel.from_config(config)
60
+ return model
61
 
62
 
63
  # ---------------------------------------------------------------------------