Test model loading
This commit is contained in:
3
.gitignore
vendored
3
.gitignore
vendored
@@ -191,4 +191,5 @@ aitk_db.db-shm
|
|||||||
/data
|
/data
|
||||||
.claude
|
.claude
|
||||||
original_repo
|
original_repo
|
||||||
.next
|
.next
|
||||||
|
testing/.model_test_outputs
|
||||||
@@ -50,10 +50,12 @@ class FakeConfig:
|
|||||||
self.patch_size = 1
|
self.patch_size = 1
|
||||||
|
|
||||||
class FakeCLIP(torch.nn.Module):
|
class FakeCLIP(torch.nn.Module):
|
||||||
def __init__(self):
|
def __init__(self, device='cuda'):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
self.device = 'cuda'
|
# the pipeline derives its execution device from this attribute;
|
||||||
|
# nn.Module.to() does not update it
|
||||||
|
self.device = device
|
||||||
self.text_model = None
|
self.text_model = None
|
||||||
self.tokenizer = None
|
self.tokenizer = None
|
||||||
self.model_max_length = 77
|
self.model_max_length = 77
|
||||||
@@ -180,8 +182,8 @@ class ChromaModel(BaseModel):
|
|||||||
self.prepare_text_encoder(text_encoder_2, dtype=dtype)
|
self.prepare_text_encoder(text_encoder_2, dtype=dtype)
|
||||||
|
|
||||||
# self.print_and_status_update("Loading CLIP")
|
# self.print_and_status_update("Loading CLIP")
|
||||||
text_encoder = FakeCLIP()
|
text_encoder = FakeCLIP(device=self.device_torch)
|
||||||
tokenizer = FakeCLIP()
|
tokenizer = FakeCLIP(device=self.device_torch)
|
||||||
text_encoder.to(self.device_torch, dtype=dtype)
|
text_encoder.to(self.device_torch, dtype=dtype)
|
||||||
|
|
||||||
self.noise_scheduler = ChromaModel.get_train_scheduler()
|
self.noise_scheduler = ChromaModel.get_train_scheduler()
|
||||||
|
|||||||
@@ -50,10 +50,12 @@ class FakeConfig:
|
|||||||
self.patch_size = 1
|
self.patch_size = 1
|
||||||
|
|
||||||
class FakeCLIP(torch.nn.Module):
|
class FakeCLIP(torch.nn.Module):
|
||||||
def __init__(self):
|
def __init__(self, device='cuda'):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
self.device = 'cuda'
|
# the pipeline derives its execution device from this attribute;
|
||||||
|
# nn.Module.to() does not update it
|
||||||
|
self.device = device
|
||||||
self.text_model = None
|
self.text_model = None
|
||||||
self.tokenizer = None
|
self.tokenizer = None
|
||||||
self.model_max_length = 77
|
self.model_max_length = 77
|
||||||
@@ -182,8 +184,8 @@ class ChromaRadianceModel(BaseModel):
|
|||||||
self.prepare_text_encoder(text_encoder_2, dtype=dtype)
|
self.prepare_text_encoder(text_encoder_2, dtype=dtype)
|
||||||
|
|
||||||
# self.print_and_status_update("Loading CLIP")
|
# self.print_and_status_update("Loading CLIP")
|
||||||
text_encoder = FakeCLIP()
|
text_encoder = FakeCLIP(device=self.device_torch)
|
||||||
tokenizer = FakeCLIP()
|
tokenizer = FakeCLIP(device=self.device_torch)
|
||||||
text_encoder.to(self.device_torch, dtype=dtype)
|
text_encoder.to(self.device_torch, dtype=dtype)
|
||||||
|
|
||||||
self.noise_scheduler = ChromaRadianceModel.get_train_scheduler()
|
self.noise_scheduler = ChromaRadianceModel.get_train_scheduler()
|
||||||
|
|||||||
286
testing/test_model_loading.py
Normal file
286
testing/test_model_loading.py
Normal file
@@ -0,0 +1,286 @@
|
|||||||
|
"""Per-arch model loading + inference smoke test (toolkit/models/v2/PLANNING.md).
|
||||||
|
|
||||||
|
Loads one registered arch through its normal loading path (the same
|
||||||
|
get_model_class -> ModelClass(...).load_model() flow training uses), runs one
|
||||||
|
small sample generation, and asserts an output file was produced.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python testing/test_model_loading.py --arch zimage # one arch, in-process
|
||||||
|
python testing/test_model_loading.py --all # every registered arch,
|
||||||
|
# one subprocess each (full
|
||||||
|
# unload between archs)
|
||||||
|
--allow-download permit hub downloads (default: HF_HUB_OFFLINE=1, so archs
|
||||||
|
whose weights are not local/cached report SKIP)
|
||||||
|
--list list registered archs
|
||||||
|
--device cuda:0
|
||||||
|
|
||||||
|
Add a new model type by adding an entry to MODEL_TESTS.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import glob
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
|
||||||
|
TOOLKIT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
sys.path.insert(0, TOOLKIT_ROOT)
|
||||||
|
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
|
# repo .env carries HF_TOKEN / HF_HOME / MODELS_PATH etc., same as run.py
|
||||||
|
load_dotenv(os.path.join(TOOLKIT_ROOT, ".env"))
|
||||||
|
|
||||||
|
OUTPUT_ROOT = os.path.join(TOOLKIT_ROOT, "testing", ".model_test_outputs")
|
||||||
|
|
||||||
|
# arch -> {"model": ModelConfig kwargs, "sample": GenerateImageConfig kwargs}
|
||||||
|
# Keep samples tiny: this asserts the load/encode/denoise/decode/save path
|
||||||
|
# works, not quality.
|
||||||
|
IMG = {"width": 512, "height": 512, "num_inference_steps": 8, "seed": 42}
|
||||||
|
VID = {"width": 256, "height": 256, "num_inference_steps": 6, "seed": 42, "num_frames": 9}
|
||||||
|
|
||||||
|
MODEL_TESTS = {
|
||||||
|
"zimage": {
|
||||||
|
"model": {"name_or_path": "Tongyi-MAI/Z-Image-Turbo"},
|
||||||
|
"sample": {**IMG, "guidance_scale": 1.0},
|
||||||
|
},
|
||||||
|
"qwen_image": {
|
||||||
|
# 20B: quantize to fit a 32GB card
|
||||||
|
"model": {"name_or_path": "Qwen/Qwen-Image", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {**IMG, "num_inference_steps": 20, "guidance_scale": 4.0},
|
||||||
|
},
|
||||||
|
"krea2": {
|
||||||
|
"model": {"name_or_path": "krea/Krea-2-Turbo", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {**IMG, "guidance_scale": 1.0},
|
||||||
|
},
|
||||||
|
"boogu_image": {
|
||||||
|
# native ~1024; 512/low-step/high-CFG degenerates to a black frame
|
||||||
|
"model": {"name_or_path": "Boogu/Boogu-Image-0.1-Base", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
|
||||||
|
},
|
||||||
|
"ernie_image": {
|
||||||
|
"model": {"name_or_path": "baidu/ERNIE-Image", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
|
||||||
|
},
|
||||||
|
"mageflow": {
|
||||||
|
"model": {"name_or_path": "microsoft/Mage-Flow-Base", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": IMG,
|
||||||
|
},
|
||||||
|
"ideogram4": {
|
||||||
|
"model": {"name_or_path": "ideogram-ai/ideogram-4-fp8", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
|
||||||
|
},
|
||||||
|
"hidream_o1": {
|
||||||
|
"model": {"name_or_path": "HiDream-ai/HiDream-O1-Image", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 28, "guidance_scale": 5.0, "seed": 42},
|
||||||
|
},
|
||||||
|
"anima": {
|
||||||
|
"model": {"name_or_path": "circlestone-labs/Anima-Base-v1.0-Diffusers"},
|
||||||
|
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.5, "seed": 42},
|
||||||
|
},
|
||||||
|
"wan21": {
|
||||||
|
"model": {"name_or_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"},
|
||||||
|
"sample": {"width": 480, "height": 480, "num_inference_steps": 20, "guidance_scale": 5.0, "seed": 42, "num_frames": 17},
|
||||||
|
},
|
||||||
|
"wan22_5b": {
|
||||||
|
"model": {"name_or_path": "Wan-AI/Wan2.2-TI2V-5B-Diffusers", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {"width": 480, "height": 480, "num_inference_steps": 20, "guidance_scale": 5.0, "seed": 42, "num_frames": 17},
|
||||||
|
},
|
||||||
|
"ltx2.3": {
|
||||||
|
# even quantized, the 22B stack does not fit a 32GB card — needs the big GPU
|
||||||
|
"model": {"name_or_path": "Lightricks/LTX-2.3/ltx-2.3-22b-dev.safetensors", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 3.0, "seed": 42, "num_frames": 25},
|
||||||
|
},
|
||||||
|
# single-file / comfy-layout archs: weights resolve under MODELS_PATH (or
|
||||||
|
# download there with --allow-download)
|
||||||
|
"chroma": {
|
||||||
|
"model": {"name_or_path": "lodestones/Chroma1-HD", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {**IMG, "num_inference_steps": 26, "guidance_scale": 4.0},
|
||||||
|
},
|
||||||
|
"flux_kontext": {
|
||||||
|
"model": {"name_or_path": "black-forest-labs/FLUX.1-Kontext-dev", "quantize": True, "quantize_te": True},
|
||||||
|
"sample": {**IMG, "num_inference_steps": 20, "guidance_scale": 2.5},
|
||||||
|
"needs_control_image": True,
|
||||||
|
},
|
||||||
|
"flux2_klein_4b": {
|
||||||
|
"model": {"name_or_path": "black-forest-labs/FLUX.2-klein-base-4B", "quantize_te": True},
|
||||||
|
"sample": {**IMG, "num_inference_steps": 25, "guidance_scale": 4.0},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
SKIP_MARKERS = (
|
||||||
|
"couldn't connect",
|
||||||
|
"offline mode",
|
||||||
|
"hf_hub_offline",
|
||||||
|
"cannot find the requested files",
|
||||||
|
"not found in cache",
|
||||||
|
"localentrynotfounderror",
|
||||||
|
"does not appear to have a file named",
|
||||||
|
"404 client error",
|
||||||
|
"entrynotfounderror",
|
||||||
|
"gatedrepoerror",
|
||||||
|
"cannot access gated repo",
|
||||||
|
"repositorynotfounderror",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def classify_error(err: BaseException) -> str:
|
||||||
|
text = f"{type(err).__name__}: {err}".lower()
|
||||||
|
if any(m in text for m in SKIP_MARKERS):
|
||||||
|
return "SKIP"
|
||||||
|
if isinstance(err, FileNotFoundError):
|
||||||
|
return "SKIP"
|
||||||
|
return "FAIL"
|
||||||
|
|
||||||
|
|
||||||
|
def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||||
|
entry = MODEL_TESTS[arch]
|
||||||
|
out_dir = os.path.join(OUTPUT_ROOT, arch.replace("/", "_").replace(":", "_"))
|
||||||
|
os.makedirs(out_dir, exist_ok=True)
|
||||||
|
for old in glob.glob(os.path.join(out_dir, "*")):
|
||||||
|
os.remove(old)
|
||||||
|
|
||||||
|
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||||
|
from toolkit.util.get_model import get_model_class
|
||||||
|
|
||||||
|
model_config = ModelConfig(arch=arch, dtype="bf16", **entry["model"])
|
||||||
|
ModelClass = get_model_class(model_config)
|
||||||
|
# get_model_class silently falls back to the legacy SD class on an
|
||||||
|
# unknown arch; that is never what a registered test wants
|
||||||
|
if getattr(ModelClass, "arch", None) not in (arch, model_config.arch):
|
||||||
|
raise ValueError(
|
||||||
|
f"arch {arch!r} resolved to {ModelClass.__name__} "
|
||||||
|
f"(arch={getattr(ModelClass, 'arch', None)!r}) — registry mismatch"
|
||||||
|
)
|
||||||
|
|
||||||
|
sampler = None
|
||||||
|
if hasattr(ModelClass, "get_train_scheduler"):
|
||||||
|
sampler = ModelClass.get_train_scheduler()
|
||||||
|
|
||||||
|
sd = ModelClass(
|
||||||
|
device=device,
|
||||||
|
model_config=model_config,
|
||||||
|
dtype="bf16",
|
||||||
|
noise_scheduler=sampler,
|
||||||
|
)
|
||||||
|
sd.load_model()
|
||||||
|
|
||||||
|
sample_kwargs = dict(entry["sample"])
|
||||||
|
if entry.get("needs_control_image"):
|
||||||
|
# edit/kontext models require a control image; a flat gray input is fine
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
ctrl_path = os.path.join(out_dir, ".ctrl.png")
|
||||||
|
Image.new(
|
||||||
|
"RGB", (sample_kwargs["width"], sample_kwargs["height"]), (128, 128, 128)
|
||||||
|
).save(ctrl_path)
|
||||||
|
sample_kwargs["ctrl_img"] = ctrl_path
|
||||||
|
gen = GenerateImageConfig(
|
||||||
|
prompt="a photo of a cat sitting on a wooden table",
|
||||||
|
output_folder=out_dir,
|
||||||
|
# the GenerateImageConfig default for output_ext is the Literal type
|
||||||
|
# alias itself; real callers always pass one
|
||||||
|
output_ext="png",
|
||||||
|
**sample_kwargs,
|
||||||
|
)
|
||||||
|
sd.generate_images([gen])
|
||||||
|
|
||||||
|
produced = [
|
||||||
|
p
|
||||||
|
for p in glob.glob(os.path.join(out_dir, "*"))
|
||||||
|
if os.path.isfile(p) and os.path.getsize(p) > 1024 and not p.endswith(".txt")
|
||||||
|
]
|
||||||
|
if not produced:
|
||||||
|
raise RuntimeError(f"no output file produced in {out_dir}")
|
||||||
|
return {"arch": arch, "status": "PASS", "outputs": produced}
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(description=__doc__)
|
||||||
|
parser.add_argument("--arch", type=str, default=None)
|
||||||
|
parser.add_argument("--all", action="store_true")
|
||||||
|
parser.add_argument("--list", action="store_true")
|
||||||
|
parser.add_argument("--device", type=str, default="cuda:0")
|
||||||
|
parser.add_argument("--allow-download", action="store_true")
|
||||||
|
parser.add_argument("--json-result", type=str, default=None)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if args.list:
|
||||||
|
for arch in MODEL_TESTS:
|
||||||
|
print(arch)
|
||||||
|
return
|
||||||
|
|
||||||
|
os.environ.setdefault("CUDA_DEVICE_ORDER", "PCI_BUS_ID")
|
||||||
|
if not args.allow_download:
|
||||||
|
os.environ.setdefault("HF_HUB_OFFLINE", "1")
|
||||||
|
|
||||||
|
if args.arch is not None:
|
||||||
|
if args.arch not in MODEL_TESTS:
|
||||||
|
raise SystemExit(
|
||||||
|
f"arch {args.arch!r} is not registered; --list shows options"
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
result = run_one(args.arch, args.device, args.allow_download)
|
||||||
|
except BaseException as err:
|
||||||
|
status = classify_error(err)
|
||||||
|
result = {"arch": args.arch, "status": status, "error": f"{type(err).__name__}: {err}"}
|
||||||
|
if status == "FAIL":
|
||||||
|
import traceback
|
||||||
|
|
||||||
|
traceback.print_exc()
|
||||||
|
if args.json_result:
|
||||||
|
with open(args.json_result, "w") as f:
|
||||||
|
json.dump(result, f)
|
||||||
|
print(f"[{result['status']}] {args.arch}" + (f" — {result.get('error', '')}" if result["status"] != "PASS" else ""))
|
||||||
|
if result["status"] == "FAIL":
|
||||||
|
sys.exit(1)
|
||||||
|
return
|
||||||
|
|
||||||
|
if not args.all:
|
||||||
|
parser.print_help()
|
||||||
|
return
|
||||||
|
|
||||||
|
# --all: one subprocess per arch so every model fully unloads (clean CUDA
|
||||||
|
# teardown) before the next loads
|
||||||
|
results = []
|
||||||
|
for arch in MODEL_TESTS:
|
||||||
|
print(f"\n===== {arch} =====")
|
||||||
|
result_path = os.path.join(OUTPUT_ROOT, f".{arch.replace('/', '_')}.result.json")
|
||||||
|
cmd = [
|
||||||
|
sys.executable,
|
||||||
|
os.path.abspath(__file__),
|
||||||
|
"--arch",
|
||||||
|
arch,
|
||||||
|
"--device",
|
||||||
|
args.device,
|
||||||
|
"--json-result",
|
||||||
|
result_path,
|
||||||
|
]
|
||||||
|
if args.allow_download:
|
||||||
|
cmd.append("--allow-download")
|
||||||
|
proc = subprocess.run(cmd, cwd=TOOLKIT_ROOT)
|
||||||
|
if os.path.exists(result_path):
|
||||||
|
with open(result_path) as f:
|
||||||
|
results.append(json.load(f))
|
||||||
|
os.remove(result_path)
|
||||||
|
else:
|
||||||
|
results.append(
|
||||||
|
{"arch": arch, "status": "FAIL", "error": f"subprocess died (exit {proc.returncode})"}
|
||||||
|
)
|
||||||
|
|
||||||
|
print("\n===== summary =====")
|
||||||
|
counts = {"PASS": 0, "FAIL": 0, "SKIP": 0}
|
||||||
|
for r in results:
|
||||||
|
counts[r["status"]] = counts.get(r["status"], 0) + 1
|
||||||
|
line = f"[{r['status']}] {r['arch']}"
|
||||||
|
if r["status"] != "PASS":
|
||||||
|
line += f" — {r.get('error', '')[:160]}"
|
||||||
|
print(line)
|
||||||
|
print(f"\n{counts['PASS']} passed, {counts['FAIL']} failed, {counts['SKIP']} skipped")
|
||||||
|
if counts["FAIL"]:
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -267,14 +267,33 @@ loads via diffusers. Nothing about sources or outputs changes yet. Suggested ord
|
|||||||
|
|
||||||
## Testing
|
## Testing
|
||||||
|
|
||||||
- [ ] `testing/` (or `tests/`) harness: for each migrated arch, load the model via
|
- [x] `testing/test_model_loading.py`: per-arch load + one small sample through
|
||||||
its v2 modules and run one small inference pass (single low-step sample; video
|
the normal training-style flow (get_model_class → load_model →
|
||||||
models at minimum frame count). One arch at a time, full unload between archs.
|
generate_images). `--arch X` runs one in-process; `--all` runs every
|
||||||
- [ ] Weights resolved through the normal resolver against `MODELS_PATH`
|
registered arch in its own subprocess (full unload between archs).
|
||||||
(GPU + local-weights test, not CI-portable at first; skip archs whose weights
|
15 archs registered so far — add each model type as it migrates.
|
||||||
are absent rather than failing).
|
- [x] Missing weights skip rather than fail: default is HF_HUB_OFFLINE=1 and
|
||||||
|
hub/file errors classify as SKIP; `--allow-download` opts into fetching.
|
||||||
|
(GPU + local-weights test, not CI-portable.)
|
||||||
|
- [x] Full sweep run 2026-08-27: 14/15 PASS (zimage, qwen_image, krea2,
|
||||||
|
boogu_image, ernie_image, ideogram4, hidream_o1, anima, wan21, wan22_5b,
|
||||||
|
chroma, flux_kontext, flux2_klein_4b, ltx2.3 — the quantized 22B ltx
|
||||||
|
stack doesn't fit 32GB, needs the 96GB card). mageflow blocked
|
||||||
|
upstream: microsoft/Mage-Flow-Base 404s on the hub (cached locally, so
|
||||||
|
it runs offline — recheck whether the repo moved/went private).
|
||||||
|
- [x] Registry carries realistic per-arch sample settings (native res, steps,
|
||||||
|
CFG) so sweep outputs are visually verifiable, not just "a file
|
||||||
|
exists". Verified: all 14 produce proper generations. Findings from
|
||||||
|
the quality pass: boogu emits a black frame below native res at
|
||||||
|
low-step/high-CFG (settings regime, present pre-restructure, not a
|
||||||
|
migration bug); chroma's FakeCLIP hardcoded device 'cuda' broke any
|
||||||
|
non-cuda:0 run (pre-existing, fixed — FakeCLIP now takes the real
|
||||||
|
device); ideogram4's fp8 release renders its own "blocked by safety
|
||||||
|
filter" card for a plain cat prompt (model behavior, not a bug —
|
||||||
|
investigate its trigger).
|
||||||
- [ ] Round-trip test per model: load → save comfy format → reload from the save →
|
- [ ] Round-trip test per model: load → save comfy format → reload from the save →
|
||||||
outputs match (bf16) / load cleanly (quantized saves).
|
outputs match (bf16) / load cleanly (quantized saves). Lands with the
|
||||||
|
Phase 2 comfy save path.
|
||||||
- [ ] Each newly migrated model adds its test in the same PR as its migration.
|
- [ ] Each newly migrated model adds its test in the same PR as its migration.
|
||||||
|
|
||||||
## TODO / look at later
|
## TODO / look at later
|
||||||
|
|||||||
Reference in New Issue
Block a user