-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy path3_train.py
More file actions
105 lines (82 loc) · 3.98 KB
/
Copy path3_train.py
File metadata and controls
105 lines (82 loc) · 3.98 KB
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
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
"""Train the PlantShade shade-generation model.
A ControlNet branch on a frozen Stable Diffusion v2.1 backbone. The model takes
a plant structure image (source) + a text prompt describing the supplementary
light position, and learns to generate the corresponding shadow map (target).
Paths are repo-relative by default and can be overridden with environment
variables (see below). Run from the repository root:
python train.py
"""
import os, sys, json, random
from datetime import datetime
ROOT = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, ROOT) # so `from share import *` works
import numpy as np
import cv2
import pytorch_lightning as pl
from torch.utils.data import Dataset, DataLoader
from pytorch_lightning.callbacks import ModelCheckpoint
from share import *
from cldm.logger import ImageLogger, LossLogger
from cldm.model import create_model, load_state_dict
def _cfg(name, default):
return os.environ.get(name, default)
# -------- configuration (override via env vars) --------
MODEL_CONFIG = _cfg("PS_MODEL_CONFIG", os.path.join(ROOT, "models", "cldm_v21.yaml"))
# ControlNet-initialised SD2.1 checkpoint (see README: tool_add_control_sd21.py)
RESUME_PATH = _cfg("PS_INIT_CKPT", os.path.join(ROOT, "models", "control_sd21_ini.ckpt"))
TRAIN_JSON = _cfg("PS_TRAIN_JSON", os.path.join(ROOT, "data", "train_split.json"))
OUT_DIR = _cfg("PS_TRAIN_OUT", os.path.join(ROOT, "outputs", "train"))
batch_size = int(_cfg("PS_BATCH", "16"))
learning_rate = float(_cfg("PS_LR", "1e-5"))
max_epochs = int(_cfg("PS_EPOCHS", "501"))
logger_freq = 300
sd_locked = True
only_mid_control = False
class PlantShadeDataset(Dataset):
def __init__(self, list_path, seed=42):
self.data = []
with open(list_path, "rt") as f:
for line in f:
line = line.strip()
if line:
self.data.append(json.loads(line))
random.seed(seed)
random.shuffle(self.data)
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
item = self.data[idx]
source = cv2.imread(item["source"])
target = cv2.imread(item["target"])
source = cv2.resize(source, (512, 512), interpolation=cv2.INTER_AREA)
target = cv2.resize(target, (512, 512), interpolation=cv2.INTER_AREA)
# OpenCV reads BGR -> RGB
source = cv2.cvtColor(source, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
target = cv2.cvtColor(target, cv2.COLOR_BGR2RGB).astype(np.float32) / 127.5 - 1.0
return dict(jpg=target, txt=item["prompt"], hint=source)
def main():
model = create_model(MODEL_CONFIG).cpu()
model.load_state_dict(load_state_dict(RESUME_PATH, location="cpu"))
model.learning_rate = learning_rate
model.sd_locked = sd_locked
model.only_mid_control = only_mid_control
dataset = PlantShadeDataset(TRAIN_JSON)
print("Training samples:", len(dataset))
dataloader = DataLoader(dataset, num_workers=0, batch_size=batch_size, shuffle=True)
run_dir = os.path.join(OUT_DIR, datetime.now().strftime("%Y-%m-%d_%H-%M-%S"))
os.makedirs(run_dir, exist_ok=True)
callbacks = [
ImageLogger(batch_frequency=logger_freq),
LossLogger(log_dir=os.path.join(run_dir, "loss_curves")),
ModelCheckpoint(dirpath=os.path.join(run_dir, "best"),
filename="best-{epoch:02d}-{train_loss_simple_step:.4f}",
monitor="train/loss_simple_step", mode="min", save_top_k=1),
ModelCheckpoint(dirpath=os.path.join(run_dir, "periodic"),
filename="epoch-{epoch:02d}", every_n_epochs=50, save_top_k=-1),
]
trainer = pl.Trainer(accelerator="gpu", devices=1, strategy="ddp", precision=32,
callbacks=callbacks, max_epochs=max_epochs)
print(f"batch_size={batch_size} lr={learning_rate} resume={RESUME_PATH}")
trainer.fit(model, train_dataloaders=dataloader)
if __name__ == "__main__":
main()