Spaces:
Running
Running
Add task/license/language/modified_after filters + /lookup endpoint; union_by_name load
Browse files
app.py
CHANGED
|
@@ -97,19 +97,25 @@ def load_table(cfg: str, table: str, has_param: bool):
|
|
| 97 |
plist = "[" + ",".join(f"'{p}'" for p in paths) + "]"
|
| 98 |
param_sel = "COALESCE(param_count,0)::BIGINT AS param_count" if has_param else "0::BIGINT AS param_count"
|
| 99 |
# slice to EMB_DIM, L2-renormalise, cast to fixed-size FLOAT[] array (nested, no correlated subquery)
|
|
|
|
|
|
|
|
|
|
| 100 |
con.execute(f"""
|
| 101 |
CREATE OR REPLACE TABLE {table} AS
|
| 102 |
SELECT id, summary,
|
| 103 |
COALESCE(likes,0)::BIGINT AS likes,
|
| 104 |
COALESCE(downloads,0)::BIGINT AS downloads,
|
| 105 |
last_modified,
|
|
|
|
|
|
|
|
|
|
| 106 |
{param_sel},
|
| 107 |
list_transform(s, x -> (x / nrm))::FLOAT[{EMB_DIM}] AS emb
|
| 108 |
FROM (
|
| 109 |
SELECT *, sqrt(list_dot_product(s, s)) AS nrm
|
| 110 |
FROM (
|
| 111 |
SELECT *, embedding[1:{EMB_DIM}] AS s
|
| 112 |
-
FROM read_parquet({plist})
|
| 113 |
WHERE embedding IS NOT NULL
|
| 114 |
QUALIFY row_number() OVER (PARTITION BY id ORDER BY last_modified DESC NULLS LAST) = 1
|
| 115 |
)
|
|
@@ -184,8 +190,15 @@ class ModelQueryResponse(BaseModel):
|
|
| 184 |
|
| 185 |
|
| 186 |
# --- core search -------------------------------------------------------------
|
| 187 |
-
def _where(min_likes, min_downloads, min_param=0, max_param=None, use_param=False
|
| 188 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 189 |
if min_likes > 0:
|
| 190 |
conds.append(f"likes >= {int(min_likes)}")
|
| 191 |
if min_downloads > 0:
|
|
@@ -196,17 +209,30 @@ def _where(min_likes, min_downloads, min_param=0, max_param=None, use_param=Fals
|
|
| 196 |
conds.append(f"param_count >= {int(min_param)}")
|
| 197 |
if max_param is not None:
|
| 198 |
conds.append(f"param_count <= {int(max_param)}")
|
| 199 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 200 |
|
| 201 |
|
| 202 |
def _vec_literal(vec):
|
| 203 |
return "[" + ",".join(repr(float(x)) for x in vec) + f"]::FLOAT[{EMB_DIM}]"
|
| 204 |
|
| 205 |
|
| 206 |
-
def _knn(table, qvec, k, where_sql, cols):
|
| 207 |
q = f"""SELECT {cols}, array_cosine_similarity(emb, {_vec_literal(qvec)}) AS sim
|
| 208 |
FROM {table} {where_sql} ORDER BY sim DESC LIMIT {int(k)}"""
|
| 209 |
-
return con.execute(q).fetchall()
|
| 210 |
|
| 211 |
|
| 212 |
def _emb_of(table, repo_id):
|
|
@@ -273,19 +299,39 @@ async def health():
|
|
| 273 |
}
|
| 274 |
|
| 275 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 276 |
@app.get("/search/datasets", response_model=QueryResponse)
|
| 277 |
-
@cache(ttl=CACHE_TTL, key="search:d:{query}:{k}:{sort_by}:{min_likes}:{min_downloads}")
|
| 278 |
async def search_datasets(
|
| 279 |
query: str,
|
| 280 |
k: int = Query(default=5, ge=1, le=100),
|
| 281 |
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending"]),
|
| 282 |
min_likes: int = Query(default=0, ge=0),
|
| 283 |
min_downloads: int = Query(default=0, ge=0),
|
|
|
|
|
|
|
|
|
|
|
|
|
| 284 |
):
|
| 285 |
try:
|
| 286 |
qvec = embed_query(f"Instruct: {DS_TASK}\nQuery:{query}")
|
| 287 |
n = k * 4 if sort_by != "similarity" else k
|
| 288 |
-
|
|
|
|
|
|
|
|
|
|
| 289 |
rows = await _sort_results(_rows_to_dataset(raw), "dataset", sort_by, k)
|
| 290 |
return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|
| 291 |
except Exception as e:
|
|
@@ -294,20 +340,27 @@ async def search_datasets(
|
|
| 294 |
|
| 295 |
|
| 296 |
@app.get("/similarity/datasets", response_model=QueryResponse)
|
| 297 |
-
@cache(ttl=CACHE_TTL, key="sim:d:{dataset_id}:{k}:{sort_by}:{min_likes}:{min_downloads}")
|
| 298 |
async def find_similar_datasets(
|
| 299 |
dataset_id: str,
|
| 300 |
k: int = Query(default=5, ge=1, le=100),
|
| 301 |
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending"]),
|
| 302 |
min_likes: int = Query(default=0, ge=0),
|
| 303 |
min_downloads: int = Query(default=0, ge=0),
|
|
|
|
|
|
|
|
|
|
|
|
|
| 304 |
):
|
| 305 |
emb = _emb_of("dataset_cards", dataset_id)
|
| 306 |
if emb is None:
|
| 307 |
raise HTTPException(status_code=404, detail=f"Dataset ID '{dataset_id}' not found")
|
| 308 |
try:
|
| 309 |
n = k * 4 if sort_by != "similarity" else k + 1
|
| 310 |
-
|
|
|
|
|
|
|
|
|
|
| 311 |
rows = [r for r in _rows_to_dataset(raw) if r["_id"] != dataset_id]
|
| 312 |
rows = await _sort_results(rows, "dataset", sort_by, k)
|
| 313 |
return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|
|
@@ -319,7 +372,7 @@ async def find_similar_datasets(
|
|
| 319 |
|
| 320 |
|
| 321 |
@app.get("/search/models", response_model=ModelQueryResponse)
|
| 322 |
-
@cache(ttl=CACHE_TTL, key="search:m:{query}:{k}:{sort_by}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}")
|
| 323 |
async def search_models(
|
| 324 |
query: str,
|
| 325 |
k: int = Query(default=5, ge=1, le=100),
|
|
@@ -328,13 +381,19 @@ async def search_models(
|
|
| 328 |
min_downloads: int = Query(default=0, ge=0),
|
| 329 |
min_param_count: int = Query(default=0, ge=0),
|
| 330 |
max_param_count: Optional[int] = Query(default=None, ge=0),
|
|
|
|
|
|
|
|
|
|
|
|
|
| 331 |
):
|
| 332 |
try:
|
| 333 |
use_param = min_param_count > 0 or max_param_count is not None
|
| 334 |
qvec = embed_query(f"search_query: {query}")
|
| 335 |
n = k * 4 if sort_by != "similarity" else k
|
| 336 |
-
|
| 337 |
-
|
|
|
|
|
|
|
| 338 |
rows = await _sort_results(_rows_to_model(raw), "model", sort_by, k)
|
| 339 |
return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|
| 340 |
except Exception as e:
|
|
@@ -343,7 +402,7 @@ async def search_models(
|
|
| 343 |
|
| 344 |
|
| 345 |
@app.get("/similarity/models", response_model=ModelQueryResponse)
|
| 346 |
-
@cache(ttl=CACHE_TTL, key="sim:m:{model_id}:{k}:{sort_by}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}")
|
| 347 |
async def find_similar_models(
|
| 348 |
model_id: str,
|
| 349 |
k: int = Query(default=5, ge=1, le=100),
|
|
@@ -352,6 +411,10 @@ async def find_similar_models(
|
|
| 352 |
min_downloads: int = Query(default=0, ge=0),
|
| 353 |
min_param_count: int = Query(default=0, ge=0),
|
| 354 |
max_param_count: Optional[int] = Query(default=None, ge=0),
|
|
|
|
|
|
|
|
|
|
|
|
|
| 355 |
):
|
| 356 |
emb = _emb_of("model_cards", model_id)
|
| 357 |
if emb is None:
|
|
@@ -359,8 +422,10 @@ async def find_similar_models(
|
|
| 359 |
try:
|
| 360 |
use_param = min_param_count > 0 or max_param_count is not None
|
| 361 |
n = k * 4 if sort_by != "similarity" else k + 1
|
| 362 |
-
|
| 363 |
-
|
|
|
|
|
|
|
| 364 |
rows = [r for r in _rows_to_model(raw) if r["_id"] != model_id]
|
| 365 |
rows = await _sort_results(rows, "model", sort_by, k)
|
| 366 |
return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|
|
|
|
| 97 |
plist = "[" + ",".join(f"'{p}'" for p in paths) + "]"
|
| 98 |
param_sel = "COALESCE(param_count,0)::BIGINT AS param_count" if has_param else "0::BIGINT AS param_count"
|
| 99 |
# slice to EMB_DIM, L2-renormalise, cast to fixed-size FLOAT[] array (nested, no correlated subquery)
|
| 100 |
+
# task/license/language are copied straight from the source cards; cast to
|
| 101 |
+
# VARCHAR because the models config's all-null `language` reads back as a
|
| 102 |
+
# NULL-typed column that would otherwise reject `=` filters at query time.
|
| 103 |
con.execute(f"""
|
| 104 |
CREATE OR REPLACE TABLE {table} AS
|
| 105 |
SELECT id, summary,
|
| 106 |
COALESCE(likes,0)::BIGINT AS likes,
|
| 107 |
COALESCE(downloads,0)::BIGINT AS downloads,
|
| 108 |
last_modified,
|
| 109 |
+
CAST(task AS VARCHAR) AS task,
|
| 110 |
+
CAST(license AS VARCHAR) AS license,
|
| 111 |
+
CAST(language AS VARCHAR) AS language,
|
| 112 |
{param_sel},
|
| 113 |
list_transform(s, x -> (x / nrm))::FLOAT[{EMB_DIM}] AS emb
|
| 114 |
FROM (
|
| 115 |
SELECT *, sqrt(list_dot_product(s, s)) AS nrm
|
| 116 |
FROM (
|
| 117 |
SELECT *, embedding[1:{EMB_DIM}] AS s
|
| 118 |
+
FROM read_parquet({plist}, union_by_name=True)
|
| 119 |
WHERE embedding IS NOT NULL
|
| 120 |
QUALIFY row_number() OVER (PARTITION BY id ORDER BY last_modified DESC NULLS LAST) = 1
|
| 121 |
)
|
|
|
|
| 190 |
|
| 191 |
|
| 192 |
# --- core search -------------------------------------------------------------
|
| 193 |
+
def _where(min_likes, min_downloads, min_param=0, max_param=None, use_param=False,
|
| 194 |
+
task=None, license=None, language=None, modified_after=None):
|
| 195 |
+
"""Build a WHERE clause + its bind params.
|
| 196 |
+
|
| 197 |
+
Numeric bounds are inlined (ints are injection-safe); the free-text metadata
|
| 198 |
+
filters (task/license/language/modified_after) are bound as `?` params.
|
| 199 |
+
Returns ``(where_sql, params)`` for ``_knn``.
|
| 200 |
+
"""
|
| 201 |
+
conds, params = [], []
|
| 202 |
if min_likes > 0:
|
| 203 |
conds.append(f"likes >= {int(min_likes)}")
|
| 204 |
if min_downloads > 0:
|
|
|
|
| 209 |
conds.append(f"param_count >= {int(min_param)}")
|
| 210 |
if max_param is not None:
|
| 211 |
conds.append(f"param_count <= {int(max_param)}")
|
| 212 |
+
if task:
|
| 213 |
+
conds.append("task = ?")
|
| 214 |
+
params.append(task)
|
| 215 |
+
if license:
|
| 216 |
+
conds.append("license = ?")
|
| 217 |
+
params.append(license)
|
| 218 |
+
if language:
|
| 219 |
+
conds.append("language = ?")
|
| 220 |
+
params.append(language)
|
| 221 |
+
if modified_after:
|
| 222 |
+
# last_modified is an ISO-8601 string; lexicographic >= is chronological.
|
| 223 |
+
conds.append("last_modified >= ?")
|
| 224 |
+
params.append(modified_after)
|
| 225 |
+
return (("WHERE " + " AND ".join(conds)) if conds else ""), params
|
| 226 |
|
| 227 |
|
| 228 |
def _vec_literal(vec):
|
| 229 |
return "[" + ",".join(repr(float(x)) for x in vec) + f"]::FLOAT[{EMB_DIM}]"
|
| 230 |
|
| 231 |
|
| 232 |
+
def _knn(table, qvec, k, where_sql, cols, params=None):
|
| 233 |
q = f"""SELECT {cols}, array_cosine_similarity(emb, {_vec_literal(qvec)}) AS sim
|
| 234 |
FROM {table} {where_sql} ORDER BY sim DESC LIMIT {int(k)}"""
|
| 235 |
+
return con.execute(q, params or []).fetchall()
|
| 236 |
|
| 237 |
|
| 238 |
def _emb_of(table, repo_id):
|
|
|
|
| 299 |
}
|
| 300 |
|
| 301 |
|
| 302 |
+
@app.get("/lookup/{repo_id:path}")
|
| 303 |
+
async def lookup(repo_id: str):
|
| 304 |
+
"""Resolve a repo id to its type in one call: {"id":..., "type": "dataset"|"model"}.
|
| 305 |
+
|
| 306 |
+
Saves clients the 404-fallthrough (try datasets, then models). 404 with a JSON
|
| 307 |
+
detail if the id is in neither index. `:path` so `org/name` ids keep their slash.
|
| 308 |
+
"""
|
| 309 |
+
for table, typ in (("dataset_cards", "dataset"), ("model_cards", "model")):
|
| 310 |
+
if con.execute(f"SELECT 1 FROM {table} WHERE id = ? LIMIT 1", [repo_id]).fetchone():
|
| 311 |
+
return {"id": repo_id, "type": typ}
|
| 312 |
+
raise HTTPException(status_code=404, detail=f"'{repo_id}' not found as a dataset or model in the index")
|
| 313 |
+
|
| 314 |
+
|
| 315 |
@app.get("/search/datasets", response_model=QueryResponse)
|
| 316 |
+
@cache(ttl=CACHE_TTL, key="search:d:{query}:{k}:{sort_by}:{min_likes}:{min_downloads}:{task}:{license}:{language}:{modified_after}")
|
| 317 |
async def search_datasets(
|
| 318 |
query: str,
|
| 319 |
k: int = Query(default=5, ge=1, le=100),
|
| 320 |
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending"]),
|
| 321 |
min_likes: int = Query(default=0, ge=0),
|
| 322 |
min_downloads: int = Query(default=0, ge=0),
|
| 323 |
+
task: Optional[str] = Query(default=None),
|
| 324 |
+
license: Optional[str] = Query(default=None),
|
| 325 |
+
language: Optional[str] = Query(default=None),
|
| 326 |
+
modified_after: Optional[str] = Query(default=None),
|
| 327 |
):
|
| 328 |
try:
|
| 329 |
qvec = embed_query(f"Instruct: {DS_TASK}\nQuery:{query}")
|
| 330 |
n = k * 4 if sort_by != "similarity" else k
|
| 331 |
+
where_sql, wparams = _where(min_likes, min_downloads,
|
| 332 |
+
task=task, license=license, language=language,
|
| 333 |
+
modified_after=modified_after)
|
| 334 |
+
raw = _knn("dataset_cards", qvec, n, where_sql, COLS, wparams)
|
| 335 |
rows = await _sort_results(_rows_to_dataset(raw), "dataset", sort_by, k)
|
| 336 |
return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|
| 337 |
except Exception as e:
|
|
|
|
| 340 |
|
| 341 |
|
| 342 |
@app.get("/similarity/datasets", response_model=QueryResponse)
|
| 343 |
+
@cache(ttl=CACHE_TTL, key="sim:d:{dataset_id}:{k}:{sort_by}:{min_likes}:{min_downloads}:{task}:{license}:{language}:{modified_after}")
|
| 344 |
async def find_similar_datasets(
|
| 345 |
dataset_id: str,
|
| 346 |
k: int = Query(default=5, ge=1, le=100),
|
| 347 |
sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending"]),
|
| 348 |
min_likes: int = Query(default=0, ge=0),
|
| 349 |
min_downloads: int = Query(default=0, ge=0),
|
| 350 |
+
task: Optional[str] = Query(default=None),
|
| 351 |
+
license: Optional[str] = Query(default=None),
|
| 352 |
+
language: Optional[str] = Query(default=None),
|
| 353 |
+
modified_after: Optional[str] = Query(default=None),
|
| 354 |
):
|
| 355 |
emb = _emb_of("dataset_cards", dataset_id)
|
| 356 |
if emb is None:
|
| 357 |
raise HTTPException(status_code=404, detail=f"Dataset ID '{dataset_id}' not found")
|
| 358 |
try:
|
| 359 |
n = k * 4 if sort_by != "similarity" else k + 1
|
| 360 |
+
where_sql, wparams = _where(min_likes, min_downloads,
|
| 361 |
+
task=task, license=license, language=language,
|
| 362 |
+
modified_after=modified_after)
|
| 363 |
+
raw = _knn("dataset_cards", emb, n, where_sql, COLS, wparams)
|
| 364 |
rows = [r for r in _rows_to_dataset(raw) if r["_id"] != dataset_id]
|
| 365 |
rows = await _sort_results(rows, "dataset", sort_by, k)
|
| 366 |
return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|
|
|
|
| 372 |
|
| 373 |
|
| 374 |
@app.get("/search/models", response_model=ModelQueryResponse)
|
| 375 |
+
@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}")
|
| 376 |
async def search_models(
|
| 377 |
query: str,
|
| 378 |
k: int = Query(default=5, ge=1, le=100),
|
|
|
|
| 381 |
min_downloads: int = Query(default=0, ge=0),
|
| 382 |
min_param_count: int = Query(default=0, ge=0),
|
| 383 |
max_param_count: Optional[int] = Query(default=None, ge=0),
|
| 384 |
+
task: Optional[str] = Query(default=None),
|
| 385 |
+
license: Optional[str] = Query(default=None),
|
| 386 |
+
language: Optional[str] = Query(default=None),
|
| 387 |
+
modified_after: Optional[str] = Query(default=None),
|
| 388 |
):
|
| 389 |
try:
|
| 390 |
use_param = min_param_count > 0 or max_param_count is not None
|
| 391 |
qvec = embed_query(f"search_query: {query}")
|
| 392 |
n = k * 4 if sort_by != "similarity" else k
|
| 393 |
+
where_sql, wparams = _where(min_likes, min_downloads, min_param_count, max_param_count, use_param,
|
| 394 |
+
task=task, license=license, language=language,
|
| 395 |
+
modified_after=modified_after)
|
| 396 |
+
raw = _knn("model_cards", qvec, n, where_sql, COLS, wparams)
|
| 397 |
rows = await _sort_results(_rows_to_model(raw), "model", sort_by, k)
|
| 398 |
return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|
| 399 |
except Exception as e:
|
|
|
|
| 402 |
|
| 403 |
|
| 404 |
@app.get("/similarity/models", response_model=ModelQueryResponse)
|
| 405 |
+
@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}")
|
| 406 |
async def find_similar_models(
|
| 407 |
model_id: str,
|
| 408 |
k: int = Query(default=5, ge=1, le=100),
|
|
|
|
| 411 |
min_downloads: int = Query(default=0, ge=0),
|
| 412 |
min_param_count: int = Query(default=0, ge=0),
|
| 413 |
max_param_count: Optional[int] = Query(default=None, ge=0),
|
| 414 |
+
task: Optional[str] = Query(default=None),
|
| 415 |
+
license: Optional[str] = Query(default=None),
|
| 416 |
+
language: Optional[str] = Query(default=None),
|
| 417 |
+
modified_after: Optional[str] = Query(default=None),
|
| 418 |
):
|
| 419 |
emb = _emb_of("model_cards", model_id)
|
| 420 |
if emb is None:
|
|
|
|
| 422 |
try:
|
| 423 |
use_param = min_param_count > 0 or max_param_count is not None
|
| 424 |
n = k * 4 if sort_by != "similarity" else k + 1
|
| 425 |
+
where_sql, wparams = _where(min_likes, min_downloads, min_param_count, max_param_count, use_param,
|
| 426 |
+
task=task, license=license, language=language,
|
| 427 |
+
modified_after=modified_after)
|
| 428 |
+
raw = _knn("model_cards", emb, n, where_sql, COLS, wparams)
|
| 429 |
rows = [r for r in _rows_to_model(raw) if r["_id"] != model_id]
|
| 430 |
rows = await _sort_results(rows, "model", sort_by, k)
|
| 431 |
return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
|