refactor: reorganize project structure by moving core modules and update import paths in API server

This commit is contained in:
2026-07-18 01:24:30 +07:00
parent abab846884
commit a82b2f6fa5
155 changed files with 25 additions and 370 deletions
+512
View File
@@ -0,0 +1,512 @@
"""
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
)