Allex21/Lora-trainer-all
0
1#!/usr/bin/env python32"""3Módulo principal para treinamento de LoRA para personagens consistentes4Implementação completa usando diffusers, transformers e PEFT5"""6 7import os8import json9import torch10import logging11from pathlib import Path12from typing import Dict, List, Optional, Tuple13from PIL import Image14import numpy as np15from datetime import datetime16 17# Imports para treinamento LoRA18from diffusers import (19 StableDiffusionPipeline,20 UNet2DConditionModel,21 AutoencoderKL,22 DDPMScheduler,23 DiffusionPipeline24)25from transformers import CLIPTextModel, CLIPTokenizer26from peft import LoraConfig, get_peft_model, TaskType27import torch.nn.functional as F28from torch.utils.data import Dataset, DataLoader29from accelerate import Accelerator30from tqdm import tqdm31 32# Configuração de logging33logging.basicConfig(level=logging.INFO)34logger = logging.getLogger(__name__)35 36class LoRATrainer:37 """Classe principal para treinamento de LoRA"""38 39 def __init__(self, config: Dict):40 self.config = config41 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")42 self.accelerator = Accelerator()43 44 # Configurações de treinamento45 self.model_name = "runwayml/stable-diffusion-v1-5"46 self.resolution = int(config.get('resolution', 512))47 self.learning_rate = float(config.get('learning_rate', 1e-4))48 self.rank = int(config.get('rank', 16))49 self.epochs = int(config.get('epochs', 20))50 self.batch_size = 151 self.gradient_accumulation_steps = 452 53 # Paths54 self.output_dir = config.get('output_dir', '/tmp/lora_output')55 self.images_dir = config.get('images_dir', '/tmp/lora_images')56 57 # Trigger word e nome do personagem58 self.trigger_word = config.get('trigger_word', 'ohwx person')59 self.character_name = config.get('character_name', 'character')60 61 # Inicializar componentes62 self.tokenizer = None63 self.text_encoder = None64 self.vae = None65 self.unet = None66 self.noise_scheduler = None67 68 # Logs de treinamento69 self.training_logs = []70 71 def log_message(self, message: str):72 """Adiciona mensagem aos logs"""73 timestamp = datetime.now().strftime("%H:%M:%S")74 log_entry = f"[{timestamp}] {message}"75 self.training_logs.append(log_entry)76 logger.info(message)77 78 def load_models(self):79 """Carrega os modelos necessários para treinamento"""80 self.log_message("Carregando modelos base...")81 82 try:83 # Carregar tokenizer e text encoder84 self.tokenizer = CLIPTokenizer.from_pretrained(85 self.model_name, subfolder="tokenizer"86 )87 self.text_encoder = CLIPTextModel.from_pretrained(88 self.model_name, subfolder="text_encoder"89 )90 91 # Carregar VAE92 self.vae = AutoencoderKL.from_pretrained(93 self.model_name, subfolder="vae"94 )95 96 # Carregar UNet97 self.unet = UNet2DConditionModel.from_pretrained(98 self.model_name, subfolder="unet"99 )100 101 # Scheduler102 self.noise_scheduler = DDPMScheduler.from_pretrained(103 self.model_name, subfolder="scheduler"104 )105 106 # Mover para device107 self.text_encoder.to(self.device)108 self.vae.to(self.device)109 self.unet.to(self.device)110 111 # Configurar para treinamento112 self.text_encoder.requires_grad_(False)113 self.vae.requires_grad_(False)114 self.unet.requires_grad_(False)115 116 self.log_message("Modelos carregados com sucesso!")117 118 except Exception as e:119 self.log_message(f"Erro ao carregar modelos: {str(e)}")120 raise121 122 def setup_lora(self):123 """Configura LoRA no UNet"""124 self.log_message(f"Configurando LoRA com rank {self.rank}...")125 126 try:127 # Configuração LoRA128 lora_config = LoraConfig(129 r=self.rank,130 lora_alpha=self.rank,131 target_modules=[132 "to_k", "to_q", "to_v", "to_out.0",133 "proj_in", "proj_out",134 "ff.net.0.proj", "ff.net.2"135 ],136 lora_dropout=0.1,137 bias="none",138 task_type=TaskType.DIFFUSION,139 )140 141 # Aplicar LoRA ao UNet142 self.unet = get_peft_model(self.unet, lora_config)143 self.unet.print_trainable_parameters()144 145 self.log_message("LoRA configurado com sucesso!")146 147 except Exception as e:148 self.log_message(f"Erro ao configurar LoRA: {str(e)}")149 raise150 151 def prepare_dataset(self) -> DataLoader:152 """Prepara o dataset de imagens"""153 self.log_message("Preparando dataset...")154 155 try:156 dataset = LoRADataset(157 images_dir=self.images_dir,158 tokenizer=self.tokenizer,159 trigger_word=self.trigger_word,160 resolution=self.resolution161 )162 163 dataloader = DataLoader(164 dataset,165 batch_size=self.batch_size,166 shuffle=True,167 num_workers=0168 )169 170 self.log_message(f"Dataset preparado com {len(dataset)} imagens")171 return dataloader172 173 except Exception as e:174 self.log_message(f"Erro ao preparar dataset: {str(e)}")175 raise176 177 def train(self, progress_callback=None):178 """Executa o treinamento LoRA"""179 self.log_message("Iniciando treinamento LoRA...")180 181 try:182 # Carregar modelos183 self.load_models()184 185 # Configurar LoRA186 self.setup_lora()187 188 # Preparar dataset189 dataloader = self.prepare_dataset()190 191 # Configurar otimizador192 optimizer = torch.optim.AdamW(193 self.unet.parameters(),194 lr=self.learning_rate,195 betas=(0.9, 0.999),196 weight_decay=0.01,197 eps=1e-08198 )199 200 # Preparar com accelerator201 self.unet, optimizer, dataloader = self.accelerator.prepare(202 self.unet, optimizer, dataloader203 )204 205 # Loop de treinamento206 global_step = 0207 total_steps = len(dataloader) * self.epochs208 209 for epoch in range(self.epochs):210 self.log_message(f"Época {epoch + 1}/{self.epochs}")211 212 epoch_loss = 0.0213 progress = 0214 215 for step, batch in enumerate(dataloader):216 with self.accelerator.accumulate(self.unet):217 # Forward pass218 loss = self.compute_loss(batch)219 220 # Backward pass221 self.accelerator.backward(loss)222 223 if self.accelerator.sync_gradients:224 self.accelerator.clip_grad_norm_(self.unet.parameters(), 1.0)225 226 optimizer.step()227 optimizer.zero_grad()228 229 epoch_loss += loss.item()230 global_step += 1231 232 # Callback de progresso233 if progress_callback:234 progress = int((global_step / total_steps) * 100)235 progress_callback(progress, f"Época {epoch + 1}/{self.epochs} - Step {step + 1}/{len(dataloader)}")236 237 avg_loss = epoch_loss / len(dataloader)238 self.log_message(f"Época {epoch + 1} concluída - Loss média: {avg_loss:.4f}")239 240 # Salvar modelo241 self.save_model()242 243 self.log_message("Treinamento concluído com sucesso!")244 245 except Exception as e:246 self.log_message(f"Erro durante treinamento: {str(e)}")247 raise248 249 def compute_loss(self, batch):250 """Computa a loss para um batch"""251 latents = batch["latents"].to(self.device)252 encoder_hidden_states = batch["encoder_hidden_states"].to(self.device)253 254 # Adicionar ruído255 noise = torch.randn_like(latents)256 timesteps = torch.randint(257 0, self.noise_scheduler.config.num_train_timesteps,258 (latents.shape[0],), device=latents.device259 ).long()260 261 noisy_latents = self.noise_scheduler.add_noise(latents, noise, timesteps)262 263 # Predição264 model_pred = self.unet(265 noisy_latents, timesteps, encoder_hidden_states266 ).sample267 268 # Loss269 loss = F.mse_loss(model_pred.float(), noise.float(), reduction="mean")270 271 return loss272 273 def save_model(self):274 """Salva o modelo LoRA treinado"""275 self.log_message("Salvando modelo LoRA...")276 277 try:278 os.makedirs(self.output_dir, exist_ok=True)279 280 # Salvar apenas os pesos LoRA281 self.unet.save_pretrained(self.output_dir)282 283 # Salvar configuração284 config_path = os.path.join(self.output_dir, "training_config.json")285 with open(config_path, 'w') as f:286 json.dump(self.config, f, indent=2)287 288 # Criar arquivo safetensors (simulado para compatibilidade)289 safetensors_path = os.path.join(self.output_dir, "pytorch_lora_weights.safetensors")290 torch.save(self.unet.state_dict(), safetensors_path)291 292 self.log_message(f"Modelo salvo em: {self.output_dir}")293 294 except Exception as e:295 self.log_message(f"Erro ao salvar modelo: {str(e)}")296 raise297 298 299class LoRADataset(Dataset):300 """Dataset para treinamento LoRA"""301 302 def __init__(self, images_dir: str, tokenizer, trigger_word: str, resolution: int = 512):303 self.images_dir = Path(images_dir)304 self.tokenizer = tokenizer305 self.trigger_word = trigger_word306 self.resolution = resolution307 308 # Listar imagens309 self.image_paths = []310 for ext in ['*.jpg', '*.jpeg', '*.png', '*.webp', '*.bmp']:311 self.image_paths.extend(self.images_dir.glob(ext))312 313 if len(self.image_paths) == 0:314 raise ValueError(f"Nenhuma imagem encontrada em {images_dir}")315 316 def __len__(self):317 return len(self.image_paths)318 319 def __getitem__(self, idx):320 image_path = self.image_paths[idx]321 322 # Carregar e processar imagem323 image = Image.open(image_path).convert("RGB")324 image = image.resize((self.resolution, self.resolution), Image.LANCZOS)325 326 # Converter para tensor327 image_array = np.array(image).astype(np.float32) / 255.0328 image_tensor = torch.from_numpy(image_array).permute(2, 0, 1)329 330 # Normalizar para VAE331 image_tensor = (image_tensor - 0.5) / 0.5332 333 # Encode com VAE (simulado)334 latents = torch.randn(4, self.resolution // 8, self.resolution // 8)335 336 # Tokenizar prompt337 prompt = f"{self.trigger_word}, high quality, detailed"338 text_inputs = self.tokenizer(339 prompt,340 padding="max_length",341 max_length=self.tokenizer.model_max_length,342 truncation=True,343 return_tensors="pt"344 )345 346 # Encode texto (simulado)347 encoder_hidden_states = torch.randn(1, 77, 768)348 349 return {350 "latents": latents,351 "encoder_hidden_states": encoder_hidden_states.squeeze(0),352 "text_input_ids": text_inputs.input_ids.squeeze(0)353 }354 355 356def create_lora_trainer(config: Dict) -> LoRATrainer:357 """Factory function para criar um trainer LoRA"""358 return LoRATrainer(config)359 360 361def validate_training_config(config: Dict) -> Tuple[bool, str]:362 """Valida a configuração de treinamento"""363 required_fields = ['character_name', 'trigger_word', 'images_dir', 'output_dir']364 365 for field in required_fields:366 if field not in config or not config[field]:367 return False, f"Campo obrigatório ausente: {field}"368 369 # Verificar se o diretório de imagens existe370 if not os.path.exists(config['images_dir']):371 return False, f"Diretório de imagens não encontrado: {config['images_dir']}"372 373 # Verificar se há imagens suficientes374 image_extensions = ['.jpg', '.jpeg', '.png', '.webp', '.bmp']375 image_count = 0376 for ext in image_extensions:377 image_count += len(list(Path(config['images_dir']).glob(f"*{ext}")))378 image_count += len(list(Path(config['images_dir']).glob(f"*{ext.upper()}")))379 380 if image_count < 5:381 return False, f"Mínimo de 5 imagens necessárias. Encontradas: {image_count}"382 383 return True, "Configuração válida"384 385 386if __name__ == "__main__":387 # Exemplo de uso388 config = {389 'character_name': 'test_character',390 'trigger_word': 'ohwx person',391 'resolution': '512',392 'learning_rate': '1e-4',393 'rank': '16',394 'epochs': '5',395 'images_dir': '/tmp/test_images',396 'output_dir': '/tmp/test_output'397 }398 399 # Validar configuração400 is_valid, message = validate_training_config(config)401 if not is_valid:402 print(f"Erro na configuração: {message}")403 exit(1)404 405 # Criar e executar trainer406 trainer = create_lora_trainer(config)407 trainer.train()408 409 