File size: 3,355 Bytes
2a3abe0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""
Regenerate data/replicate_models.json from Replicate's live "text-to-image"
collection. Run this to pick up newly released models:

    python scripts/refresh_replicate_models.py

Needs REPLICATE_API_TOKEN (env or .env). For each model in the collection it
inspects the input schema and keeps only pure text-to-image models -- ones
that take a `prompt` and do NOT require an input image (that filters out
editing / img2img / controlnet models). It records the square-image input
each model understands (aspect_ratio vs width/height) so every audit uses a
consistent 1:1 framing.
"""
import json
import os
import sys
from pathlib import Path

PROJ = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(PROJ))

CONFIG_PATH = PROJ / "data" / "replicate_models.json"
IMAGE_INPUT_KEYS = {"image", "input_image", "image_input", "images", "subject", "mask", "control_image"}


def _slug(owner, name, taken):
    key = name.lower().replace("/", "-")
    if key in taken:
        key = f"{owner}-{name}".lower().replace("/", "-")
    return key


def build_config():
    # load .env if present
    env = PROJ / ".env"
    if env.exists():
        for line in env.read_text().splitlines():
            line = line.strip()
            if line and not line.startswith("#") and "=" in line:
                k, _, v = line.partition("=")
                os.environ.setdefault(k.strip(), v.strip().strip('"').strip("'"))

    import replicate
    client = replicate.Client(api_token=os.environ["REPLICATE_API_TOKEN"])
    col = client.collections.get("text-to-image")

    models, taken = [], set()
    skipped = []
    for m in col.models:
        mid = f"{m.owner}/{m.name}"
        try:
            ver = m.latest_version
            schema = ver.openapi_schema["components"]["schemas"]["Input"]
            props = schema.get("properties", {})
            required = set(schema.get("required", []))
        except Exception as e:
            skipped.append((mid, f"schema:{e}"))
            continue

        if "prompt" not in props:
            skipped.append((mid, "no prompt"))
            continue
        if required & IMAGE_INPUT_KEYS:
            skipped.append((mid, "requires image (editing model)"))
            continue

        # Square-framing input the model understands (consistent across models).
        extra = {}
        if "aspect_ratio" in props:
            extra["aspect_ratio"] = "1:1"
        elif "width" in props and "height" in props:
            extra["width"] = 1024
            extra["height"] = 1024

        key = _slug(m.owner, m.name, taken)
        taken.add(key)
        # Pin the version: works for official AND community models (bare
        # owner/name 404s for community ones), and records exactly which
        # model version an audit used (reproducibility). `label` stays the
        # readable owner/name.
        run_id = f"{mid}:{ver.id}" if getattr(ver, "id", None) else mid
        models.append({"key": key, "label": mid, "id": run_id, "input": extra})

    models.sort(key=lambda x: x["label"])
    CONFIG_PATH.write_text(json.dumps(models, indent=2))
    print(f"Wrote {len(models)} text-to-image models to {CONFIG_PATH}")
    print(f"Skipped {len(skipped)} (not pure text-to-image):")
    for mid, why in skipped:
        print(f"  - {mid}: {why}")


if __name__ == "__main__":
    build_config()