Spaces:
Running
Running
Minimum model size: 50K parameters
Browse files- players.py +6 -0
players.py
CHANGED
|
@@ -12,6 +12,7 @@ from huggingface_hub import HfApi, hf_hub_download
|
|
| 12 |
from huggingface_hub.utils import GatedRepoError, RepositoryNotFoundError
|
| 13 |
|
| 14 |
MAX_PARAMS = int(os.environ.get("MAX_PARAMS", 250_000_000))
|
|
|
|
| 15 |
ALLOW_REMOTE_CODE = os.environ.get("ALLOW_REMOTE_CODE", "1") == "1"
|
| 16 |
MAX_CACHED_MODELS = int(os.environ.get("MAX_CACHED_MODELS", 6))
|
| 17 |
BATCH_SIZE = 16
|
|
@@ -255,6 +256,8 @@ def precheck(model_id: str) -> dict:
|
|
| 255 |
est = None
|
| 256 |
if getattr(info, "safetensors", None) and getattr(info.safetensors, "total", None):
|
| 257 |
est = int(info.safetensors.total)
|
|
|
|
|
|
|
| 258 |
else:
|
| 259 |
dtype = str(cfg.get("dtype") or cfg.get("torch_dtype") or "float32")
|
| 260 |
est = int(sum(use.values()) / (2 if ("16" in dtype) else 4))
|
|
@@ -292,6 +295,9 @@ def load_player(model_id: str, meta: dict) -> LMPlayer:
|
|
| 292 |
if n_params > MAX_PARAMS:
|
| 293 |
del model
|
| 294 |
raise ModelRejected(f"`{model_id}` has {fmt_params(n_params)} parameters; the limit is {fmt_params(MAX_PARAMS)}.")
|
|
|
|
|
|
|
|
|
|
| 295 |
player = LMPlayer(model_id, meta["sha"], model, tok, n_params, meta["custom_code"])
|
| 296 |
# smoke test: one forward pass on a real prompt
|
| 297 |
try:
|
|
|
|
| 12 |
from huggingface_hub.utils import GatedRepoError, RepositoryNotFoundError
|
| 13 |
|
| 14 |
MAX_PARAMS = int(os.environ.get("MAX_PARAMS", 250_000_000))
|
| 15 |
+
MIN_PARAMS = int(os.environ.get("MIN_PARAMS", 50_000))
|
| 16 |
ALLOW_REMOTE_CODE = os.environ.get("ALLOW_REMOTE_CODE", "1") == "1"
|
| 17 |
MAX_CACHED_MODELS = int(os.environ.get("MAX_CACHED_MODELS", 6))
|
| 18 |
BATCH_SIZE = 16
|
|
|
|
| 256 |
est = None
|
| 257 |
if getattr(info, "safetensors", None) and getattr(info.safetensors, "total", None):
|
| 258 |
est = int(info.safetensors.total)
|
| 259 |
+
if est < MIN_PARAMS:
|
| 260 |
+
raise ModelRejected(f"`{model_id}` has only ~{fmt_params(est)} parameters; the minimum is {fmt_params(MIN_PARAMS)}.")
|
| 261 |
else:
|
| 262 |
dtype = str(cfg.get("dtype") or cfg.get("torch_dtype") or "float32")
|
| 263 |
est = int(sum(use.values()) / (2 if ("16" in dtype) else 4))
|
|
|
|
| 295 |
if n_params > MAX_PARAMS:
|
| 296 |
del model
|
| 297 |
raise ModelRejected(f"`{model_id}` has {fmt_params(n_params)} parameters; the limit is {fmt_params(MAX_PARAMS)}.")
|
| 298 |
+
if n_params < MIN_PARAMS:
|
| 299 |
+
del model
|
| 300 |
+
raise ModelRejected(f"`{model_id}` has only {fmt_params(n_params)} parameters; the minimum is {fmt_params(MIN_PARAMS)}.")
|
| 301 |
player = LMPlayer(model_id, meta["sha"], model, tok, n_params, meta["custom_code"])
|
| 302 |
# smoke test: one forward pass on a real prompt
|
| 303 |
try:
|