Spaces:
Running
Running
File size: 43,863 Bytes
a10a183 d0a25b4 a10a183 da95441 a10a183 e36772a a10a183 9ac7863 a10a183 e48acc2 a10a183 e48acc2 a10a183 e48acc2 a10a183 04d67e9 a10a183 d0a25b4 a10a183 d0a25b4 a10a183 34b0d20 9ac7863 34b0d20 9ac7863 34b0d20 9ac7863 a10a183 e36772a a10a183 e36772a a10a183 e48acc2 e36772a e48acc2 a10a183 e36772a e48acc2 e36772a e48acc2 e36772a e48acc2 a10a183 b3a08a4 d0a25b4 b3a08a4 d0a25b4 b3a08a4 d0a25b4 b3a08a4 d0a25b4 b3a08a4 a10a183 34b0d20 9ac7863 34b0d20 9ac7863 34b0d20 9ac7863 34b0d20 9ac7863 d0a25b4 04aaad8 d0a25b4 9ac7863 d0a25b4 9ac7863 d0a25b4 9ac7863 d0a25b4 9ac7863 d0a25b4 9ac7863 d0a25b4 a10a183 7528321 a10a183 d0a25b4 a10a183 d0a25b4 a10a183 e36772a a10a183 e36772a a10a183 e36772a a10a183 e36772a a10a183 e36772a a10a183 9ac7863 a10a183 e48acc2 a10a183 9ac7863 a10a183 7528321 a10a183 e48acc2 9ac7863 a10a183 d0a25b4 e48acc2 9ac7863 d0a25b4 a10a183 e48acc2 a10a183 7528321 a10a183 e48acc2 a10a183 e48acc2 a10a183 9ac7863 a10a183 7528321 a10a183 e48acc2 9ac7863 a10a183 d0a25b4 e48acc2 9ac7863 d0a25b4 a10a183 e48acc2 a10a183 7528321 a10a183 e48acc2 a10a183 e48acc2 a10a183 e36772a a10a183 e36772a a10a183 e36772a a10a183 e36772a a10a183 0ac08e4 ae9ddcd 0ac08e4 da95441 b3a08a4 da95441 a10a183 | 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 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 | """hub-search-v3 backend — FastAPI + in-memory DuckDB brute-force vector search.
Drop-in compatible with the old davanstrien/huggingface-datasets-search-v2 API
(same routes / params / response shapes) so the librarian-bots front-end can be
repointed with no code change. Replaces the ChromaDB boot-time HNSW rebuild
(which timed out) with an in-memory DuckDB FLOAT[dim] table loaded once at boot.
See bench: duckdb-vss-bench-2026-07-18. Query encoding replicates the old
backend so query vectors stay aligned with the seed document embeddings.
"""
import logging
import os
import re
import time
from contextlib import asynccontextmanager
from typing import List, Optional
import duckdb
import httpx
from cashews import cache
from fastapi import FastAPI, Header, HTTPException, Query
from fastapi.middleware.cors import CORSMiddleware
from huggingface_hub import HfApi, hf_hub_download, list_repo_files, login
from pydantic import BaseModel
os.environ.setdefault("HF_HUB_ENABLE_HF_TRANSFER", "1")
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
logger = logging.getLogger("hub-search-v3")
# --- config -----------------------------------------------------------------
EMBEDDING_MODEL = "Qwen/Qwen3-Embedding-0.6B"
EMBEDDINGS_REPO = os.getenv("EMBEDDINGS_REPO", "davanstrien/search-v2-embeddings")
EMB_DIM = int(os.getenv("EMB_DIM", "256")) # MRL truncation dim; 1024 = full
FULL_DIM = 1024
THREADS = int(os.getenv("DUCKDB_THREADS", "2"))
CACHE_TTL = "24h"
TRENDING_CACHE_TTL = "1h"
DS_TASK = "Given a search query, retrieve relevant model and dataset summaries that match the query. "
# ISO-639-3/-2B codes that appear in source cards alongside their two-letter
# equivalents (e.g. datasets tagged `fra` vs `fr`); a two-letter filter matches both.
LANG_EQUIV = {
"en": ["eng"], "fr": ["fra", "fre"], "de": ["deu", "ger"], "es": ["spa"],
"zh": ["zho", "chi"], "pt": ["por"], "ru": ["rus"], "ja": ["jpn"],
"ko": ["kor"], "ar": ["ara"], "it": ["ita"], "nl": ["nld", "dut"],
"pl": ["pol"], "tr": ["tur"], "vi": ["vie"], "hi": ["hin"], "sv": ["swe"],
}
# config dir names inside the embeddings repo (datasets vs models slices)
DATASET_CFG = os.getenv("DATASET_CONFIG", "dataset_cards")
MODEL_CFG = os.getenv("MODEL_CONFIG", "model_cards")
# prebuilt full-card-text FTS database (hybrid search channel); optional at boot
FTS_FILE = os.getenv("FTS_FILE", "_fts/card-text.duckdb")
def config_files(cfg: str) -> List[str]:
"""Discover the parquet shards for a config dir (shard count varies across seed/v3)."""
files = [f for f in list_repo_files(EMBEDDINGS_REPO, repo_type="dataset")
if f.startswith(f"{cfg}/") and f.endswith(".parquet")]
if not files:
raise RuntimeError(f"no parquet files found under {cfg}/ in {EMBEDDINGS_REPO}")
return sorted(files)
HF_TOKEN = os.getenv("HF_TOKEN")
if HF_TOKEN:
login(token=HF_TOKEN)
cache.setup("mem://", size_limit="1gb")
# in-memory DuckDB (single connection, read-only after boot)
con = duckdb.connect(":memory:")
con.execute(f"SET threads={THREADS}")
_query_model = None
STATE = {"ready": False, "boot_seconds": None, "counts": {}, "revision": None,
"dim": EMB_DIM, "model": EMBEDDING_MODEL}
def get_query_model():
global _query_model
if _query_model is None:
from sentence_transformers import SentenceTransformer
logger.info(f"Loading query model {EMBEDDING_MODEL} on CPU")
_query_model = SentenceTransformer(EMBEDDING_MODEL, device="cpu")
return _query_model
def embed_query(text: str) -> List[float]:
"""Replicate the old backend's query encoding (prompt_name='query')."""
model = get_query_model()
vec = model.encode(text, prompt_name="query", normalize_embeddings=False)
vec = vec[:EMB_DIM]
# L2 normalise (matches the stored-vector normalisation; makes cosine stable)
import numpy as np
n = np.linalg.norm(vec)
if n > 0:
vec = vec / n
return vec.astype(float).tolist()
def load_table(cfg: str, table: str, has_param: bool):
"""Load one config into an in-memory FLOAT[EMB_DIM] table (MRL trunc + renorm).
Filters NULL embeddings (v3 refusal rows have no embedding). Discovers shards
dynamically so the seed and the ~1.17M-row search-v3 dataset both work.
"""
paths = [hf_hub_download(EMBEDDINGS_REPO, f, repo_type="dataset") for f in config_files(cfg)]
plist = "[" + ",".join(f"'{p}'" for p in paths) + "]"
param_sel = "COALESCE(param_count,0)::BIGINT AS param_count" if has_param else "0::BIGINT AS param_count"
# slice to EMB_DIM, L2-renormalise, cast to fixed-size FLOAT[] array (nested, no correlated subquery)
# task/license/language are copied straight from the source cards; cast to
# VARCHAR because the models config's all-null `language` reads back as a
# NULL-typed column that would otherwise reject `=` filters at query time.
con.execute(f"""
CREATE OR REPLACE TABLE {table} AS
SELECT id, summary,
COALESCE(likes,0)::BIGINT AS likes,
COALESCE(downloads,0)::BIGINT AS downloads,
last_modified,
CAST(task AS VARCHAR) AS task,
CAST(license AS VARCHAR) AS license,
CAST(language AS VARCHAR) AS language,
{param_sel},
list_transform(s, x -> (x / nrm))::FLOAT[{EMB_DIM}] AS emb
FROM (
SELECT *, sqrt(list_dot_product(s, s)) AS nrm
FROM (
SELECT *, embedding[1:{EMB_DIM}] AS s
FROM read_parquet({plist}, union_by_name=True)
WHERE embedding IS NOT NULL
QUALIFY row_number() OVER (PARTITION BY id ORDER BY last_modified DESC NULLS LAST) = 1
)
)
""")
_refresh_metadata(cfg, table)
n = con.execute(f"SELECT count(*) FROM {table}").fetchone()[0]
STATE["counts"][table] = n
logger.info(f"loaded {table} (from {cfg}): {n:,} rows @ {EMB_DIM}d")
def _refresh_metadata(cfg: str, table: str) -> None:
"""Overwrite likes/downloads from the pipeline's metadata snapshot, if present.
Corpus rows carry the counts a repo had when its card was last summarised, and
cards rarely change — so a repo can sit at 0 downloads long after it is widely
used. Ranking now leans on those counts, so refresh them from the small
(id, downloads, likes) file the delta pipeline publishes each run. Absent file
or unreadable snapshot is not fatal: the corpus values still work, just stale.
"""
try:
path = hf_hub_download(EMBEDDINGS_REPO, f"_state/metadata/{cfg}.parquet", repo_type="dataset")
except Exception as e:
logger.warning(f"{table}: no metadata snapshot ({e}); using corpus counts")
return
try:
n = con.execute(f"""
UPDATE {table} AS t
SET downloads = COALESCE(m.downloads, 0)::BIGINT,
likes = COALESCE(m.likes, 0)::BIGINT
FROM read_parquet('{path}') AS m
WHERE m.id = t.id
""").fetchone()
logger.info(f"{table}: refreshed metadata for {(n[0] if n else 0):,} rows")
except Exception as e:
logger.warning(f"{table}: metadata refresh failed ({e}); using corpus counts")
@asynccontextmanager
async def lifespan(app: FastAPI):
t0 = time.time()
logger.info(f"boot: loading model + embeddings (dim={EMB_DIM})")
get_query_model() # load model up front so first query isn't slow
try:
one = hf_hub_download(EMBEDDINGS_REPO, config_files(MODEL_CFG)[0], repo_type="dataset")
cols = con.execute(f"DESCRIBE SELECT * FROM read_parquet('{one}')").fetchall()
has_param = any(c[0] == "param_count" for c in cols)
except Exception as e:
logger.warning(f"param_count detection failed ({e}); assuming absent")
has_param = False
load_table(DATASET_CFG, "dataset_cards", has_param=False)
load_table(MODEL_CFG, "model_cards", has_param=has_param)
# Optional: open the prebuilt card-text FTS db on its OWN connection.
# (Not ATTACHed: the fts match_bm25 macro references its internal schema
# unqualified, which breaks across catalog boundaries.) Any failure
# degrades hybrid search to vector-only rather than blocking boot.
global fts_con
try:
fts_path = hf_hub_download(EMBEDDINGS_REPO, FTS_FILE, repo_type="dataset")
fts_con = duckdb.connect(fts_path, read_only=True)
fts_con.execute("INSTALL fts; LOAD fts;")
STATE["fts_cards"] = fts_con.execute("SELECT count(*) FROM cards").fetchone()[0]
logger.info(f"card-text FTS attached: {STATE['fts_cards']:,} cards")
except Exception as e:
fts_con = None
STATE["fts_cards"] = None
logger.warning(f"card-text FTS unavailable ({e}); hybrid=vector-only")
try:
info = HfApi().dataset_info(EMBEDDINGS_REPO)
STATE["revision"] = info.sha[:8] if info.sha else None
STATE["last_modified"] = str(getattr(info, "last_modified", None))
except Exception:
pass
STATE["boot_seconds"] = round(time.time() - t0, 1)
STATE["ready"] = True
logger.info(f"boot complete in {STATE['boot_seconds']}s; ready")
yield
await cache.close()
app = FastAPI(title="hub-search-v3", lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
allow_origin_regex=r"https://.*\.(hf\.space|huggingface\.co)",
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# --- response models (identical to old backend) -----------------------------
class QueryResult(BaseModel):
dataset_id: str
similarity: float
summary: str
likes: int
downloads: int
task: Optional[str] = None
license: Optional[str] = None
language: Optional[str] = None
last_modified: Optional[str] = None
class QueryResponse(BaseModel):
results: List[QueryResult]
class ModelQueryResult(BaseModel):
model_id: str
similarity: float
summary: str
likes: int
downloads: int
param_count: Optional[int] = None
task: Optional[str] = None
license: Optional[str] = None
language: Optional[str] = None
last_modified: Optional[str] = None
class ModelQueryResponse(BaseModel):
results: List[ModelQueryResult]
# --- core search -------------------------------------------------------------
def _where(min_likes, min_downloads, min_param=0, max_param=None, use_param=False,
task=None, license=None, language=None, modified_after=None):
"""Build a WHERE clause + its bind params.
Numeric bounds are inlined (ints are injection-safe); the free-text metadata
filters (task/license/language/modified_after) are bound as `?` params.
task/language are comma-joined multi-value strings in the source cards, so
they match by token membership, not string equality; two-letter language
codes also match their ISO-639-3 equivalents (LANG_EQUIV).
Returns ``(where_sql, params)`` for ``_knn``.
"""
conds, params = [], []
if min_likes > 0:
conds.append(f"likes >= {int(min_likes)}")
if min_downloads > 0:
conds.append(f"downloads >= {int(min_downloads)}")
if use_param:
conds.append("param_count > 0")
if min_param > 0:
conds.append(f"param_count >= {int(min_param)}")
if max_param is not None:
conds.append(f"param_count <= {int(max_param)}")
token_match = "list_contains(string_split(replace(lower({col}), ' ', ''), ','), ?)"
if task:
conds.append(token_match.format(col="task"))
params.append(task.lower())
if license:
conds.append("license = ?")
params.append(license)
if language:
variants = [language.lower()] + LANG_EQUIV.get(language.lower(), [])
ors = " OR ".join(token_match.format(col="language") for _ in variants)
conds.append(f"({ors})")
params.extend(variants)
if modified_after:
# last_modified is an ISO-8601 string; lexicographic >= is chronological.
conds.append("last_modified >= ?")
params.append(modified_after)
return (("WHERE " + " AND ".join(conds)) if conds else ""), params
def _vec_literal(vec):
return "[" + ",".join(repr(float(x)) for x in vec) + f"]::FLOAT[{EMB_DIM}]"
DUP_SIM = 0.985 # near-identical cards (mirrors/re-uploads) collapse to the top-ranked copy
DUP_OVERFETCH = 6 # candidates per requested row, so dup-collapse can be backfilled
_DL_IDX = 3 # position of `downloads` within COLS (id, summary, likes, downloads, ...)
def _knn(table, qvec, k, where_sql, cols, params=None, return_emb=False):
"""Top-k by cosine, with near-duplicate collapse.
Overfetches DUP_OVERFETCH x, then greedily drops rows whose embedding is
~identical (cos >= DUP_SIM) to a higher-ranked kept row — mirror repos
otherwise eat adjacent result slots. When a duplicate is more downloaded than
the copy already kept, it replaces it in place: the group keeps its ranking
position but is represented by the canonical repo rather than whichever
mirror happened to score a hair higher.
The overfetch has to exceed the dup rate or collapse silently returns fewer
than k (popular model families are dominated by near-identical conversion
cards); the LIMIT is cheap because the full cosine scan dominates either way.
Returns (cols..., sim) tuples; with return_emb=True, (cols..., emb, sim).
"""
import numpy as np
q = f"""SELECT {cols}, emb, array_cosine_similarity(emb, {_vec_literal(qvec)}) AS sim
FROM {table} {where_sql} ORDER BY sim DESC LIMIT {int(k) * DUP_OVERFETCH}"""
rows = con.execute(q, params or []).fetchall()
kept, kept_vecs = [], []
for r in rows:
v = np.asarray(r[-2], dtype=np.float32)
if kept_vecs:
sims = np.stack(kept_vecs) @ v
j = int(np.argmax(sims))
if float(sims[j]) >= DUP_SIM:
# Same card, different repo: keep whichever copy people actually use.
if (r[_DL_IDX] or 0) > (kept[j][_DL_IDX] or 0):
kept[j] = r if return_emb else r[:-2] + (r[-1],)
kept_vecs[j] = v
continue
kept.append(r if return_emb else r[:-2] + (r[-1],))
kept_vecs.append(v)
if len(kept) >= int(k):
break
return kept
def _emb_of(table, repo_id):
row = con.execute(f"SELECT emb FROM {table} WHERE id = ?", [repo_id]).fetchone()
return row[0] if row else None
fts_con = None
def _bm25_ids(kind: str, text: str, n: int) -> list:
"""Top-n repo ids by BM25 over full card text (empty if FTS unavailable)."""
if fts_con is None:
return []
rows = fts_con.execute(
"""SELECT id FROM (
SELECT id, fts_main_cards.match_bm25(doc_id, ?) AS s
FROM cards WHERE repo_type = ?
) WHERE s IS NOT NULL ORDER BY s DESC LIMIT ?""",
[text, kind, int(n)],
).fetchall()
return [r[0] for r in rows]
# A query that is short and has no spaces is a repo name, not a description
# ("iSAID", "rtdetr"). Summary embeddings cannot retrieve those: the summary of
# ariG23498/iSAID never says "iSAID", and "iSAID" alone retrieves Icelandic
# census records via ID-ish subword tokens. Matching the id directly is the only
# channel that can find them.
_NAME_MAX_TOKENS = 3 # "iSAID instance segmentation" still carries a name token
_NAME_MIN_CHARS = 4 # shorter tokens ("ocr", "sam") match far too much
# Equal to the vector channel. RRF scores a rank-r hit at weight/(60+r), so the
# vector pool (k*10) reaches ~1/140 at its tail while a name-only hit sits at
# weight/60. At 0.5 such a hit loses to every vector result down to rank ~60 and so
# can never enter the top k — which is the entire point of the channel. False
# positives are held back at the trigger (identifier-shaped tokens, matched at an
# id boundary) rather than by de-weighting the whole channel.
NAME_CHANNEL_WEIGHT = 1.0
# Hyphenated descriptors that pass the identifier test but name no repo.
_NAME_STOPWORDS = frozenset({"real-time", "state-of-the-art", "open-source", "fine-tuned"})
def _is_identifier_like(tok: str) -> bool:
"""True for tokens shaped like a repo/model name rather than an English word.
Names carry a digit (yolov10), an internal capital (iSAID, RF-DETR), or a
separator (rf-detr). Ordinary description words ("herbarium", "element") carry
none of these, and matching those against ids returns download-ranked noise.
"""
if tok.lower() in _NAME_STOPWORDS:
return False
return bool(
any(c.isdigit() for c in tok)
or any(c in "-_." for c in tok)
or re.search(r"[a-z][A-Z]|^[A-Z]{2,}$", tok)
)
def _name_tokens(query: str) -> list:
"""Tokens from the query that plausibly name a repo rather than describe one."""
if len(query.strip()) > 60 or len(query.split()) > _NAME_MAX_TOKENS:
return []
return [
tok
for raw in re.split(r"[\s,]+", query.strip())
if len(tok := raw.strip("\"'()[]")) >= _NAME_MIN_CHARS and _is_identifier_like(tok)
]
def _name_ids(table, tokens, n, where_sql, wparams):
"""Top-n ids whose name contains any of these tokens, most-downloaded first.
The token must start at an id boundary (start, or after / _ . -). A bare
substring match puts `RASSAISAID/finance-deepseek-prompts` above the real
`ariG23498/iSAID` for the query "iSAID". The trailing side is deliberately left
open so name variants still match — "yolov10" has to find `jameslahm/yolov10x`.
Ordering by downloads (not cosine) is deliberate: this channel exists to find
the repo actually named, and RRF only uses the rank, so the canonical copy of
a name should lead its own channel.
"""
if not tokens:
return []
cond = f"AND {where_sql[6:]}" if where_sql else ""
ors = " OR ".join("regexp_matches(id, ?, 'i')" for _ in tokens)
pats = [f"(^|[/_.-]){re.escape(t)}" for t in tokens]
rows = con.execute(
f"""SELECT id FROM {table}
WHERE ({ors}) {cond}
ORDER BY downloads DESC LIMIT {int(n)}""",
pats + list(wparams or []),
).fetchall()
return [r[0] for r in rows]
def _fuse(table, qvec, raw, extra_ids, k, where_sql, wparams, weight=1.0):
"""RRF-fuse the vector KNN rows with a second channel's id list.
Channel-only hits are fetched from the in-memory table (with their true cosine
similarity) and must pass the same metadata filters as the vector channel.
`weight` scales the second channel's contribution, so a speculative channel can
rescue a miss without outvoting the vector ranking.
Returns rows in fused order, same tuple shape as _knn output.
"""
if not extra_ids:
return raw
rrf: dict = {}
for rank, rid in enumerate(r[0] for r in raw):
rrf[rid] = rrf.get(rid, 0.0) + 1.0 / (60 + rank)
for rank, rid in enumerate(extra_ids):
rrf[rid] = rrf.get(rid, 0.0) + weight / (60 + rank)
order = sorted(rrf, key=rrf.get, reverse=True)
known = {r[0]: r for r in raw}
missing = [rid for rid in order if rid not in known][: k * 2]
if missing:
ph = ",".join("?" for _ in missing)
cond = f"AND {where_sql[6:]}" if where_sql else ""
extra = con.execute(
f"""SELECT {COLS}, array_cosine_similarity(emb, {_vec_literal(qvec)}) AS sim
FROM {table} WHERE id IN ({ph}) {cond}""",
missing + list(wparams or []),
).fetchall()
known.update({r[0]: r for r in extra})
return [known[rid] for rid in order if rid in known]
def _hybrid_fuse(table, kind, query, qvec, raw, k, where_sql, wparams):
"""RRF-fuse vector KNN rows with a BM25 channel over full card text."""
return _fuse(table, qvec, raw, _bm25_ids(kind, query, k * 2), k, where_sql, wparams)
async def _sort_results(rows, id_field, sort_by, k):
"""rows: list of dicts with keys id, similarity, summary, likes, downloads, [param_count]."""
if sort_by == "trending":
scores = {}
tasks = [get_trending_score(r["_id"], id_field) for r in rows]
import asyncio
vals = await asyncio.gather(*tasks)
scores = {r["_id"]: v for r, v in zip(rows, vals)}
rows.sort(key=lambda r: scores.get(r["_id"], 0), reverse=True)
rows = rows[:k]
elif sort_by in ("likes", "downloads"):
rows.sort(key=lambda r: r[sort_by], reverse=True)
rows = rows[:k]
elif sort_by == "updated":
# last_modified is an ISO-8601 string; lexicographic sort is chronological
rows.sort(key=lambda r: r.get("last_modified") or "", reverse=True)
rows = rows[:k]
else:
rows = _apply_popularity_prior(rows)[:k]
return rows
# Optional log-scaled download term, added to cosine to break ties toward the copy
# people actually use (a 0-download mirror otherwise outranks the canonical repo,
# since their cards read the same).
#
# OFF by default, deliberately. On hf-find's eval/queries.jsonl it looks like a
# clear win (mean recall 0.363 -> 0.435, zero-download slots 28.3% -> 6.2% at 0.08),
# but that eval's expected repos were chosen *because* they are canonical, so it
# cannot help but reward a popularity prior. A blind 3-judge A/B over 24 unseen
# queries split by domain instead:
#
# mainstream queries prior 14 - 2 baseline
# niche / historical / GLAM prior 3 - 5 baseline
#
# The mechanism: a 150k-download repo gets the full weight while a 50-download
# exact match gets ~a third of it, so any similarity gap under ~0.054 flips — and
# niche queries live in exactly that regime. Concretely, at 0.08 the query "OCR for
# historical printed text" ranks generic dots.ocr-1.5 above
# wjbmattingly/LightOnOCR-2-1B-cultural-heritage-english.
#
# So this is a domain trade, not a free improvement. Enable per-deployment with
# POPULARITY_WEIGHT=0.08 if the audience is mainstream model/dataset discovery.
POPULARITY_WEIGHT = float(os.getenv("POPULARITY_WEIGHT", "0"))
def _rerank_pool(k: int, sort_by: str) -> int:
"""Candidate rows to pull before re-ranking down to k.
The relevance path re-ranks (popularity prior, optional name fusion) and gains
from a deeper pool: on eval/queries.jsonl mean recall is 0.363 at k*1, 0.399 at
k*4 and 0.435 at k*10, flat thereafter. The cosine scan touches every row
regardless, so a larger LIMIT is close to free.
The explicit sort modes keep the old k*4: they sort hard by likes/downloads, so
widening the pool would trade relevance for popularity rather than break ties.
"""
return k * (10 if sort_by == "similarity" else 4)
def _apply_popularity_prior(rows):
"""Re-rank by similarity + a log-scaled, per-result-set-normalised download term."""
if POPULARITY_WEIGHT <= 0 or not rows:
return rows
import math
mx = max((r.get("downloads") or 0) for r in rows)
if mx <= 0:
return rows
denom = math.log1p(mx)
return sorted(
rows,
key=lambda r: r["similarity"] + POPULARITY_WEIGHT * (math.log1p(r.get("downloads") or 0) / denom),
reverse=True,
)
def _rows_to_dataset(raw):
out = []
for id_, summary, likes, downloads, param, task, license_, language, last_modified, sim in raw:
out.append({"_id": id_, "dataset_id": id_, "similarity": float(sim),
"summary": summary, "likes": int(likes), "downloads": int(downloads),
"task": task, "license": license_, "language": language,
"last_modified": last_modified})
return out
def _rows_to_model(raw):
out = []
for id_, summary, likes, downloads, param, task, license_, language, last_modified, sim in raw:
out.append({"_id": id_, "model_id": id_, "similarity": float(sim),
"summary": summary, "likes": int(likes), "downloads": int(downloads),
"param_count": int(param),
"task": task, "license": license_, "language": language,
"last_modified": last_modified})
return out
META_COLS = "task, license, language, CAST(last_modified AS VARCHAR) AS last_modified"
COLS = f"id, summary, likes, downloads, param_count, {META_COLS}"
# --- routes ------------------------------------------------------------------
@app.get("/")
async def root():
from fastapi.responses import RedirectResponse
return RedirectResponse(url="/docs")
@app.get("/health")
async def health():
return {
"ready": STATE["ready"],
"boot_seconds": STATE["boot_seconds"],
"counts": STATE["counts"],
"index_revision": STATE.get("revision"),
"index_last_modified": STATE.get("last_modified"),
"fts_cards": STATE.get("fts_cards"),
"embedding_model": STATE["model"],
"dim": STATE["dim"],
"embeddings_repo": EMBEDDINGS_REPO,
}
@app.get("/lookup/{repo_id:path}")
async def lookup(repo_id: str):
"""Resolve a repo id to its type in one call: {"id":..., "type": "dataset"|"model"}.
Saves clients the 404-fallthrough (try datasets, then models). 404 with a JSON
detail if the id is in neither index. `:path` so `org/name` ids keep their slash.
"""
for table, typ in (("dataset_cards", "dataset"), ("model_cards", "model")):
if con.execute(f"SELECT 1 FROM {table} WHERE id = ? LIMIT 1", [repo_id]).fetchone():
return {"id": repo_id, "type": typ}
raise HTTPException(status_code=404, detail=f"'{repo_id}' not found as a dataset or model in the index")
@app.get("/search/datasets", response_model=QueryResponse)
@cache(ttl=CACHE_TTL, key="search:d:{query}:{k}:{sort_by}:{min_likes}:{min_downloads}:{task}:{license}:{language}:{modified_after}:{hybrid}")
async def search_datasets(
query: str,
k: int = Query(default=5, ge=1, le=100),
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
min_likes: int = Query(default=0, ge=0),
min_downloads: int = Query(default=0, ge=0),
task: Optional[str] = Query(default=None),
license: Optional[str] = Query(default=None),
language: Optional[str] = Query(default=None),
modified_after: Optional[str] = Query(default=None),
hybrid: bool = Query(default=False),
):
try:
qvec = embed_query(f"Instruct: {DS_TASK}\nQuery:{query}")
n = _rerank_pool(k, sort_by)
where_sql, wparams = _where(min_likes, min_downloads,
task=task, license=license, language=language,
modified_after=modified_after)
raw = _knn("dataset_cards", qvec, n, where_sql, COLS, wparams)
if hybrid:
raw = _hybrid_fuse("dataset_cards", "datasets", query, qvec, raw, k, where_sql, wparams)
elif (name_toks := _name_tokens(query)):
raw = _fuse("dataset_cards", qvec, raw,
_name_ids("dataset_cards", name_toks, k, where_sql, wparams),
k, where_sql, wparams, weight=NAME_CHANNEL_WEIGHT)
rows = await _sort_results(_rows_to_dataset(raw), "dataset", sort_by, k)
return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
except Exception as e:
logger.error(f"search datasets: {e}")
raise HTTPException(status_code=500, detail="Search failed")
@app.get("/similarity/datasets", response_model=QueryResponse)
@cache(ttl=CACHE_TTL, key="sim:d:{dataset_id}:{k}:{sort_by}:{min_likes}:{min_downloads}:{task}:{license}:{language}:{modified_after}")
async def find_similar_datasets(
dataset_id: str,
k: int = Query(default=5, ge=1, le=100),
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
min_likes: int = Query(default=0, ge=0),
min_downloads: int = Query(default=0, ge=0),
task: Optional[str] = Query(default=None),
license: Optional[str] = Query(default=None),
language: Optional[str] = Query(default=None),
modified_after: Optional[str] = Query(default=None),
):
emb = _emb_of("dataset_cards", dataset_id)
if emb is None:
raise HTTPException(status_code=404, detail=f"Dataset ID '{dataset_id}' not found")
try:
n = k * 4 if sort_by != "similarity" else k + 1
where_sql, wparams = _where(min_likes, min_downloads,
task=task, license=license, language=language,
modified_after=modified_after)
raw = _knn("dataset_cards", emb, n, where_sql, COLS, wparams)
rows = [r for r in _rows_to_dataset(raw) if r["_id"] != dataset_id]
rows = await _sort_results(rows, "dataset", sort_by, k)
return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
except HTTPException:
raise
except Exception as e:
logger.error(f"similarity datasets: {e}")
raise HTTPException(status_code=500, detail="Similarity search failed")
@app.get("/search/models", response_model=ModelQueryResponse)
@cache(ttl=CACHE_TTL, key="search:m:{query}:{k}:{sort_by}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}:{task}:{license}:{language}:{modified_after}:{hybrid}")
async def search_models(
query: str,
k: int = Query(default=5, ge=1, le=100),
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
min_likes: int = Query(default=0, ge=0),
min_downloads: int = Query(default=0, ge=0),
min_param_count: int = Query(default=0, ge=0),
max_param_count: Optional[int] = Query(default=None, ge=0),
task: Optional[str] = Query(default=None),
license: Optional[str] = Query(default=None),
language: Optional[str] = Query(default=None),
modified_after: Optional[str] = Query(default=None),
hybrid: bool = Query(default=False),
):
try:
use_param = min_param_count > 0 or max_param_count is not None
qvec = embed_query(f"search_query: {query}")
n = _rerank_pool(k, sort_by)
where_sql, wparams = _where(min_likes, min_downloads, min_param_count, max_param_count, use_param,
task=task, license=license, language=language,
modified_after=modified_after)
raw = _knn("model_cards", qvec, n, where_sql, COLS, wparams)
if hybrid:
raw = _hybrid_fuse("model_cards", "models", query, qvec, raw, k, where_sql, wparams)
elif (name_toks := _name_tokens(query)):
raw = _fuse("model_cards", qvec, raw,
_name_ids("model_cards", name_toks, k, where_sql, wparams),
k, where_sql, wparams, weight=NAME_CHANNEL_WEIGHT)
rows = await _sort_results(_rows_to_model(raw), "model", sort_by, k)
return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
except Exception as e:
logger.error(f"search models: {e}")
raise HTTPException(status_code=500, detail="Model search failed")
@app.get("/similarity/models", response_model=ModelQueryResponse)
@cache(ttl=CACHE_TTL, key="sim:m:{model_id}:{k}:{sort_by}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}:{task}:{license}:{language}:{modified_after}")
async def find_similar_models(
model_id: str,
k: int = Query(default=5, ge=1, le=100),
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
min_likes: int = Query(default=0, ge=0),
min_downloads: int = Query(default=0, ge=0),
min_param_count: int = Query(default=0, ge=0),
max_param_count: Optional[int] = Query(default=None, ge=0),
task: Optional[str] = Query(default=None),
license: Optional[str] = Query(default=None),
language: Optional[str] = Query(default=None),
modified_after: Optional[str] = Query(default=None),
):
emb = _emb_of("model_cards", model_id)
if emb is None:
raise HTTPException(status_code=404, detail=f"Model ID '{model_id}' not found")
try:
use_param = min_param_count > 0 or max_param_count is not None
n = k * 4 if sort_by != "similarity" else k + 1
where_sql, wparams = _where(min_likes, min_downloads, min_param_count, max_param_count, use_param,
task=task, license=license, language=language,
modified_after=modified_after)
raw = _knn("model_cards", emb, n, where_sql, COLS, wparams)
rows = [r for r in _rows_to_model(raw) if r["_id"] != model_id]
rows = await _sort_results(rows, "model", sort_by, k)
return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
except HTTPException:
raise
except Exception as e:
logger.error(f"similarity models: {e}")
raise HTTPException(status_code=500, detail="Model similarity search failed")
# --- trending (proxies HF API, joins summaries from the in-memory table) -----
@cache(ttl=TRENDING_CACHE_TTL, key="tscore:{item_type}:{item_id}")
async def get_trending_score(item_id: str, item_type: str) -> float:
try:
endpoint = "models" if item_type == "model" else "datasets"
async with httpx.AsyncClient(timeout=10) as c:
r = await c.get(f"https://huggingface.co/api/{endpoint}/{item_id}?expand=trendingScore")
r.raise_for_status()
return r.json().get("trendingScore", 0)
except Exception:
return 0
@app.get("/trending/datasets", response_model=QueryResponse)
@cache(ttl=TRENDING_CACHE_TTL, key="trend:d:{limit}:{min_likes}:{min_downloads}")
async def trending_datasets(
limit: int = Query(default=10, ge=1, le=100),
min_likes: int = Query(default=0, ge=0),
min_downloads: int = Query(default=0, ge=0),
):
try:
async with httpx.AsyncClient(timeout=15) as c:
r = await c.get("https://huggingface.co/api/datasets?sort=trendingScore&limit=200")
r.raise_for_status()
items = r.json()
except Exception as e:
logger.error(f"trending datasets fetch: {e}")
raise HTTPException(status_code=502, detail="Failed to fetch trending datasets")
items = [d for d in items if d.get("likes", 0) >= min_likes and d.get("downloads", 0) >= min_downloads]
ids = [d["id"] for d in items[: limit * 3]]
if not ids:
return QueryResponse(results=[])
placeholders = ",".join("?" for _ in ids)
found = {r[0]: r for r in con.execute(
f"SELECT id, summary, likes, downloads, {META_COLS} FROM dataset_cards WHERE id IN ({placeholders})", ids
).fetchall()}
out = []
for d in items:
row = found.get(d["id"])
if row:
out.append(QueryResult(dataset_id=d["id"], similarity=1.0, summary=row[1],
likes=d.get("likes", 0), downloads=d.get("downloads", 0),
task=row[4], license=row[5], language=row[6],
last_modified=row[7]))
if len(out) >= limit:
break
return QueryResponse(results=out)
@app.get("/trending/models", response_model=ModelQueryResponse)
@cache(ttl=TRENDING_CACHE_TTL, key="trend:m:{limit}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}")
async def trending_models(
limit: int = Query(default=10, ge=1, le=100),
min_likes: int = Query(default=0, ge=0),
min_downloads: int = Query(default=0, ge=0),
min_param_count: int = Query(default=0, ge=0),
max_param_count: Optional[int] = Query(default=None, ge=0),
):
try:
async with httpx.AsyncClient(timeout=15) as c:
r = await c.get("https://huggingface.co/api/models?sort=trendingScore&limit=200")
r.raise_for_status()
items = r.json()
except Exception as e:
logger.error(f"trending models fetch: {e}")
raise HTTPException(status_code=502, detail="Failed to fetch trending models")
items = [m for m in items if m.get("likes", 0) >= min_likes and m.get("downloads", 0) >= min_downloads]
ids = [m["id"] for m in items[: limit * 3]]
if not ids:
return ModelQueryResponse(results=[])
placeholders = ",".join("?" for _ in ids)
found = {r[0]: r for r in con.execute(
f"SELECT id, summary, likes, downloads, param_count, {META_COLS} FROM model_cards WHERE id IN ({placeholders})", ids
).fetchall()}
use_param = min_param_count > 0 or max_param_count is not None
out = []
for m in items:
row = found.get(m["id"])
if not row:
continue
param = int(row[4])
if use_param:
if param == 0:
continue
if min_param_count > 0 and param < min_param_count:
continue
if max_param_count is not None and param > max_param_count:
continue
out.append(ModelQueryResult(model_id=m["id"], similarity=1.0, summary=row[1],
likes=m.get("likes", 0), downloads=m.get("downloads", 0),
param_count=param,
task=row[5], license=row[6], language=row[7],
last_modified=row[8]))
if len(out) >= limit:
break
return ModelQueryResponse(results=out)
# --- embedding primitives (personal-feeds prototype + compare tooling) -------
class EmbBatchReq(BaseModel):
repo_type: str
ids: List[str]
@app.post("/embeddings_batch")
async def embeddings_batch(req: EmbBatchReq):
"""Return stored embeddings for up to 200 repo ids (e.g. a user's likes)."""
if req.repo_type not in ("datasets", "models"):
raise HTTPException(status_code=422, detail="repo_type must be datasets|models")
table = "dataset_cards" if req.repo_type == "datasets" else "model_cards"
ids = req.ids[:200]
if not ids:
return {"embeddings": {}}
placeholders = ",".join("?" for _ in ids)
rows = con.execute(
f"SELECT id, emb FROM {table} WHERE id IN ({placeholders})", ids
).fetchall()
return {"embeddings": {r[0]: [float(x) for x in r[1]] for r in rows}}
@app.get("/embed_query")
@cache(ttl=CACHE_TTL, key="embq:{repo_type}:{text}")
async def embed_query_route(
text: str = Query(min_length=3, max_length=200),
repo_type: str = Query(default="datasets"),
):
"""Embed a free-text topic phrase with the serving query encoder.
repo_type picks the query prefix the corpus was embedded against
(datasets use the instruct prompt, models the search_query prefix).
"""
if repo_type not in ("datasets", "models"):
raise HTTPException(status_code=422, detail="repo_type must be datasets|models")
prompt = (f"Instruct: {DS_TASK}\nQuery:{text}" if repo_type == "datasets"
else f"search_query: {text}")
return {"embedding": embed_query(prompt)}
class VectorSearchReq(BaseModel):
repo_type: str
vector: List[float]
k: int = 20
exclude_ids: List[str] = []
include_embeddings: bool = False
@app.post("/search_vector")
async def search_vector(req: VectorSearchReq):
"""KNN search with a caller-supplied (already 256-d, normalised) vector.
Powers relevance-feedback loops: refine a preference vector client-side
from votes, then pull fresh results for it. exclude_ids drops already-seen
items server-side so refinement surfaces new material.
"""
if req.repo_type not in ("datasets", "models"):
raise HTTPException(status_code=422, detail="repo_type must be datasets|models")
if len(req.vector) != EMB_DIM:
raise HTTPException(status_code=422, detail=f"vector must be {EMB_DIM}-d")
table = "dataset_cards" if req.repo_type == "datasets" else "model_cards"
k = max(1, min(int(req.k), 100))
n = k + len(req.exclude_ids[:300])
raw = _knn(table, req.vector, n, "", COLS, return_emb=req.include_embeddings)
excl = set(req.exclude_ids[:300])
out = []
for row in raw:
emb_val = row[-2] if req.include_embeddings else None
(id_, summary, likes, downloads, param, task, license_, language, last_modified) = row[:9]
if id_ in excl:
continue
item = {"id": id_, "summary": summary, "likes": int(likes),
"downloads": int(downloads), "param_count": int(param),
"task": task, "license": license_, "language": language,
"last_modified": last_modified, "similarity": float(row[-1])}
if req.include_embeddings:
item["embedding"] = [float(x) for x in emb_val]
out.append(item)
if len(out) >= k:
break
return {"results": out}
# --- per-user feeds store (OAuth token -> whoami -> one JSON doc per user) ---
FEEDS_STORE_REPO = os.getenv("FEEDS_STORE_REPO", "davanstrien/hub-feeds-store")
FEEDS_DOC_LIMIT = 300_000 # bytes
@cache(ttl="10m", key="whoami:{token}")
async def _resolve_user(token: str) -> str:
async with httpx.AsyncClient(timeout=10) as c:
r = await c.get("https://huggingface.co/api/whoami-v2",
headers={"Authorization": f"Bearer {token}"})
if r.status_code != 200:
raise HTTPException(status_code=401, detail="Invalid or expired HF token")
name = r.json().get("name")
if not name:
raise HTTPException(status_code=401, detail="Could not resolve username")
return name
async def _auth_user(authorization: str) -> str:
if not authorization or not authorization.startswith("Bearer "):
raise HTTPException(status_code=401, detail="Missing Bearer token")
return await _resolve_user(authorization[7:])
@app.get("/feeds_store")
async def feeds_load(authorization: str = Header(default=None)):
user = await _auth_user(authorization)
import json as _json
try:
path = hf_hub_download(FEEDS_STORE_REPO, f"users/{user}.json",
repo_type="dataset", force_download=True)
with open(path) as f:
return {"user": user, "doc": _json.load(f)}
except Exception:
return {"user": user, "doc": None}
class FeedsSaveReq(BaseModel):
doc: dict
@app.put("/feeds_store")
async def feeds_save(req: FeedsSaveReq, authorization: str = Header(default=None)):
user = await _auth_user(authorization)
import json as _json
payload = _json.dumps(req.doc).encode()
if len(payload) > FEEDS_DOC_LIMIT:
raise HTTPException(status_code=413, detail="Feeds doc too large")
HfApi().upload_file(path_or_fileobj=payload, path_in_repo=f"users/{user}.json",
repo_id=FEEDS_STORE_REPO, repo_type="dataset",
commit_message=f"feeds sync: {user}")
return {"ok": True, "user": user, "bytes": len(payload)}
# --- suggest (autocomplete; old backend never implemented it — bonus) --------
@app.get("/suggest/{repo_type}")
@cache(ttl="1h", key="suggest:{repo_type}:{q}")
async def suggest(repo_type: str, q: str):
if len(q) < 2:
return {"suggestions": []}
table = "dataset_cards" if repo_type == "datasets" else "model_cards"
rows = con.execute(
f"SELECT id FROM {table} WHERE id ILIKE ? ORDER BY downloads DESC LIMIT 10",
[f"%{q}%"],
).fetchall()
return {"suggestions": [r[0] for r in rows]}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=7860)
|