File size: 1,652 Bytes
fbaf630
 
 
 
 
 
f2dc03b
fbaf630
f2dc03b
fbaf630
 
 
 
 
 
 
 
 
 
 
 
 
f2dc03b
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
# src/syntax_pred/hf_weights.py
from __future__ import annotations
import os, glob
from typing import Dict, List, Iterable
from huggingface_hub import snapshot_download

def _collect(dirpath: str, patterns: Iterable[str] = ("*.pt", "*.ckpt")) -> List[str]:
    out: List[str] = []
    for pat in patterns:
        out.extend(glob.glob(os.path.join(dirpath, pat)))
    return sorted(out)

def fetch_weights(repo_id: str,
                  allow_patterns: Iterable[str] = ("left/*.pt","left/*.ckpt","right/*.pt","right/*.ckpt")
                 ) -> Dict[str, List[str]]:
    cache_dir = snapshot_download(repo_id=repo_id, allow_patterns=list(allow_patterns))
    left_dir  = os.path.join(cache_dir, "left")
    right_dir = os.path.join(cache_dir, "right")
    return {
        "left": _collect(left_dir)  if os.path.isdir(left_dir)  else [],
        "right": _collect(right_dir) if os.path.isdir(right_dir) else [],
    }

def fetch_classifier_weight(repo_id: str,
                            subdir: str = "classifier",
                            file_patterns: Iterable[str] = ("*.pt", "*.ckpt")) -> str:
    """
    Скачивает веса классификатора из HF-репозитория.
    Возвращает путь к первому (отсортированному) найденному файлу или "".
    """
    allow = [f"{subdir}/{p}" for p in file_patterns]
    cache_dir = snapshot_download(repo_id=repo_id, allow_patterns=allow)
    target_dir = os.path.join(cache_dir, subdir)
    files = _collect(target_dir, patterns=file_patterns) if os.path.isdir(target_dir) else []
    return files[0] if files else ""