CoolFace
Apppublic

Singularity666/Magix

sourceHugging Facemitupdated 2y agoView on Hugging Face
1likes
main.py98 linesDownload Raw Back to root
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)