""" Cloud Removal Module - Hệ thống xử lý mây độc lập Cung cấp nhiều phương pháp khử mây cho dữ liệu Sentinel-2 """ import numpy as np import xarray as xr from typing import Tuple, Optional, Dict from sklearn.neighbors import KNeighborsRegressor from sklearn.ensemble import RandomForestRegressor import warnings from pathlib import Path warnings.filterwarnings('ignore') class CloudRemovalStrategy: """Base class cho các chiến lược xử lý mây""" def __init__(self, name: str, description: str): self.name = name self.description = description def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: """ Xử lý mây và trả về dữ liệu đã được làm sạch Returns: Tuple[xr.Dataset, Dict]: (cleaned_data, metadata) """ raise NotImplementedError class ClassicStrategy(CloudRemovalStrategy): """ Chiến lược cổ điển 3 bước: 1. Temporal interpolation (ffill + bfill) 2. Median compositing (nếu >= 3 scenes) 3. Spatial interpolation (nearest neighbor) """ def __init__(self): super().__init__( name="classic", description="3-step classical approach: temporal → median → spatial interpolation" ) def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: metadata = { 'method': self.name, 'steps_applied': [] } # Apply mask for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].where(~cloud_mask) # Step 1: Temporal Interpolation for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].ffill(dim='time').bfill(dim='time') metadata['steps_applied'].append('temporal_interpolation') # Step 2: Median Compositing (if >= 3 time steps) if len(s2_data.time) >= 3: for band in s2_data.data_vars: if band != "SCL": median_composite = s2_data[band].median(dim='time', skipna=True) s2_data[band] = s2_data[band].fillna(median_composite) metadata['steps_applied'].append('median_compositing') # Step 3: Spatial Interpolation for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].interpolate_na(dim='x', method='nearest', fill_value='extrapolate') s2_data[band] = s2_data[band].interpolate_na(dim='y', method='nearest', fill_value='extrapolate') metadata['steps_applied'].append('spatial_interpolation') # Final fallback for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].fillna(0) return s2_data, metadata class NoRemovalStrategy(CloudRemovalStrategy): """Không xử lý mây - giữ nguyên dữ liệu gốc, chỉ fill NaN bằng 0""" def __init__(self): super().__init__( name="none", description="No cloud removal - keep original data with NaN filled as 0" ) def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: metadata = { 'method': self.name, 'steps_applied': ['none'], 'note': 'No cloud removal applied, only NaN filling' } # Chỉ fill NaN bằng 0, không apply cloud mask for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].fillna(0) return s2_data, metadata class TemporalOnlyStrategy(CloudRemovalStrategy): """Chỉ sử dụng temporal interpolation - nhanh nhất, phù hợp khi có nhiều time steps""" def __init__(self): super().__init__( name="temporal_only", description="Temporal interpolation only - fast, good for time series with many scenes" ) def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: metadata = { 'method': self.name, 'steps_applied': ['temporal_interpolation'] } # Apply mask for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].where(~cloud_mask) # Temporal interpolation for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].ffill(dim='time').bfill(dim='time') s2_data[band] = s2_data[band].fillna(0) return s2_data, metadata class MedianCompositeStrategy(CloudRemovalStrategy): """Ưu tiên median composite - tốt nhất cho giảm noise""" def __init__(self): super().__init__( name="median_composite", description="Median composite priority - best for noise reduction" ) def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: metadata = { 'method': self.name, 'steps_applied': ['median_compositing', 'spatial_interpolation'] } # Apply mask for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].where(~cloud_mask) # Direct median composite for band in s2_data.data_vars: if band != "SCL": median_composite = s2_data[band].median(dim='time', skipna=True) # Fill all NaN with median s2_data[band] = s2_data[band].fillna(median_composite) # Spatial interpolation for remaining gaps for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].interpolate_na(dim='x', method='nearest') s2_data[band] = s2_data[band].interpolate_na(dim='y', method='nearest') s2_data[band] = s2_data[band].fillna(0) return s2_data, metadata class MLInpaintingStrategy(CloudRemovalStrategy): """ Machine Learning Inpainting - sử dụng KNN hoặc Random Forest Học từ pixels hợp lệ để dự đoán pixels bị mây """ def __init__(self, ml_model: str = "knn"): """ Args: ml_model: 'knn' hoặc 'rf' (random forest) """ super().__init__( name=f"ml_inpainting_{ml_model}", description=f"ML-based cloud removal using {ml_model.upper()} - learns from valid pixels" ) self.ml_model = ml_model def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: metadata = { 'method': self.name, 'ml_model': self.ml_model, 'steps_applied': [] } # Apply mask for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].where(~cloud_mask) # ML inpainting cho từng time step for time_idx in range(len(s2_data.time)): # Get all bands for this time step bands_data = [] band_names = [] for band in s2_data.data_vars: if band != "SCL": band_data = s2_data[band].isel(time=time_idx).values bands_data.append(band_data.flatten()) band_names.append(band) if not bands_data: continue # Stack bands: shape (n_pixels, n_bands) X_all = np.column_stack(bands_data) # Find valid (non-NaN) and invalid (NaN) pixels valid_mask = ~np.isnan(X_all).any(axis=1) if valid_mask.sum() < 10: # Not enough training data continue X_valid = X_all[valid_mask] X_invalid_indices = np.where(~valid_mask)[0] if len(X_invalid_indices) == 0: # No clouds continue # Prepare features: use spatial coordinates + spectral values y_coords, x_coords = np.meshgrid( np.arange(s2_data.dims['y']), np.arange(s2_data.dims['x']), indexing='ij' ) coords_flat = np.column_stack([y_coords.flatten(), x_coords.flatten()]) # Train ML model on valid pixels X_train = coords_flat[valid_mask] y_train = X_valid try: if self.ml_model == "knn": model = KNeighborsRegressor(n_neighbors=min(5, len(X_train)), weights='distance') else: # random forest model = RandomForestRegressor(n_estimators=10, max_depth=10, random_state=42, n_jobs=-1) model.fit(X_train, y_train) # Predict invalid pixels X_test = coords_flat[X_invalid_indices] predictions = model.predict(X_test) # Fill predictions back X_all[X_invalid_indices] = predictions # Reshape and update dataset for band_idx, band in enumerate(band_names): filled_data = X_all[:, band_idx].reshape(s2_data.dims['y'], s2_data.dims['x']) s2_data[band].values[time_idx] = filled_data metadata['steps_applied'].append(f'ml_inpainting_time_{time_idx}') except Exception as e: print(f"[ML INPAINTING] Error at time {time_idx}: {e}") continue # Final cleanup for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].fillna(0) return s2_data, metadata class DeepInpaintingStrategy(CloudRemovalStrategy): """ Deep Learning Inpainting - sử dụng U-Net CNN Phức tạp hơn nhưng cho kết quả tốt nhất với large cloud gaps Note: Yêu cầu pretrained model (train bằng train_cloud_removal.py) """ def __init__(self, model_path: Optional[str] = None): super().__init__( name="deep_inpainting", description="Deep Learning U-Net based cloud removal - best quality for large gaps" ) self.model_path = model_path or "model_train/cloud_removal_unet_best.pth" self.model = None self.device = None # Try to load model if provided if model_path or Path(self.model_path).exists(): try: import torch import torch.nn as nn # Load checkpoint checkpoint = torch.load(self.model_path, map_location='cpu') # Recreate U-Net architecture from train_cloud_removal import UNet self.model = UNet( in_channels=checkpoint.get('in_channels', 4), out_channels=checkpoint.get('out_channels', 4) ) self.model.load_state_dict(checkpoint['model_state_dict']) # Set device self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') self.model = self.model.to(self.device) self.model.eval() print(f"[DEEP INPAINTING] Loaded U-Net model from {self.model_path}") print(f"[DEEP INPAINTING] Using device: {self.device}") except Exception as e: print(f"[DEEP INPAINTING] Could not load model: {e}") self.model = None def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: metadata = { 'method': self.name, 'has_model': self.model is not None, 'steps_applied': [] } # Apply mask for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].where(~cloud_mask) if self.model is None: # Fallback to classical method print("[DEEP INPAINTING] No model available, falling back to median composite") for band in s2_data.data_vars: if band != "SCL": median_composite = s2_data[band].median(dim='time', skipna=True) s2_data[band] = s2_data[band].fillna(median_composite) s2_data[band] = s2_data[band].interpolate_na(dim='x', method='nearest') s2_data[band] = s2_data[band].interpolate_na(dim='y', method='nearest') s2_data[band] = s2_data[band].fillna(0) metadata['steps_applied'].append('fallback_median') else: # Use U-Net for cloud removal print("[DEEP INPAINTING] Applying U-Net cloud removal...") import torch try: # Process each time step for time_idx in range(len(s2_data.time)): # Get bands for this time step (B02, B03, B04, B08) bands_to_process = ['B02', 'B03', 'B04', 'B08'] available_bands = [b for b in bands_to_process if b in s2_data.data_vars] if len(available_bands) < 4: print(f"[DEEP INPAINTING] Warning: Not all required bands available, skipping time {time_idx}") continue # Stack bands [C, H, W] input_bands = [] for band in available_bands: band_data = s2_data[band].isel(time=time_idx).values.astype(np.float32) # Normalize to [0, 1] (S2 values are typically 0-10000) band_data = np.clip(band_data / 10000.0, 0, 1) input_bands.append(band_data) input_array = np.stack(input_bands, axis=0) # [C, H, W] # Convert to tensor and add batch dimension input_tensor = torch.from_numpy(input_array).unsqueeze(0).to(self.device) if hasattr(self.model, 'encoder'): # Custom UNet from train_cloud_removal.py expected_channels = self.model.encoder[0].double_conv[0].in_channels elif hasattr(self.model, 'inc') and hasattr(self.model.inc.double_conv[0], 'in_channels'): expected_channels = self.model.inc.double_conv[0].in_channels elif hasattr(self.model, 'conv1') and hasattr(self.model.conv1, 'in_channels'): expected_channels = self.model.conv1.in_channels else: expected_channels = 6 if expected_channels > input_tensor.shape[1]: pad_channels = expected_channels - input_tensor.shape[1] padding = torch.zeros(1, pad_channels, *input_tensor.shape[2:]).to(self.device) input_tensor = torch.cat([input_tensor, padding], dim=1) # Run through U-Net with torch.no_grad(): output_tensor = self.model(input_tensor) # Convert back to numpy output_array = output_tensor[0].cpu().numpy() # [C, H, W] # Denormalize back to original scale output_array = output_array * 10000.0 # Update dataset with cleaned data for i, band in enumerate(available_bands): s2_data[band].values[time_idx] = output_array[i] metadata['steps_applied'].append(f'unet_time_{time_idx}') print(f"[DEEP INPAINTING] Processed {len(s2_data.time)} time steps with U-Net") except Exception as e: print(f"[DEEP INPAINTING] Error during inference: {e}") # Fallback to classical method for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].ffill(dim='time').bfill(dim='time') s2_data[band] = s2_data[band].fillna(0) metadata['steps_applied'].append('unet_error_fallback') return s2_data, metadata class HybridStrategy(CloudRemovalStrategy): """ Hybrid Strategy - kết hợp Classical + ML 1. Classical temporal interpolation (nhanh) 2. ML inpainting cho gaps còn lại (chất lượng cao) 3. Spatial interpolation (cleanup) """ def __init__(self): super().__init__( name="hybrid", description="Hybrid classical + ML - balanced speed and quality" ) def remove_clouds(self, s2_data: xr.Dataset, cloud_mask: xr.DataArray) -> Tuple[xr.Dataset, Dict]: metadata = { 'method': self.name, 'steps_applied': [] } # Apply mask for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].where(~cloud_mask) # Step 1: Temporal interpolation (fast) for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].ffill(dim='time').bfill(dim='time') metadata['steps_applied'].append('temporal_interpolation') # Step 2: Check remaining NaN percentage nan_count = 0 total_count = 0 for band in s2_data.data_vars: if band != "SCL": nan_count += np.isnan(s2_data[band].values).sum() total_count += s2_data[band].values.size nan_percentage = (nan_count / total_count * 100) if total_count > 0 else 0 # Step 3: ML inpainting if still significant gaps (>5%) if nan_percentage > 5.0: print(f"[HYBRID] {nan_percentage:.1f}% NaN remaining, applying ML inpainting...") ml_strategy = MLInpaintingStrategy(ml_model="knn") s2_data, ml_meta = ml_strategy.remove_clouds(s2_data, cloud_mask) metadata['steps_applied'].extend(['ml_inpainting_knn']) metadata['nan_before_ml'] = nan_percentage else: # Step 4: Spatial interpolation for small gaps for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].interpolate_na(dim='x', method='nearest') s2_data[band] = s2_data[band].interpolate_na(dim='y', method='nearest') metadata['steps_applied'].append('spatial_interpolation') # Final cleanup for band in s2_data.data_vars: if band != "SCL": s2_data[band] = s2_data[band].fillna(0) return s2_data, metadata # ============ FACTORY & UTILITIES ============ def get_available_methods() -> Dict[str, str]: """Trả về dictionary của tất cả methods có sẵn""" return { "none": "No cloud removal - keep original data (fastest, may have cloud artifacts)", "classic": "3-step classical: temporal → median → spatial (default, balanced)", "temporal_only": "Temporal interpolation only (fast, needs many scenes)", "median_composite": "Median composite priority (best noise reduction)", "ml_knn": "ML K-Nearest Neighbors inpainting (good quality, medium speed)", "ml_rf": "ML Random Forest inpainting (high quality, slower)", "deep": "Deep Learning CNN inpainting (best quality, requires model)", "hybrid": "Hybrid classical + ML (balanced speed & quality)" } def create_cloud_removal_strategy(method: str = "classic", **kwargs) -> CloudRemovalStrategy: """ Factory function để tạo strategy từ tên method Args: method: Tên method ("classic", "temporal_only", "median_composite", "ml_knn", "ml_rf", "deep", "hybrid") **kwargs: Additional parameters cho specific strategies Returns: CloudRemovalStrategy instance """ method = method.lower() if method == "none": return NoRemovalStrategy() elif method == "classic": return ClassicStrategy() elif method == "temporal_only": return TemporalOnlyStrategy() elif method == "median_composite": return MedianCompositeStrategy() elif method == "ml_knn": return MLInpaintingStrategy(ml_model="knn") elif method == "ml_rf": return MLInpaintingStrategy(ml_model="rf") elif method == "deep": model_path = kwargs.get('model_path', None) return DeepInpaintingStrategy(model_path=model_path) elif method == "hybrid": return HybridStrategy() else: print(f"[CLOUD REMOVAL] Unknown method '{method}', using 'classic'") return ClassicStrategy() def process_cloud_removal( s2_data: xr.Dataset, method: str = "classic", verbose: bool = True, **kwargs ) -> Tuple[xr.Dataset, Dict]: """ Main entry point cho cloud removal Args: s2_data: Sentinel-2 dataset với SCL band method: Cloud removal method name verbose: Print progress messages **kwargs: Additional parameters Returns: Tuple[xr.Dataset, Dict]: (cleaned_data, metadata) """ if verbose: print(f"[CLOUD REMOVAL] Using method: {method}") # Detect clouds from SCL if "SCL" not in s2_data: if verbose: print("[CLOUD REMOVAL] Warning: No SCL band, cannot mask clouds") return s2_data, {'method': 'none', 'warning': 'no_scl_band'} scl = s2_data["SCL"] # Create comprehensive cloud mask cloud_mask = (scl == 3) | (scl == 8) | (scl == 9) | (scl == 10) | (scl == 11) invalid_mask = (scl == 0) | (scl == 1) full_mask = cloud_mask | invalid_mask # Calculate coverage total_pixels = full_mask.size masked_pixels = int(full_mask.sum().values) cloud_coverage_percent = (masked_pixels / total_pixels * 100) if total_pixels > 0 else 0 if verbose: print(f"[CLOUD REMOVAL] Cloud coverage: {cloud_coverage_percent:.1f}%") print(f"[CLOUD REMOVAL] Masked pixels: {masked_pixels:,}/{total_pixels:,}") # Create strategy and process strategy = create_cloud_removal_strategy(method, **kwargs) cleaned_data, metadata = strategy.remove_clouds(s2_data.copy(deep=True), full_mask) # Add coverage info to metadata metadata['cloud_coverage_percent'] = float(cloud_coverage_percent) metadata['masked_pixels'] = masked_pixels metadata['total_pixels'] = total_pixels if verbose: print(f"[CLOUD REMOVAL] Completed using {metadata['method']}") print(f"[CLOUD REMOVAL] Steps: {', '.join(metadata['steps_applied'])}") return cleaned_data, metadata # ============ TESTING & COMPARISON ============ def compare_methods(s2_data: xr.Dataset, methods: list = None) -> Dict: """ So sánh các methods khác nhau trên cùng dữ liệu Args: s2_data: Sentinel-2 dataset methods: List of method names to compare (default: all) Returns: Dict: Comparison results """ if methods is None: methods = ["classic", "temporal_only", "median_composite", "ml_knn", "hybrid"] results = {} for method in methods: try: print(f"\n{'='*60}") print(f"Testing: {method}") print(f"{'='*60}") cleaned_data, metadata = process_cloud_removal(s2_data, method=method, verbose=True) # Calculate remaining NaN nan_count = sum(np.isnan(cleaned_data[band].values).sum() for band in cleaned_data.data_vars if band != "SCL") total_count = sum(cleaned_data[band].values.size for band in cleaned_data.data_vars if band != "SCL") results[method] = { 'metadata': metadata, 'remaining_nan_percent': (nan_count / total_count * 100) if total_count > 0 else 0, 'success': True } except Exception as e: results[method] = { 'error': str(e), 'success': False } print(f"[ERROR] {method}: {e}") return results