CoolFace
Apppublic

Allex21/Lora-trainer-all

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
lora_trainer.py409 linesDownload Raw Back to root
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