Team Ai
Apppublic

Milcho/ControlNet-Guidance

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
tutorial_train.py36 linesDownload Raw Back to root
1from share import *2 3import pytorch_lightning as pl4from torch.utils.data import DataLoader5from tutorial_dataset import MyDataset6from cldm.logger import ImageLogger7from cldm.model import create_model, load_state_dict8 9 10# Configs11resume_path = './models/control_sd15_ini.ckpt'12batch_size = 413logger_freq = 30014learning_rate = 1e-515sd_locked = True16only_mid_control = False17 18 19# First use cpu to load models. Pytorch Lightning will automatically move it to GPUs.20model = create_model('./models/cldm_v15.yaml').cpu()21model.load_state_dict(load_state_dict(resume_path, location='cpu'))22model.learning_rate = learning_rate23model.sd_locked = sd_locked24model.only_mid_control = only_mid_control25 26 27# Misc28dataset = MyDataset()29dataloader = DataLoader(dataset, num_workers=0, batch_size=batch_size, shuffle=True)30logger = ImageLogger(batch_frequency=logger_freq)31trainer = pl.Trainer(gpus=1, precision=32, callbacks=[logger])32 33 34# Train!35trainer.fit(model, dataloader)36