""" 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)