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)