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%
9.6 KB · 237 lines python
Raw Blame History
1from __future__ import annotations23import types4import typing5from typing import Any67from pydantic import BaseModel, Field, model_validator89OPENNESS = ("open-weights", "open-source", "proprietary", "restricted", "unknown")10MODALITIES = ("text", "image", "audio", "video", "code", "embedding", "3d", "multimodal")111213def _is_list_type(tp: Any) -> bool:14    origin = typing.get_origin(tp)15    if origin is list:16        return True17    if origin in (typing.Union, types.UnionType):18        return any(_is_list_type(a) for a in typing.get_args(tp))19    return False202122def _is_type(tp: Any, target: type) -> bool:23    if tp is target:24        return True25    if typing.get_origin(tp) in (typing.Union, types.UnionType):26        return any(a is target for a in typing.get_args(tp))27    return False282930class Tolerant(BaseModel):31    """LLM outputs are messy: null for lists, "yes"/"no" for booleans, "70B" for integers. Coerce instead of failing (and escalating)."""3233    @model_validator(mode="before")34    @classmethod35    def _coerce(cls, data: Any) -> Any:36        if not isinstance(data, dict):37            return data38        out = dict(data)39        for name, field in cls.model_fields.items():40            if name not in out:41                continue42            v = out[name]43            tp = field.annotation44            if _is_list_type(tp):45                if v is None or v == "" or v == "null":46                    out[name] = []47                elif isinstance(v, (str, int, float)):48                    out[name] = [x.strip() for x in str(v).split(",") if x.strip()] if isinstance(v, str) else [v]49                elif isinstance(v, dict):50                    out[name] = [str(x) for x in v.values() if x not in (None, "")]51            elif _is_type(tp, bool) and isinstance(v, str):52                low = v.strip().lower()53                out[name] = True if low in ("yes", "true", "supported", "y") else False if low in ("no", "false", "unsupported", "n", "none") else None54            elif _is_type(tp, int) and isinstance(v, str):55                from aiatlas.sdk.extract.numbers import parse_context_length, parse_param_count5657                parsed = parse_param_count(v) if "param" in name else parse_context_length(v)58                out[name] = parsed if parsed is not None else None59            elif _is_type(tp, float) and isinstance(v, str):60                import re6162                m = re.search(r"-?\d+(?:\.\d+)?", v.replace(",", ""))63                out[name] = float(m.group(0)) if m else None64            elif _is_type(tp, str) and isinstance(v, (int, float)) and not isinstance(v, bool):65                out[name] = str(v)66            elif v == "null":67                out[name] = None68        return out697071class DocumentClassification(Tolerant):72    doc_class: str = Field(description="one of: model_release, model_update, pricing, research_paper, company_news, product_launch, "73                                       "framework_release, dataset_release, benchmark, hardware, regulation, incident, funding, acquisition, other")74    is_primary_source: bool | None = None75    mentioned_models: list[str] = Field(default_factory=list)76    mentioned_companies: list[str] = Field(default_factory=list)77    language: str | None = None78    relevance: float = Field(default=0.5, ge=0, le=1, description="relevance to the AI ecosystem")798081class ModelPassport(Tolerant):82    name: str | None = None83    developer: str | None = Field(default=None, description="organization that developed the model")84    family: str | None = None85    version: str | None = None86    release_date: str | None = Field(default=None, description="ISO date if stated (YYYY-MM-DD, YYYY-MM or YYYY)")87    status: str | None = Field(default=None, description="available | preview | deprecated | retired | announced")88    openness: str | None = Field(default=None, description=" | ".join(OPENNESS))89    license: str | None = None90    architecture: str | None = Field(default=None, description="e.g. transformer decoder, MoE, diffusion, state-space")91    parameter_count: int | None = Field(default=None, description="total parameters, as an integer (e.g. 70000000000)")92    active_parameter_count: int | None = None93    is_moe: bool | None = None94    modalities_input: list[str] = Field(default_factory=list)95    modalities_output: list[str] = Field(default_factory=list)96    context_length: int | None = Field(default=None, description="tokens")97    max_output_tokens: int | None = None98    knowledge_cutoff: str | None = None99    languages: list[str] = Field(default_factory=list)100    tool_calling: bool | None = None101    structured_output: bool | None = None102    reasoning: bool | None = None103    vision: bool | None = None104    audio: bool | None = None105    fine_tuning_available: bool | None = None106    tokenizer: str | None = None107    training_data_notes: str | None = None108    predecessor: str | None = None109    base_model: str | None = None110    quantizations: list[str] = Field(default_factory=list)111    paper_url: str | None = None112    model_card_url: str | None = None113    repository_url: str | None = None114    official_page_url: str | None = None115    hardware_requirements: str | None = None116    safety_notes: str | None = None117    evidence: list[str] = Field(default_factory=list, description="short quotes supporting the most important fields")118119120class PriceLine(Tolerant):121    model: str122    provider_model_id: str | None = None123    input_per_mtok: float | None = Field(default=None, description="USD per 1M input tokens")124    output_per_mtok: float | None = None125    cached_input_per_mtok: float | None = None126    cache_write_per_mtok: float | None = None127    batch_input_per_mtok: float | None = None128    batch_output_per_mtok: float | None = None129    per_image: float | None = None130    context_length: int | None = None131    max_output_tokens: int | None = None132    notes: str | None = None133134135class PricingExtraction(Tolerant):136    provider: str | None = None137    currency: str = "USD"138    effective_date: str | None = None139    prices: list[PriceLine] = Field(default_factory=list)140141142class CompanyPassport(Tolerant):143    name: str | None = None144    legal_name: str | None = None145    country: str | None = Field(default=None, description="ISO 3166-1 alpha-2 if determinable")146    headquarters: str | None = None147    founded: str | None = None148    founders: list[str] = Field(default_factory=list)149    leadership: list[str] = Field(default_factory=list)150    website: str | None = None151    description: str | None = None152    products: list[str] = Field(default_factory=list)153    models: list[str] = Field(default_factory=list)154    investors: list[str] = Field(default_factory=list)155    parent_company: str | None = None156    subsidiaries: list[str] = Field(default_factory=list)157    employee_count: int | None = None158    funding_total_usd: float | None = None159160161class PaperPassport(Tolerant):162    title: str | None = None163    authors: list[str] = Field(default_factory=list)164    affiliations: list[str] = Field(default_factory=list)165    date: str | None = None166    field: str | None = None167    summary: str | None = Field(default=None, description="2-3 sentence factual summary")168    methods: list[str] = Field(default_factory=list)169    models: list[str] = Field(default_factory=list, description="models introduced or evaluated")170    datasets: list[str] = Field(default_factory=list)171    benchmarks: list[str] = Field(default_factory=list)172    key_claims: list[str] = Field(default_factory=list)173    results: list[str] = Field(default_factory=list)174    limitations: list[str] = Field(default_factory=list)175    code_url: str | None = None176177178class BenchmarkRow(Tolerant):179    model: str180    score: float181    metric: str | None = None182    config: str | None = None183184185class BenchmarkResultExtraction(Tolerant):186    benchmark: str | None = None187    metric: str | None = None188    higher_is_better: bool | None = True189    evaluated_at: str | None = None190    rows: list[BenchmarkRow] = Field(default_factory=list)191192193class HardwareSpec(Tolerant):194    name: str | None = None195    manufacturer: str | None = None196    kind: str | None = Field(default=None, description="gpu | cpu | npu | tpu | asic | soc | accelerator | system")197    architecture: str | None = None198    release_date: str | None = None199    memory_gb: float | None = None200    memory_type: str | None = None201    memory_bandwidth_gbs: float | None = None202    compute_fp16_tflops: float | None = None203    compute_fp8_tflops: float | None = None204    compute_int8_tops: float | None = None205    tdp_watts: float | None = None206    form_factor: str | None = None207    price_usd: float | None = None208    interconnect: str | None = None209210211class ReleaseAnnouncement(Tolerant):212    kind: str | None = Field(default=None, description="new_model | model_update | pricing_change | deprecation | new_product | framework_release | research | other")213    title: str | None = None214    date: str | None = None215    organization: str | None = None216    models: list[str] = Field(default_factory=list)217    summary: str | None = None218    facts: list[str] = Field(default_factory=list, description="atomic factual statements (one per item)")219    model_passport: ModelPassport | None = None220    pricing: PricingExtraction | None = None221222223TASKS: dict[str, tuple[type[BaseModel], str]] = {224    "classify": (DocumentClassification, "small"),225    "model_passport": (ModelPassport, "medium"),226    "pricing": (PricingExtraction, "medium"),227    "company_passport": (CompanyPassport, "medium"),228    "paper_passport": (PaperPassport, "medium"),229    "benchmark_results": (BenchmarkResultExtraction, "medium"),230    "hardware_spec": (HardwareSpec, "medium"),231    "release_announcement": (ReleaseAnnouncement, "medium"),232}233234235def schema_for(task: str) -> tuple[type[BaseModel], str]:236    return TASKS[task]237