Download dima/dmap.py from jcandane/DIMAX: direct link, hf CLI and curl.
- Browser
- Download file 23.4 kB
-
https://huggingface.co/jcandane/DIMAX/resolve/main/dima/dmap.py
- Command line
-
hf download hf://jcandane/DIMAX/dima/dmap.py
-
curl -L -o dmap.py https://huggingface.co/jcandane/DIMAX/resolve/main/dima/dmap.py
23.4 kB
| # src/dima/dmap.py | |
| from __future__ import annotations | |
| import json | |
| from typing import Any, Dict, Optional, Tuple, Union | |
| import numpy as np | |
| import scipy.sparse as sp | |
| from scipy.sparse.linalg import eigsh, LinearOperator, lobpcg | |
| from huggingface_hub import HfApi, HfFolder, upload_file, hf_hub_download | |
| from .ann import ANNBackend, make_ann | |
| from .utils import median_eps_from_knn_d2 | |
| def k_ideal(d: int, N: int) -> int: | |
| """ | |
| Heuristic for kNN graph size in diffusion maps. | |
| Stable default: grows slowly with N and linearly with d. | |
| """ | |
| d = int(max(1, d)) | |
| N = int(max(2, N)) | |
| k = int(np.ceil(2.0 * d * np.log2(N))) | |
| return int(min(max(8, k), N - 1)) | |
| def _sqdist_ab(A: np.ndarray, B: np.ndarray) -> np.ndarray: | |
| """ | |
| Squared Euclidean distances between rows: | |
| A: (a,d), B: (b,d) -> D2: (a,b) | |
| """ | |
| A = np.asarray(A, dtype=np.float64) | |
| B = np.asarray(B, dtype=np.float64) | |
| A2 = np.sum(A * A, axis=1, keepdims=True) | |
| B2 = np.sum(B * B, axis=1, keepdims=True).T | |
| G = A @ B.T | |
| return np.maximum(A2 + B2 - 2.0 * G, 0.0) | |
| class DMAP: | |
| """ | |
| Diffusion Maps encoder with Nystr枚m out-of-sample extension. | |
| Notation (arrays named by indices): | |
| R_iX: reference ambient data | |
| K_ij: kernel on graph edges (sparse CSR) | |
| q_i = 危_j K_ij | |
| q伪_i = (q_i)^伪 | |
| K伪_ij = K_ij / (q伪_i q伪_j) | |
| d_i = 危_j K伪_ij | |
| A_ij = K伪_ij / sqrt(d_i d_j) (symmetric) | |
| eigsh(A) -> 位_x, u_ix | |
| 蠄_ix = u_ix / sqrt(d_i) | |
| R_ix = (位_x)^t 蠄_ix | |
| Nystr枚m OOS for novel ambient R_aX: | |
| K_ai = exp(-尾 * D2_ai / 蔚) | |
| q_a = 危_i K_ai, q伪_a = (q_a)^伪 | |
| K伪_ai = K_ai / (q伪_a q伪_i) | |
| d_a = 危_i K伪_ai | |
| P_ai = K伪_ai / d_a | |
| R_ax = 危_i P_ai * (R_ix / 位_x) | |
| This implementation uses a kNN graph, sparse eigensolver, and | |
| provides optional dense refinement via streaming LOBPCG. | |
| """ | |
| def __init__( | |
| self, | |
| R_iX: np.ndarray, | |
| *, | |
| d: int = 32, | |
| beta: float = 1.0, | |
| alpha: float = 0.0, | |
| t: float = 1.0, | |
| k: Optional[int] = None, | |
| eps: Optional[float] = None, | |
| eps_use_kth: bool = True, | |
| eps_mul: float = 1.0, | |
| drop_trivial: bool = True, | |
| seed: int = 0, | |
| sym: str = "max", | |
| dtype: Any = np.float32, | |
| ann_backend: ANNBackend = "auto", | |
| ann_params: Optional[Dict[str, Any]] = None, | |
| n_jobs: int = -1, | |
| # dense refinement (O(N^2) compute, streaming memory) | |
| refine_dense: bool = False, | |
| stream_block: int = 4096, | |
| lobpcg_maxiter: int = 3, | |
| lobpcg_tol: float = 1e-6, | |
| use_symmetry: bool = True, | |
| # allow unicode kwargs (尾, 伪, 蔚, 蔚_mul, 蔚_use_kth, ...) | |
| **kwargs: Any, | |
| ): | |
| # ---- map unicode kwargs -> ascii ---- | |
| if "尾" in kwargs: | |
| beta = kwargs.pop("尾") | |
| if "伪" in kwargs: | |
| alpha = kwargs.pop("伪") | |
| if "蔚" in kwargs: | |
| eps = kwargs.pop("蔚") | |
| if "蔚_use_kth" in kwargs: | |
| eps_use_kth = kwargs.pop("蔚_use_kth") | |
| if "蔚_mul" in kwargs: | |
| eps_mul = kwargs.pop("蔚_mul") | |
| if "sym" in kwargs: | |
| sym = kwargs.pop("sym") | |
| if kwargs: | |
| raise TypeError(f"Unexpected kwargs: {sorted(kwargs.keys())}") | |
| self.d = int(d) | |
| self.k = int(k_ideal(self.d, int(np.asarray(R_iX).shape[0])) if k is None else int(k)) | |
| self.beta = float(beta) | |
| self.alpha = float(alpha) | |
| self.t = float(t) | |
| self.drop_trivial = bool(drop_trivial) | |
| self.seed = int(seed) | |
| self.sym = str(sym) | |
| self.dtype = dtype | |
| # unicode aliases (so older code + packers can find them) | |
| self.尾 = self.beta | |
| self.伪 = self.alpha | |
| self.refine_dense = bool(refine_dense) | |
| self.stream_block = int(stream_block) | |
| self.lobpcg_maxiter = int(lobpcg_maxiter) | |
| self.lobpcg_tol = float(lobpcg_tol) | |
| self.use_symmetry = bool(use_symmetry) | |
| rng = np.random.default_rng(self.seed) | |
| # reference data | |
| R_iX = np.asarray(R_iX) | |
| if R_iX.ndim != 2: | |
| raise ValueError(f"R_iX must be 2D array, got shape {R_iX.shape}") | |
| self.R_iX = np.ascontiguousarray(R_iX.astype(self.dtype, copy=False)) | |
| Nref, D = self.R_iX.shape | |
| self.Nref = int(Nref) | |
| self.D = int(D) | |
| if not (2 <= self.k < Nref): | |
| raise ValueError(f"Invalid k={self.k} for Nref={Nref}") | |
| # ANN index (ambient) | |
| self.ann, self.ann_backend = make_ann(ann_backend, ann_params=ann_params, n_jobs=n_jobs) | |
| self.ann.build(self.R_iX) | |
| # kNN of reference points | |
| j_iK, D2_iK = self.ann.search(self.R_iX, self.k) # (Nref,k) | |
| # eps selection | |
| if eps is None: | |
| self.eps = float(median_eps_from_knn_d2(D2_iK, use_kth=bool(eps_use_kth)) * float(eps_mul)) | |
| else: | |
| self.eps = float(eps) | |
| self.蔚 = self.eps # unicode alias | |
| # ------------------------- | |
| # 1) Build sparse graph and warm-start eigensolver | |
| # ------------------------- | |
| K_iK = np.exp(-self.beta * (D2_iK / self.eps)).astype(np.float64, copy=False) | |
| indptr = (np.arange(Nref + 1, dtype=np.int64) * self.k) | |
| indices = j_iK.reshape(-1).astype(np.int64, copy=False) | |
| data = K_iK.reshape(-1) | |
| K_ij = sp.csr_matrix((data, indices, indptr), shape=(Nref, Nref), dtype=np.float64) | |
| # symmetrize | |
| if self.sym == "max": | |
| K_ij = K_ij.maximum(K_ij.T) | |
| elif self.sym == "mean": | |
| K_ij = (K_ij + K_ij.T) * 0.5 | |
| else: | |
| raise ValueError(f"Unknown sym={self.sym!r}") | |
| # degrees q_i and qalpha_i (warm) | |
| q_i_warm = np.asarray(K_ij.sum(axis=1)).ravel() | |
| q_i_warm = np.maximum(q_i_warm, 1e-30) | |
| qalpha_i_warm = np.maximum(np.power(q_i_warm, self.alpha), 1e-30) | |
| Qinv = sp.diags(1.0 / qalpha_i_warm, format="csr") | |
| Kalpha_ij = Qinv @ K_ij @ Qinv | |
| d_i_warm = np.asarray(Kalpha_ij.sum(axis=1)).ravel() | |
| d_i_warm = np.maximum(d_i_warm, 1e-30) | |
| Dinv_sqrt = sp.diags(1.0 / np.sqrt(d_i_warm), format="csr") | |
| A_ij = Dinv_sqrt @ Kalpha_ij @ Dinv_sqrt | |
| nev = self.d + (1 if self.drop_trivial else 0) | |
| v0 = rng.normal(size=Nref).astype(np.float64) | |
| lam0, u0 = eigsh(A_ij, k=nev, which="LA", v0=v0) | |
| ord0 = np.argsort(lam0)[::-1] | |
| lam0 = lam0[ord0] | |
| u0 = u0[:, ord0] | |
| # LOBPCG warm-start block (orthonormalize) | |
| X0, _ = np.linalg.qr(u0.astype(np.float64, copy=False)) | |
| # ------------------------- | |
| # 2) Optional refinement: streaming dense LOBPCG on dense PSD operator | |
| # ------------------------- | |
| if self.refine_dense: | |
| try: | |
| # Rebuild dense operator using streaming blocks without allocating full N^2 matrix | |
| X = self.R_iX.astype(np.float64, copy=False) | |
| X2 = np.sum(X * X, axis=1, keepdims=True) | |
| # compute q_i, qalpha_i, d_i for dense kernel operator: | |
| # q_i = 危_j K_ij, K_ij = exp(-尾*||xi-xj||^2 / eps) | |
| # K伪_ij = K_ij/(q伪_i q伪_j) | |
| # d_i = 危_j K伪_ij | |
| # Step A: q_i | |
| ones = np.ones((Nref, 1), dtype=np.float64) | |
| q_i = self._K_matmat_dense(X, X2, ones).ravel() | |
| q_i = np.maximum(q_i, 1e-30) | |
| qalpha_i = np.maximum(np.power(q_i, self.alpha), 1e-30) | |
| # Step B: d_i | |
| inv_qalpha = 1.0 / qalpha_i | |
| V = (ones * inv_qalpha[:, None]) # (N,1) actually N x 1 | |
| tmp = self._K_matmat_dense(X, X2, V).ravel() # 危_j K_ij * inv_qalpha_j | |
| d_i = inv_qalpha * tmp # 危_j K_ij/(q伪_i q伪_j) | |
| d_i = np.maximum(d_i, 1e-30) | |
| # Build symmetric operator A(v) = D^{-1/2} Q^{-1} K Q^{-1} D^{-1/2} v | |
| inv_sqrt_d = 1.0 / np.sqrt(d_i) | |
| inv_qalpha = 1.0 / qalpha_i | |
| def matvec(v: np.ndarray) -> np.ndarray: | |
| v = v.astype(np.float64, copy=False).reshape(-1, 1) # (N,1) | |
| w = v * inv_sqrt_d[:, None] | |
| w = w * inv_qalpha[:, None] | |
| y = self._K_matmat_dense(X, X2, w) | |
| y = y * inv_qalpha[:, None] | |
| y = y * inv_sqrt_d[:, None] | |
| return y.ravel() | |
| Aop = LinearOperator((Nref, Nref), matvec=matvec, dtype=np.float64) | |
| # LOBPCG refine using warm-start X0 | |
| lam, u = lobpcg(Aop, X0, largest=True, maxiter=self.lobpcg_maxiter, tol=self.lobpcg_tol) | |
| ord1 = np.argsort(lam)[::-1] | |
| lam = lam[ord1] | |
| u = u[:, ord1] | |
| self.q_i = q_i.astype(np.float64, copy=False) | |
| self.qalpha_i = qalpha_i.astype(np.float64, copy=False) | |
| self.d_i = d_i.astype(np.float64, copy=False) | |
| except Exception: | |
| # fallback to warm start if refinement fails | |
| lam, u = lam0, u0 | |
| self.q_i = q_i_warm.astype(np.float64, copy=False) | |
| self.qalpha_i = qalpha_i_warm.astype(np.float64, copy=False) | |
| self.d_i = d_i_warm.astype(np.float64, copy=False) | |
| else: | |
| lam, u = lam0, u0 | |
| self.q_i = q_i_warm.astype(np.float64, copy=False) | |
| self.qalpha_i = qalpha_i_warm.astype(np.float64, copy=False) | |
| self.d_i = d_i_warm.astype(np.float64, copy=False) | |
| # provide unicode aliases for packers / older code | |
| self.q伪_i = self.qalpha_i | |
| self.位_x = lam.astype(np.float64, copy=False) | |
| # psi and drop trivial | |
| psi = u / np.sqrt(self.d_i)[:, None] | |
| if self.drop_trivial: | |
| lam = lam[1:] | |
| psi = psi[:, 1:] | |
| u = u[:, 1:] | |
| # diffusion coords | |
| R_ix = psi * (lam ** self.t)[None, :] | |
| # store | |
| self.位_x = lam.astype(np.float64, copy=False) # (d,) | |
| self.u_ix = u.astype(np.float64, copy=False) # (Nref,d) | |
| self.蠄_ix = psi.astype(np.float64, copy=False) # (Nref,d) | |
| self.R_ix = R_ix.astype(np.float64, copy=False) # (Nref,d) | |
| self.蟺_i = (self.d_i / self.d_i.sum()).astype(np.float64, copy=False) | |
| # for Nystr枚m: R_ix / 位_x | |
| self.R_over_位_ix = (self.R_ix / self.位_x[None, :]).astype(np.float64, copy=False) | |
| # --------- Dense kernel streaming utilities ---------- | |
| def _rbf_block(self, Xi: np.ndarray, Xj: np.ndarray, X2i: np.ndarray, X2j: np.ndarray) -> np.ndarray: | |
| # ||xi-xj||^2 = xi^2 + xj^2 - 2 xi路xj | |
| G = Xi @ Xj.T | |
| D2 = np.maximum(X2i + X2j.T - 2.0 * G, 0.0) | |
| return np.exp(-self.beta * (D2 / self.eps)) | |
| def _K_matmat_dense(self, X: np.ndarray, X2: np.ndarray, V: np.ndarray) -> np.ndarray: | |
| """ | |
| Compute (dense) K @ V without forming K explicitly, using streaming blocks. | |
| X: (N,D), X2: (N,1), V: (N,m) -> out: (N,m) | |
| """ | |
| X = np.asarray(X, dtype=np.float64) | |
| X2 = np.asarray(X2, dtype=np.float64) | |
| V = np.asarray(V, dtype=np.float64) | |
| N = X.shape[0] | |
| bs = self.stream_block | |
| out = np.zeros((N, V.shape[1]), dtype=np.float64) | |
| if not self.use_symmetry: | |
| for i0 in range(0, N, bs): | |
| i1 = min(N, i0 + bs) | |
| Xi = X[i0:i1] | |
| X2i = X2[i0:i1] | |
| acc = np.zeros((i1 - i0, V.shape[1]), dtype=np.float64) | |
| for j0 in range(0, N, bs): | |
| j1 = min(N, j0 + bs) | |
| Xj = X[j0:j1] | |
| X2j = X2[j0:j1] | |
| Kij = self._rbf_block(Xi, Xj, X2i, X2j) | |
| acc += Kij @ V[j0:j1] | |
| out[i0:i1] = acc | |
| return out | |
| # symmetric tiling | |
| for i0 in range(0, N, bs): | |
| i1 = min(N, i0 + bs) | |
| Xi = X[i0:i1] | |
| X2i = X2[i0:i1] | |
| Vi = V[i0:i1] | |
| # diagonal tile | |
| Kii = self._rbf_block(Xi, Xi, X2i, X2i) | |
| out[i0:i1] += Kii @ Vi | |
| for j0 in range(i1, N, bs): | |
| j1 = min(N, j0 + bs) | |
| Xj = X[j0:j1] | |
| X2j = X2[j0:j1] | |
| Vj = V[j0:j1] | |
| Kij = self._rbf_block(Xi, Xj, X2i, X2j) | |
| out[i0:i1] += Kij @ Vj | |
| out[j0:j1] += Kij.T @ Vi | |
| return out | |
| # --------- Nystr枚m embedding ---------- | |
| def __call__(self, R_aX: Union[np.ndarray, list], *, batch_size: Optional[int] = None) -> np.ndarray: | |
| R_aX = np.asarray(R_aX) | |
| single = (R_aX.ndim == 1) | |
| if single: | |
| R_aX = R_aX[None, :] | |
| R_aX = np.ascontiguousarray(R_aX.astype(self.dtype, copy=False)) | |
| if batch_size is None: | |
| R_ax = self._embed(R_aX) | |
| else: | |
| bs = int(batch_size) | |
| out = [] | |
| for s in range(0, R_aX.shape[0], bs): | |
| out.append(self._embed(R_aX[s:s + bs])) | |
| R_ax = np.vstack(out) | |
| return R_ax[0] if single else R_ax | |
| def _embed(self, R_aX: np.ndarray) -> np.ndarray: | |
| # kNN for novel points | |
| j_aK, D2_aK = self.ann.search(R_aX, self.k) # (a,k) | |
| K_ai = np.exp(-self.beta * (D2_aK.astype(np.float64) / self.eps)) # (a,k) | |
| q_a = np.maximum(K_ai.sum(axis=1), 1e-30) | |
| qalpha_a = np.maximum(np.power(q_a, self.alpha), 1e-30) | |
| qalpha_i = np.maximum(self.qalpha_i[j_aK], 1e-30) | |
| Kalpha_ai = K_ai / (qalpha_a[:, None] * qalpha_i) | |
| d_a = np.maximum(Kalpha_ai.sum(axis=1), 1e-30) | |
| P_ai = Kalpha_ai / d_a[:, None] | |
| R_over = self.R_over_位_ix[j_aK, :] # (a,k,d) | |
| R_ax = (P_ai[:, :, None] * R_over).sum(axis=1) | |
| return R_ax | |
| # ------------------------- | |
| # Serialization helpers | |
| # ------------------------- | |
| def _np_dtype_str(x) -> str: | |
| """Best-effort numpy dtype string for serialization.""" | |
| try: | |
| return str(np.dtype(x)) | |
| except Exception: | |
| return "float32" | |
| def state_dict(self) -> Dict[str, Any]: | |
| """ | |
| Minimal, inference-sufficient state for Nystr枚m out-of-sample embedding. | |
| The resulting dict is intentionally compatible with `FrozenDMAP` in `dima.py` | |
| (keys: R_iX, qalpha_i, R_over_lam_ix, k, beta, alpha, eps, dtype). | |
| """ | |
| R_over = getattr(self, "R_over_位_ix", None) | |
| if R_over is None: | |
| # Backward/alternate name safety | |
| R_over = getattr(self, "R_over_lam_ix", None) | |
| if R_over is None: | |
| raise AttributeError("DMAP object has no Nystr枚m matrix `R_over_位_ix` (did training finish?).") | |
| state: Dict[str, Any] = dict( | |
| k=int(self.k), | |
| beta=float(self.beta), | |
| alpha=float(self.alpha), | |
| eps=float(self.eps), | |
| dtype=self._np_dtype_str(getattr(self, "dtype", "float32")), | |
| # arrays | |
| R_iX=np.asarray(self.R_iX), | |
| qalpha_i=np.asarray(self.qalpha_i, dtype=np.float64), | |
| R_over_lam_ix=np.asarray(R_over, dtype=np.float64), | |
| # small meta (optional) | |
| meta=dict( | |
| Nref=int(np.asarray(self.R_iX).shape[0]), | |
| D=int(np.asarray(self.R_iX).shape[1]), | |
| d=int(np.asarray(R_over).shape[1]), | |
| t=float(getattr(self, "t", 1.0)), | |
| drop_trivial=bool(getattr(self, "drop_trivial", False)), | |
| sym=str(getattr(self, "sym", "max")), | |
| ), | |
| ) | |
| return state | |
| def save_local(self, weights_file: str = "dmap.msgpack", config_file: Optional[str] = "dmap_config.json") -> None: | |
| """ | |
| Save DMAP weights (and optionally a lightweight JSON config) locally. | |
| - `weights_file`: binary msgpack with arrays (via flax.serialization) | |
| - `config_file`: JSON with scalar metadata only (no large arrays) | |
| """ | |
| # local import keeps `dima` usable without flax unless you call this | |
| from flax import serialization as flax_ser # type: ignore | |
| import json as _json | |
| state = self.state_dict() | |
| blob = flax_ser.msgpack_serialize(state) | |
| with open(weights_file, "wb") as f: | |
| f.write(blob) | |
| if config_file is not None: | |
| cfg = dict( | |
| k=int(state["k"]), | |
| beta=float(state["beta"]), | |
| alpha=float(state["alpha"]), | |
| eps=float(state["eps"]), | |
| dtype=str(state.get("dtype", "float32")), | |
| **(state.get("meta", {}) or {}), | |
| ) | |
| with open(config_file, "w") as f: | |
| _json.dump(cfg, f, indent=2) | |
| return None | |
| def from_state( | |
| cls, | |
| state: Dict[str, Any], | |
| *, | |
| ann_backend: ANNBackend = "auto", | |
| ann_params: Optional[Dict[str, Any]] = None, | |
| n_jobs: int = -1, | |
| ) -> "DMAP": | |
| """ | |
| Rehydrate a DMAP object from `state_dict()` output WITHOUT retraining. | |
| This bypasses `__init__` and rebuilds only what is needed for out-of-sample Nystr枚m embedding: | |
| reference points, normalization factors, Nystr枚m matrix, and ANN index. | |
| """ | |
| # allow passing a full DIMA state dict | |
| if "encoder" in state and isinstance(state["encoder"], dict): | |
| state = state["encoder"] # type: ignore[assignment] | |
| # tolerate unicode key variants | |
| k = int(state["k"]) | |
| beta = float(state.get("beta", state.get("尾"))) | |
| alpha = float(state.get("alpha", state.get("伪"))) | |
| eps = float(state.get("eps", state.get("蔚"))) | |
| dtype = np.dtype(state.get("dtype", "float32")) | |
| R_iX = np.ascontiguousarray(np.asarray(state["R_iX"]).astype(dtype, copy=False)) | |
| qalpha_i = np.asarray(state.get("qalpha_i", state.get("q伪_i")), dtype=np.float64) | |
| R_over = state.get("R_over_lam_ix", state.get("R_over_位_ix")) | |
| if R_over is None: | |
| raise KeyError("State must contain 'R_over_lam_ix' (or 'R_over_位_ix').") | |
| R_over = np.asarray(R_over, dtype=np.float64) | |
| obj = cls.__new__(cls) # bypass __init__ | |
| # required for embedding | |
| obj.k = k | |
| obj.beta = beta | |
| obj.alpha = alpha | |
| obj.eps = eps | |
| obj.dtype = dtype | |
| obj.R_iX = R_iX | |
| obj.qalpha_i = qalpha_i | |
| obj.q伪_i = obj.qalpha_i # alias | |
| obj.R_over_位_ix = R_over | |
| obj.R_over_lam_ix = obj.R_over_位_ix # alias for external code | |
| obj.尾 = obj.beta | |
| obj.伪 = obj.alpha | |
| obj.蔚 = obj.eps | |
| # optional metadata for introspection | |
| meta = state.get("meta", {}) if isinstance(state.get("meta", {}), dict) else {} | |
| obj.d = int(meta.get("d", R_over.shape[1])) | |
| obj.t = float(meta.get("t", 1.0)) | |
| obj.drop_trivial = bool(meta.get("drop_trivial", False)) | |
| obj.sym = str(meta.get("sym", "max")) | |
| # rebuild ANN index | |
| obj.ann, obj.ann_backend = make_ann(ann_backend, ann_params=ann_params, n_jobs=n_jobs) | |
| obj.ann.build(obj.R_iX) | |
| return obj | |
| def load_local( | |
| cls, | |
| weights_file: str = "dmap.msgpack", | |
| *, | |
| ann_backend: ANNBackend = "auto", | |
| ann_params: Optional[Dict[str, Any]] = None, | |
| n_jobs: int = -1, | |
| ) -> "DMAP": | |
| """Load DMAP weights saved by `save_local()`.""" | |
| from flax import serialization as flax_ser # type: ignore | |
| with open(weights_file, "rb") as f: | |
| state = flax_ser.msgpack_restore(f.read()) | |
| return cls.from_state(state, ann_backend=ann_backend, ann_params=ann_params, n_jobs=n_jobs) | |
| def upload_to_huggingface( | |
| self, | |
| repo_id: str, | |
| *, | |
| weights_file: str = "dmap.msgpack", | |
| config_file: str = "dmap_config.json", | |
| token: Optional[str] = None, | |
| repo_type: str = "model", | |
| revision: Optional[str] = None, | |
| ) -> None: | |
| """ | |
| Upload DMAP weights to the Hugging Face Hub. | |
| This follows the same pattern as `DIMA.save_hf`: | |
| - saves locally first (msgpack + JSON) | |
| - creates the repo if needed | |
| - uploads both files | |
| """ | |
| try: | |
| from huggingface_hub import HfApi, HfFolder, upload_file # type: ignore | |
| except Exception as e: | |
| raise RuntimeError("huggingface_hub not installed. Install extras: `pip install dima[hf]`.") from e | |
| self.save_local(weights_file=weights_file, config_file=config_file) | |
| if token is None: | |
| token = HfFolder.get_token() | |
| if token is None: | |
| raise RuntimeError("No HF token found. Run `huggingface-cli login`, or pass `token=...`.") | |
| import os as _os | |
| api = HfApi() | |
| api.create_repo(repo_id=repo_id, repo_type=repo_type, exist_ok=True, token=token) | |
| wf = _os.path.basename(weights_file) | |
| cf = _os.path.basename(config_file) | |
| upload_file( | |
| path_or_fileobj=weights_file, | |
| path_in_repo=wf, | |
| repo_id=repo_id, | |
| token=token, | |
| repo_type=repo_type, | |
| revision=revision, | |
| ) | |
| upload_file( | |
| path_or_fileobj=config_file, | |
| path_in_repo=cf, | |
| repo_id=repo_id, | |
| token=token, | |
| repo_type=repo_type, | |
| revision=revision, | |
| ) | |
| return None | |
| def download_from_huggingface( | |
| cls, | |
| repo_id: str, | |
| *, | |
| weights_file: str = "dmap.msgpack", | |
| ann_backend: ANNBackend = "auto", | |
| ann_params: Optional[Dict[str, Any]] = None, | |
| n_jobs: int = -1, | |
| token: Optional[str] = None, | |
| repo_type: str = "model", | |
| revision: Optional[str] = None, | |
| ) -> "DMAP": | |
| """ | |
| Download DMAP weights from Hugging Face Hub and rehydrate a DMAP instance. | |
| Note: despite the name, this *downloads from* the Hub. | |
| """ | |
| try: | |
| from huggingface_hub import hf_hub_download # type: ignore | |
| except Exception as e: | |
| raise RuntimeError("huggingface_hub not installed. Install extras: `pip install dima[hf]`.") from e | |
| path = hf_hub_download( | |
| repo_id=repo_id, | |
| filename=weights_file, | |
| token=token, | |
| repo_type=repo_type, | |
| revision=revision, | |
| ) | |
| return cls.load_local(path, ann_backend=ann_backend, ann_params=ann_params, n_jobs=n_jobs) | |
| __all__ = ["DMAP", "k_ideal"] |