modelspec-dev 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- api/__init__.py +0 -0
- api/class_fit.py +334 -0
- api/classes.py +557 -0
- api/ranking/__init__.py +12 -0
- api/ranking/engine.py +1943 -0
- cli/__init__.py +0 -0
- cli/modelspec/__init__.py +0 -0
- cli/modelspec/cli.py +1819 -0
- cli/modelspec/commands/__init__.py +0 -0
- cli/modelspec/decide_cmd.py +333 -0
- cli/modelspec/offline.py +623 -0
- cli/modelspec/snapshot.py +698 -0
- cli/modelspec/snapshot_build_cmd.py +49 -0
- cli/modelspec/verify_cmd.py +125 -0
- cli/modelspec/vocab_cmd.py +204 -0
- cli/modelspec/vocabulary_cache.py +54 -0
- decision/__init__.py +13 -0
- decision/capability.py +872 -0
- decision/computed.py +125 -0
- decision/contract.py +1575 -0
- decision/engine.py +238 -0
- decision/excluded.py +34 -0
- decision/explain.py +908 -0
- decision/filter.py +796 -0
- decision/model.py +438 -0
- decision/normalise.py +604 -0
- decision/optimise.py +320 -0
- decision/registry.py +717 -0
- decision/relax.py +132 -0
- decision/resolve.py +111 -0
- decision/schema.py +21 -0
- decision/snapshot.py +1483 -0
- decision/sources.py +544 -0
- decision/templates.py +134 -0
- decision/verify.py +1745 -0
- decision/vocabulary.py +433 -0
- modelspec_dev-0.1.0.dist-info/METADATA +101 -0
- modelspec_dev-0.1.0.dist-info/RECORD +63 -0
- modelspec_dev-0.1.0.dist-info/WHEEL +4 -0
- modelspec_dev-0.1.0.dist-info/entry_points.txt +2 -0
- modelspec_dev-0.1.0.dist-info/licenses/LICENSE +43 -0
- modelspec_dev-0.1.0.dist-info/licenses/LICENSE-DATA +428 -0
- pipeline/__init__.py +0 -0
- pipeline/class_export.py +172 -0
- pipeline/hardware.py +434 -0
- pipeline/hosts.py +247 -0
- pipeline/load.py +224 -0
- pipeline/ranking.py +551 -0
- registry/domains.yaml +130 -0
- registry/facets.yaml +888 -0
- registry/harnesses.yaml +79 -0
- registry/providers.yaml +354 -0
- registry/sources.yaml +3059 -0
- registry/templates.yaml +166 -0
- schema/__init__.py +0 -0
- schema/applicability.py +147 -0
- schema/benchmark.py +175 -0
- schema/benchmark_eligibility.py +304 -0
- schema/card.py +1463 -0
- schema/enrichment.py +162 -0
- schema/enums.py +327 -0
- schema/graph.py +406 -0
- schema/suppliers.py +72 -0
cli/modelspec/cli.py
ADDED
|
@@ -0,0 +1,1819 @@
|
|
|
1
|
+
"""ModelSpec CLI — explore, search, compare, and rank AI models.
|
|
2
|
+
|
|
3
|
+
Entry point: ``modelspec`` (registered in pyproject.toml).
|
|
4
|
+
Graph commands use FalkorDB on localhost:6382.
|
|
5
|
+
Community commands (gaps, research, contribute, validate) work offline from YAML files.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import json
|
|
11
|
+
import subprocess
|
|
12
|
+
import sys
|
|
13
|
+
from datetime import date
|
|
14
|
+
from pathlib import Path
|
|
15
|
+
from typing import Any, Optional
|
|
16
|
+
|
|
17
|
+
import typer
|
|
18
|
+
from falkordb import FalkorDB
|
|
19
|
+
from rich.console import Console
|
|
20
|
+
from rich.panel import Panel
|
|
21
|
+
from rich.progress import Progress, SpinnerColumn, TextColumn
|
|
22
|
+
from rich.table import Table
|
|
23
|
+
from rich.text import Text
|
|
24
|
+
from rich.tree import Tree
|
|
25
|
+
|
|
26
|
+
from . import offline as _offline # noqa: E402
|
|
27
|
+
|
|
28
|
+
# ───────────────────────────────────────────────────────────────
|
|
29
|
+
# App setup
|
|
30
|
+
# ───────────────────────────────────────────────────────────────
|
|
31
|
+
|
|
32
|
+
app = typer.Typer(
|
|
33
|
+
cls=_offline.ContractGroup,
|
|
34
|
+
name="modelspec",
|
|
35
|
+
help="ModelSpec — explore, search, compare, and rank AI models.",
|
|
36
|
+
no_args_is_help=True,
|
|
37
|
+
rich_markup_mode="rich",
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
# The offline path: answers from a local snapshot of the published export, with
|
|
41
|
+
# no database and no network. This is what dpf calls.
|
|
42
|
+
app.add_typer(_offline.app, name="offline")
|
|
43
|
+
app.add_typer(_offline.snapshot_app, name="snapshot")
|
|
44
|
+
|
|
45
|
+
# The decision contract (MODEL-135). Parses and validates a spec; the engine
|
|
46
|
+
# behind it lands in MODEL-141/142/145.
|
|
47
|
+
from . import decide_cmd as _decide_cmd # noqa: E402
|
|
48
|
+
|
|
49
|
+
app.command("decide")(_decide_cmd.decide)
|
|
50
|
+
|
|
51
|
+
from . import vocab_cmd as _vocab_cmd # noqa: E402
|
|
52
|
+
|
|
53
|
+
app.command("vocab", cls=_offline.ContractCommand)(_vocab_cmd.vocab)
|
|
54
|
+
|
|
55
|
+
# The decision snapshot (MODEL-138): a new subcommand beside `fetch` and `status`.
|
|
56
|
+
from . import snapshot_build_cmd as _snapshot_build_cmd # noqa: E402
|
|
57
|
+
|
|
58
|
+
_offline.snapshot_app.command("build", cls=_offline.ContractCommand)(_snapshot_build_cmd.build)
|
|
59
|
+
|
|
60
|
+
# Two-key verification (MODEL-140): re-reads queued values from their sources.
|
|
61
|
+
from . import verify_cmd as _verify_cmd # noqa: E402
|
|
62
|
+
|
|
63
|
+
app.add_typer(_verify_cmd.app, name="verify")
|
|
64
|
+
|
|
65
|
+
console = Console()
|
|
66
|
+
|
|
67
|
+
# ───────────────────────────────────────────────────────────────
|
|
68
|
+
# Graph connection helper
|
|
69
|
+
# ───────────────────────────────────────────────────────────────
|
|
70
|
+
|
|
71
|
+
_FALKORDB_HOST = "localhost"
|
|
72
|
+
_FALKORDB_PORT = 6382
|
|
73
|
+
_GRAPH_NAME = "modelspec"
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _get_graph():
|
|
77
|
+
"""Connect to FalkorDB and return the modelspec graph handle."""
|
|
78
|
+
try:
|
|
79
|
+
db = FalkorDB(host=_FALKORDB_HOST, port=_FALKORDB_PORT)
|
|
80
|
+
return db.select_graph(_GRAPH_NAME)
|
|
81
|
+
except Exception as exc:
|
|
82
|
+
console.print(f"[bold red]Error:[/] Could not connect to FalkorDB at {_FALKORDB_HOST}:{_FALKORDB_PORT}")
|
|
83
|
+
console.print("[dim]The graph commands need a local FalkorDB. For an answer without one, "
|
|
84
|
+
"use `modelspec snapshot fetch` then `modelspec offline rank`.[/]")
|
|
85
|
+
console.print(f" {exc}")
|
|
86
|
+
raise typer.Exit(1)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
# ───────────────────────────────────────────────────────────────
|
|
90
|
+
# Formatting helpers
|
|
91
|
+
# ───────────────────────────────────────────────────────────────
|
|
92
|
+
|
|
93
|
+
def _fmt_params(n: int | float | None) -> str:
|
|
94
|
+
"""Format parameter count to human-readable string."""
|
|
95
|
+
if n is None:
|
|
96
|
+
return "-"
|
|
97
|
+
n = int(n)
|
|
98
|
+
if n >= 1_000_000_000_000:
|
|
99
|
+
return f"{n / 1_000_000_000_000:.1f}T"
|
|
100
|
+
if n >= 1_000_000_000:
|
|
101
|
+
return f"{n / 1_000_000_000:.1f}B"
|
|
102
|
+
if n >= 1_000_000:
|
|
103
|
+
return f"{n / 1_000_000:.0f}M"
|
|
104
|
+
return f"{n:,}"
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def _fmt_cost(v: float | None) -> str:
|
|
108
|
+
"""Format cost per million tokens."""
|
|
109
|
+
if v is None:
|
|
110
|
+
return "-"
|
|
111
|
+
if v == 0:
|
|
112
|
+
return "[green]free[/]"
|
|
113
|
+
return f"${v:.2f}"
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def _fmt_elo(v: float | None) -> str:
|
|
117
|
+
if v is None:
|
|
118
|
+
return "-"
|
|
119
|
+
return f"{v:.0f}"
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def _fmt_float(v: float | None, suffix: str = "") -> str:
|
|
123
|
+
if v is None:
|
|
124
|
+
return "-"
|
|
125
|
+
return f"{v:.1f}{suffix}"
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def _fmt_int(v: int | None) -> str:
|
|
129
|
+
if v is None:
|
|
130
|
+
return "-"
|
|
131
|
+
return f"{v:,}"
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def _status_color(status: str | None) -> str:
|
|
135
|
+
colors = {
|
|
136
|
+
"active": "green",
|
|
137
|
+
"beta": "yellow",
|
|
138
|
+
"alpha": "yellow",
|
|
139
|
+
"preview": "cyan",
|
|
140
|
+
"deprecated": "red",
|
|
141
|
+
"sunset": "dim red",
|
|
142
|
+
}
|
|
143
|
+
if not status:
|
|
144
|
+
return "white"
|
|
145
|
+
return colors.get(status.lower(), "white")
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def _tier_style(tier: str | None) -> str:
|
|
149
|
+
if not tier:
|
|
150
|
+
return "-"
|
|
151
|
+
styles = {
|
|
152
|
+
"tier-1": "[bold green]tier-1[/]",
|
|
153
|
+
"tier-2": "[yellow]tier-2[/]",
|
|
154
|
+
"tier-3": "[dim]tier-3[/]",
|
|
155
|
+
}
|
|
156
|
+
return styles.get(tier, tier)
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def _bool_icon(v: Any) -> str:
|
|
160
|
+
if v is True:
|
|
161
|
+
return "[green]Y[/]"
|
|
162
|
+
if v is False:
|
|
163
|
+
return "[dim]-[/]"
|
|
164
|
+
return "-"
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def _node_props(node) -> dict[str, Any]:
|
|
168
|
+
"""Extract properties dict from a FalkorDB node."""
|
|
169
|
+
if hasattr(node, "properties"):
|
|
170
|
+
return node.properties
|
|
171
|
+
return {}
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
_UNPUBLISHED_MODEL_KEYS = ("applicable_field_coverage", "card_completeness")
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def _published_model_props(node) -> dict[str, Any]:
|
|
178
|
+
"""Model node properties minus the internal coverage statistic.
|
|
179
|
+
|
|
180
|
+
A stale FalkorDB may still hold the key from before MODEL-74 stopped
|
|
181
|
+
writing it. Drop it here so ``info`` cannot republish it.
|
|
182
|
+
"""
|
|
183
|
+
props = dict(_node_props(node))
|
|
184
|
+
for key in _UNPUBLISHED_MODEL_KEYS:
|
|
185
|
+
props.pop(key, None)
|
|
186
|
+
return props
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def _edge_props(edge) -> dict[str, Any]:
|
|
190
|
+
"""Extract properties dict from a FalkorDB edge."""
|
|
191
|
+
if hasattr(edge, "properties"):
|
|
192
|
+
return edge.properties
|
|
193
|
+
return {}
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
# ───────────────────────────────────────────────────────────────
|
|
197
|
+
# Offline YAML helpers (used by gaps, research, contribute, validate)
|
|
198
|
+
# ───────────────────────────────────────────────────────────────
|
|
199
|
+
|
|
200
|
+
# Resolve project root relative to this file:
|
|
201
|
+
# cli/modelspec/cli.py -> project root is ../../..
|
|
202
|
+
_PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
|
203
|
+
_MODELS_DIR = _PROJECT_ROOT / "models"
|
|
204
|
+
|
|
205
|
+
# Lazy import — schema.card lives at project root level
|
|
206
|
+
_ModelCard = None # type: ignore[assignment]
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def _get_model_card_class():
|
|
210
|
+
"""Import ModelCard lazily so graph-only commands don't need the schema on sys.path."""
|
|
211
|
+
global _ModelCard
|
|
212
|
+
if _ModelCard is not None:
|
|
213
|
+
return _ModelCard
|
|
214
|
+
root = str(_PROJECT_ROOT)
|
|
215
|
+
if root not in sys.path:
|
|
216
|
+
sys.path.insert(0, root)
|
|
217
|
+
from schema.card import ModelCard
|
|
218
|
+
_ModelCard = ModelCard
|
|
219
|
+
return ModelCard
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def _discover_card_files(
|
|
223
|
+
provider: str | None = None,
|
|
224
|
+
) -> list[Path]:
|
|
225
|
+
"""Find all .md model card files, optionally filtered by provider."""
|
|
226
|
+
if provider:
|
|
227
|
+
provider_dir = _MODELS_DIR / provider
|
|
228
|
+
if not provider_dir.is_dir():
|
|
229
|
+
return []
|
|
230
|
+
return sorted(provider_dir.glob("*.md"))
|
|
231
|
+
return sorted(_MODELS_DIR.glob("*/*.md"))
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
def _load_cards(
|
|
235
|
+
provider: str | None = None,
|
|
236
|
+
model_type: str | None = None,
|
|
237
|
+
) -> list[Any]:
|
|
238
|
+
"""Load and parse all model cards, with optional filters."""
|
|
239
|
+
ModelCard = _get_model_card_class()
|
|
240
|
+
cards = []
|
|
241
|
+
for path in _discover_card_files(provider=provider):
|
|
242
|
+
try:
|
|
243
|
+
card = ModelCard.from_yaml_file(path)
|
|
244
|
+
if model_type and card.identity.model_type:
|
|
245
|
+
if card.identity.model_type.value != model_type:
|
|
246
|
+
continue
|
|
247
|
+
elif model_type and not card.identity.model_type:
|
|
248
|
+
continue
|
|
249
|
+
card._source_path = path # stash for later use
|
|
250
|
+
cards.append(card)
|
|
251
|
+
except Exception:
|
|
252
|
+
# Skip unparseable cards silently in bulk load
|
|
253
|
+
continue
|
|
254
|
+
return cards
|
|
255
|
+
|
|
256
|
+
|
|
257
|
+
# Key benchmark fields that matter most
|
|
258
|
+
_KEY_BENCHMARKS = [
|
|
259
|
+
"humaneval", "gpqa_diamond", "mmlu_pro", "arena_elo_overall",
|
|
260
|
+
"swe_bench_verified", "math_500", "aime_2025", "live_code_bench",
|
|
261
|
+
"ifeval", "mmmu", "mteb_overall",
|
|
262
|
+
]
|
|
263
|
+
|
|
264
|
+
# Major providers get a higher importance multiplier
|
|
265
|
+
_MAJOR_PROVIDERS = {
|
|
266
|
+
"openai": 3.0, "anthropic": 3.0, "google": 3.0, "meta": 2.5,
|
|
267
|
+
"mistral": 2.0, "deepseek": 2.0, "qwen": 2.0, "xai": 2.0,
|
|
268
|
+
"cohere": 1.5, "microsoft": 1.5, "nvidia": 1.5, "amazon": 1.5,
|
|
269
|
+
}
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def _compute_gap_info(card: Any) -> dict[str, Any]:
|
|
273
|
+
"""Analyze a card for missing data and compute a priority score."""
|
|
274
|
+
missing: list[str] = []
|
|
275
|
+
|
|
276
|
+
# Benchmarks
|
|
277
|
+
bench = card.benchmarks
|
|
278
|
+
filled_bench = bench.filled_count()
|
|
279
|
+
if filled_bench == 0:
|
|
280
|
+
missing.append("benchmarks (none)")
|
|
281
|
+
elif filled_bench < 5:
|
|
282
|
+
missing.append(f"benchmarks ({filled_bench} filled)")
|
|
283
|
+
|
|
284
|
+
# Cost
|
|
285
|
+
if card.cost.input is None:
|
|
286
|
+
missing.append("cost.input")
|
|
287
|
+
if card.cost.output is None:
|
|
288
|
+
missing.append("cost.output")
|
|
289
|
+
|
|
290
|
+
# Architecture
|
|
291
|
+
if card.architecture.total_parameters is None:
|
|
292
|
+
missing.append("total_parameters")
|
|
293
|
+
|
|
294
|
+
# Context window (None means missing; 0 is a valid value for non-text models)
|
|
295
|
+
if card.modalities.text.context_window is None:
|
|
296
|
+
missing.append("context_window")
|
|
297
|
+
|
|
298
|
+
# Capabilities — check if the key sub-sections have overall tier set
|
|
299
|
+
cap = card.capabilities
|
|
300
|
+
cap_filled = False
|
|
301
|
+
for section_name in ("coding", "reasoning", "tool_use"):
|
|
302
|
+
sub = getattr(cap, section_name)
|
|
303
|
+
if sub.overall is not None:
|
|
304
|
+
cap_filled = True
|
|
305
|
+
break
|
|
306
|
+
# Also check language and creative via their own key fields
|
|
307
|
+
if not cap_filled and card.capabilities.language.multilingual:
|
|
308
|
+
cap_filled = True
|
|
309
|
+
if not cap_filled and card.capabilities.creative.writing is not None:
|
|
310
|
+
cap_filled = True
|
|
311
|
+
if not cap_filled:
|
|
312
|
+
missing.append("capabilities")
|
|
313
|
+
|
|
314
|
+
# Platform availability
|
|
315
|
+
platforms = card.availability.platforms_available()
|
|
316
|
+
if len(platforms) == 0:
|
|
317
|
+
missing.append("availability")
|
|
318
|
+
|
|
319
|
+
# Model importance heuristic
|
|
320
|
+
provider = card.identity.provider
|
|
321
|
+
importance = _MAJOR_PROVIDERS.get(provider, 1.0)
|
|
322
|
+
|
|
323
|
+
# Bump importance for models with community signals
|
|
324
|
+
if card.adoption.huggingface_downloads and card.adoption.huggingface_downloads > 100_000:
|
|
325
|
+
importance *= 1.5
|
|
326
|
+
arena_elo = card.benchmarks.scores.get("arena_elo_overall")
|
|
327
|
+
if arena_elo and arena_elo > 1200:
|
|
328
|
+
importance *= 1.3
|
|
329
|
+
|
|
330
|
+
priority = importance * len(missing)
|
|
331
|
+
|
|
332
|
+
return {
|
|
333
|
+
"model_id": card.identity.model_id,
|
|
334
|
+
"display_name": card.identity.display_name,
|
|
335
|
+
"applicable_field_coverage": card.applicable_field_coverage,
|
|
336
|
+
"missing": missing,
|
|
337
|
+
"missing_count": len(missing),
|
|
338
|
+
"priority": round(priority, 1),
|
|
339
|
+
"provider": provider,
|
|
340
|
+
}
|
|
341
|
+
|
|
342
|
+
|
|
343
|
+
# ───────────────────────────────────────────────────────────────
|
|
344
|
+
# 1. modelspec info <model_id>
|
|
345
|
+
# ───────────────────────────────────────────────────────────────
|
|
346
|
+
|
|
347
|
+
@app.command()
|
|
348
|
+
def info(
|
|
349
|
+
model_id: str = typer.Argument(..., help="Model ID, e.g. qwen/qwen3-30b-a3b"),
|
|
350
|
+
format: Optional[str] = typer.Option(None, "--format", "-f", help="Output format: json"),
|
|
351
|
+
) -> None:
|
|
352
|
+
"""Show the full model card for a given model ID."""
|
|
353
|
+
graph = _get_graph()
|
|
354
|
+
|
|
355
|
+
# Fetch model node
|
|
356
|
+
result = graph.query(
|
|
357
|
+
"MATCH (m:Model {id: $model_id}) RETURN m",
|
|
358
|
+
{"model_id": model_id},
|
|
359
|
+
)
|
|
360
|
+
if not result.result_set:
|
|
361
|
+
console.print(f"[bold red]Model not found:[/] {model_id}")
|
|
362
|
+
raise typer.Exit(1)
|
|
363
|
+
|
|
364
|
+
m = _published_model_props(result.result_set[0][0])
|
|
365
|
+
|
|
366
|
+
# Fetch provider
|
|
367
|
+
prov_result = graph.query(
|
|
368
|
+
"MATCH (m:Model {id: $mid})-[:MADE_BY]->(p:Provider) RETURN p",
|
|
369
|
+
{"mid": model_id},
|
|
370
|
+
)
|
|
371
|
+
provider = _node_props(prov_result.result_set[0][0]) if prov_result.result_set else {}
|
|
372
|
+
|
|
373
|
+
# Fetch capabilities
|
|
374
|
+
cap_result = graph.query(
|
|
375
|
+
"MATCH (m:Model {id: $mid})-[r:HAS_CAPABILITY]->(c:Capability) "
|
|
376
|
+
"RETURN c.id, c.name, r.tier ORDER BY c.id",
|
|
377
|
+
{"mid": model_id},
|
|
378
|
+
)
|
|
379
|
+
|
|
380
|
+
# Fetch benchmarks
|
|
381
|
+
bench_result = graph.query(
|
|
382
|
+
"MATCH (m:Model {id: $mid})-[r:SCORED_ON]->(b:Benchmark) "
|
|
383
|
+
"RETURN b.id, b.name, r.value ORDER BY b.id",
|
|
384
|
+
{"mid": model_id},
|
|
385
|
+
)
|
|
386
|
+
|
|
387
|
+
# Fetch hardware fits
|
|
388
|
+
hw_result = graph.query(
|
|
389
|
+
"MATCH (m:Model {id: $mid})-[r:FITS_ON]->(h:Hardware) "
|
|
390
|
+
"RETURN h.id, r",
|
|
391
|
+
{"mid": model_id},
|
|
392
|
+
)
|
|
393
|
+
|
|
394
|
+
# Fetch platforms
|
|
395
|
+
plat_result = graph.query(
|
|
396
|
+
"MATCH (m:Model {id: $mid})-[:AVAILABLE_ON]->(p:Platform) "
|
|
397
|
+
"RETURN p.id, p.display_name",
|
|
398
|
+
{"mid": model_id},
|
|
399
|
+
)
|
|
400
|
+
|
|
401
|
+
# Fetch license
|
|
402
|
+
lic_result = graph.query(
|
|
403
|
+
"MATCH (m:Model {id: $mid})-[:LICENSED_AS]->(l:License) RETURN l",
|
|
404
|
+
{"mid": model_id},
|
|
405
|
+
)
|
|
406
|
+
license_info = _node_props(lic_result.result_set[0][0]) if lic_result.result_set else {}
|
|
407
|
+
|
|
408
|
+
# Fetch tags
|
|
409
|
+
tag_result = graph.query(
|
|
410
|
+
"MATCH (m:Model {id: $mid})-[:TAGGED_WITH]->(t:Tag) RETURN t.id",
|
|
411
|
+
{"mid": model_id},
|
|
412
|
+
)
|
|
413
|
+
tags = [row[0] for row in tag_result.result_set]
|
|
414
|
+
|
|
415
|
+
# ── JSON output ────────────────────────────────────────────
|
|
416
|
+
if format == "json":
|
|
417
|
+
data = {
|
|
418
|
+
"model": m,
|
|
419
|
+
"provider": provider,
|
|
420
|
+
"license": license_info,
|
|
421
|
+
"tags": tags,
|
|
422
|
+
"capabilities": [
|
|
423
|
+
{"id": r[0], "name": r[1], "tier": r[2]}
|
|
424
|
+
for r in cap_result.result_set
|
|
425
|
+
],
|
|
426
|
+
"benchmarks": [
|
|
427
|
+
{"id": r[0], "name": r[1], "value": r[2]}
|
|
428
|
+
for r in bench_result.result_set
|
|
429
|
+
],
|
|
430
|
+
"hardware": [
|
|
431
|
+
{"id": r[0], **_edge_props(r[1])}
|
|
432
|
+
for r in hw_result.result_set
|
|
433
|
+
],
|
|
434
|
+
"platforms": [
|
|
435
|
+
{"id": r[0], "name": r[1]}
|
|
436
|
+
for r in plat_result.result_set
|
|
437
|
+
],
|
|
438
|
+
}
|
|
439
|
+
console.print_json(json.dumps(data, default=str))
|
|
440
|
+
return
|
|
441
|
+
|
|
442
|
+
# ── Rich output ────────────────────────────────────────────
|
|
443
|
+
|
|
444
|
+
# Identity panel
|
|
445
|
+
status = m.get("status", "")
|
|
446
|
+
status_clr = _status_color(status)
|
|
447
|
+
identity_lines = [
|
|
448
|
+
f"[bold]{m.get('display_name', model_id)}[/]",
|
|
449
|
+
f" ID: {m.get('id', model_id)}",
|
|
450
|
+
f" Provider: {provider.get('display_name', m.get('id', '').split('/')[0])}",
|
|
451
|
+
f" Type: {m.get('model_type', '-')}",
|
|
452
|
+
f" Status: [{status_clr}]{status}[/]",
|
|
453
|
+
f" Released: {m.get('release_date', '-')}",
|
|
454
|
+
f" Family: {m.get('family', '-')}",
|
|
455
|
+
f" Country: {m.get('origin_country', '-')}",
|
|
456
|
+
f" Open: {_bool_icon(m.get('open_weights'))}",
|
|
457
|
+
]
|
|
458
|
+
if tags:
|
|
459
|
+
identity_lines.append(f" Tags: {', '.join(tags)}")
|
|
460
|
+
|
|
461
|
+
console.print(Panel("\n".join(identity_lines), title="Identity", border_style="blue"))
|
|
462
|
+
|
|
463
|
+
# Architecture
|
|
464
|
+
arch_lines = []
|
|
465
|
+
if m.get("architecture_type"):
|
|
466
|
+
arch_lines.append(f" Architecture: {m['architecture_type']}")
|
|
467
|
+
if m.get("total_parameters"):
|
|
468
|
+
active = m.get("active_parameters")
|
|
469
|
+
params_str = _fmt_params(m["total_parameters"])
|
|
470
|
+
if active:
|
|
471
|
+
params_str += f" ({_fmt_params(active)} active)"
|
|
472
|
+
arch_lines.append(f" Parameters: {params_str}")
|
|
473
|
+
if m.get("context_window"):
|
|
474
|
+
arch_lines.append(f" Context: {_fmt_int(m['context_window'])} tokens")
|
|
475
|
+
if m.get("max_input"):
|
|
476
|
+
arch_lines.append(f" Max input: {_fmt_int(m['max_input'])} tokens")
|
|
477
|
+
if m.get("max_output"):
|
|
478
|
+
arch_lines.append(f" Max output: {_fmt_int(m['max_output'])} tokens")
|
|
479
|
+
if m.get("embedding_dimensions"):
|
|
480
|
+
arch_lines.append(f" Embedding: {_fmt_int(m['embedding_dimensions'])} dims")
|
|
481
|
+
if arch_lines:
|
|
482
|
+
console.print(Panel("\n".join(arch_lines), title="Architecture", border_style="cyan"))
|
|
483
|
+
|
|
484
|
+
# Capabilities
|
|
485
|
+
if cap_result.result_set:
|
|
486
|
+
cap_table = Table(title="Capabilities", show_header=True, header_style="bold magenta")
|
|
487
|
+
cap_table.add_column("Capability", style="white", min_width=30)
|
|
488
|
+
cap_table.add_column("Tier", justify="center")
|
|
489
|
+
for row in cap_result.result_set:
|
|
490
|
+
cap_table.add_row(row[1] or row[0], _tier_style(row[2]))
|
|
491
|
+
console.print(cap_table)
|
|
492
|
+
|
|
493
|
+
# Benchmarks
|
|
494
|
+
if bench_result.result_set:
|
|
495
|
+
bench_table = Table(title="Benchmarks", show_header=True, header_style="bold yellow")
|
|
496
|
+
bench_table.add_column("Benchmark", style="white", min_width=25)
|
|
497
|
+
bench_table.add_column("Score", justify="right", style="bold")
|
|
498
|
+
for row in bench_result.result_set:
|
|
499
|
+
bench_table.add_row(row[1] or row[0], _fmt_float(row[2]))
|
|
500
|
+
console.print(bench_table)
|
|
501
|
+
|
|
502
|
+
# Hardware fit
|
|
503
|
+
if hw_result.result_set:
|
|
504
|
+
hw_table = Table(title="Hardware Fit", show_header=True, header_style="bold green")
|
|
505
|
+
hw_table.add_column("Hardware", style="white", min_width=25)
|
|
506
|
+
hw_table.add_column("Quant", justify="center")
|
|
507
|
+
hw_table.add_column("VRAM/RAM", justify="right")
|
|
508
|
+
hw_table.add_column("tok/s", justify="right")
|
|
509
|
+
hw_table.add_column("TTFT", justify="right")
|
|
510
|
+
hw_table.add_column("Engine", justify="center")
|
|
511
|
+
for row in hw_result.result_set:
|
|
512
|
+
hw_id = row[0]
|
|
513
|
+
ep = _edge_props(row[1])
|
|
514
|
+
hw_table.add_row(
|
|
515
|
+
hw_id.replace("_", " ").title(),
|
|
516
|
+
ep.get("quantization", "-"),
|
|
517
|
+
_fmt_float(ep.get("vram_usage_gb"), " GB"),
|
|
518
|
+
_fmt_float(ep.get("tokens_per_sec"), ""),
|
|
519
|
+
_fmt_float(ep.get("ttft_ms"), " ms"),
|
|
520
|
+
ep.get("inference_engine", "-"),
|
|
521
|
+
)
|
|
522
|
+
console.print(hw_table)
|
|
523
|
+
|
|
524
|
+
# Availability
|
|
525
|
+
if plat_result.result_set:
|
|
526
|
+
plat_names = [row[1] or row[0] for row in plat_result.result_set]
|
|
527
|
+
console.print(Panel(
|
|
528
|
+
" " + " | ".join(plat_names),
|
|
529
|
+
title="Available On",
|
|
530
|
+
border_style="green",
|
|
531
|
+
))
|
|
532
|
+
|
|
533
|
+
# Cost
|
|
534
|
+
cost_lines = []
|
|
535
|
+
if m.get("cost_input") is not None:
|
|
536
|
+
cost_lines.append(f" Input: {_fmt_cost(m['cost_input'])}/M tokens")
|
|
537
|
+
if m.get("cost_output") is not None:
|
|
538
|
+
cost_lines.append(f" Output: {_fmt_cost(m['cost_output'])}/M tokens")
|
|
539
|
+
if cost_lines:
|
|
540
|
+
console.print(Panel("\n".join(cost_lines), title="Cost", border_style="yellow"))
|
|
541
|
+
|
|
542
|
+
# License
|
|
543
|
+
if license_info:
|
|
544
|
+
lic_lines = [f" License: {license_info.get('name', '-')}"]
|
|
545
|
+
if license_info.get("commercial_ok"):
|
|
546
|
+
# A permission string since MODEL-77 ("allowed" / "restricted" /
|
|
547
|
+
# "prohibited" / "unspecified" / "withheld"), never a boolean.
|
|
548
|
+
lic_lines.append(f" Commercial: {license_info['commercial_ok']}")
|
|
549
|
+
console.print(Panel("\n".join(lic_lines), title="License", border_style="dim"))
|
|
550
|
+
|
|
551
|
+
|
|
552
|
+
# ───────────────────────────────────────────────────────────────
|
|
553
|
+
# 2. modelspec search
|
|
554
|
+
# ───────────────────────────────────────────────────────────────
|
|
555
|
+
|
|
556
|
+
@app.command()
|
|
557
|
+
def search(
|
|
558
|
+
type: Optional[str] = typer.Option(None, "--type", "-t", help="Model type, e.g. llm-chat"),
|
|
559
|
+
hardware: Optional[str] = typer.Option(None, "--hardware", "-hw", help="Hardware ID filter (FITS_ON)"),
|
|
560
|
+
license: Optional[str] = typer.Option(None, "--license", "-l", help="License type filter"),
|
|
561
|
+
origin: Optional[str] = typer.Option(None, "--origin", help="Origin country code, e.g. US, CN"),
|
|
562
|
+
open_weights: Optional[bool] = typer.Option(None, "--open-weights/--closed-weights", help="Filter by open weights"),
|
|
563
|
+
min_params: Optional[int] = typer.Option(None, "--min-params", help="Minimum total parameters"),
|
|
564
|
+
max_params: Optional[int] = typer.Option(None, "--max-params", help="Maximum total parameters"),
|
|
565
|
+
query: Optional[str] = typer.Option(None, "--query", "-q", help="Text search on display_name"),
|
|
566
|
+
format: Optional[str] = typer.Option(None, "--format", "-f", help="Output format: json"),
|
|
567
|
+
) -> None:
|
|
568
|
+
"""Search models with filters."""
|
|
569
|
+
graph = _get_graph()
|
|
570
|
+
|
|
571
|
+
# Build dynamic WHERE clauses
|
|
572
|
+
where_clauses: list[str] = []
|
|
573
|
+
params: dict[str, Any] = {}
|
|
574
|
+
|
|
575
|
+
match_prefix = "MATCH (m:Model)"
|
|
576
|
+
extra_matches: list[str] = []
|
|
577
|
+
|
|
578
|
+
if type:
|
|
579
|
+
where_clauses.append("m.model_type = $model_type")
|
|
580
|
+
params["model_type"] = type
|
|
581
|
+
|
|
582
|
+
if origin:
|
|
583
|
+
where_clauses.append("m.origin_country = $origin")
|
|
584
|
+
params["origin"] = origin.upper()
|
|
585
|
+
|
|
586
|
+
if open_weights is not None:
|
|
587
|
+
where_clauses.append("m.open_weights = $open_weights")
|
|
588
|
+
params["open_weights"] = open_weights
|
|
589
|
+
|
|
590
|
+
if min_params is not None:
|
|
591
|
+
where_clauses.append("m.total_parameters >= $min_params")
|
|
592
|
+
params["min_params"] = min_params
|
|
593
|
+
|
|
594
|
+
if max_params is not None:
|
|
595
|
+
where_clauses.append("m.total_parameters <= $max_params")
|
|
596
|
+
params["max_params"] = max_params
|
|
597
|
+
|
|
598
|
+
if query:
|
|
599
|
+
where_clauses.append("toLower(m.display_name) CONTAINS toLower($query_text)")
|
|
600
|
+
params["query_text"] = query
|
|
601
|
+
|
|
602
|
+
if hardware:
|
|
603
|
+
extra_matches.append(f"MATCH (m)-[:FITS_ON]->(h:Hardware {{id: $hw_id}})")
|
|
604
|
+
params["hw_id"] = hardware
|
|
605
|
+
|
|
606
|
+
if license:
|
|
607
|
+
extra_matches.append(f"MATCH (m)-[:LICENSED_AS]->(l:License {{id: $lic_id}})")
|
|
608
|
+
params["lic_id"] = license
|
|
609
|
+
|
|
610
|
+
cypher = match_prefix
|
|
611
|
+
if extra_matches:
|
|
612
|
+
cypher += " " + " ".join(extra_matches)
|
|
613
|
+
if where_clauses:
|
|
614
|
+
cypher += " WHERE " + " AND ".join(where_clauses)
|
|
615
|
+
cypher += (
|
|
616
|
+
" RETURN m.id, m.display_name, m.model_type, m.total_parameters, "
|
|
617
|
+
"m.arena_elo_overall, m.cost_input, m.cost_output, m.status, "
|
|
618
|
+
"m.open_weights, m.origin_country "
|
|
619
|
+
"ORDER BY m.arena_elo_overall DESC"
|
|
620
|
+
)
|
|
621
|
+
|
|
622
|
+
result = graph.query(cypher, params)
|
|
623
|
+
|
|
624
|
+
if format == "json":
|
|
625
|
+
rows = []
|
|
626
|
+
for r in result.result_set:
|
|
627
|
+
rows.append({
|
|
628
|
+
"id": r[0], "name": r[1], "type": r[2],
|
|
629
|
+
"parameters": r[3], "arena_elo": r[4],
|
|
630
|
+
"cost_input": r[5], "cost_output": r[6],
|
|
631
|
+
"status": r[7], "open_weights": r[8],
|
|
632
|
+
"origin": r[9],
|
|
633
|
+
})
|
|
634
|
+
console.print_json(json.dumps(rows, default=str))
|
|
635
|
+
return
|
|
636
|
+
|
|
637
|
+
if not result.result_set:
|
|
638
|
+
console.print("[yellow]No models found matching your filters.[/]")
|
|
639
|
+
return
|
|
640
|
+
|
|
641
|
+
table = Table(
|
|
642
|
+
title=f"Models ({len(result.result_set)} found)",
|
|
643
|
+
show_header=True,
|
|
644
|
+
header_style="bold cyan",
|
|
645
|
+
)
|
|
646
|
+
table.add_column("Model", style="bold white", min_width=20)
|
|
647
|
+
table.add_column("Type", style="dim")
|
|
648
|
+
table.add_column("Params", justify="right")
|
|
649
|
+
table.add_column("ELO", justify="right")
|
|
650
|
+
table.add_column("$/M in", justify="right")
|
|
651
|
+
table.add_column("$/M out", justify="right")
|
|
652
|
+
table.add_column("Status", justify="center")
|
|
653
|
+
table.add_column("Open", justify="center")
|
|
654
|
+
table.add_column("Origin", justify="center")
|
|
655
|
+
|
|
656
|
+
for r in result.result_set:
|
|
657
|
+
status = r[7] or ""
|
|
658
|
+
table.add_row(
|
|
659
|
+
r[1] or r[0],
|
|
660
|
+
r[2] or "-",
|
|
661
|
+
_fmt_params(r[3]),
|
|
662
|
+
_fmt_elo(r[4]),
|
|
663
|
+
_fmt_cost(r[5]),
|
|
664
|
+
_fmt_cost(r[6]),
|
|
665
|
+
f"[{_status_color(status)}]{status}[/]",
|
|
666
|
+
_bool_icon(r[8]),
|
|
667
|
+
r[9] or "-",
|
|
668
|
+
)
|
|
669
|
+
|
|
670
|
+
console.print(table)
|
|
671
|
+
|
|
672
|
+
|
|
673
|
+
# ───────────────────────────────────────────────────────────────
|
|
674
|
+
# 3. modelspec compare <model_ids>
|
|
675
|
+
# ───────────────────────────────────────────────────────────────
|
|
676
|
+
|
|
677
|
+
@app.command()
|
|
678
|
+
def compare(
|
|
679
|
+
model_ids: list[str] = typer.Argument(..., help="2-4 model IDs to compare"),
|
|
680
|
+
format: Optional[str] = typer.Option(None, "--format", "-f", help="Output format: json"),
|
|
681
|
+
) -> None:
|
|
682
|
+
"""Side-by-side comparison of 2-4 models."""
|
|
683
|
+
if len(model_ids) < 2 or len(model_ids) > 4:
|
|
684
|
+
console.print("[bold red]Error:[/] Provide 2 to 4 model IDs.")
|
|
685
|
+
raise typer.Exit(1)
|
|
686
|
+
|
|
687
|
+
graph = _get_graph()
|
|
688
|
+
|
|
689
|
+
models: list[dict[str, Any]] = []
|
|
690
|
+
all_benchmarks: dict[str, dict[str, float | None]] = {} # bench_id -> {model_id: value}
|
|
691
|
+
all_caps: dict[str, dict[str, str | None]] = {} # cap_id -> {model_id: tier}
|
|
692
|
+
all_hw: dict[str, dict[str, dict]] = {} # hw_id -> {model_id: props}
|
|
693
|
+
|
|
694
|
+
for mid in model_ids:
|
|
695
|
+
# Model node
|
|
696
|
+
result = graph.query("MATCH (m:Model {id: $mid}) RETURN m", {"mid": mid})
|
|
697
|
+
if not result.result_set:
|
|
698
|
+
console.print(f"[bold red]Model not found:[/] {mid}")
|
|
699
|
+
raise typer.Exit(1)
|
|
700
|
+
m = _published_model_props(result.result_set[0][0])
|
|
701
|
+
models.append(m)
|
|
702
|
+
|
|
703
|
+
# Benchmarks
|
|
704
|
+
bench_r = graph.query(
|
|
705
|
+
"MATCH (m:Model {id: $mid})-[r:SCORED_ON]->(b:Benchmark) "
|
|
706
|
+
"RETURN b.id, b.name, r.value",
|
|
707
|
+
{"mid": mid},
|
|
708
|
+
)
|
|
709
|
+
for row in bench_r.result_set:
|
|
710
|
+
bid = row[0]
|
|
711
|
+
if bid not in all_benchmarks:
|
|
712
|
+
all_benchmarks[bid] = {"_name": row[1]}
|
|
713
|
+
all_benchmarks[bid][mid] = row[2]
|
|
714
|
+
|
|
715
|
+
# Capabilities (section-level only: coding, reasoning, tool_use, etc.)
|
|
716
|
+
cap_r = graph.query(
|
|
717
|
+
"MATCH (m:Model {id: $mid})-[r:HAS_CAPABILITY]->(c:Capability) "
|
|
718
|
+
"WHERE NOT c.id CONTAINS ':' "
|
|
719
|
+
"RETURN c.id, c.name, r.tier",
|
|
720
|
+
{"mid": mid},
|
|
721
|
+
)
|
|
722
|
+
for row in cap_r.result_set:
|
|
723
|
+
cid = row[0]
|
|
724
|
+
if cid not in all_caps:
|
|
725
|
+
all_caps[cid] = {"_name": row[1]}
|
|
726
|
+
all_caps[cid][mid] = row[2]
|
|
727
|
+
|
|
728
|
+
# Hardware
|
|
729
|
+
hw_r = graph.query(
|
|
730
|
+
"MATCH (m:Model {id: $mid})-[r:FITS_ON]->(h:Hardware) RETURN h.id, r",
|
|
731
|
+
{"mid": mid},
|
|
732
|
+
)
|
|
733
|
+
for row in hw_r.result_set:
|
|
734
|
+
hid = row[0]
|
|
735
|
+
if hid not in all_hw:
|
|
736
|
+
all_hw[hid] = {}
|
|
737
|
+
all_hw[hid][mid] = _edge_props(row[1])
|
|
738
|
+
|
|
739
|
+
if format == "json":
|
|
740
|
+
data = {
|
|
741
|
+
"models": [{
|
|
742
|
+
"id": m.get("id"),
|
|
743
|
+
"display_name": m.get("display_name"),
|
|
744
|
+
"model_type": m.get("model_type"),
|
|
745
|
+
"total_parameters": m.get("total_parameters"),
|
|
746
|
+
"arena_elo_overall": m.get("arena_elo_overall"),
|
|
747
|
+
"cost_input": m.get("cost_input"),
|
|
748
|
+
"cost_output": m.get("cost_output"),
|
|
749
|
+
} for m in models],
|
|
750
|
+
"benchmarks": all_benchmarks,
|
|
751
|
+
"capabilities": all_caps,
|
|
752
|
+
"hardware": all_hw,
|
|
753
|
+
}
|
|
754
|
+
console.print_json(json.dumps(data, default=str))
|
|
755
|
+
return
|
|
756
|
+
|
|
757
|
+
names = [m.get("display_name", m.get("id", "?")) for m in models]
|
|
758
|
+
|
|
759
|
+
# ── Overview table ──────────────────────────────────────
|
|
760
|
+
overview = Table(title="Model Comparison", show_header=True, header_style="bold cyan")
|
|
761
|
+
overview.add_column("Attribute", style="bold", min_width=18)
|
|
762
|
+
for name in names:
|
|
763
|
+
overview.add_column(name, justify="center", min_width=16)
|
|
764
|
+
|
|
765
|
+
attrs = [
|
|
766
|
+
("Type", "model_type"),
|
|
767
|
+
("Parameters", None),
|
|
768
|
+
("Context", "context_window"),
|
|
769
|
+
("Arena ELO", "arena_elo_overall"),
|
|
770
|
+
("Cost (in)", "cost_input"),
|
|
771
|
+
("Cost (out)", "cost_output"),
|
|
772
|
+
("Open Weights", "open_weights"),
|
|
773
|
+
("Origin", "origin_country"),
|
|
774
|
+
("Status", "status"),
|
|
775
|
+
]
|
|
776
|
+
|
|
777
|
+
for label, key in attrs:
|
|
778
|
+
cells: list[str] = []
|
|
779
|
+
for m in models:
|
|
780
|
+
if label == "Parameters":
|
|
781
|
+
active = m.get("active_parameters")
|
|
782
|
+
total = m.get("total_parameters")
|
|
783
|
+
s = _fmt_params(total)
|
|
784
|
+
if active:
|
|
785
|
+
s += f"\n({_fmt_params(active)} active)"
|
|
786
|
+
cells.append(s)
|
|
787
|
+
elif key == "arena_elo_overall":
|
|
788
|
+
cells.append(_fmt_elo(m.get(key)))
|
|
789
|
+
elif key in ("cost_input", "cost_output"):
|
|
790
|
+
cells.append(_fmt_cost(m.get(key)))
|
|
791
|
+
elif key == "context_window":
|
|
792
|
+
cells.append(_fmt_int(m.get(key)))
|
|
793
|
+
elif key == "open_weights":
|
|
794
|
+
cells.append(_bool_icon(m.get(key)))
|
|
795
|
+
elif key == "status":
|
|
796
|
+
v = m.get(key, "")
|
|
797
|
+
cells.append(f"[{_status_color(v)}]{v}[/]")
|
|
798
|
+
else:
|
|
799
|
+
cells.append(str(m.get(key, "-") or "-"))
|
|
800
|
+
overview.add_row(label, *cells)
|
|
801
|
+
|
|
802
|
+
console.print(overview)
|
|
803
|
+
|
|
804
|
+
# ── Capabilities comparison ─────────────────────────────
|
|
805
|
+
if all_caps:
|
|
806
|
+
cap_table = Table(title="Capabilities", show_header=True, header_style="bold magenta")
|
|
807
|
+
cap_table.add_column("Capability", style="bold", min_width=18)
|
|
808
|
+
for name in names:
|
|
809
|
+
cap_table.add_column(name, justify="center", min_width=16)
|
|
810
|
+
for cid in sorted(all_caps.keys()):
|
|
811
|
+
row_data = all_caps[cid]
|
|
812
|
+
cap_table.add_row(
|
|
813
|
+
row_data.get("_name", cid),
|
|
814
|
+
*[_tier_style(row_data.get(mid)) for mid in model_ids],
|
|
815
|
+
)
|
|
816
|
+
console.print(cap_table)
|
|
817
|
+
|
|
818
|
+
# ── Benchmark comparison (highlight winner) ─────────────
|
|
819
|
+
if all_benchmarks:
|
|
820
|
+
bench_table = Table(title="Benchmarks", show_header=True, header_style="bold yellow")
|
|
821
|
+
bench_table.add_column("Benchmark", style="bold", min_width=22)
|
|
822
|
+
for name in names:
|
|
823
|
+
bench_table.add_column(name, justify="right", min_width=16)
|
|
824
|
+
|
|
825
|
+
for bid in sorted(all_benchmarks.keys()):
|
|
826
|
+
row_data = all_benchmarks[bid]
|
|
827
|
+
values = {mid: row_data.get(mid) for mid in model_ids}
|
|
828
|
+
numeric_vals = [v for v in values.values() if v is not None]
|
|
829
|
+
# Lower is better for some benchmarks (like FID), but most are higher-is-better
|
|
830
|
+
lower_is_better = bid in ("fid", "wer_librispeech", "api_latency_p50_ms", "api_latency_p99_ms", "api_ttft_ms")
|
|
831
|
+
if numeric_vals:
|
|
832
|
+
best = min(numeric_vals) if lower_is_better else max(numeric_vals)
|
|
833
|
+
else:
|
|
834
|
+
best = None
|
|
835
|
+
|
|
836
|
+
cells: list[str] = []
|
|
837
|
+
for mid in model_ids:
|
|
838
|
+
v = values.get(mid)
|
|
839
|
+
if v is None:
|
|
840
|
+
cells.append("-")
|
|
841
|
+
elif v == best and len(numeric_vals) > 1:
|
|
842
|
+
cells.append(f"[bold green]{v:.1f}[/]")
|
|
843
|
+
else:
|
|
844
|
+
cells.append(f"{v:.1f}")
|
|
845
|
+
|
|
846
|
+
bench_table.add_row(row_data.get("_name", bid), *cells)
|
|
847
|
+
|
|
848
|
+
console.print(bench_table)
|
|
849
|
+
|
|
850
|
+
# ── Hardware fit comparison ──────────────────────────────
|
|
851
|
+
if all_hw:
|
|
852
|
+
hw_table = Table(title="Hardware Fit", show_header=True, header_style="bold green")
|
|
853
|
+
hw_table.add_column("Hardware", style="bold", min_width=22)
|
|
854
|
+
for name in names:
|
|
855
|
+
hw_table.add_column(name, justify="center", min_width=16)
|
|
856
|
+
|
|
857
|
+
for hid in sorted(all_hw.keys()):
|
|
858
|
+
hw_data = all_hw[hid]
|
|
859
|
+
cells: list[str] = []
|
|
860
|
+
for mid in model_ids:
|
|
861
|
+
props = hw_data.get(mid)
|
|
862
|
+
if props:
|
|
863
|
+
tps = props.get("tokens_per_sec")
|
|
864
|
+
quant = props.get("quantization", "?")
|
|
865
|
+
s = f"[green]Y[/] {quant}"
|
|
866
|
+
if tps:
|
|
867
|
+
s += f"\n{tps:.0f} tok/s"
|
|
868
|
+
cells.append(s)
|
|
869
|
+
else:
|
|
870
|
+
cells.append("[dim]-[/]")
|
|
871
|
+
hw_table.add_row(hid.replace("_", " ").title(), *cells)
|
|
872
|
+
|
|
873
|
+
console.print(hw_table)
|
|
874
|
+
|
|
875
|
+
|
|
876
|
+
# ───────────────────────────────────────────────────────────────
|
|
877
|
+
# 4. modelspec rank --use-case <USE_CASE>
|
|
878
|
+
# ───────────────────────────────────────────────────────────────
|
|
879
|
+
|
|
880
|
+
# Benchmark weights per use case
|
|
881
|
+
_USE_CASE_WEIGHTS: dict[str, dict[str, float]] = {
|
|
882
|
+
"coding": {
|
|
883
|
+
"humaneval": 2.0, "swe_bench_verified": 3.0, "live_code_bench": 2.0,
|
|
884
|
+
"aider_polyglot": 2.0, "arena_elo_coding": 2.0, "arena_elo_overall": 1.0,
|
|
885
|
+
},
|
|
886
|
+
"reasoning": {
|
|
887
|
+
"gpqa_diamond": 2.0, "math_500": 2.0, "aime_2025": 2.0,
|
|
888
|
+
"bbh": 1.5, "arena_elo_overall": 1.0, "arena_elo_math": 1.5,
|
|
889
|
+
},
|
|
890
|
+
"chat": {
|
|
891
|
+
"arena_elo_overall": 3.0, "arena_elo_style_control": 1.5,
|
|
892
|
+
"mt_bench": 1.5, "alpaca_eval": 1.0, "ifeval": 1.5,
|
|
893
|
+
},
|
|
894
|
+
"embedding": {
|
|
895
|
+
"mteb_overall": 3.0, "mteb_retrieval": 2.0, "mteb_classification": 1.5,
|
|
896
|
+
"beir": 1.5,
|
|
897
|
+
},
|
|
898
|
+
"agentic": {
|
|
899
|
+
"swe_bench_verified": 2.0, "swe_bench_agent": 3.0, "tau_bench": 2.0,
|
|
900
|
+
"web_arena": 2.0, "arena_elo_overall": 1.0,
|
|
901
|
+
},
|
|
902
|
+
"general": {
|
|
903
|
+
"arena_elo_overall": 2.0, "gpqa_diamond": 1.0, "humaneval": 1.0,
|
|
904
|
+
"math_500": 1.0, "swe_bench_verified": 1.0, "mmlu_pro": 1.0,
|
|
905
|
+
},
|
|
906
|
+
}
|
|
907
|
+
|
|
908
|
+
|
|
909
|
+
@app.command()
|
|
910
|
+
def rank(
|
|
911
|
+
use_case: str = typer.Option(..., "--use-case", "-u", help="Use case: coding, reasoning, chat, embedding, agentic, general"),
|
|
912
|
+
hardware: Optional[str] = typer.Option(None, "--hardware", "-hw", help="Restrict to models fitting this hardware"),
|
|
913
|
+
license: Optional[str] = typer.Option(None, "--license", "-l", help="Restrict to this license type"),
|
|
914
|
+
top: int = typer.Option(20, "--top", "-n", help="Show top N models"),
|
|
915
|
+
format: Optional[str] = typer.Option(None, "--format", "-f", help="Output format: json"),
|
|
916
|
+
) -> None:
|
|
917
|
+
"""Rank models by use case, scored on relevant benchmarks."""
|
|
918
|
+
graph = _get_graph()
|
|
919
|
+
|
|
920
|
+
weights = _USE_CASE_WEIGHTS.get(use_case.lower())
|
|
921
|
+
if not weights:
|
|
922
|
+
valid = ", ".join(_USE_CASE_WEIGHTS.keys())
|
|
923
|
+
console.print(f"[bold red]Unknown use case:[/] {use_case}")
|
|
924
|
+
console.print(f" Valid options: {valid}")
|
|
925
|
+
raise typer.Exit(1)
|
|
926
|
+
|
|
927
|
+
# Build match
|
|
928
|
+
cypher = "MATCH (m:Model)"
|
|
929
|
+
extra: list[str] = []
|
|
930
|
+
params: dict[str, Any] = {}
|
|
931
|
+
|
|
932
|
+
if hardware:
|
|
933
|
+
extra.append("MATCH (m)-[:FITS_ON]->(h:Hardware {id: $hw_id})")
|
|
934
|
+
params["hw_id"] = hardware
|
|
935
|
+
if license:
|
|
936
|
+
extra.append("MATCH (m)-[:LICENSED_AS]->(l:License {id: $lic_id})")
|
|
937
|
+
params["lic_id"] = license
|
|
938
|
+
|
|
939
|
+
if extra:
|
|
940
|
+
cypher += " " + " ".join(extra)
|
|
941
|
+
|
|
942
|
+
cypher += (
|
|
943
|
+
" OPTIONAL MATCH (m)-[r:SCORED_ON]->(b:Benchmark) "
|
|
944
|
+
"RETURN m.id, m.display_name, m.model_type, m.arena_elo_overall, "
|
|
945
|
+
"m.cost_input, m.total_parameters, collect([b.id, r.value])"
|
|
946
|
+
)
|
|
947
|
+
|
|
948
|
+
result = graph.query(cypher, params)
|
|
949
|
+
|
|
950
|
+
# Score each model
|
|
951
|
+
scored: list[dict[str, Any]] = []
|
|
952
|
+
for row in result.result_set:
|
|
953
|
+
mid, name, mtype, elo, cost_in, params_val, bench_pairs = row
|
|
954
|
+
bench_map: dict[str, float] = {}
|
|
955
|
+
for pair in bench_pairs:
|
|
956
|
+
if pair[0] and pair[1] is not None:
|
|
957
|
+
bench_map[pair[0]] = pair[1]
|
|
958
|
+
|
|
959
|
+
total_score = 0.0
|
|
960
|
+
total_weight = 0.0
|
|
961
|
+
score_parts: dict[str, float] = {}
|
|
962
|
+
for bench_id, weight in weights.items():
|
|
963
|
+
val = bench_map.get(bench_id)
|
|
964
|
+
if val is not None:
|
|
965
|
+
# Normalize: most benchmarks 0-100, ELO ~1000-1400
|
|
966
|
+
if "elo" in bench_id:
|
|
967
|
+
normalized = (val - 1000) / 400 * 100 # map 1000-1400 to 0-100
|
|
968
|
+
else:
|
|
969
|
+
normalized = val
|
|
970
|
+
total_score += normalized * weight
|
|
971
|
+
total_weight += weight
|
|
972
|
+
score_parts[bench_id] = val
|
|
973
|
+
|
|
974
|
+
final_score = (total_score / total_weight) if total_weight > 0 else 0.0
|
|
975
|
+
|
|
976
|
+
scored.append({
|
|
977
|
+
"id": mid,
|
|
978
|
+
"name": name,
|
|
979
|
+
"type": mtype,
|
|
980
|
+
"score": round(final_score, 1),
|
|
981
|
+
"elo": elo,
|
|
982
|
+
"cost_input": cost_in,
|
|
983
|
+
"params": params_val,
|
|
984
|
+
"benchmarks_matched": len(score_parts),
|
|
985
|
+
"score_parts": score_parts,
|
|
986
|
+
})
|
|
987
|
+
|
|
988
|
+
scored.sort(key=lambda x: x["score"], reverse=True)
|
|
989
|
+
scored = scored[:top]
|
|
990
|
+
|
|
991
|
+
if format == "json":
|
|
992
|
+
console.print_json(json.dumps(scored, default=str))
|
|
993
|
+
return
|
|
994
|
+
|
|
995
|
+
if not scored:
|
|
996
|
+
console.print("[yellow]No models found.[/]")
|
|
997
|
+
return
|
|
998
|
+
|
|
999
|
+
table = Table(
|
|
1000
|
+
title=f"Rankings: {use_case} (top {top})",
|
|
1001
|
+
show_header=True,
|
|
1002
|
+
header_style="bold cyan",
|
|
1003
|
+
)
|
|
1004
|
+
table.add_column("#", justify="right", style="dim", width=4)
|
|
1005
|
+
table.add_column("Model", style="bold white", min_width=25)
|
|
1006
|
+
table.add_column("Type", style="dim")
|
|
1007
|
+
table.add_column("Score", justify="right", style="bold yellow")
|
|
1008
|
+
table.add_column("ELO", justify="right")
|
|
1009
|
+
table.add_column("$/M in", justify="right")
|
|
1010
|
+
table.add_column("Params", justify="right")
|
|
1011
|
+
table.add_column("Benchmarks", justify="center")
|
|
1012
|
+
|
|
1013
|
+
for i, entry in enumerate(scored, 1):
|
|
1014
|
+
rank_str = str(i)
|
|
1015
|
+
if i == 1:
|
|
1016
|
+
rank_str = f"[bold green]{i}[/]"
|
|
1017
|
+
elif i <= 3:
|
|
1018
|
+
rank_str = f"[yellow]{i}[/]"
|
|
1019
|
+
|
|
1020
|
+
table.add_row(
|
|
1021
|
+
rank_str,
|
|
1022
|
+
entry["name"] or entry["id"],
|
|
1023
|
+
entry["type"] or "-",
|
|
1024
|
+
f"{entry['score']:.1f}",
|
|
1025
|
+
_fmt_elo(entry["elo"]),
|
|
1026
|
+
_fmt_cost(entry["cost_input"]),
|
|
1027
|
+
_fmt_params(entry["params"]),
|
|
1028
|
+
f"{entry['benchmarks_matched']}/{len(weights)}",
|
|
1029
|
+
)
|
|
1030
|
+
|
|
1031
|
+
console.print(table)
|
|
1032
|
+
|
|
1033
|
+
# Show which benchmarks contributed
|
|
1034
|
+
bench_names = ", ".join(weights.keys())
|
|
1035
|
+
console.print(f"\n[dim]Scoring benchmarks: {bench_names}[/]")
|
|
1036
|
+
|
|
1037
|
+
|
|
1038
|
+
# ───────────────────────────────────────────────────────────────
|
|
1039
|
+
# 5. modelspec stats
|
|
1040
|
+
# ───────────────────────────────────────────────────────────────
|
|
1041
|
+
|
|
1042
|
+
@app.command()
|
|
1043
|
+
def stats(
|
|
1044
|
+
format: Optional[str] = typer.Option(None, "--format", "-f", help="Output format: json"),
|
|
1045
|
+
) -> None:
|
|
1046
|
+
"""Show database overview: node counts, edge counts, type breakdown."""
|
|
1047
|
+
graph = _get_graph()
|
|
1048
|
+
|
|
1049
|
+
# Node counts by label
|
|
1050
|
+
labels_result = graph.query("CALL db.labels()")
|
|
1051
|
+
node_counts: dict[str, int] = {}
|
|
1052
|
+
for row in labels_result.result_set:
|
|
1053
|
+
label = row[0]
|
|
1054
|
+
count_r = graph.query(f"MATCH (n:{label}) RETURN count(n)")
|
|
1055
|
+
node_counts[label] = count_r.result_set[0][0]
|
|
1056
|
+
|
|
1057
|
+
# Edge counts by type
|
|
1058
|
+
rel_result = graph.query("CALL db.relationshipTypes()")
|
|
1059
|
+
edge_counts: dict[str, int] = {}
|
|
1060
|
+
for row in rel_result.result_set:
|
|
1061
|
+
rtype = row[0]
|
|
1062
|
+
count_r = graph.query(f"MATCH ()-[r:{rtype}]->() RETURN count(r)")
|
|
1063
|
+
edge_counts[rtype] = count_r.result_set[0][0]
|
|
1064
|
+
|
|
1065
|
+
# Type breakdown
|
|
1066
|
+
type_result = graph.query(
|
|
1067
|
+
"MATCH (m:Model) "
|
|
1068
|
+
"RETURN m.model_type, count(m) "
|
|
1069
|
+
"ORDER BY count(m) DESC"
|
|
1070
|
+
)
|
|
1071
|
+
|
|
1072
|
+
if format == "json":
|
|
1073
|
+
data = {
|
|
1074
|
+
"node_counts": node_counts,
|
|
1075
|
+
"edge_counts": edge_counts,
|
|
1076
|
+
"type_breakdown": [
|
|
1077
|
+
{"type": r[0], "count": r[1]}
|
|
1078
|
+
for r in type_result.result_set
|
|
1079
|
+
],
|
|
1080
|
+
}
|
|
1081
|
+
console.print_json(json.dumps(data, default=str))
|
|
1082
|
+
return
|
|
1083
|
+
|
|
1084
|
+
total_nodes = sum(node_counts.values())
|
|
1085
|
+
total_edges = sum(edge_counts.values())
|
|
1086
|
+
|
|
1087
|
+
# Header
|
|
1088
|
+
console.print(Panel(
|
|
1089
|
+
f"[bold]{total_nodes}[/] nodes | [bold]{total_edges}[/] edges | "
|
|
1090
|
+
f"[bold]{node_counts.get('Model', 0)}[/] models",
|
|
1091
|
+
title="ModelSpec Database",
|
|
1092
|
+
border_style="blue",
|
|
1093
|
+
))
|
|
1094
|
+
|
|
1095
|
+
# Node counts
|
|
1096
|
+
node_table = Table(title="Nodes by Label", show_header=True, header_style="bold cyan")
|
|
1097
|
+
node_table.add_column("Label", style="bold", min_width=15)
|
|
1098
|
+
node_table.add_column("Count", justify="right")
|
|
1099
|
+
for label, count in sorted(node_counts.items(), key=lambda x: -x[1]):
|
|
1100
|
+
node_table.add_row(label, str(count))
|
|
1101
|
+
console.print(node_table)
|
|
1102
|
+
|
|
1103
|
+
# Edge counts
|
|
1104
|
+
edge_table = Table(title="Edges by Type", show_header=True, header_style="bold magenta")
|
|
1105
|
+
edge_table.add_column("Relationship", style="bold", min_width=20)
|
|
1106
|
+
edge_table.add_column("Count", justify="right")
|
|
1107
|
+
for rtype, count in sorted(edge_counts.items(), key=lambda x: -x[1]):
|
|
1108
|
+
edge_table.add_row(rtype, str(count))
|
|
1109
|
+
console.print(edge_table)
|
|
1110
|
+
|
|
1111
|
+
# Model type breakdown
|
|
1112
|
+
if type_result.result_set:
|
|
1113
|
+
type_table = Table(title="Models by Type", show_header=True, header_style="bold green")
|
|
1114
|
+
type_table.add_column("Type", style="bold", min_width=20)
|
|
1115
|
+
type_table.add_column("Count", justify="right")
|
|
1116
|
+
for row in type_result.result_set:
|
|
1117
|
+
type_table.add_row(row[0] or "(untyped)", str(row[1]))
|
|
1118
|
+
console.print(type_table)
|
|
1119
|
+
|
|
1120
|
+
|
|
1121
|
+
# ───────────────────────────────────────────────────────────────
|
|
1122
|
+
# 6. modelspec hardware <hardware_id>
|
|
1123
|
+
# ───────────────────────────────────────────────────────────────
|
|
1124
|
+
|
|
1125
|
+
@app.command()
|
|
1126
|
+
def hardware(
|
|
1127
|
+
hardware_id: str = typer.Argument(..., help="Hardware ID, e.g. macbook_air_m4_24gb"),
|
|
1128
|
+
format: Optional[str] = typer.Option(None, "--format", "-f", help="Output format: json"),
|
|
1129
|
+
) -> None:
|
|
1130
|
+
"""Show models that fit on a specific hardware device."""
|
|
1131
|
+
graph = _get_graph()
|
|
1132
|
+
|
|
1133
|
+
result = graph.query(
|
|
1134
|
+
"MATCH (m:Model)-[r:FITS_ON]->(h:Hardware {id: $hw_id}) "
|
|
1135
|
+
"RETURN m.id, m.display_name, m.model_type, m.arena_elo_overall, "
|
|
1136
|
+
"m.total_parameters, m.cost_input, r "
|
|
1137
|
+
"ORDER BY m.arena_elo_overall DESC",
|
|
1138
|
+
{"hw_id": hardware_id},
|
|
1139
|
+
)
|
|
1140
|
+
|
|
1141
|
+
if not result.result_set:
|
|
1142
|
+
# Check if the hardware exists at all
|
|
1143
|
+
hw_check = graph.query(
|
|
1144
|
+
"MATCH (h:Hardware {id: $hw_id}) RETURN h.id", {"hw_id": hardware_id}
|
|
1145
|
+
)
|
|
1146
|
+
if not hw_check.result_set:
|
|
1147
|
+
# List available hardware
|
|
1148
|
+
all_hw = graph.query("MATCH (h:Hardware) RETURN h.id ORDER BY h.id")
|
|
1149
|
+
if all_hw.result_set:
|
|
1150
|
+
hw_list = ", ".join(r[0] for r in all_hw.result_set)
|
|
1151
|
+
console.print(f"[bold red]Hardware not found:[/] {hardware_id}")
|
|
1152
|
+
console.print(f" Available: {hw_list}")
|
|
1153
|
+
else:
|
|
1154
|
+
console.print(f"[bold red]No hardware profiles in the database.[/]")
|
|
1155
|
+
else:
|
|
1156
|
+
console.print(f"[yellow]No models fit on {hardware_id}.[/]")
|
|
1157
|
+
return
|
|
1158
|
+
|
|
1159
|
+
if format == "json":
|
|
1160
|
+
rows = []
|
|
1161
|
+
for r in result.result_set:
|
|
1162
|
+
ep = _edge_props(r[6])
|
|
1163
|
+
rows.append({
|
|
1164
|
+
"id": r[0], "name": r[1], "type": r[2],
|
|
1165
|
+
"arena_elo": r[3], "parameters": r[4], "cost_input": r[5],
|
|
1166
|
+
**ep,
|
|
1167
|
+
})
|
|
1168
|
+
console.print_json(json.dumps(rows, default=str))
|
|
1169
|
+
return
|
|
1170
|
+
|
|
1171
|
+
table = Table(
|
|
1172
|
+
title=f"Models for {hardware_id.replace('_', ' ').title()}",
|
|
1173
|
+
show_header=True,
|
|
1174
|
+
header_style="bold green",
|
|
1175
|
+
)
|
|
1176
|
+
table.add_column("Model", style="bold white", min_width=25)
|
|
1177
|
+
table.add_column("Type", style="dim")
|
|
1178
|
+
table.add_column("ELO", justify="right")
|
|
1179
|
+
table.add_column("Params", justify="right")
|
|
1180
|
+
table.add_column("Quant", justify="center")
|
|
1181
|
+
table.add_column("VRAM/RAM", justify="right")
|
|
1182
|
+
table.add_column("tok/s", justify="right")
|
|
1183
|
+
table.add_column("TTFT", justify="right")
|
|
1184
|
+
table.add_column("Engine", justify="center")
|
|
1185
|
+
|
|
1186
|
+
for r in result.result_set:
|
|
1187
|
+
ep = _edge_props(r[6])
|
|
1188
|
+
table.add_row(
|
|
1189
|
+
r[1] or r[0],
|
|
1190
|
+
r[2] or "-",
|
|
1191
|
+
_fmt_elo(r[3]),
|
|
1192
|
+
_fmt_params(r[4]),
|
|
1193
|
+
ep.get("quantization", "-"),
|
|
1194
|
+
_fmt_float(ep.get("vram_usage_gb"), " GB"),
|
|
1195
|
+
_fmt_float(ep.get("tokens_per_sec"), ""),
|
|
1196
|
+
_fmt_float(ep.get("ttft_ms"), " ms"),
|
|
1197
|
+
ep.get("inference_engine", "-"),
|
|
1198
|
+
)
|
|
1199
|
+
|
|
1200
|
+
console.print(table)
|
|
1201
|
+
|
|
1202
|
+
|
|
1203
|
+
# ═══════════════════════════════════════════════════════════════
|
|
1204
|
+
# Community contribution commands (offline, YAML-based)
|
|
1205
|
+
# ═══════════════════════════════════════════════════════════════
|
|
1206
|
+
|
|
1207
|
+
|
|
1208
|
+
# ───────────────────────────────────────────────────────────────
|
|
1209
|
+
# 7. modelspec gaps — find data gaps
|
|
1210
|
+
# ───────────────────────────────────────────────────────────────
|
|
1211
|
+
|
|
1212
|
+
@app.command()
|
|
1213
|
+
def gaps(
|
|
1214
|
+
type: Optional[str] = typer.Option(None, "--type", "-t", help="Filter by model_type (e.g. llm-chat, vlm)"),
|
|
1215
|
+
provider: Optional[str] = typer.Option(None, "--provider", "-p", help="Filter by provider slug"),
|
|
1216
|
+
top: int = typer.Option(20, "--top", "-n", help="Show top N models with gaps"),
|
|
1217
|
+
field: Optional[str] = typer.Option(None, "--field", help="Show gaps for a specific field (benchmarks, cost, capabilities, context_window, availability, total_parameters)"),
|
|
1218
|
+
) -> None:
|
|
1219
|
+
"""Find models with the most missing data, ranked by priority.
|
|
1220
|
+
|
|
1221
|
+
Works offline from YAML model cards in models/.
|
|
1222
|
+
"""
|
|
1223
|
+
cards = _load_cards(provider=provider, model_type=type)
|
|
1224
|
+
if not cards:
|
|
1225
|
+
console.print("[yellow]No model cards found matching your filters.[/]")
|
|
1226
|
+
raise typer.Exit(1)
|
|
1227
|
+
|
|
1228
|
+
gap_rows = []
|
|
1229
|
+
for card in cards:
|
|
1230
|
+
info = _compute_gap_info(card)
|
|
1231
|
+
# If --field is given, only include models missing that field
|
|
1232
|
+
if field:
|
|
1233
|
+
matches = [m for m in info["missing"] if field.lower() in m.lower()]
|
|
1234
|
+
if not matches:
|
|
1235
|
+
continue
|
|
1236
|
+
if info["missing_count"] == 0:
|
|
1237
|
+
continue
|
|
1238
|
+
gap_rows.append(info)
|
|
1239
|
+
|
|
1240
|
+
# Sort by priority descending
|
|
1241
|
+
gap_rows.sort(key=lambda x: x["priority"], reverse=True)
|
|
1242
|
+
gap_rows = gap_rows[:top]
|
|
1243
|
+
|
|
1244
|
+
if not gap_rows:
|
|
1245
|
+
console.print("[green]All model cards look complete for the selected criteria.[/]")
|
|
1246
|
+
return
|
|
1247
|
+
|
|
1248
|
+
table = Table(
|
|
1249
|
+
title=f"Data Gaps (top {len(gap_rows)} of {len(cards)} models)",
|
|
1250
|
+
show_header=True,
|
|
1251
|
+
header_style="bold cyan",
|
|
1252
|
+
)
|
|
1253
|
+
table.add_column("#", justify="right", style="dim", width=4)
|
|
1254
|
+
table.add_column("Model", style="bold white", min_width=30)
|
|
1255
|
+
table.add_column("Provider", style="dim", min_width=10)
|
|
1256
|
+
table.add_column("Coverage", justify="right", min_width=9)
|
|
1257
|
+
table.add_column("Missing Fields", style="yellow", min_width=35)
|
|
1258
|
+
table.add_column("Priority", justify="right", style="bold magenta")
|
|
1259
|
+
|
|
1260
|
+
for i, row in enumerate(gap_rows, 1):
|
|
1261
|
+
pct = row["applicable_field_coverage"]
|
|
1262
|
+
if pct >= 50:
|
|
1263
|
+
comp_str = f"[green]{pct:.1f}%[/]"
|
|
1264
|
+
elif pct >= 25:
|
|
1265
|
+
comp_str = f"[yellow]{pct:.1f}%[/]"
|
|
1266
|
+
else:
|
|
1267
|
+
comp_str = f"[red]{pct:.1f}%[/]"
|
|
1268
|
+
|
|
1269
|
+
missing_str = ", ".join(row["missing"][:5])
|
|
1270
|
+
if len(row["missing"]) > 5:
|
|
1271
|
+
missing_str += f" (+{len(row['missing']) - 5} more)"
|
|
1272
|
+
|
|
1273
|
+
table.add_row(
|
|
1274
|
+
str(i),
|
|
1275
|
+
row["display_name"] or row["model_id"],
|
|
1276
|
+
row["provider"],
|
|
1277
|
+
comp_str,
|
|
1278
|
+
missing_str,
|
|
1279
|
+
f"{row['priority']:.1f}",
|
|
1280
|
+
)
|
|
1281
|
+
|
|
1282
|
+
console.print(table)
|
|
1283
|
+
|
|
1284
|
+
# Summary
|
|
1285
|
+
total_gaps = sum(r["missing_count"] for r in gap_rows)
|
|
1286
|
+
console.print(f"\n[dim]Total gap fields across shown models: {total_gaps}[/]")
|
|
1287
|
+
console.print("[dim]Run [bold]modelspec research <model_id>[/bold] to auto-fill data from HuggingFace.[/]")
|
|
1288
|
+
|
|
1289
|
+
|
|
1290
|
+
# ───────────────────────────────────────────────────────────────
|
|
1291
|
+
# 8. modelspec research <model_id> — auto-research a model
|
|
1292
|
+
# ───────────────────────────────────────────────────────────────
|
|
1293
|
+
|
|
1294
|
+
def _find_card_path(model_id: str) -> Path | None:
|
|
1295
|
+
"""Find the YAML file for a given model_id like 'openai/gpt-4o'."""
|
|
1296
|
+
parts = model_id.split("/", 1)
|
|
1297
|
+
if len(parts) != 2:
|
|
1298
|
+
return None
|
|
1299
|
+
provider, slug = parts
|
|
1300
|
+
path = _MODELS_DIR / provider / f"{slug}.md"
|
|
1301
|
+
if path.exists():
|
|
1302
|
+
return path
|
|
1303
|
+
# Try case-insensitive match
|
|
1304
|
+
provider_dir = _MODELS_DIR / provider
|
|
1305
|
+
if provider_dir.is_dir():
|
|
1306
|
+
for f in provider_dir.glob("*.md"):
|
|
1307
|
+
if f.stem.lower() == slug.lower():
|
|
1308
|
+
return f
|
|
1309
|
+
return None
|
|
1310
|
+
|
|
1311
|
+
|
|
1312
|
+
def _fetch_huggingface_data(model_id: str) -> dict[str, Any]:
|
|
1313
|
+
"""Fetch model metadata from HuggingFace Hub API."""
|
|
1314
|
+
import httpx
|
|
1315
|
+
|
|
1316
|
+
url = f"https://huggingface.co/api/models/{model_id}"
|
|
1317
|
+
try:
|
|
1318
|
+
resp = httpx.get(url, timeout=15.0, follow_redirects=True)
|
|
1319
|
+
if resp.status_code == 200:
|
|
1320
|
+
return resp.json()
|
|
1321
|
+
return {}
|
|
1322
|
+
except Exception:
|
|
1323
|
+
return {}
|
|
1324
|
+
|
|
1325
|
+
|
|
1326
|
+
def _apply_hf_updates(card: Any, hf_data: dict[str, Any]) -> dict[str, tuple[Any, Any]]:
|
|
1327
|
+
"""Apply HuggingFace data to a card, only filling None/empty fields.
|
|
1328
|
+
|
|
1329
|
+
Returns a dict of {field_path: (old_value, new_value)} for changes made.
|
|
1330
|
+
"""
|
|
1331
|
+
changes: dict[str, tuple[Any, Any]] = {}
|
|
1332
|
+
|
|
1333
|
+
# Downloads
|
|
1334
|
+
downloads = hf_data.get("downloads")
|
|
1335
|
+
if downloads and not card.adoption.huggingface_downloads:
|
|
1336
|
+
old = card.adoption.huggingface_downloads
|
|
1337
|
+
card.adoption.huggingface_downloads = downloads
|
|
1338
|
+
changes["adoption.huggingface_downloads"] = (old, downloads)
|
|
1339
|
+
|
|
1340
|
+
# Likes
|
|
1341
|
+
likes = hf_data.get("likes")
|
|
1342
|
+
if likes and not card.adoption.huggingface_likes:
|
|
1343
|
+
old = card.adoption.huggingface_likes
|
|
1344
|
+
card.adoption.huggingface_likes = likes
|
|
1345
|
+
changes["adoption.huggingface_likes"] = (old, likes)
|
|
1346
|
+
|
|
1347
|
+
# Tags -> identity.tags (only if empty)
|
|
1348
|
+
hf_tags = hf_data.get("tags", [])
|
|
1349
|
+
if hf_tags and not card.identity.tags:
|
|
1350
|
+
card.identity.tags = hf_tags[:20] # cap at 20 tags
|
|
1351
|
+
changes["identity.tags"] = ([], hf_tags[:20])
|
|
1352
|
+
|
|
1353
|
+
# Pipeline tag
|
|
1354
|
+
pipeline_tag = hf_data.get("pipeline_tag")
|
|
1355
|
+
if pipeline_tag and not card.identity.pipeline_tag:
|
|
1356
|
+
old = card.identity.pipeline_tag
|
|
1357
|
+
card.identity.pipeline_tag = pipeline_tag
|
|
1358
|
+
changes["identity.pipeline_tag"] = (old, pipeline_tag)
|
|
1359
|
+
|
|
1360
|
+
# License
|
|
1361
|
+
license_str = hf_data.get("cardData", {}).get("license") if isinstance(hf_data.get("cardData"), dict) else None
|
|
1362
|
+
if not license_str:
|
|
1363
|
+
# Try top-level tags for license
|
|
1364
|
+
for tag in hf_tags:
|
|
1365
|
+
if tag.startswith("license:"):
|
|
1366
|
+
license_str = tag.split(":", 1)[1]
|
|
1367
|
+
break
|
|
1368
|
+
|
|
1369
|
+
# Parameter count from safetensors
|
|
1370
|
+
safetensors = hf_data.get("safetensors")
|
|
1371
|
+
if isinstance(safetensors, dict):
|
|
1372
|
+
total = safetensors.get("total")
|
|
1373
|
+
if total and card.architecture.total_parameters is None:
|
|
1374
|
+
card.architecture.total_parameters = total
|
|
1375
|
+
changes["architecture.total_parameters"] = (None, total)
|
|
1376
|
+
|
|
1377
|
+
# Base model from tags
|
|
1378
|
+
for tag in hf_tags:
|
|
1379
|
+
if tag.startswith("base_model:"):
|
|
1380
|
+
base = tag.split(":", 1)[1]
|
|
1381
|
+
if not card.lineage.base_model:
|
|
1382
|
+
card.lineage.base_model = base
|
|
1383
|
+
changes["lineage.base_model"] = ("", base)
|
|
1384
|
+
break
|
|
1385
|
+
|
|
1386
|
+
# Library name
|
|
1387
|
+
library = hf_data.get("library_name")
|
|
1388
|
+
if library and not card.lineage.library_name:
|
|
1389
|
+
card.lineage.library_name = library
|
|
1390
|
+
changes["lineage.library_name"] = ("", library)
|
|
1391
|
+
|
|
1392
|
+
# HuggingFace URL in sources
|
|
1393
|
+
hf_id = hf_data.get("id") or hf_data.get("modelId")
|
|
1394
|
+
if hf_id and not card.sources.huggingface_url:
|
|
1395
|
+
url = f"https://huggingface.co/{hf_id}"
|
|
1396
|
+
card.sources.huggingface_url = url
|
|
1397
|
+
changes["sources.huggingface_url"] = ("", url)
|
|
1398
|
+
|
|
1399
|
+
# HuggingFace platform availability
|
|
1400
|
+
if hf_id and not card.availability.huggingface.available:
|
|
1401
|
+
card.availability.huggingface.available = True
|
|
1402
|
+
card.availability.huggingface.model_id = hf_id
|
|
1403
|
+
changes["availability.huggingface.available"] = (False, True)
|
|
1404
|
+
|
|
1405
|
+
# Update last_scraped timestamp
|
|
1406
|
+
card.sources.last_scraped_huggingface = str(date.today())
|
|
1407
|
+
|
|
1408
|
+
return changes
|
|
1409
|
+
|
|
1410
|
+
|
|
1411
|
+
@app.command()
|
|
1412
|
+
def research(
|
|
1413
|
+
model_id: str = typer.Argument(..., help="Model ID, e.g. meta/llama-3.1-8b-instruct"),
|
|
1414
|
+
source: str = typer.Option("all", "--source", "-s", help="Data source: huggingface, all"),
|
|
1415
|
+
dry_run: bool = typer.Option(False, "--dry-run", help="Show what would change without writing"),
|
|
1416
|
+
) -> None:
|
|
1417
|
+
"""Auto-research a model: fetch data from HuggingFace and update its card.
|
|
1418
|
+
|
|
1419
|
+
Only fills empty/None fields — never overwrites existing data.
|
|
1420
|
+
"""
|
|
1421
|
+
ModelCard = _get_model_card_class()
|
|
1422
|
+
|
|
1423
|
+
card_path = _find_card_path(model_id)
|
|
1424
|
+
if not card_path:
|
|
1425
|
+
console.print(f"[bold red]Card not found:[/] {model_id}")
|
|
1426
|
+
console.print(f" Expected at: models/{model_id.replace('/', '/')}.md")
|
|
1427
|
+
raise typer.Exit(1)
|
|
1428
|
+
|
|
1429
|
+
card = ModelCard.from_yaml_file(card_path)
|
|
1430
|
+
all_changes: dict[str, tuple[Any, Any]] = {}
|
|
1431
|
+
|
|
1432
|
+
# HuggingFace
|
|
1433
|
+
if source in ("huggingface", "all"):
|
|
1434
|
+
# For open-weights models, use model_id directly;
|
|
1435
|
+
# for proprietary, try the HF URL from sources
|
|
1436
|
+
hf_id = model_id
|
|
1437
|
+
if card.sources.huggingface_url:
|
|
1438
|
+
# Extract id from URL like https://huggingface.co/meta-llama/...
|
|
1439
|
+
url_parts = card.sources.huggingface_url.rstrip("/").split("huggingface.co/")
|
|
1440
|
+
if len(url_parts) == 2:
|
|
1441
|
+
hf_id = url_parts[1]
|
|
1442
|
+
|
|
1443
|
+
with Progress(
|
|
1444
|
+
SpinnerColumn(),
|
|
1445
|
+
TextColumn("[progress.description]{task.description}"),
|
|
1446
|
+
console=console,
|
|
1447
|
+
) as progress:
|
|
1448
|
+
progress.add_task(f"Fetching HuggingFace data for {hf_id}...", total=None)
|
|
1449
|
+
hf_data = _fetch_huggingface_data(hf_id)
|
|
1450
|
+
|
|
1451
|
+
if hf_data:
|
|
1452
|
+
hf_changes = _apply_hf_updates(card, hf_data)
|
|
1453
|
+
all_changes.update(hf_changes)
|
|
1454
|
+
if not hf_changes:
|
|
1455
|
+
console.print(f"[dim]HuggingFace: no new data to fill (existing fields already populated).[/]")
|
|
1456
|
+
else:
|
|
1457
|
+
console.print(f"[yellow]HuggingFace: no data found for {hf_id}[/]")
|
|
1458
|
+
|
|
1459
|
+
if not all_changes:
|
|
1460
|
+
console.print(f"\n[yellow]No changes to make for {model_id}.[/]")
|
|
1461
|
+
return
|
|
1462
|
+
|
|
1463
|
+
# Show diff
|
|
1464
|
+
diff_lines = []
|
|
1465
|
+
for field_path, (old_val, new_val) in sorted(all_changes.items()):
|
|
1466
|
+
old_display = repr(old_val) if old_val not in (None, "", [], False) else "[dim]empty[/]"
|
|
1467
|
+
new_display = repr(new_val)
|
|
1468
|
+
# Truncate long values
|
|
1469
|
+
if len(str(new_display)) > 60:
|
|
1470
|
+
new_display = str(new_display)[:57] + "..."
|
|
1471
|
+
diff_lines.append(f" [red]- {field_path}: {old_display}[/]")
|
|
1472
|
+
diff_lines.append(f" [green]+ {field_path}: {new_display}[/]")
|
|
1473
|
+
|
|
1474
|
+
console.print(Panel(
|
|
1475
|
+
"\n".join(diff_lines),
|
|
1476
|
+
title=f"Changes for {model_id} ({len(all_changes)} fields)",
|
|
1477
|
+
border_style="cyan",
|
|
1478
|
+
))
|
|
1479
|
+
|
|
1480
|
+
if dry_run:
|
|
1481
|
+
console.print(f"\n[yellow]Dry run — no files modified.[/]")
|
|
1482
|
+
console.print(f" Run without --dry-run to write changes to {card_path.relative_to(_PROJECT_ROOT)}")
|
|
1483
|
+
return
|
|
1484
|
+
|
|
1485
|
+
# Write updated card
|
|
1486
|
+
yaml_content = card.to_yaml()
|
|
1487
|
+
card_path.write_text(yaml_content, encoding="utf-8")
|
|
1488
|
+
console.print(f"\n[green]Updated {len(all_changes)} fields in {card_path.relative_to(_PROJECT_ROOT)}[/]")
|
|
1489
|
+
|
|
1490
|
+
|
|
1491
|
+
# ───────────────────────────────────────────────────────────────
|
|
1492
|
+
# 9. modelspec contribute — submit a PR
|
|
1493
|
+
# ───────────────────────────────────────────────────────────────
|
|
1494
|
+
|
|
1495
|
+
def _run_cmd(cmd: list[str], check: bool = True, capture: bool = True) -> subprocess.CompletedProcess:
|
|
1496
|
+
"""Run a subprocess command, returning the result."""
|
|
1497
|
+
return subprocess.run(
|
|
1498
|
+
cmd,
|
|
1499
|
+
cwd=str(_PROJECT_ROOT),
|
|
1500
|
+
capture_output=capture,
|
|
1501
|
+
text=True,
|
|
1502
|
+
check=check,
|
|
1503
|
+
)
|
|
1504
|
+
|
|
1505
|
+
|
|
1506
|
+
@app.command()
|
|
1507
|
+
def contribute(
|
|
1508
|
+
message: Optional[str] = typer.Option(None, "--message", "-m", help="Commit/PR message describing your changes"),
|
|
1509
|
+
) -> None:
|
|
1510
|
+
"""Submit your model card changes as a pull request to turbobeest/modelspec.
|
|
1511
|
+
|
|
1512
|
+
Requires the GitHub CLI (gh) to be installed and authenticated.
|
|
1513
|
+
"""
|
|
1514
|
+
# 1. Check gh CLI
|
|
1515
|
+
try:
|
|
1516
|
+
_run_cmd(["gh", "auth", "status"])
|
|
1517
|
+
except (FileNotFoundError, subprocess.CalledProcessError):
|
|
1518
|
+
console.print("[bold red]Error:[/] GitHub CLI (gh) is not installed or not authenticated.")
|
|
1519
|
+
console.print(" Install: https://cli.github.com/")
|
|
1520
|
+
console.print(" Auth: gh auth login")
|
|
1521
|
+
raise typer.Exit(1)
|
|
1522
|
+
|
|
1523
|
+
# 2. Check for changes
|
|
1524
|
+
diff_result = _run_cmd(["git", "diff", "--name-only"], check=False)
|
|
1525
|
+
staged_result = _run_cmd(["git", "diff", "--name-only", "--cached"], check=False)
|
|
1526
|
+
untracked_result = _run_cmd(["git", "ls-files", "--others", "--exclude-standard", "models/"], check=False)
|
|
1527
|
+
|
|
1528
|
+
changed_files = set()
|
|
1529
|
+
for output in (diff_result.stdout, staged_result.stdout, untracked_result.stdout):
|
|
1530
|
+
for line in output.strip().splitlines():
|
|
1531
|
+
line = line.strip()
|
|
1532
|
+
if line and line.startswith("models/") and line.endswith(".md"):
|
|
1533
|
+
changed_files.add(line)
|
|
1534
|
+
|
|
1535
|
+
if not changed_files:
|
|
1536
|
+
console.print("[yellow]No modified or new model cards found.[/]")
|
|
1537
|
+
console.print(" Run [bold]modelspec research <model_id>[/] first to enrich a model card.")
|
|
1538
|
+
raise typer.Exit(0)
|
|
1539
|
+
|
|
1540
|
+
console.print(f"[bold]Found {len(changed_files)} changed model card(s):[/]")
|
|
1541
|
+
for f in sorted(changed_files):
|
|
1542
|
+
console.print(f" [green]+[/] {f}")
|
|
1543
|
+
|
|
1544
|
+
# 3. Determine upstream vs fork
|
|
1545
|
+
remote_result = _run_cmd(["git", "remote", "-v"], check=False)
|
|
1546
|
+
is_fork = "turbobeest/modelspec" not in remote_result.stdout
|
|
1547
|
+
|
|
1548
|
+
# 4. Get username
|
|
1549
|
+
user_result = _run_cmd(["gh", "api", "user", "--jq", ".login"], check=False)
|
|
1550
|
+
username = user_result.stdout.strip() or "contributor"
|
|
1551
|
+
|
|
1552
|
+
# 5. Build branch name and message
|
|
1553
|
+
today = date.today().strftime("%Y%m%d")
|
|
1554
|
+
model_names = []
|
|
1555
|
+
for f in sorted(changed_files):
|
|
1556
|
+
parts = Path(f).stem
|
|
1557
|
+
model_names.append(parts)
|
|
1558
|
+
|
|
1559
|
+
summary = "-".join(model_names[:3])
|
|
1560
|
+
if len(model_names) > 3:
|
|
1561
|
+
summary += f"-and-{len(model_names) - 3}-more"
|
|
1562
|
+
branch_name = f"contrib/{username}/{today}-{summary}"
|
|
1563
|
+
|
|
1564
|
+
if not message:
|
|
1565
|
+
message = f"Data enrichment: {', '.join(model_names[:5])}"
|
|
1566
|
+
if len(model_names) > 5:
|
|
1567
|
+
message += f" and {len(model_names) - 5} more"
|
|
1568
|
+
|
|
1569
|
+
# 6. Load before/after completeness for PR body
|
|
1570
|
+
ModelCard = _get_model_card_class()
|
|
1571
|
+
pr_body_lines = ["## Summary", "", f"Updated {len(changed_files)} model card(s):", ""]
|
|
1572
|
+
for f in sorted(changed_files):
|
|
1573
|
+
card_path = _PROJECT_ROOT / f
|
|
1574
|
+
try:
|
|
1575
|
+
card = ModelCard.from_yaml_file(card_path)
|
|
1576
|
+
pr_body_lines.append(
|
|
1577
|
+
f"- **{card.identity.display_name}** (`{card.identity.model_id}`): "
|
|
1578
|
+
f"{card.applicable_field_coverage:.1f}% applicable field coverage"
|
|
1579
|
+
)
|
|
1580
|
+
except Exception:
|
|
1581
|
+
pr_body_lines.append(f"- `{f}`")
|
|
1582
|
+
|
|
1583
|
+
pr_body_lines.extend([
|
|
1584
|
+
"",
|
|
1585
|
+
"## Details",
|
|
1586
|
+
"",
|
|
1587
|
+
f"Fields enriched via `modelspec research`.",
|
|
1588
|
+
"",
|
|
1589
|
+
"## Test plan",
|
|
1590
|
+
"",
|
|
1591
|
+
"- [ ] `modelspec validate` passes",
|
|
1592
|
+
"- [ ] Spot-check updated fields against source",
|
|
1593
|
+
])
|
|
1594
|
+
|
|
1595
|
+
pr_body = "\n".join(pr_body_lines)
|
|
1596
|
+
|
|
1597
|
+
# 7. Create branch, commit, push, open PR
|
|
1598
|
+
try:
|
|
1599
|
+
_run_cmd(["git", "checkout", "-b", branch_name])
|
|
1600
|
+
console.print(f" Created branch: [bold]{branch_name}[/]")
|
|
1601
|
+
except subprocess.CalledProcessError:
|
|
1602
|
+
console.print(f"[bold red]Error:[/] Could not create branch {branch_name}")
|
|
1603
|
+
raise typer.Exit(1)
|
|
1604
|
+
|
|
1605
|
+
try:
|
|
1606
|
+
for f in sorted(changed_files):
|
|
1607
|
+
_run_cmd(["git", "add", f])
|
|
1608
|
+
_run_cmd(["git", "commit", "-m", message])
|
|
1609
|
+
console.print(f" Committed: {message}")
|
|
1610
|
+
except subprocess.CalledProcessError as exc:
|
|
1611
|
+
console.print(f"[bold red]Error committing:[/] {exc.stderr}")
|
|
1612
|
+
raise typer.Exit(1)
|
|
1613
|
+
|
|
1614
|
+
try:
|
|
1615
|
+
if is_fork:
|
|
1616
|
+
_run_cmd(["git", "push", "-u", "origin", branch_name])
|
|
1617
|
+
else:
|
|
1618
|
+
_run_cmd(["git", "push", "-u", "origin", branch_name])
|
|
1619
|
+
console.print(" Pushed to origin.")
|
|
1620
|
+
except subprocess.CalledProcessError as exc:
|
|
1621
|
+
console.print(f"[bold red]Error pushing:[/] {exc.stderr}")
|
|
1622
|
+
raise typer.Exit(1)
|
|
1623
|
+
|
|
1624
|
+
# Open PR
|
|
1625
|
+
try:
|
|
1626
|
+
pr_result = _run_cmd([
|
|
1627
|
+
"gh", "pr", "create",
|
|
1628
|
+
"--repo", "turbobeest/modelspec",
|
|
1629
|
+
"--title", f"Data enrichment: {summary}",
|
|
1630
|
+
"--body", pr_body,
|
|
1631
|
+
])
|
|
1632
|
+
pr_url = pr_result.stdout.strip()
|
|
1633
|
+
console.print(f"\n[bold green]Pull request created![/]")
|
|
1634
|
+
console.print(f" {pr_url}")
|
|
1635
|
+
except subprocess.CalledProcessError as exc:
|
|
1636
|
+
console.print(f"[bold red]Error creating PR:[/] {exc.stderr}")
|
|
1637
|
+
raise typer.Exit(1)
|
|
1638
|
+
|
|
1639
|
+
|
|
1640
|
+
# ───────────────────────────────────────────────────────────────
|
|
1641
|
+
# 10. modelspec validate — validate all cards
|
|
1642
|
+
# ───────────────────────────────────────────────────────────────
|
|
1643
|
+
|
|
1644
|
+
@app.command()
|
|
1645
|
+
def validate(
|
|
1646
|
+
fix: bool = typer.Option(False, "--fix", help="Attempt to fix common issues"),
|
|
1647
|
+
) -> None:
|
|
1648
|
+
"""Validate all model cards against the schema and report issues.
|
|
1649
|
+
|
|
1650
|
+
Works offline from YAML model cards in models/.
|
|
1651
|
+
"""
|
|
1652
|
+
ModelCard = _get_model_card_class()
|
|
1653
|
+
|
|
1654
|
+
card_files = _discover_card_files()
|
|
1655
|
+
if not card_files:
|
|
1656
|
+
console.print("[yellow]No model cards found in models/.[/]")
|
|
1657
|
+
raise typer.Exit(1)
|
|
1658
|
+
|
|
1659
|
+
valid_cards = []
|
|
1660
|
+
errors: list[tuple[Path, str]] = []
|
|
1661
|
+
fixed: list[tuple[Path, str]] = []
|
|
1662
|
+
|
|
1663
|
+
with Progress(
|
|
1664
|
+
SpinnerColumn(),
|
|
1665
|
+
TextColumn("[progress.description]{task.description}"),
|
|
1666
|
+
console=console,
|
|
1667
|
+
) as progress:
|
|
1668
|
+
task = progress.add_task(f"Validating {len(card_files)} cards...", total=len(card_files))
|
|
1669
|
+
|
|
1670
|
+
for path in card_files:
|
|
1671
|
+
try:
|
|
1672
|
+
card = ModelCard.from_yaml_file(path)
|
|
1673
|
+
valid_cards.append(card)
|
|
1674
|
+
|
|
1675
|
+
# Check for common issues even on parseable cards
|
|
1676
|
+
issues = []
|
|
1677
|
+
|
|
1678
|
+
# Missing required identity fields
|
|
1679
|
+
if not card.identity.display_name:
|
|
1680
|
+
issues.append("missing display_name")
|
|
1681
|
+
if not card.identity.provider:
|
|
1682
|
+
issues.append("missing provider")
|
|
1683
|
+
|
|
1684
|
+
# Enum mismatches (already caught by Pydantic, but double check)
|
|
1685
|
+
if card.identity.model_type is None:
|
|
1686
|
+
issues.append("no model_type set")
|
|
1687
|
+
|
|
1688
|
+
if issues and fix:
|
|
1689
|
+
changed = False
|
|
1690
|
+
if not card.identity.display_name:
|
|
1691
|
+
# Derive from model_id
|
|
1692
|
+
card.identity.display_name = card.identity.model_id.split("/")[-1].replace("-", " ").title()
|
|
1693
|
+
changed = True
|
|
1694
|
+
if changed:
|
|
1695
|
+
yaml_content = card.to_yaml()
|
|
1696
|
+
path.write_text(yaml_content, encoding="utf-8")
|
|
1697
|
+
fixed.append((path, "; ".join(issues)))
|
|
1698
|
+
|
|
1699
|
+
if issues and not fix:
|
|
1700
|
+
for issue in issues:
|
|
1701
|
+
errors.append((path, f"warning: {issue}"))
|
|
1702
|
+
|
|
1703
|
+
except Exception as exc:
|
|
1704
|
+
err_msg = str(exc)
|
|
1705
|
+
# Truncate long error messages
|
|
1706
|
+
if len(err_msg) > 120:
|
|
1707
|
+
err_msg = err_msg[:117] + "..."
|
|
1708
|
+
errors.append((path, err_msg))
|
|
1709
|
+
|
|
1710
|
+
if fix:
|
|
1711
|
+
# Try to fix by re-loading raw YAML and patching
|
|
1712
|
+
try:
|
|
1713
|
+
import yaml
|
|
1714
|
+
content = path.read_text(encoding="utf-8")
|
|
1715
|
+
parts = content.split("---", 2)
|
|
1716
|
+
if len(parts) >= 3:
|
|
1717
|
+
data = yaml.safe_load(parts[1]) or {}
|
|
1718
|
+
# Common fix: remove invalid enum values
|
|
1719
|
+
if "status" in data and data["status"] not in (
|
|
1720
|
+
"active", "beta", "alpha", "deprecated", "sunset", "preview"
|
|
1721
|
+
):
|
|
1722
|
+
data["status"] = "active"
|
|
1723
|
+
yaml_str = yaml.dump(data, default_flow_style=False, sort_keys=False, allow_unicode=True)
|
|
1724
|
+
new_content = f"---\n{yaml_str}---\n\n{parts[2].strip()}"
|
|
1725
|
+
path.write_text(new_content, encoding="utf-8")
|
|
1726
|
+
# Try parsing again
|
|
1727
|
+
card = ModelCard.from_yaml_file(path)
|
|
1728
|
+
valid_cards.append(card)
|
|
1729
|
+
fixed.append((path, "re-serialized with fixes"))
|
|
1730
|
+
except Exception:
|
|
1731
|
+
pass
|
|
1732
|
+
|
|
1733
|
+
progress.advance(task)
|
|
1734
|
+
|
|
1735
|
+
# Report errors
|
|
1736
|
+
if errors:
|
|
1737
|
+
err_table = Table(
|
|
1738
|
+
title=f"Validation Issues ({len(errors)})",
|
|
1739
|
+
show_header=True,
|
|
1740
|
+
header_style="bold red",
|
|
1741
|
+
)
|
|
1742
|
+
err_table.add_column("File", style="dim", min_width=40)
|
|
1743
|
+
err_table.add_column("Issue", style="yellow")
|
|
1744
|
+
for path, msg in errors[:50]:
|
|
1745
|
+
rel_path = str(path.relative_to(_PROJECT_ROOT))
|
|
1746
|
+
err_table.add_row(rel_path, msg)
|
|
1747
|
+
if len(errors) > 50:
|
|
1748
|
+
console.print(f"[dim]... and {len(errors) - 50} more issues[/]")
|
|
1749
|
+
console.print(err_table)
|
|
1750
|
+
|
|
1751
|
+
if fixed:
|
|
1752
|
+
fix_table = Table(
|
|
1753
|
+
title=f"Fixed ({len(fixed)})",
|
|
1754
|
+
show_header=True,
|
|
1755
|
+
header_style="bold green",
|
|
1756
|
+
)
|
|
1757
|
+
fix_table.add_column("File", style="dim", min_width=40)
|
|
1758
|
+
fix_table.add_column("Fix Applied", style="green")
|
|
1759
|
+
for path, msg in fixed:
|
|
1760
|
+
rel_path = str(path.relative_to(_PROJECT_ROOT))
|
|
1761
|
+
fix_table.add_row(rel_path, msg)
|
|
1762
|
+
console.print(fix_table)
|
|
1763
|
+
|
|
1764
|
+
# Applicable field coverage distribution
|
|
1765
|
+
if valid_cards:
|
|
1766
|
+
completeness_values = [c.applicable_field_coverage for c in valid_cards]
|
|
1767
|
+
avg = sum(completeness_values) / len(completeness_values)
|
|
1768
|
+
high = len([v for v in completeness_values if v >= 50])
|
|
1769
|
+
mid = len([v for v in completeness_values if 25 <= v < 50])
|
|
1770
|
+
low = len([v for v in completeness_values if v < 25])
|
|
1771
|
+
|
|
1772
|
+
dist_table = Table(
|
|
1773
|
+
title="Applicable field coverage",
|
|
1774
|
+
show_header=True,
|
|
1775
|
+
header_style="bold cyan",
|
|
1776
|
+
)
|
|
1777
|
+
dist_table.add_column("Range", style="bold", min_width=15)
|
|
1778
|
+
dist_table.add_column("Count", justify="right")
|
|
1779
|
+
dist_table.add_column("Bar", min_width=30)
|
|
1780
|
+
|
|
1781
|
+
max_count = max(high, mid, low, 1)
|
|
1782
|
+
bar_width = 30
|
|
1783
|
+
|
|
1784
|
+
dist_table.add_row(
|
|
1785
|
+
"[green]>= 50%[/]",
|
|
1786
|
+
str(high),
|
|
1787
|
+
"[green]" + "#" * int(high / max_count * bar_width) + "[/]",
|
|
1788
|
+
)
|
|
1789
|
+
dist_table.add_row(
|
|
1790
|
+
"[yellow]25-49%[/]",
|
|
1791
|
+
str(mid),
|
|
1792
|
+
"[yellow]" + "#" * int(mid / max_count * bar_width) + "[/]",
|
|
1793
|
+
)
|
|
1794
|
+
dist_table.add_row(
|
|
1795
|
+
"[red]< 25%[/]",
|
|
1796
|
+
str(low),
|
|
1797
|
+
"[red]" + "#" * int(low / max_count * bar_width) + "[/]",
|
|
1798
|
+
)
|
|
1799
|
+
console.print(dist_table)
|
|
1800
|
+
|
|
1801
|
+
console.print(
|
|
1802
|
+
f"\n[bold]Summary:[/] {len(valid_cards)} valid / {len(card_files)} total cards | "
|
|
1803
|
+
f"Average applicable field coverage: [bold]{avg:.1f}%[/]"
|
|
1804
|
+
)
|
|
1805
|
+
|
|
1806
|
+
if errors and not fix:
|
|
1807
|
+
console.print("\n[dim]Run with --fix to attempt automatic fixes.[/]")
|
|
1808
|
+
|
|
1809
|
+
|
|
1810
|
+
# ───────────────────────────────────────────────────────────────
|
|
1811
|
+
# Entry point (for Typer)
|
|
1812
|
+
# ───────────────────────────────────────────────────────────────
|
|
1813
|
+
|
|
1814
|
+
def main() -> None:
|
|
1815
|
+
app()
|
|
1816
|
+
|
|
1817
|
+
|
|
1818
|
+
if __name__ == "__main__":
|
|
1819
|
+
main()
|