Singularity666/Magix
1
1import os2import shutil3import json4import torch5import random6from pathlib import Path7from torch.utils.data import Dataset8from torchvision import transforms9from diffusers import StableDiffusionPipeline, DDIMScheduler, UNet2DConditionModel, AutoencoderKL, DDPMScheduler10from transformers import CLIPTextModel, CLIPTokenizer11from accelerate import Accelerator12from tqdm.auto import tqdm13from PIL import Image14 15class CustomDataset(Dataset):16 def __init__(self, data_dir, prompt, tokenizer, size=512, center_crop=False):17 self.data_dir = Path(data_dir)18 self.prompt = prompt19 self.tokenizer = tokenizer20 self.size = size21 self.center_crop = center_crop22 23 self.image_transforms = transforms.Compose([24 transforms.Resize(size, interpolation=transforms.InterpolationMode.BILINEAR),25 transforms.CenterCrop(size) if center_crop else transforms.RandomCrop(size),26 transforms.ToTensor(),27 transforms.Normalize([0.5], [0.5])28 ])29 30 self.images = [f for f in self.data_dir.iterdir() if f.is_file() and not str(f).endswith(".txt")]31 32 def __len__(self):33 return len(self.images)34 35 def __getitem__(self, idx):36 image_path = self.images[idx]37 image = Image.open(image_path)38 if not image.mode == "RGB":39 image = image.convert("RGB")40 41 image = self.image_transforms(image)42 prompt_ids = self.tokenizer(43 self.prompt, padding="max_length", truncation=True, max_length=self.tokenizer.model_max_length44 ).input_ids45 46 return {"image": image, "prompt_ids": prompt_ids}47 48def fine_tune_model(instance_data_dir, instance_prompt, model_name, output_dir, seed=1337, resolution=512, train_batch_size=1, max_train_steps=800):49 # Setup50 accelerator = Accelerator()51 set_seed(seed)52 53 tokenizer = CLIPTokenizer.from_pretrained(model_name)54 text_encoder = CLIPTextModel.from_pretrained(model_name)55 vae = AutoencoderKL.from_pretrained(model_name)56 unet = UNet2DConditionModel.from_pretrained(model_name)57 noise_scheduler = DDPMScheduler.from_pretrained(model_name, subfolder="scheduler")58 59 dataset = CustomDataset(instance_data_dir, instance_prompt, tokenizer, resolution)60 dataloader = torch.utils.data.DataLoader(dataset, batch_size=train_batch_size, shuffle=True)61 62 optimizer = torch.optim.AdamW(unet.parameters(), lr=1e-6)63 64 unet, optimizer, dataloader = accelerator.prepare(unet, optimizer, dataloader)65 vae.to(accelerator.device)66 text_encoder.to(accelerator.device)67 68 global_step = 069 for step, batch in tqdm(enumerate(dataloader), total=max_train_steps):70 latents = vae.encode(batch["image"].to(accelerator.device)).latent_dist.sample() * 0.1821571 noise = torch.randn_like(latents)72 timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (latents.shape[0],), device=latents.device).long()73 noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)74 encoder_hidden_states = text_encoder(batch["prompt_ids"].to(accelerator.device))[0]75 76 model_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample77 78 loss = torch.nn.functional.mse_loss(model_pred.float(), noise.float(), reduction="mean")79 accelerator.backward(loss)80 81 optimizer.step()82 optimizer.zero_grad()83 global_step += 184 if global_step >= max_train_steps:85 break86 87 # Save model88 unet = accelerator.unwrap_model(unet)89 unet.save_pretrained(output_dir)90 vae.save_pretrained(output_dir)91 text_encoder.save_pretrained(output_dir)92 tokenizer.save_pretrained(output_dir)93 94def set_seed(seed):95 random.seed(seed)96 torch.manual_seed(seed)97 if torch.cuda.is_available():98 torch.cuda.manual_seed_all(seed)