Files
remote-sensing/cloud_removal.py
T
2026-03-07 17:14:00 +07:00

629 lines
24 KiB
Python

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