Allex21/Tre
0
1import argparse2import os3import torch4from diffusers import StableDiffusionPipeline5from peft import LoraConfig, get_peft_model6from torch.utils.data import Dataset, DataLoader7from PIL import Image8from torchvision import transforms9 10# Dataset customizado11class ImageDataset(Dataset):12 def __init__(self, folder, size=512):13 self.files = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith((".png", ".jpg", ".jpeg"))]14 self.transform = transforms.Compose([15 transforms.Resize((size, size)),16 transforms.ToTensor()17 ])18 19 def __len__(self):20 return len(self.files)21 22 def __getitem__(self, idx):23 img = Image.open(self.files[idx]).convert("RGB")24 return self.transform(img)25 26def main(args):27 # Carrega modelo base28 model_id = "runwayml/stable-diffusion-v1-5"29 pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=torch.float16).to("cuda")30 31 # Configuração do LoRA32 lora_config = LoraConfig(33 r=args.rank,34 lora_alpha=16,35 target_modules=["to_q", "to_v"],36 lora_dropout=0.1,37 bias="none",38 task_type="CAUSAL_LM",39 )40 pipe.unet = get_peft_model(pipe.unet, lora_config)41 42 # Dataset43 dataset = ImageDataset(args.images_dir)44 dataloader = DataLoader(dataset, batch_size=1, shuffle=True)45 46 # Otimizador47 optimizer = torch.optim.AdamW(pipe.unet.parameters(), lr=args.learning_rate)48 49 # Loop de treino50 for epoch in range(args.num_epochs):51 for batch in dataloader:52 batch = batch.to("cuda")53 noise = torch.randn_like(batch)54 optimizer.zero_grad()55 loss = pipe.unet(batch, noise)["loss"]56 loss.backward()57 optimizer.step()58 print(f"✅ Epoch {epoch+1}/{args.num_epochs} finalizado.")59 60 # Salvar LoRA61 os.makedirs(args.output_dir, exist_ok=True)62 torch.save(pipe.unet.state_dict(), os.path.join(args.output_dir, "lora.safetensors"))63 print("✅ Treinamento concluído. Arquivo salvo em lora.safetensors")64 65if __name__ == "__main__":66 parser = argparse.ArgumentParser()67 parser.add_argument("--images_dir", type=str, required=True)68 parser.add_argument("--output_dir", type=str, required=True)69 parser.add_argument("--learning_rate", type=float, default=1e-4)70 parser.add_argument("--num_epochs", type=int, default=10)71 parser.add_argument("--rank", type=int, default=4)72 args = parser.parse_args()73 main(args)