AI_Development_Automation_Manager / tests /test_ddpm_adapter.py
SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
2.81 kB
from __future__ import annotations
import io
import threading
from pathlib import Path
from adam.config import ConfigManager
from adam.executor import ToolContext
from adam.registry import ToolSpec
from adam.tools import ddpm_adapter
def _context(root: Path, logs: list[str]) -> ToolContext:
running = threading.Event()
running.set()
return ToolContext(
root=root,
job_id="DDPMTEST",
tool=ToolSpec("ddpm_trainer", "DDPM Trainer", "test", "Training", "train_ddpm"),
cancel_event=threading.Event(),
run_event=running,
progress_callback=lambda *_args: None,
log_callback=logs.append,
)
def test_resolution_change_branches_from_pipeline_instead_of_resuming_checkpoint(tmp_path: Path, monkeypatch) -> None:
trainer_root = tmp_path / "DDPM"
model = trainer_root / "output" / "Anime"
checkpoint = model / "checkpoint-40"
dataset = tmp_path / "dataset"
for folder in (checkpoint / "unet", model / "unet", model / "scheduler", dataset):
folder.mkdir(parents=True)
(trainer_root / "train.py").write_text("# fake trainer", encoding="utf-8")
(model / "model_index.json").write_text("{}", encoding="utf-8")
(checkpoint / "optimizer.bin").write_bytes(b"optimizer")
(checkpoint / "scheduler.bin").write_bytes(b"scheduler")
(checkpoint / "unet" / "diffusion_pytorch_model.safetensors").write_bytes(b"weights")
(checkpoint / "unet" / "config.json").write_text('{"sample_size": 64}', encoding="utf-8")
for index in range(2):
(dataset / f"{index}.png").write_bytes(b"image")
ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}})
monkeypatch.setattr(ddpm_adapter.importlib.util, "find_spec", lambda _name: object())
commands: list[list[str]] = []
class FakeProcess:
stdout = io.StringIO("PROGRESS_JSON:{\"event\":\"done\"}\n")
returncode = 0
def poll(self):
return 0
def terminate(self):
self.returncode = -15
def fake_popen(command, **_kwargs):
commands.append(command)
return FakeProcess()
monkeypatch.setattr(ddpm_adapter.subprocess, "Popen", fake_popen)
logs: list[str] = []
result = ddpm_adapter.train_ddpm(
_context(tmp_path, logs),
dataset_dir=str(dataset),
model_name="Anime",
epochs=5,
output_dir=str(model),
resume_from=str(checkpoint),
resolution=256,
)
command = commands[0]
assert "--pretrained_model_path" in command
assert str(model.resolve()) in command
assert "--resume_from_checkpoint" not in command
assert result["output_folder"] != str(model.resolve())
assert any("Changing DDPM resolution from 64px to 256px" in line for line in logs)