Files
remote-sensing/train_cloud_removal.py
T

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
)