ArchEnhancer / backend /services /training_engine.py
Aguilar Elizondo
Initial commit: Architecture AI Enhancer v1.0.0
6bfa765
Raw
History Blame Contribute Delete
12.6 kB
"""
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)