Download tests/test_ddpm_adapter.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 2.81 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/tests/test_ddpm_adapter.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager/tests/test_ddpm_adapter.py
-
curl -L -o test_ddpm_adapter.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/tests/test_ddpm_adapter.py
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) | |