cck-0702's picture
Clean commit without binary (image) files
c8c00f0
Raw
History Blame Contribute Delete
18.9 kB
import os
import json
import torch
import argparse
import torch.nn.functional as F
import random
import time
from torch.optim import AdamW
from tqdm import tqdm
from diffusers import DDPMScheduler
from model import DefectFillModel, USE_MODELSCOPE
from data_loader import get_data_loaders
from utils import save_checkpoint, load_checkpoint
# TensorBoard support
from torch.utils.tensorboard import SummaryWriter
import datetime
def generate_seed_from_timestamp():
"""Generates a random seed based on the current timestamp"""
return int(time.time() * 1000) % (2**31)
def train(args):
# Set up device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
if torch.cuda.is_available():
print(f"GPU Name: {torch.cuda.get_device_name(0)}")
else:
print("CUDA is NOT available. The model will run on CPU, causing severe slowdown.")
# Create output directory structure
os.makedirs(args.output_dir, exist_ok=True)
checkpoints_dir = os.path.join(args.output_dir, "checkpoints")
tensorboard_dir = os.path.join(args.output_dir, "tensorboard")
os.makedirs(checkpoints_dir, exist_ok=True)
os.makedirs(tensorboard_dir, exist_ok=True)
# Create log file in the output directory
log_file_path = os.path.join(args.output_dir, "train_log.txt")
log_file = open(log_file_path, "a")
log_file.write(f"\n\n{'='*60}\n")
log_file.write(f"Training started at {datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n")
log_file.write(f"Object class: {args.object_class}\n")
log_file.write(f"Defect type: {args.defect_type if args.defect_type else 'all'}\n")
log_file.write(f"Config name: {args.config_name}\n")
log_file.write(f"Lambda defect: {args.lambda_defect}\n")
log_file.write(f"Lambda obj: {args.lambda_obj}\n")
log_file.write(f"Lambda attn: {args.lambda_attn}\n")
log_file.write(f"Alpha (obj branch bg weight): {args.alpha}\n")
log_file.write(f"Gradient accumulation steps: {args.gradient_accumulation_steps}\n")
log_file.write(f"Random seed: {args.seed}\n")
log_file.write(f"{'='*60}\n\n")
# Save training configuration to JSON
train_config = {
"timestamp": datetime.datetime.now().strftime('%Y-%m-%dT%H:%M:%S'),
"object_class": args.object_class,
"defect_type": args.defect_type if args.defect_type else "all",
"config_name": args.config_name,
"lambda_defect": args.lambda_defect,
"lambda_obj": args.lambda_obj,
"lambda_attn": args.lambda_attn,
"alpha": args.alpha,
"batch_size": args.batch_size,
"max_train_steps": args.max_train_steps,
"gradient_accumulation_steps": args.gradient_accumulation_steps,
"lora_rank": args.lora_rank,
"lora_alpha": args.lora_alpha,
"text_encoder_lr": args.text_encoder_lr,
"unet_lr": args.unet_lr,
"lr_warmup_steps": args.lr_warmup_steps,
"save_steps": args.save_steps,
"seed": args.seed
}
config_path = os.path.join(args.output_dir, "train_config.json")
with open(config_path, "w") as f:
json.dump(train_config, f, indent=4)
print(f"Training config saved to: {config_path}")
# Initialize TensorBoard writer
writer = SummaryWriter(tensorboard_dir)
# Load data
train_loader, test_loader = get_data_loaders(
root_dir=args.data_dir,
object_class=args.object_class,
batch_size=args.batch_size,
defect_type=args.defect_type,
dilate_mask=args.dilate_mask,
mask_kernel_size=args.mask_kernel_size
)
# Initialize model
model = DefectFillModel(
device=device,
lora_rank=args.lora_rank,
lora_alpha=args.lora_alpha,
seed=args.seed
)
# Set up optimizers with specific learning rates for Text Encoder and UNet
text_encoder_params = [p for n, p in model.pipeline.text_encoder.named_parameters() if "lora" in n]
unet_params = [p for n, p in model.pipeline.unet.named_parameters() if "lora" in n]
optimizer = AdamW([
{"params": text_encoder_params, "lr": args.text_encoder_lr},
{"params": unet_params, "lr": args.unet_lr}
])
# Save original LRs for warmup calculation
base_lrs = [args.text_encoder_lr, args.unet_lr]
# Set up noise scheduler (handling ModelScope vs HuggingFace)
hf_model_id = "sd2-community/stable-diffusion-2-inpainting"
if USE_MODELSCOPE:
try:
from modelscope import snapshot_download
print(f"[ModelScope] Downloading scheduler: {hf_model_id}")
local_model_path = snapshot_download(hf_model_id)
noise_scheduler = DDPMScheduler.from_pretrained(local_model_path, subfolder="scheduler")
except ImportError:
print("[Warning] modelscope not installed, falling back to HuggingFace...")
noise_scheduler = DDPMScheduler.from_pretrained(hf_model_id, subfolder="scheduler")
else:
noise_scheduler = DDPMScheduler.from_pretrained(hf_model_id, subfolder="scheduler")
# Resume from checkpoint if specified
start_step = 0
if args.resume_from:
start_step = load_checkpoint(model, optimizer, args.resume_from)
print(f"Resuming from step {start_step}")
log_file.write(f"Resuming from step {start_step}\n")
# Set models to training mode
model.pipeline.unet.train()
model.pipeline.text_encoder.train()
total_steps = args.max_train_steps
progress_bar = tqdm(range(start_step, total_steps), desc="Training Progress")
global_step = start_step
accumulation_step = 0
# Add this BEFORE the while loop
with torch.no_grad():
clean_latents_cache = {}
for batch in train_loader:
images = batch["image"].to(device, dtype=torch.float16)
is_defect = batch["is_defect"]
defect_samples = torch.nonzero(is_defect).squeeze(1)
if len(defect_samples) > 0:
defect_images = images[defect_samples]
# Cache original latents
latents = model.pipeline.vae.encode(defect_images).latent_dist.sample() * model.pipeline.vae.config.scaling_factor
clean_latents_cache[tuple(defect_samples.cpu().tolist())] = latents.detach()
while global_step < total_steps:
for batch in train_loader:
if global_step >= total_steps:
break
# t_start = time.time()
# Move data to device
images = batch["image"].to(device, dtype=torch.float16)
masks = batch["mask"].to(device, dtype=torch.float16)
backgrounds = batch["background"].to(device, dtype=torch.float16)
adjusted_masks = batch["adjusted_mask"].to(device, dtype=torch.float16)
is_defect = batch["is_defect"]
# Ensure we only process defective samples
defect_samples = torch.nonzero(is_defect).squeeze(1)
if len(defect_samples) == 0:
continue # Skip batch if no defects present
# Extract defect-only samples
defect_images = images[defect_samples]
defect_masks = masks[defect_samples]
defect_backgrounds = backgrounds[defect_samples]
defect_adjusted_masks = adjusted_masks[defect_samples]
object_classes = [batch["object_class"][i] for i in defect_samples]
# t_data = time.time()
# Extract defect type from file path for specific prompting
defect_types = []
for i in defect_samples:
if hasattr(train_loader.dataset, 'images') and i < len(train_loader.dataset.images):
img_path = train_loader.dataset.images[i]
parts = img_path.split(os.sep)
for j, part in enumerate(parts):
if part == "defective" and j + 1 < len(parts):
defect_types.append(parts[j + 1])
break
else:
defect_types.append("defect")
else:
defect_types.append("defect")
# Learning rate warmup
if global_step < args.lr_warmup_steps:
lr_scale = min(1.0, (global_step + 1) / args.lr_warmup_steps)
for i, param_group in enumerate(optimizer.param_groups):
param_group["lr"] = base_lrs[i] * lr_scale
# Reset attention maps
if hasattr(model, 'attention_maps'):
model.attention_maps = {}
# ========== PHASE 1: Defect Branch (Defect Texture Learning) ==========
# Using <defect> learnable token for concept isolation
# t_text = time.time()
defect_prompts = [f"A photo of {model.placeholder_token}" for _ in range(len(defect_samples))]
text_embeddings = model.get_text_embeddings(defect_prompts, enable_grad=True)
# ========== PHASE 1 & 2 OPTIMIZED VAE & MASK PASS ==========
# 1. Single VAE encoding for original defect images
with torch.no_grad():
# Instead of re-encoding, load from cache
latents = clean_latents_cache[tuple(defect_samples.cpu().tolist())]
# t_vae = time.time()
if len(defect_masks.shape) == 3:
defect_masks = defect_masks.unsqueeze(1)
# 2. Fully Vectorized Random Mask Generation directly on GPU (0 Python Loops)
B, _, H, W = defect_images.shape
K = 15 # Number of random boxes
# Create broadcastable coordinate grids (1, 1, H, 1) and (1, 1, 1, W)
grid_y = torch.arange(H, device=device).view(1, 1, H, 1)
grid_x = torch.arange(W, device=device).view(1, 1, 1, W)
# Sample all box parameters for all batch items and box counts simultaneously: shape (B, K, 1, 1)
rh = torch.randint(max(1, int(H * 0.03)), max(2, int(H * 0.25)), (B, K, 1, 1), device=device)
rw = torch.randint(max(1, int(W * 0.03)), max(2, int(W * 0.25)), (B, K, 1, 1), device=device)
ry = torch.randint(0, max(1, H - 20), (B, K, 1, 1), device=device)
rx = torch.randint(0, max(1, W - 20), (B, K, 1, 1), device=device)
# Compute rectangle coverage across all (B, K, H, W) elements in a single broadcasted operation
in_box = (grid_y >= ry) & (grid_y < (ry + rh)) & (grid_x >= rx) & (grid_x < (rx + rw))
# Merge all K boxes per image using logical OR (.any), shape -> (B, 1, H, W)
random_masks = in_box.any(dim=1, keepdim=True).to(dtype=defect_images.dtype)
# 3. Batched VAE Encoding (Combines Phase 1 and Phase 2 into 1 pass)
with torch.no_grad():
masked_images = defect_images * (1 - defect_masks)
rand_masked_images = defect_images * (1 - random_masks)
combined_masked = torch.cat([masked_images, rand_masked_images], dim=0)
combined_latents = model.pipeline.vae.encode(combined_masked).latent_dist.sample() * model.pipeline.vae.config.scaling_factor
masked_image_latents, random_masked_image_latents = torch.chunk(combined_latents, 2, dim=0)
# 4. Latent interpolation for masks
mask_latents = F.interpolate(defect_masks, size=(latents.shape[2], latents.shape[3]))
random_mask_latents = F.interpolate(random_masks, size=(latents.shape[2], latents.shape[3]))
# ========== Phase 1 Forward Pass ==========
noise = torch.randn_like(latents)
timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (latents.shape[0],), device=device)
noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)
outputs = model(
noisy_latents=noisy_latents,
masked_image_latents=masked_image_latents,
mask_latents=mask_latents,
timesteps=timesteps,
encoder_hidden_states=text_embeddings
)
# t_forward = time.time()
# print(f"[TIMING RESULTS]")
# print(f" β”œβ”€ Data Fetch & Filter : {t_data - t_start:.2f}s")
# print(f" β”œβ”€ Text Encoding : {t_text - t_data:.2f}s")
# print(f" β”œβ”€ VAE Encoding : {t_vae - t_text:.2f}s")
# print(f" └─ UNet Forward & Loss : {t_forward - t_vae:.2f}s")
noise_pred = outputs["noise_pred"]
defect_loss = model.compute_defect_loss(noise_pred, noise, mask_latents)
attention_loss = outputs.get("attention_loss", torch.tensor(0.0, device=device))
# ========== Phase 2 Forward Pass ==========
obj_prompts = [f"A {obj_class} with {model.placeholder_token}" for obj_class in object_classes]
obj_text_embeddings = model.get_text_embeddings(obj_prompts, enable_grad=True)
obj_noise = torch.randn_like(latents)
obj_timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (latents.shape[0],), device=device)
obj_noisy_latents = noise_scheduler.add_noise(latents, obj_noise, obj_timesteps)
obj_outputs = model(
noisy_latents=obj_noisy_latents,
masked_image_latents=random_masked_image_latents,
mask_latents=random_mask_latents,
timesteps=obj_timesteps,
encoder_hidden_states=obj_text_embeddings
)
object_loss = model.compute_object_loss(obj_outputs["noise_pred"], obj_noise, random_mask_latents, alpha=args.alpha)
# Total Loss Calculation
total_loss = args.lambda_defect * defect_loss + args.lambda_obj * object_loss + args.lambda_attn * attention_loss
total_loss = total_loss / args.gradient_accumulation_steps
# NaN Check
if torch.isnan(total_loss):
print(f"Warning: NaN loss detected at step {global_step}")
log_file.write(f"Warning: NaN loss at step {global_step}\n")
optimizer.zero_grad()
continue
# t_before_backward = time.time()
total_loss.backward()
# t_after_backward = time.time()
# print(f" └─ Backward Pass Time : {t_after_backward - t_before_backward:.2f}s")
accumulation_step += 1
# Optimization step
if accumulation_step >= args.gradient_accumulation_steps:
optimizer.step()
optimizer.zero_grad()
accumulation_step = 0
progress_bar.update(1)
global_step += 1
# Logging and TensorBoard updates
writer.add_scalar("Loss/Defect", defect_loss.item(), global_step)
writer.add_scalar("Loss/Object", object_loss.item(), global_step)
writer.add_scalar("Loss/Attention", attention_loss.item(), global_step)
writer.add_scalar("Loss/Total", total_loss.item() * args.gradient_accumulation_steps, global_step)
if global_step % 10 == 0:
for i, param_group in enumerate(optimizer.param_groups):
writer.add_scalar(f"LearningRate/group{i}", param_group["lr"], global_step)
# Periodic checkpointing
if global_step % args.save_steps == 0 or global_step == total_steps:
checkpoint_path = os.path.join(checkpoints_dir, f"checkpoint_{global_step}.pt")
save_checkpoint(model, optimizer, global_step, checkpoint_path)
log_file.write(f"Checkpoint saved at step {global_step}\n")
# Save final model
final_checkpoint_path = os.path.join(checkpoints_dir, "checkpoint_final.pt")
save_checkpoint(model, optimizer, global_step, final_checkpoint_path)
print(f"Final model saved to: {final_checkpoint_path}")
writer.close()
log_file.close()
return model
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Train DefectFill model")
# Paths
parser.add_argument("--data_dir", type=str, required=True, help="Path to MVTec AD dataset")
parser.add_argument("--object_class", type=str, required=True, help="Object class to train on")
parser.add_argument("--output_dir", type=str, default="./output", help="Directory to save models")
# Loss Weights
parser.add_argument("--lambda_defect", type=float, default=0.5, help="Defect loss weight (L_def)")
parser.add_argument("--lambda_obj", type=float, default=0.2, help="Object integrity loss weight (L_obj)")
parser.add_argument("--lambda_attn", type=float, default=0.05, help="Attention loss weight (L_attn)")
parser.add_argument("--alpha", type=float, default=0.3, help="Background weight for object branch")
# Training Config
parser.add_argument("--config_name", type=str, default="base", help="Experiment name (base/tex/obj)")
parser.add_argument("--defect_type", type=str, default=None, help="Specific defect type to train on")
parser.add_argument("--batch_size", type=int, default=2)
parser.add_argument("--lora_rank", type=int, default=8)
parser.add_argument("--lora_alpha", type=int, default=16)
parser.add_argument("--text_encoder_lr", type=float, default=4e-5)
parser.add_argument("--unet_lr", type=float, default=2e-4)
parser.add_argument("--max_train_steps", type=int, default=2000)
parser.add_argument("--lr_warmup_steps", type=int, default=100)
parser.add_argument("--save_steps", type=int, default=500)
parser.add_argument("--seed", type=int, default=-1, help="-1 for timestamp-based seed")
parser.add_argument("--resume_from", type=str, default=None)
parser.add_argument("--gradient_accumulation_steps", type=int, default=2)
parser.add_argument("--dilate_mask", type=str, default="False", help="Whether to dilate masks (True/False)")
parser.add_argument("--mask_kernel_size", type=int, default=3, help="Size of dilation kernel (must be odd, e.g. 3, 5, 7)")
args = parser.parse_args()
# Helper to convert string "True" to boolean
args.dilate_mask = args.dilate_mask.lower() == "true"
# Seed setup
if args.seed == -1:
args.seed = generate_seed_from_timestamp()
torch.manual_seed(args.seed)
random.seed(args.seed)
train(args)