davanstrien HF Staff commited on
Commit
e48acc2
·
verified ·
1 Parent(s): 8b7c589

Add task/license/language/modified_after filters + /lookup endpoint; union_by_name load

Browse files
Files changed (1) hide show
  1. app.py +81 -16
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
- conds = []
 
 
 
 
 
 
 
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
- return ("WHERE " + " AND ".join(conds)) if conds else ""
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- raw = _knn("dataset_cards", qvec, n, _where(min_likes, min_downloads), COLS)
 
 
 
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
- raw = _knn("dataset_cards", emb, n, _where(min_likes, min_downloads), COLS)
 
 
 
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
- raw = _knn("model_cards", qvec, n,
337
- _where(min_likes, min_downloads, min_param_count, max_param_count, use_param), COLS)
 
 
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
- raw = _knn("model_cards", emb, n,
363
- _where(min_likes, min_downloads, min_param_count, max_param_count, use_param), COLS)
 
 
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])