forked from Haidra-Org/AI-Horde
-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathsimple_model_test.py
More file actions
130 lines (109 loc) · 4.13 KB
/
Copy pathsimple_model_test.py
File metadata and controls
130 lines (109 loc) · 4.13 KB
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
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
#!/usr/bin/env python3
# SPDX-FileCopyrightText: 2026 AI Power Grid
#
# SPDX-License-Identifier: AGPL-3.0-or-later
"""
Simple test to debug model validation issues
- Downloads the model reference JSON directly (no horde imports)
- Rebuilds the stable_diffusion_names set like server does
- Simulates ImageWorker.parse_models() accept/reject logic
- Can simulate server failure (e.g., empty reference) to reproduce 400 BadRequest
"""
import json
import os
import sys
import requests
# Defaults
REF_URL_DEFAULT = "https://raw.githubusercontent.com/AIPowerGrid/grid-image-model-reference/main/stable_diffusion.json"
DIFFUSERS_URL_DEFAULT = "https://raw.githubusercontent.com/AIPowerGrid/grid-image-model-reference/main/diffusers.json"
ACCEPTED_BASELINES = {
"stable diffusion 1",
"stable diffusion 2",
"stable diffusion 2 512",
"stable_diffusion_xl",
"stable_cascade",
"flux_1",
}
def download_reference() -> dict:
ref_url = os.getenv("HORDE_IMAGE_COMPVIS_REFERENCE", REF_URL_DEFAULT)
diff_url = os.getenv("HORDE_IMAGE_DIFFUSERS_REFERENCE", DIFFUSERS_URL_DEFAULT)
print(f"Reference URL: {ref_url}")
print(f"Diffusers URL: {diff_url}")
ref, diff = {}, {}
resp = requests.get(ref_url, timeout=10)
resp.raise_for_status()
ref = resp.json()
try:
resp2 = requests.get(diff_url, timeout=10)
resp2.raise_for_status()
diff = resp2.json()
except Exception as e:
print(f"Note: could not load diffusers reference ({e}) - continuing with compvis only")
ref.update(diff)
return ref
def build_stable_names(reference: dict) -> set[str]:
stable = set()
for name, info in reference.items():
baseline = info.get("baseline")
if baseline in ACCEPTED_BASELINES:
stable.add(name)
return stable
def simulate_parse_models(
worker_models: list[str],
stable_names: set[str],
user_special=False,
user_customizer=False,
testing_models: set[str] = None,
) -> set[str]:
if testing_models is None:
testing_models = set()
accepted = set()
for model in worker_models:
parts = model.split("::")
if user_special and len(parts) == 2:
accepted.add(model)
elif (model in stable_names) or user_customizer or (model in testing_models):
accepted.add(model)
else:
print(f"Rejecting unknown model: {model}")
return accepted
def run_scenario(worker_models: list[str], simulate_ref_failure: bool = False):
print("=== Scenario ===")
print(f"Worker advertises: {worker_models}")
reference = {}
stable_names = set()
try:
if not simulate_ref_failure:
reference = download_reference()
stable_names = build_stable_names(reference)
print(f"Reference models: {len(reference)} | stable_names: {len(stable_names)}")
else:
print("Simulating server reference failure (empty stable_names)")
reference = {}
stable_names = set()
except Exception as e:
print(f"Reference download/parse failed: {e}")
stable_names = set() # replicate failure path on server
accepted = simulate_parse_models(worker_models, stable_names)
if not accepted:
print("RESULT: 400 BadRequest -> 'Unfortunately we cannot accept workers serving unrecognised models at this time'")
else:
print(f"RESULT: OK. Accepted models: {sorted(list(accepted))}")
def main():
# Use the exact models from your failing worker payload
default_worker_models = [
"FLUX.1-dev-Kontext-fp8-scaled",
"FLUX.1-schnell",
"Chroma",
]
if len(sys.argv) > 1:
try:
default_worker_models = json.loads(sys.argv[1])
except Exception as e:
print(f"Could not parse argv[1] as JSON list, using defaults. Error: {e}")
print("\n=== Test: Normal (should succeed if reference loads and names match) ===")
run_scenario(default_worker_models, simulate_ref_failure=False)
print("\n=== Test: Simulated server failure (reproduces 400) ===")
run_scenario(default_worker_models, simulate_ref_failure=True)
if __name__ == "__main__":
main()