Spaces:
Sleeping
Sleeping
Commit ·
e30d0ec
1
Parent(s): dc9f449
Materialize Spaces meta parameters during load
Browse files- src/evaluator.py +20 -0
- 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 |
-
|
| 54 |
except TypeError:
|
| 55 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
# ---------------------------------------------------------------------------
|