SPB Git forge

spb/ai-atlas

Public
41commits 1branches 0releases
4.6 MBsize
maindefault branch
12 days agolast push
HTML 77.2% TypeScript 10.5% Python 9.6% JavaScript 2.5%
12.3 KB · 167 lines python
Raw Blame History
1"""/compare?ids=a,b,… — side-by-side of 2–6 entities of the same type with per-type dimension sets and provenance.23API 1.1: benchmark dimensions are keyed by (benchmark, canonical metric, config_key) and only appear when EVERY compared model has a current4result in that comparability group; `comparability` labels each benchmark dimension; `diff_only=1` keeps only differing dimensions;5`mode=models|providers|hardware|companies|frameworks` asserts the expected type."""6from __future__ import annotations78from typing import Any910from fastapi import APIRouter, Query, Request1112from aiatlas.api.common import (13    COMPANY_TYPES,14    PRICE_COLS,15    PRICE_FROM,16    ApiError,17    cached,18    csv,19    enrich_provenance,20    entity_summary,21    price_row,22    resolve_entity,23)24from aiatlas.db import connection, fetch_all25from aiatlas.ontology.benchmarks import TRUST_LABELS, comparability26from aiatlas.services.frontier import best_per_model, config_summary, group_rows, load_results2728router = APIRouter(prefix="/api/v1/compare", tags=["compare"])2930D = lambda key, label, kind="text", unit=None, source="attr": {"key": key, "label": label, "kind": kind, **({"unit": unit} if unit else {}), "source": source}3132DIMENSIONS: dict[str, list[dict[str, Any]]] = {33    "model": [D("parameter_count", "Parameters", "number", "params"), D("active_parameter_count", "Active parameters", "number", "params"),34              D("context_length", "Context window", "number", "tokens"), D("max_output_tokens", "Max output", "number", "tokens"), D("openness", "Openness"),35              D("license", "License"), D("modalities", "Modalities", "list"), D("release_date", "Release date", "date"), D("knowledge_cutoff", "Knowledge cutoff", "date"),36              D("status", "Status"), D("family", "Family"), D("architecture", "Architecture"), D("reasoning", "Reasoning", "bool"), D("tool_calling", "Tool calling", "bool"),37              D("best_input_per_mtok", "Best input price", "number", "USD / 1M tokens", "prices"), D("best_output_per_mtok", "Best output price", "number", "USD / 1M tokens", "prices"),38              D("provider_count", "Providers", "number", None, "prices")],39    "provider": [D("website", "Website"), D("pricing_url", "Pricing page"), D("model_count", "Models priced", "number", None, "prices"),40                 D("min_input_per_mtok", "Min input price", "number", "USD / 1M tokens", "prices"), D("min_output_per_mtok", "Min output price", "number", "USD / 1M tokens", "prices"),41                 D("features", "Features", "list", None, "prices")],42    "hardware": [D("kind", "Kind"), D("manufacturer", "Manufacturer"), D("memory_gb", "Memory", "number", "GB"), D("memory_bandwidth_gbs", "Memory bandwidth", "number", "GB/s"),43                 D("tdp_watts", "TDP", "number", "W"), D("architecture", "Architecture"), D("memory_type", "Memory type"), D("runtimes", "Runtimes", "list"), D("release_date", "Release", "date")],44    "framework": [D("latest_version", "Latest version"), D("latest_release_at", "Latest release", "date"), D("license", "License"), D("language", "Language"),45                  D("metric.stars", "GitHub stars", "number"), D("metric.forks", "Forks", "number"), D("repository_url", "Repository")],46    "company": [D("country", "Country"), D("founded", "Founded", "date"), D("headquarters", "Headquarters"), D("org_kind", "Kind"), D("website", "Website"),47                D("model_count", "Models", "number", None, "graph"), D("paper_count", "Papers", "number", None, "graph")],48    "benchmark": [D("category", "Category"), D("metric", "Metric"), D("unit", "Unit"), D("task", "Task"), D("result_count", "Results", "number", None, "results")],49}50MODES = {"models": "model", "providers": "provider", "hardware": "hardware", "companies": "company", "frameworks": "framework", "benchmarks": "benchmark"}51TYPE_ALIASES = {"library": "framework", "runtime": "framework", "artifact": "model"}525354def _norm(v: Any) -> Any:55    if isinstance(v, list):56        return sorted(str(x) for x in v)57    if isinstance(v, float) and v.is_integer():58        return int(v)59    return v606162async def compare_entities(keys: list[str], *, diff_only: bool = False, mode: str | None = None) -> dict[str, Any]:63    if not 2 <= len(keys) <= 6:64        raise ApiError(400, "ids must list between 2 and 6 entities")65    if mode and mode not in MODES:66        raise ApiError(400, f"mode must be one of {', '.join(MODES)}")67    async with connection() as conn:68        rows = [await resolve_entity(conn, k) for k in keys]69        types = {TYPE_ALIASES.get(r["entity_type"], r["entity_type"]) for r in rows}70        etype = TYPE_ALIASES.get(rows[0]["entity_type"], rows[0]["entity_type"])71        if etype in COMPANY_TYPES:72            etype = "company"73            types = {"company"}74        if len(types) != 1:75            raise ApiError(400, f"all entities must share one type, got {', '.join(sorted(types))}")76        if mode and MODES[mode] != etype:77            raise ApiError(400, f"mode={mode} expects {MODES[mode]} entities, got {etype}")78        dims = [dict(d) for d in DIMENSIONS.get(etype, [])]79        eids = [r["id"] for r in rows]80        items: list[dict[str, Any]] = []81        extra: dict[str, dict[str, Any]] = {eid: {} for eid in eids}82        prices_by: dict[str, list[dict[str, Any]]] = {eid: [] for eid in eids}83        results_by: dict[str, list[dict[str, Any]]] = {eid: [] for eid in eids}84        comp: dict[str, dict[str, Any]] = {}85        if etype == "model":86            for p in await fetch_all(conn, f"select {PRICE_COLS} from {PRICE_FROM} where p.model_id = any(cast(:ids as text[])) and p.valid_to is null order by p.input_per_mtok nulls last", ids=eids):87                prices_by[p["m_id"]].append(price_row(p))88            for eid, plist in prices_by.items():89                ins = [x["input_per_mtok"] for x in plist if x["input_per_mtok"] is not None and x["input_per_mtok"] > 0]90                outs = [x["output_per_mtok"] for x in plist if x["output_per_mtok"] is not None and x["output_per_mtok"] > 0]91                extra[eid] = {"best_input_per_mtok": min(ins) if ins else None, "best_output_per_mtok": min(outs) if outs else None,92                              "provider_count": len({x["provider"]["id"] for x in plist if x["provider"]})}93            res = await load_results(conn, model_ids=eids, canonical_models_only=False)94            for r in res:95                results_by[r["model_id"]].append({"id": r["id"], "benchmark": {"id": r["benchmark_id"], "slug": r["benchmark_slug"], "name": r["benchmark_name"]}, "score": r["score"],96                                                  "metric": r["metric_canonical"], "unit": r.get("unit"), "higher_is_better": r.get("higher_is_better"), "config": config_summary(r.get("config")),97                                                  "config_key": r["config_key"], "trust_level": r["trust_level"], "evaluated_at": r.get("evaluated_at"), "observed_at": r["observed_at"],98                                                  "source_url": r.get("source_url"), "tier": r.get("tier")})99            for g in sorted(group_rows(res).values(), key=lambda g: (g["rows"][0]["benchmark_name"] or "", g["label"])):100                if g["models"] < set(eids) and not set(eids) <= g["models"]:101                    continue102                best = {r["model_id"]: r for r in best_per_model(g["rows"], g["higher_is_better"])}103                if not all(eid in best for eid in eids):104                    continue105                first = g["rows"][0]106                key = f"bench:{first['benchmark_slug']}:{g['metric']}:{g['config_key']}"107                ref = best[eids[0]]108                worst_level, reasons = "comparable", []109                for eid in eids[1:]:110                    level, why = comparability(ref.get("config"), best[eid].get("config"), ref.get("metric"), best[eid].get("metric"))111                    if level == "not-comparable" or (level == "partially-comparable" and worst_level == "comparable"):112                        worst_level = level113                    reasons += [w for w in why if w not in reasons and level != "comparable"]114                dims.append({"key": key, "label": f"{first['benchmark_name']} · {g['label']}", "kind": "number", "unit": first.get("unit"), "source": "results",115                             "higher_is_better": g["higher_is_better"], "benchmark": first["benchmark_slug"], "metric": g["metric"], "config_key": g["config_key"],116                             "comparability": worst_level, "trust_levels": sorted({best[e]["trust_level"] for e in eids})})117                comp[key] = {"level": worst_level, "reasons": reasons or ["same variant, metric and evaluation conditions"],118                             "trust": {e: {"level": best[e]["trust_level"], "label": TRUST_LABELS.get(best[e]["trust_level"], best[e]["trust_level"])} for e in eids}}119                for eid in eids:120                    extra[eid][key] = best[eid]["score"]121        elif etype == "provider":122            agg = await fetch_all(conn, """select provider_id, count(distinct model_id) as model_count, min(nullif(input_per_mtok, 0)) as min_input_per_mtok,123                                           min(nullif(output_per_mtok, 0)) as min_output_per_mtok, jsonb_agg(distinct k.key) filter (where k.key is not null) as features124                                           from prices p left join lateral jsonb_object_keys(p.features) k(key) on true125                                           where provider_id = any(cast(:ids as text[])) and valid_to is null group by 1""", ids=eids)126            for a in agg:127                extra[a["provider_id"]] = {"model_count": int(a["model_count"]), "min_input_per_mtok": a["min_input_per_mtok"], "min_output_per_mtok": a["min_output_per_mtok"], "features": a["features"] or []}128        elif etype == "company":129            agg = await fetch_all(conn, """select e.id, (select count(*) from entities m where m.organization_id = e.id and m.entity_type = 'model' and m.merged_into is null) as model_count,130                                           (select count(*) from entities m where m.organization_id = e.id and m.entity_type = 'paper' and m.merged_into is null) as paper_count131                                           from entities e where e.id = any(cast(:ids as text[]))""", ids=eids)132            for a in agg:133                extra[a["id"]] = {"model_count": int(a["model_count"]), "paper_count": int(a["paper_count"])}134        elif etype == "benchmark":135            agg = await fetch_all(conn, "select benchmark_id, count(*) as n from benchmark_results where benchmark_id = any(cast(:ids as text[])) and valid_to is null group by 1", ids=eids)136            for a in agg:137                extra[a["benchmark_id"]] = {"result_count": int(a["n"])}138        await enrich_provenance(conn, *(r.get("provenance") for r in rows))139        for r in rows:140            attrs = r.get("attributes") or {}141            prov = r.get("provenance") or {}142            values = {d["key"]: (extra[r["id"]].get(d["key"]) if d.get("source") != "attr" else attrs.get(d["key"])) for d in dims}143            item: dict[str, Any] = {"entity": entity_summary(r), "values": values, "provenance": {k: prov[k] for k in values if k in prov}}144            if etype == "model":145                item["prices"] = prices_by[r["id"]]146                item["results"] = results_by[r["id"]]147            items.append(item)148    if diff_only:149        keep = [d for d in dims if len({repr(_norm(it["values"].get(d["key"]))) for it in items}) > 1]150        kept = {d["key"] for d in keep}151        dims = keep152        for it in items:153            it["values"] = {k: v for k, v in it["values"].items() if k in kept}154            it["provenance"] = {k: v for k, v in it["provenance"].items() if k in kept}155    return {"entity_type": etype, "dimensions": dims, "items": items, "comparability": comp, "diff_only": diff_only,156            "note": "Benchmark dimensions appear only when every compared model has a current result in the same comparability group (benchmark × metric × config_key)."}157158159@router.get("")160@cached(300)161async def compare(request: Request, ids: str = Query(..., description="2–6 slugs or ids, comma-separated"), diff_only: int = Query(0, ge=0, le=1),162                  mode: str | None = Query(None, description="models|providers|hardware|companies|frameworks|benchmarks")) -> dict[str, Any]:163    return await compare_entities(csv(ids), diff_only=bool(diff_only), mode=mode)164165166__all__ = ["DIMENSIONS", "compare_entities", "router"]167