File size: 2,807 Bytes
1ac9744
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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)