678 lines
23 KiB
Python
Executable File
678 lines
23 KiB
Python
Executable File
"""
|
|
Train Cloud Removal Model using SEN12MS-CR Dataset
|
|
Huấn luyện model Deep Learning để khử mây từ ảnh Sentinel-2
|
|
|
|
Dataset: SEN12MS-CR (Sentinel-12 Multi-Seasonal Cloud Removal)
|
|
- Input: S2 cloudy images (ảnh Sentinel-2 bị mây)
|
|
- Target: S2 clean images (ảnh Sentinel-2 sạch)
|
|
- Optional: S1 SAR data (radar data không bị ảnh hưởng bởi mây)
|
|
"""
|
|
|
|
import os
|
|
import sys
|
|
import numpy as np
|
|
import torch
|
|
import torch.nn as nn
|
|
import torch.optim as optim
|
|
from torch.utils.data import Dataset, DataLoader
|
|
from pathlib import Path
|
|
import matplotlib.pyplot as plt
|
|
from tqdm import tqdm
|
|
import math
|
|
from torch.utils.tensorboard import SummaryWriter
|
|
|
|
# Add winter_dataset to path
|
|
sys.path.insert(0, str(Path(__file__).parent / "winter_dataset"))
|
|
from sen12ms_cr_dataLoader import SEN12MSCRDataset, Seasons, S1Bands, S2Bands
|
|
|
|
|
|
# ============ DATASET WRAPPER ============
|
|
|
|
class CloudRemovalDataset(Dataset):
|
|
"""
|
|
PyTorch Dataset wrapper cho SEN12MS-CR
|
|
Input: S2 cloudy + S1 (optional)
|
|
Target: S2 clean
|
|
"""
|
|
|
|
def __init__(self, base_dir, season=Seasons.WINTER, use_s1=True,
|
|
s2_bands=S2Bands.ALL, normalize=True):
|
|
"""
|
|
Args:
|
|
base_dir: Đường dẫn đến thư mục chứa dữ liệu
|
|
season: Mùa (SPRING, SUMMER, FALL, WINTER)
|
|
use_s1: Có sử dụng dữ liệu S1 (radar) không
|
|
s2_bands: Các band S2 cần dùng
|
|
normalize: Normalize dữ liệu về [0, 1]
|
|
"""
|
|
self.dataset = SEN12MSCRDataset(base_dir)
|
|
self.season = season
|
|
self.use_s1 = use_s1
|
|
self.s2_bands = s2_bands
|
|
self.normalize = normalize
|
|
|
|
# Lấy tất cả scene và patch IDs
|
|
season_ids = self.dataset.get_season_ids(season)
|
|
|
|
# Tạo list of (scene_id, patch_id) pairs
|
|
self.samples = []
|
|
for scene_id, patch_ids in season_ids.items():
|
|
for patch_id in patch_ids:
|
|
self.samples.append((scene_id, patch_id))
|
|
|
|
# Get band count
|
|
n_s2_bands = len(s2_bands.value) if hasattr(s2_bands, 'value') else len(s2_bands)
|
|
|
|
print(f"[DATASET] Loaded {len(self.samples)} samples from {season.value}")
|
|
print(f"[DATASET] Use S1: {use_s1}, S2 bands: {n_s2_bands}")
|
|
|
|
def __len__(self):
|
|
return len(self.samples)
|
|
|
|
def __getitem__(self, idx):
|
|
scene_id, patch_id = self.samples[idx]
|
|
|
|
# Load triplet: S1, S2 clean, S2 cloudy
|
|
s1, s2_clean, s2_cloudy, bounds = self.dataset.get_s1s2s2cloudy_triplet(
|
|
self.season,
|
|
scene_id,
|
|
patch_id,
|
|
s1_bands=S1Bands.ALL if self.use_s1 else S1Bands.NONE,
|
|
s2_bands=self.s2_bands,
|
|
s2cloudy_bands=self.s2_bands
|
|
)
|
|
|
|
# Normalize to [0, 1] if needed
|
|
if self.normalize:
|
|
s2_clean = s2_clean.astype(np.float32) / 10000.0 # S2 values are in [0, 10000]
|
|
s2_cloudy = s2_cloudy.astype(np.float32) / 10000.0
|
|
if self.use_s1:
|
|
# S1 values need different normalization (dB scale)
|
|
s1 = (s1.astype(np.float32) + 30) / 50.0 # Normalize from [-30, 20] to [0, 1]
|
|
s1 = np.clip(s1, 0, 1)
|
|
|
|
# Convert to torch tensors
|
|
s2_clean = torch.from_numpy(s2_clean).float()
|
|
s2_cloudy = torch.from_numpy(s2_cloudy).float()
|
|
|
|
# Input: S2 cloudy + S1 (if enabled)
|
|
if self.use_s1:
|
|
s1 = torch.from_numpy(s1).float()
|
|
input_data = torch.cat([s2_cloudy, s1], dim=0)
|
|
else:
|
|
input_data = s2_cloudy
|
|
|
|
return input_data, s2_clean
|
|
|
|
|
|
# ============ U-NET ARCHITECTURE ============
|
|
|
|
class DoubleConv(nn.Module):
|
|
"""(Conv2d -> BatchNorm -> ReLU) x 2"""
|
|
|
|
def __init__(self, in_channels, out_channels):
|
|
super().__init__()
|
|
self.double_conv = nn.Sequential(
|
|
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
|
|
nn.BatchNorm2d(out_channels),
|
|
nn.ReLU(inplace=True),
|
|
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
|
|
nn.BatchNorm2d(out_channels),
|
|
nn.ReLU(inplace=True)
|
|
)
|
|
|
|
def forward(self, x):
|
|
return self.double_conv(x)
|
|
|
|
|
|
class UNet(nn.Module):
|
|
"""
|
|
U-Net architecture cho cloud removal
|
|
Input: S2 cloudy (+ S1 optional) [B, C_in, H, W]
|
|
Output: S2 clean [B, C_out, H, W]
|
|
"""
|
|
|
|
def __init__(self, in_channels, out_channels, features=[64, 128, 256, 512]):
|
|
super().__init__()
|
|
self.encoder = nn.ModuleList()
|
|
self.decoder = nn.ModuleList()
|
|
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
|
|
|
|
# Encoder (downsampling)
|
|
for feature in features:
|
|
self.encoder.append(DoubleConv(in_channels, feature))
|
|
in_channels = feature
|
|
|
|
# Bottleneck
|
|
self.bottleneck = DoubleConv(features[-1], features[-1] * 2)
|
|
|
|
# Decoder (upsampling)
|
|
for feature in reversed(features):
|
|
self.decoder.append(
|
|
nn.ConvTranspose2d(feature * 2, feature, kernel_size=2, stride=2)
|
|
)
|
|
self.decoder.append(DoubleConv(feature * 2, feature))
|
|
|
|
# Final output layer
|
|
self.final_conv = nn.Conv2d(features[0], out_channels, kernel_size=1)
|
|
|
|
def forward(self, x):
|
|
skip_connections = []
|
|
|
|
# Encoder
|
|
for encode in self.encoder:
|
|
x = encode(x)
|
|
skip_connections.append(x)
|
|
x = self.pool(x)
|
|
|
|
# Bottleneck
|
|
x = self.bottleneck(x)
|
|
|
|
# Decoder
|
|
skip_connections = skip_connections[::-1]
|
|
|
|
for idx in range(0, len(self.decoder), 2):
|
|
x = self.decoder[idx](x) # Upsample
|
|
skip_connection = skip_connections[idx // 2]
|
|
|
|
# Handle size mismatch
|
|
if x.shape != skip_connection.shape:
|
|
x = nn.functional.interpolate(x, size=skip_connection.shape[2:])
|
|
|
|
concat_skip = torch.cat((skip_connection, x), dim=1)
|
|
x = self.decoder[idx + 1](concat_skip) # Double conv
|
|
|
|
return self.final_conv(x)
|
|
|
|
|
|
# ============ TRAINING FUNCTIONS ============
|
|
|
|
def calculate_psnr(img1, img2, max_value=1.0):
|
|
"""
|
|
Tính PSNR (Peak Signal-to-Noise Ratio) giữa 2 ảnh
|
|
|
|
Args:
|
|
img1: Tensor [B, C, H, W] hoặc [C, H, W]
|
|
img2: Tensor [B, C, H, W] hoặc [C, H, W]
|
|
max_value: Giá trị pixel tối đa (thường là 1.0 cho ảnh đã normalize)
|
|
|
|
Returns:
|
|
psnr_value: PSNR tính bằng dB
|
|
"""
|
|
mse = torch.mean((img1 - img2) ** 2)
|
|
if mse == 0:
|
|
return float('inf')
|
|
psnr = 20 * math.log10(max_value) - 10 * torch.log10(mse)
|
|
return psnr.item()
|
|
|
|
|
|
def calculate_baseline_psnr(dataloader, device):
|
|
"""
|
|
Tính Cận dưới (Baseline PSNR) - PSNR giữa ảnh cloudy và ảnh clean
|
|
Đây là giá trị tham chiếu để đánh giá model có tốt hơn không làm gì không.
|
|
|
|
Args:
|
|
dataloader: Validation dataloader
|
|
device: 'cuda' hoặc 'cpu'
|
|
|
|
Returns:
|
|
baseline_psnr: PSNR trung bình giữa input cloudy và target clean
|
|
"""
|
|
total_psnr = 0
|
|
num_batches = 0
|
|
|
|
print("\n📏 Đang tính Cận dưới (Baseline PSNR)...")
|
|
print(" Baseline = PSNR(s2_cloudy, s2_clean)")
|
|
|
|
with torch.no_grad():
|
|
for inputs, targets in tqdm(dataloader, desc="Calculating Baseline"):
|
|
inputs = inputs.to(device)
|
|
targets = targets.to(device)
|
|
|
|
# Lấy phần S2 cloudy từ input
|
|
# Input có thể là [S2_cloudy] hoặc [S2_cloudy, S1]
|
|
# Output channels = target channels
|
|
num_s2_bands = targets.shape[1]
|
|
s2_cloudy = inputs[:, :num_s2_bands, :, :]
|
|
|
|
# Tính PSNR giữa cloudy và clean
|
|
psnr = calculate_psnr(s2_cloudy, targets)
|
|
total_psnr += psnr
|
|
num_batches += 1
|
|
|
|
baseline_psnr = total_psnr / num_batches
|
|
print(f"\n✅ Cận dưới (Baseline PSNR): {baseline_psnr:.2f} dB")
|
|
print(f" → Đây là 'thanh thước đo' - model phải vượt qua giá trị này!")
|
|
|
|
return baseline_psnr
|
|
|
|
|
|
def train_epoch(model, dataloader, criterion, optimizer, device):
|
|
"""Train for one epoch"""
|
|
model.train()
|
|
total_loss = 0
|
|
|
|
pbar = tqdm(dataloader, desc="Training")
|
|
for batch_idx, (inputs, targets) in enumerate(pbar):
|
|
inputs = inputs.to(device)
|
|
targets = targets.to(device)
|
|
|
|
# Forward pass
|
|
optimizer.zero_grad()
|
|
outputs = model(inputs)
|
|
loss = criterion(outputs, targets)
|
|
|
|
# Backward pass
|
|
loss.backward()
|
|
optimizer.step()
|
|
|
|
total_loss += loss.item()
|
|
pbar.set_postfix({'loss': loss.item()})
|
|
|
|
return total_loss / len(dataloader)
|
|
|
|
|
|
def validate(model, dataloader, criterion, device):
|
|
"""Validate model và tính PSNR"""
|
|
model.eval()
|
|
total_loss = 0
|
|
total_psnr = 0
|
|
|
|
with torch.no_grad():
|
|
for inputs, targets in tqdm(dataloader, desc="Validation"):
|
|
inputs = inputs.to(device)
|
|
targets = targets.to(device)
|
|
|
|
outputs = model(inputs)
|
|
loss = criterion(outputs, targets)
|
|
total_loss += loss.item()
|
|
|
|
# Tính PSNR giữa output (fake_s2) và target (s2_clean)
|
|
psnr = calculate_psnr(outputs, targets)
|
|
total_psnr += psnr
|
|
|
|
avg_loss = total_loss / len(dataloader)
|
|
avg_psnr = total_psnr / len(dataloader)
|
|
|
|
return avg_loss, avg_psnr
|
|
|
|
|
|
def visualize_results(model, dataset, device, num_samples=3):
|
|
"""Visualize cloud removal results"""
|
|
model.eval()
|
|
|
|
fig, axes = plt.subplots(num_samples, 3, figsize=(15, 5 * num_samples))
|
|
|
|
with torch.no_grad():
|
|
for i in range(num_samples):
|
|
idx = np.random.randint(0, len(dataset))
|
|
input_data, target = dataset[idx]
|
|
|
|
input_data = input_data.unsqueeze(0).to(device)
|
|
output = model(input_data)
|
|
|
|
# Convert to numpy
|
|
input_rgb = input_data[0, :3, :, :].cpu().numpy().transpose(1, 2, 0)
|
|
target_rgb = target[:3, :, :].cpu().numpy().transpose(1, 2, 0)
|
|
output_rgb = output[0, :3, :, :].cpu().numpy().transpose(1, 2, 0)
|
|
|
|
# Clip to [0, 1]
|
|
input_rgb = np.clip(input_rgb * 3, 0, 1) # Enhance for visualization
|
|
target_rgb = np.clip(target_rgb * 3, 0, 1)
|
|
output_rgb = np.clip(output_rgb * 3, 0, 1)
|
|
|
|
if num_samples == 1:
|
|
axes[0].imshow(input_rgb)
|
|
axes[0].set_title("Input (Cloudy)")
|
|
axes[0].axis('off')
|
|
|
|
axes[1].imshow(output_rgb)
|
|
axes[1].set_title("Output (Predicted)")
|
|
axes[1].axis('off')
|
|
|
|
axes[2].imshow(target_rgb)
|
|
axes[2].set_title("Target (Clean)")
|
|
axes[2].axis('off')
|
|
else:
|
|
axes[i, 0].imshow(input_rgb)
|
|
axes[i, 0].set_title(f"Sample {i+1}: Input (Cloudy)")
|
|
axes[i, 0].axis('off')
|
|
|
|
axes[i, 1].imshow(output_rgb)
|
|
axes[i, 1].set_title(f"Sample {i+1}: Output (Predicted)")
|
|
axes[i, 1].axis('off')
|
|
|
|
axes[i, 2].imshow(target_rgb)
|
|
axes[i, 2].set_title(f"Sample {i+1}: Target (Clean)")
|
|
axes[i, 2].axis('off')
|
|
|
|
plt.tight_layout()
|
|
return fig
|
|
|
|
|
|
# ============ MAIN TRAINING SCRIPT ============
|
|
|
|
def train_cloud_removal_model(
|
|
data_dir="winter_dataset",
|
|
use_s1=True,
|
|
batch_size=8,
|
|
num_epochs=50,
|
|
learning_rate=1e-4,
|
|
device="cuda" if torch.cuda.is_available() else "cpu",
|
|
save_dir="model_train",
|
|
status_callback=None
|
|
):
|
|
"""
|
|
Train cloud removal model
|
|
|
|
Args:
|
|
data_dir: Thư mục chứa dữ liệu SEN12MS-CR
|
|
use_s1: Có sử dụng S1 radar data không
|
|
batch_size: Batch size
|
|
num_epochs: Số epochs
|
|
learning_rate: Learning rate
|
|
device: 'cuda' hoặc 'cpu'
|
|
save_dir: Thư mục lưu model
|
|
status_callback: Callback function(epoch, train_loss, val_loss, val_psnr, baseline_psnr)
|
|
"""
|
|
|
|
print("=" * 70)
|
|
print("🌥️ CLOUD REMOVAL MODEL TRAINING")
|
|
print("=" * 70)
|
|
print(f"Data directory: {data_dir}")
|
|
print(f"Use S1 (SAR): {use_s1}")
|
|
print(f"Device: {device}")
|
|
print(f"Batch size: {batch_size}")
|
|
print(f"Epochs: {num_epochs}")
|
|
print(f"Learning rate: {learning_rate}")
|
|
print("=" * 70)
|
|
|
|
# Create dataset
|
|
print("\n📂 Loading dataset...")
|
|
|
|
# Use RGB + NIR bands for training (B02, B03, B04, B08)
|
|
s2_bands = [S2Bands.B02, S2Bands.B03, S2Bands.B04, S2Bands.B08]
|
|
|
|
dataset = CloudRemovalDataset(
|
|
base_dir=data_dir,
|
|
season=Seasons.WINTER,
|
|
use_s1=use_s1,
|
|
s2_bands=s2_bands,
|
|
normalize=True
|
|
)
|
|
|
|
# Split train/val
|
|
train_size = int(0.8 * len(dataset))
|
|
val_size = len(dataset) - train_size
|
|
train_dataset, val_dataset = torch.utils.data.random_split(
|
|
dataset, [train_size, val_size]
|
|
)
|
|
|
|
print(f"Train samples: {len(train_dataset)}")
|
|
print(f"Val samples: {len(val_dataset)}")
|
|
|
|
# Create dataloaders
|
|
train_loader = DataLoader(
|
|
train_dataset,
|
|
batch_size=batch_size,
|
|
shuffle=True,
|
|
num_workers=4,
|
|
pin_memory=True if device == "cuda" else False
|
|
)
|
|
|
|
val_loader = DataLoader(
|
|
val_dataset,
|
|
batch_size=batch_size,
|
|
shuffle=False,
|
|
num_workers=4,
|
|
pin_memory=True if device == "cuda" else False
|
|
)
|
|
|
|
# Create model
|
|
print("\n🏗️ Creating U-Net model...")
|
|
in_channels = len(s2_bands) + (2 if use_s1 else 0) # S2 + S1 (VV, VH)
|
|
out_channels = len(s2_bands)
|
|
|
|
model = UNet(in_channels=in_channels, out_channels=out_channels)
|
|
model = model.to(device)
|
|
|
|
print(f"Input channels: {in_channels}")
|
|
print(f"Output channels: {out_channels}")
|
|
print(f"Model parameters: {sum(p.numel() for p in model.parameters()):,}")
|
|
|
|
# Loss and optimizer
|
|
criterion = nn.L1Loss() # MAE loss
|
|
optimizer = optim.Adam(model.parameters(), lr=learning_rate)
|
|
scheduler = optim.lr_scheduler.ReduceLROnPlateau(
|
|
optimizer, mode='min', factor=0.5, patience=5
|
|
)
|
|
|
|
# Tính Baseline PSNR trước khi bắt đầu train
|
|
print("\n" + "="*70)
|
|
print("📏 TÍNH CẬN DƯỚI (BASELINE) - CHỈ TÍNH 1 LẦN DUY NHẤT")
|
|
print("="*70)
|
|
baseline_psnr = calculate_baseline_psnr(val_loader, device)
|
|
print("="*70)
|
|
|
|
# Setup Tensorboard
|
|
writer = SummaryWriter(log_dir=str(Path(save_dir) / "runs" / "cloud_removal"))
|
|
print(f"\n📊 Tensorboard logging: {Path(save_dir) / 'runs' / 'cloud_removal'}")
|
|
print(f" Chạy lệnh: tensorboard --logdir={Path(save_dir) / 'runs'}")
|
|
|
|
# Training loop
|
|
print("\n🚀 Starting training...")
|
|
best_val_loss = float('inf')
|
|
best_val_psnr = 0
|
|
train_losses = []
|
|
val_losses = []
|
|
val_psnrs = []
|
|
|
|
for epoch in range(num_epochs):
|
|
print(f"\n{'='*70}")
|
|
print(f"Epoch {epoch + 1}/{num_epochs}")
|
|
print(f"{'='*70}")
|
|
|
|
# Train
|
|
train_loss = train_epoch(model, train_loader, criterion, optimizer, device)
|
|
train_losses.append(train_loss)
|
|
|
|
# Validate
|
|
val_loss, val_psnr = validate(model, val_loader, criterion, device)
|
|
val_losses.append(val_loss)
|
|
val_psnrs.append(val_psnr)
|
|
|
|
# Update learning rate
|
|
scheduler.step(val_loss)
|
|
|
|
# Tensorboard logging
|
|
writer.add_scalar('Loss/train', train_loss, epoch)
|
|
writer.add_scalar('Loss/val', val_loss, epoch)
|
|
writer.add_scalar('PSNR/model', val_psnr, epoch)
|
|
writer.add_scalar('PSNR/baseline', baseline_psnr, epoch)
|
|
writer.add_scalar('PSNR/improvement', val_psnr - baseline_psnr, epoch)
|
|
writer.add_scalar('Learning_Rate', optimizer.param_groups[0]['lr'], epoch)
|
|
|
|
# Tính độ cải thiện so với baseline
|
|
improvement = val_psnr - baseline_psnr
|
|
improvement_percent = (improvement / baseline_psnr) * 100 if baseline_psnr > 0 else 0
|
|
|
|
print(f"\nEpoch {epoch + 1} Summary:")
|
|
print(f" Train Loss: {train_loss:.6f}")
|
|
print(f" Val Loss: {val_loss:.6f}")
|
|
print(f" ")
|
|
print(f" 📊 PSNR Comparison:")
|
|
print(f" ├─ Model PSNR: {val_psnr:.2f} dB")
|
|
print(f" ├─ Baseline PSNR: {baseline_psnr:.2f} dB (cận dưới)")
|
|
print(f" └─ Improvement: {improvement:+.2f} dB ({improvement_percent:+.1f}%)")
|
|
if improvement > 0:
|
|
print(f" ✅ Model đang TỐT HƠN baseline!")
|
|
else:
|
|
print(f" ⚠️ Model chưa vượt qua baseline")
|
|
|
|
# Call status callback if provided
|
|
if status_callback:
|
|
try:
|
|
status_callback(epoch + 1, train_loss, val_loss, val_psnr, baseline_psnr)
|
|
except Exception as cb_err:
|
|
print(f"[WARNING] Status callback error: {cb_err}")
|
|
|
|
# Save best model (based on PSNR)
|
|
if val_psnr > best_val_psnr:
|
|
best_val_psnr = val_psnr
|
|
best_val_loss = val_loss
|
|
save_path = Path(save_dir) / "cloud_removal_unet_best.pth"
|
|
save_path.parent.mkdir(exist_ok=True)
|
|
|
|
torch.save({
|
|
'epoch': epoch,
|
|
'model_state_dict': model.state_dict(),
|
|
'optimizer_state_dict': optimizer.state_dict(),
|
|
'train_loss': train_loss,
|
|
'val_loss': val_loss,
|
|
'val_psnr': val_psnr,
|
|
'baseline_psnr': baseline_psnr,
|
|
'use_s1': use_s1,
|
|
'in_channels': in_channels,
|
|
'out_channels': out_channels,
|
|
# Hyperparameters for reproducibility
|
|
'hyperparameters': {
|
|
'data_dir': data_dir,
|
|
'use_s1': use_s1,
|
|
'batch_size': batch_size,
|
|
'num_epochs': num_epochs,
|
|
'learning_rate': learning_rate,
|
|
'device': device,
|
|
'model_type': 'unet',
|
|
'architecture': 'U-Net',
|
|
'optimizer': 'Adam',
|
|
'criterion': 'L1Loss',
|
|
'scheduler': 'ReduceLROnPlateau',
|
|
'train_size': len(train_dataset),
|
|
'val_size': len(val_dataset),
|
|
'test_size': len(val_dataset) / (len(train_dataset) + len(val_dataset)),
|
|
's2_bands': len(s2_bands),
|
|
'feature_mode': 'image',
|
|
'random_seed': 42 # Default seed
|
|
}
|
|
}, save_path)
|
|
|
|
print(f" 💾 Saved best model (highest PSNR): {save_path}")
|
|
|
|
# Visualize every 10 epochs
|
|
if (epoch + 1) % 10 == 0:
|
|
print("\n📊 Generating visualizations...")
|
|
fig = visualize_results(model, val_dataset, device, num_samples=3)
|
|
|
|
viz_path = Path(save_dir) / f"cloud_removal_epoch_{epoch+1}.png"
|
|
fig.savefig(viz_path, dpi=150, bbox_inches='tight')
|
|
plt.close(fig)
|
|
|
|
print(f" 💾 Saved visualization: {viz_path}")
|
|
|
|
# Close Tensorboard writer
|
|
writer.close()
|
|
|
|
# Plot training curves
|
|
print("\n📈 Plotting training curves...")
|
|
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))
|
|
|
|
# Loss curves
|
|
ax1.plot(train_losses, label='Train Loss', linewidth=2)
|
|
ax1.plot(val_losses, label='Val Loss', linewidth=2)
|
|
ax1.set_xlabel('Epoch', fontsize=12)
|
|
ax1.set_ylabel('Loss (MAE)', fontsize=12)
|
|
ax1.set_title('Loss Curves', fontsize=14, fontweight='bold')
|
|
ax1.legend(fontsize=10)
|
|
ax1.grid(True, alpha=0.3)
|
|
|
|
# PSNR curves với baseline
|
|
epochs = range(len(val_psnrs))
|
|
ax2.plot(epochs, val_psnrs, label='Model PSNR', linewidth=2, color='green')
|
|
ax2.axhline(y=baseline_psnr, color='red', linestyle='--', linewidth=2,
|
|
label=f'Baseline PSNR ({baseline_psnr:.2f} dB)')
|
|
ax2.fill_between(epochs, baseline_psnr, val_psnrs,
|
|
where=[p > baseline_psnr for p in val_psnrs],
|
|
alpha=0.3, color='green', label='Model > Baseline')
|
|
ax2.fill_between(epochs, baseline_psnr, val_psnrs,
|
|
where=[p <= baseline_psnr for p in val_psnrs],
|
|
alpha=0.3, color='red', label='Model < Baseline')
|
|
ax2.set_xlabel('Epoch', fontsize=12)
|
|
ax2.set_ylabel('PSNR (dB)', fontsize=12)
|
|
ax2.set_title('PSNR vs Baseline (Cận dưới)', fontsize=14, fontweight='bold')
|
|
ax2.legend(fontsize=10)
|
|
ax2.grid(True, alpha=0.3)
|
|
|
|
plt.tight_layout()
|
|
curve_path = Path(save_dir) / "training_curves.png"
|
|
fig.savefig(curve_path, dpi=150, bbox_inches='tight')
|
|
plt.close(fig)
|
|
|
|
print(f" 💾 Saved training curves: {curve_path}")
|
|
|
|
# Final summary
|
|
print("\n" + "=" * 70)
|
|
print("✅ TRAINING COMPLETED!")
|
|
print("=" * 70)
|
|
print(f"Best validation loss: {best_val_loss:.6f}")
|
|
print(f"Best validation PSNR: {best_val_psnr:.2f} dB")
|
|
print(f"Baseline PSNR: {baseline_psnr:.2f} dB")
|
|
print(f"Final improvement: {best_val_psnr - baseline_psnr:+.2f} dB")
|
|
print(f"Model saved to: {Path(save_dir) / 'cloud_removal_unet_best.pth'}")
|
|
print(f"Tensorboard logs: {Path(save_dir) / 'runs' / 'cloud_removal'}")
|
|
print("=" * 70)
|
|
|
|
return model, train_losses, val_losses, val_psnrs, baseline_psnr
|
|
|
|
|
|
# ============ INFERENCE FUNCTION ============
|
|
|
|
def apply_cloud_removal(model_path, cloudy_image, s1_data=None, device="cuda"):
|
|
"""
|
|
Áp dụng model để khử mây cho một ảnh
|
|
|
|
Args:
|
|
model_path: Đường dẫn đến model đã train
|
|
cloudy_image: Ảnh S2 bị mây [C, H, W]
|
|
s1_data: Dữ liệu S1 (optional) [2, H, W]
|
|
device: 'cuda' hoặc 'cpu'
|
|
|
|
Returns:
|
|
cleaned_image: Ảnh đã khử mây [C, H, W]
|
|
"""
|
|
# Load model
|
|
checkpoint = torch.load(model_path, map_location=device)
|
|
|
|
model = UNet(
|
|
in_channels=checkpoint['in_channels'],
|
|
out_channels=checkpoint['out_channels']
|
|
)
|
|
model.load_state_dict(checkpoint['model_state_dict'])
|
|
model = model.to(device)
|
|
model.eval()
|
|
|
|
# Prepare input
|
|
input_tensor = torch.from_numpy(cloudy_image).float().unsqueeze(0).to(device)
|
|
|
|
if checkpoint['use_s1'] and s1_data is not None:
|
|
s1_tensor = torch.from_numpy(s1_data).float().unsqueeze(0).to(device)
|
|
input_tensor = torch.cat([input_tensor, s1_tensor], dim=1)
|
|
|
|
# Inference
|
|
with torch.no_grad():
|
|
output = model(input_tensor)
|
|
|
|
cleaned_image = output[0].cpu().numpy()
|
|
|
|
return cleaned_image
|
|
|
|
|
|
if __name__ == "__main__":
|
|
# Train model
|
|
model, train_losses, val_losses = train_cloud_removal_model(
|
|
data_dir="winter_dataset",
|
|
use_s1=True,
|
|
batch_size=8,
|
|
num_epochs=50,
|
|
learning_rate=1e-4
|
|
)
|