Spaces:
Runtime error
Runtime error
| """ | |
| Training Engine for LoRA Models | |
| This module handles the complete training pipeline for custom LoRA adapters: | |
| - Dataset preparation and loading | |
| - Model initialization | |
| - Training loop with checkpointing | |
| - Model saving and validation | |
| """ | |
| import logging | |
| import os | |
| from pathlib import Path | |
| from typing import Optional | |
| import torch | |
| from torch.utils.data import Dataset, DataLoader | |
| from PIL import Image | |
| from diffusers import ( | |
| StableDiffusionXLPipeline, | |
| AutoencoderKL, | |
| UNet2DConditionModel | |
| ) | |
| from transformers import CLIPTextModel, CLIPTokenizer | |
| from accelerate import Accelerator | |
| from accelerate.utils import set_seed | |
| from peft import LoraConfig, get_peft_model | |
| import torch.nn.functional as F | |
| from config import settings | |
| logger = logging.getLogger(__name__) | |
| class ArchitectureDataset(Dataset): | |
| """ | |
| Custom dataset for architectural image pairs | |
| Loads input-target pairs for LoRA training. | |
| """ | |
| def __init__(self, input_dir: Path, target_dir: Path, size: int = 1024): | |
| """ | |
| Initialize dataset | |
| Args: | |
| input_dir: Directory containing input images | |
| target_dir: Directory containing target images | |
| size: Target image size | |
| """ | |
| self.input_dir = input_dir | |
| self.target_dir = target_dir | |
| self.size = size | |
| # Load all pairs | |
| self.pairs = self._load_pairs() | |
| logger.info(f"Loaded {len(self.pairs)} training pairs") | |
| def _load_pairs(self): | |
| """ | |
| Match input and target images by their pair ID | |
| Returns: | |
| List of (input_path, target_path) tuples | |
| """ | |
| pairs = [] | |
| # Get all input files | |
| input_files = sorted(self.input_dir.glob("*_input.*")) | |
| for input_file in input_files: | |
| # Extract pair ID | |
| pair_id = input_file.stem.replace("_input", "") | |
| # Find matching target file | |
| target_files = list(self.target_dir.glob(f"{pair_id}_target.*")) | |
| if target_files: | |
| pairs.append((input_file, target_files[0])) | |
| else: | |
| logger.warning(f"No target found for input: {input_file}") | |
| return pairs | |
| def __len__(self): | |
| return len(self.pairs) | |
| def __getitem__(self, idx): | |
| """ | |
| Get a training pair | |
| Returns: | |
| Dictionary with 'input' and 'target' tensors | |
| """ | |
| input_path, target_path = self.pairs[idx] | |
| # Load images | |
| input_img = Image.open(input_path).convert("RGB") | |
| target_img = Image.open(target_path).convert("RGB") | |
| # Resize to target size | |
| input_img = input_img.resize((self.size, self.size), Image.LANCZOS) | |
| target_img = target_img.resize((self.size, self.size), Image.LANCZOS) | |
| # Convert to tensors and normalize to [-1, 1] | |
| input_tensor = torch.from_numpy( | |
| (torch.tensor(input_img).float() / 127.5 - 1.0).numpy() | |
| ).permute(2, 0, 1) | |
| target_tensor = torch.from_numpy( | |
| (torch.tensor(target_img).float() / 127.5 - 1.0).numpy() | |
| ).permute(2, 0, 1) | |
| return { | |
| "input": input_tensor, | |
| "target": target_tensor, | |
| "input_path": str(input_path), | |
| "target_path": str(target_path) | |
| } | |
| def prepare_lora_config(rank: int = 8) -> LoraConfig: | |
| """ | |
| Create LoRA configuration | |
| Args: | |
| rank: LoRA rank (dimensionality of adaptation matrices) | |
| Returns: | |
| LoraConfig object | |
| """ | |
| return LoraConfig( | |
| r=rank, | |
| lora_alpha=rank, # Often set equal to rank | |
| target_modules=[ | |
| "to_q", "to_k", "to_v", "to_out.0", # Attention layers | |
| "proj_in", "proj_out", # Projections | |
| "ff.net.0.proj", "ff.net.2" # Feed-forward layers | |
| ], | |
| lora_dropout=0.1, | |
| bias="none", | |
| task_type="CAUSAL_LM" | |
| ) | |
| def train_lora_model( | |
| train_steps: int = settings.TRAIN_STEPS, | |
| learning_rate: float = settings.LEARNING_RATE, | |
| lora_rank: int = settings.LORA_RANK, | |
| batch_size: int = settings.BATCH_SIZE, | |
| lock_file: Optional[Path] = None | |
| ): | |
| """ | |
| Main training function for LoRA model | |
| This function orchestrates the complete training process: | |
| 1. Initialize accelerator and models | |
| 2. Prepare dataset and dataloader | |
| 3. Training loop with gradient accumulation | |
| 4. Save final model | |
| Args: | |
| train_steps: Total number of training steps | |
| learning_rate: Learning rate for optimizer | |
| lora_rank: Rank for LoRA adaptation | |
| batch_size: Batch size for training | |
| lock_file: Optional lock file to remove after training | |
| """ | |
| try: | |
| logger.info("=" * 60) | |
| logger.info("Starting LoRA Training") | |
| logger.info("=" * 60) | |
| logger.info(f"Configuration:") | |
| logger.info(f" - Train Steps: {train_steps}") | |
| logger.info(f" - Learning Rate: {learning_rate}") | |
| logger.info(f" - LoRA Rank: {lora_rank}") | |
| logger.info(f" - Batch Size: {batch_size}") | |
| logger.info(f" - Base Model: {settings.BASE_MODEL}") | |
| # Initialize accelerator | |
| accelerator = Accelerator( | |
| mixed_precision=settings.MIXED_PRECISION, | |
| gradient_accumulation_steps=settings.GRADIENT_ACCUMULATION_STEPS, | |
| log_with="tensorboard", | |
| project_dir=str(settings.LOGS_DIR) | |
| ) | |
| # Set random seed for reproducibility | |
| set_seed(42) | |
| # Load dataset | |
| logger.info("Loading training dataset...") | |
| dataset = ArchitectureDataset( | |
| input_dir=settings.INPUT_DIR, | |
| target_dir=settings.TARGET_DIR, | |
| size=settings.TARGET_RESOLUTION | |
| ) | |
| if len(dataset) == 0: | |
| raise ValueError("No training data found!") | |
| dataloader = DataLoader( | |
| dataset, | |
| batch_size=batch_size, | |
| shuffle=True, | |
| num_workers=0, # Set to 0 for Windows compatibility | |
| pin_memory=True | |
| ) | |
| # Load base model components | |
| logger.info("Loading Stable Diffusion XL components...") | |
| # Load UNet (the main model we'll adapt) | |
| unet = UNet2DConditionModel.from_pretrained( | |
| settings.BASE_MODEL, | |
| subfolder="unet", | |
| torch_dtype=torch.float16 if accelerator.device.type == "cuda" else torch.float32 | |
| ) | |
| # Load VAE for encoding images | |
| vae = AutoencoderKL.from_pretrained( | |
| settings.BASE_MODEL, | |
| subfolder="vae", | |
| torch_dtype=torch.float16 if accelerator.device.type == "cuda" else torch.float32 | |
| ) | |
| vae.requires_grad_(False) # Freeze VAE | |
| # Load text encoder | |
| text_encoder = CLIPTextModel.from_pretrained( | |
| settings.BASE_MODEL, | |
| subfolder="text_encoder" | |
| ) | |
| text_encoder.requires_grad_(False) # Freeze text encoder | |
| tokenizer = CLIPTokenizer.from_pretrained( | |
| settings.BASE_MODEL, | |
| subfolder="tokenizer" | |
| ) | |
| # Apply LoRA to UNet | |
| logger.info(f"Applying LoRA with rank {lora_rank}...") | |
| lora_config = prepare_lora_config(rank=lora_rank) | |
| unet = get_peft_model(unet, lora_config) | |
| unet.print_trainable_parameters() | |
| # Setup optimizer | |
| optimizer = torch.optim.AdamW( | |
| unet.parameters(), | |
| lr=learning_rate, | |
| betas=(0.9, 0.999), | |
| weight_decay=1e-2, | |
| eps=1e-8 | |
| ) | |
| # Prepare with accelerator | |
| unet, optimizer, dataloader = accelerator.prepare( | |
| unet, optimizer, dataloader | |
| ) | |
| vae = vae.to(accelerator.device) | |
| text_encoder = text_encoder.to(accelerator.device) | |
| # Training loop | |
| logger.info("Starting training loop...") | |
| global_step = 0 | |
| progress_bar = range(train_steps) | |
| # Encode prompt for all samples | |
| prompt = settings.DEFAULT_PROMPT | |
| text_inputs = tokenizer( | |
| prompt, | |
| padding="max_length", | |
| max_length=tokenizer.model_max_length, | |
| truncation=True, | |
| return_tensors="pt" | |
| ) | |
| text_embeddings = text_encoder(text_inputs.input_ids.to(accelerator.device))[0] | |
| unet.train() | |
| for epoch in range(100): # Large number, will break when steps reached | |
| for batch in dataloader: | |
| with accelerator.accumulate(unet): | |
| # Encode images to latent space | |
| with torch.no_grad(): | |
| latents_input = vae.encode(batch["input"].to(accelerator.device)).latent_dist.sample() | |
| latents_target = vae.encode(batch["target"].to(accelerator.device)).latent_dist.sample() | |
| # Scale latents | |
| latents_input = latents_input * vae.config.scaling_factor | |
| latents_target = latents_target * vae.config.scaling_factor | |
| # Sample random timestep | |
| timesteps = torch.randint( | |
| 0, 1000, (latents_input.shape[0],), | |
| device=accelerator.device | |
| ).long() | |
| # Add noise to target latents | |
| noise = torch.randn_like(latents_target) | |
| noisy_latents = latents_target # Simplified for img2img | |
| # Predict noise | |
| model_pred = unet( | |
| noisy_latents, | |
| timesteps, | |
| text_embeddings.repeat(latents_input.shape[0], 1, 1) | |
| ).sample | |
| # Calculate loss (MSE between prediction and target) | |
| loss = F.mse_loss(model_pred.float(), latents_target.float(), reduction="mean") | |
| # Backpropagation | |
| accelerator.backward(loss) | |
| if accelerator.sync_gradients: | |
| accelerator.clip_grad_norm_(unet.parameters(), settings.MAX_GRAD_NORM) | |
| optimizer.step() | |
| optimizer.zero_grad() | |
| if accelerator.sync_gradients: | |
| global_step += 1 | |
| if global_step % 50 == 0: | |
| logger.info(f"Step {global_step}/{train_steps} - Loss: {loss.item():.4f}") | |
| # Save checkpoint | |
| if global_step % settings.SAVE_STEPS == 0: | |
| logger.info(f"Saving checkpoint at step {global_step}") | |
| save_path = settings.LORA_DIR / f"checkpoint_{global_step}" | |
| accelerator.unwrap_model(unet).save_pretrained(save_path) | |
| # Check if training complete | |
| if global_step >= train_steps: | |
| break | |
| if global_step >= train_steps: | |
| break | |
| # Save final model | |
| logger.info("Training completed! Saving final model...") | |
| final_path = settings.LORA_DIR / settings.LORA_MODEL_NAME.replace(".safetensors", "") | |
| accelerator.unwrap_model(unet).save_pretrained(final_path) | |
| # Also save as safetensors | |
| from safetensors.torch import save_file | |
| state_dict = accelerator.unwrap_model(unet).state_dict() | |
| safetensors_path = settings.LORA_DIR / settings.LORA_MODEL_NAME | |
| save_file(state_dict, safetensors_path) | |
| logger.info(f"Model saved to: {safetensors_path}") | |
| logger.info("=" * 60) | |
| logger.info("Training Complete!") | |
| logger.info("=" * 60) | |
| except Exception as e: | |
| logger.error(f"Training failed: {e}", exc_info=True) | |
| raise | |
| finally: | |
| # Remove lock file | |
| if lock_file and lock_file.exists(): | |
| lock_file.unlink() | |
| logger.info("Training lock removed") | |
| if __name__ == "__main__": | |
| # Test training | |
| train_lora_model(train_steps=100) | |