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