mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-29 15:31:25 -07:00
fix: Check if a model is gated / accessible
This commit is contained in:
@@ -13,8 +13,12 @@ from urllib.request import urlopen
|
|||||||
|
|
||||||
import cpuinfo
|
import cpuinfo
|
||||||
import torch
|
import torch
|
||||||
from huggingface_hub import HfApi, hf_hub_download
|
from huggingface_hub import HfApi, get_token, hf_hub_download
|
||||||
from huggingface_hub.utils import disable_progress_bars, enable_progress_bars
|
from huggingface_hub.utils import (
|
||||||
|
GatedRepoError,
|
||||||
|
disable_progress_bars,
|
||||||
|
enable_progress_bars,
|
||||||
|
)
|
||||||
from questionary import Choice
|
from questionary import Choice
|
||||||
from rich.table import Table
|
from rich.table import Table
|
||||||
|
|
||||||
@@ -32,11 +36,14 @@ def collect_reproducibles(path: str):
|
|||||||
)
|
)
|
||||||
print()
|
print()
|
||||||
|
|
||||||
api = HfApi()
|
token = get_token()
|
||||||
|
token_arg = token or False
|
||||||
|
api = HfApi(token=token_arg)
|
||||||
|
|
||||||
models = api.list_models(
|
models = api.list_models(
|
||||||
filter=["heretic", "reproducible"],
|
filter=["heretic", "reproducible"],
|
||||||
sort="created_at",
|
sort="created_at",
|
||||||
|
expand=["gated", "tags"],
|
||||||
)
|
)
|
||||||
|
|
||||||
found = 0
|
found = 0
|
||||||
@@ -51,6 +58,14 @@ def collect_reproducibles(path: str):
|
|||||||
if model.tags is not None and "gguf" in model.tags:
|
if model.tags is not None and "gguf" in model.tags:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if model.gated is not False:
|
||||||
|
if not token:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
api.auth_check(model.id, repo_type="model")
|
||||||
|
except GatedRepoError:
|
||||||
|
continue
|
||||||
|
|
||||||
print(f"[bold]{model.id}[/]...", end="")
|
print(f"[bold]{model.id}[/]...", end="")
|
||||||
|
|
||||||
user, repository = model.id.split("/")
|
user, repository = model.id.split("/")
|
||||||
@@ -83,6 +98,7 @@ def collect_reproducibles(path: str):
|
|||||||
cache_path = hf_hub_download(
|
cache_path = hf_hub_download(
|
||||||
model.id,
|
model.id,
|
||||||
"reproduce/reproduce.json",
|
"reproduce/reproduce.json",
|
||||||
|
token=token_arg,
|
||||||
)
|
)
|
||||||
|
|
||||||
file_path.parent.mkdir(parents=True, exist_ok=True)
|
file_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|||||||
Reference in New Issue
Block a user