""" 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 # 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 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""" model.eval() total_loss = 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() return total_loss / len(dataloader) 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" ): """ 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 """ 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 ) # Training loop print("\n🚀 Starting training...") best_val_loss = float('inf') train_losses = [] val_losses = [] 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 = validate(model, val_loader, criterion, device) val_losses.append(val_loss) # Update learning rate scheduler.step(val_loss) print(f"\nEpoch {epoch + 1} Summary:") print(f" Train Loss: {train_loss:.6f}") print(f" Val Loss: {val_loss:.6f}") # Save best model if val_loss < best_val_loss: 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, 'use_s1': use_s1, 'in_channels': in_channels, 'out_channels': out_channels }, save_path) print(f" 💾 Saved best model: {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}") # Plot training curves print("\n📈 Plotting training curves...") fig, ax = plt.subplots(figsize=(10, 6)) ax.plot(train_losses, label='Train Loss') ax.plot(val_losses, label='Val Loss') ax.set_xlabel('Epoch') ax.set_ylabel('Loss (MAE)') ax.set_title('Cloud Removal Training Progress') ax.legend() ax.grid(True) 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"Model saved to: {Path(save_dir) / 'cloud_removal_unet_best.pth'}") print("=" * 70) return model, train_losses, val_losses # ============ 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 )