hoàn thành tính cận trên và cận dưới của tất cả các thuật toán
This commit is contained in:
Regular → Executable
+185
-20
@@ -18,6 +18,8 @@ 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"))
|
||||
@@ -185,6 +187,66 @@ class UNet(nn.Module):
|
||||
|
||||
# ============ 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()
|
||||
@@ -211,9 +273,10 @@ def train_epoch(model, dataloader, criterion, optimizer, device):
|
||||
|
||||
|
||||
def validate(model, dataloader, criterion, device):
|
||||
"""Validate model"""
|
||||
"""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"):
|
||||
@@ -223,8 +286,15 @@ def validate(model, dataloader, criterion, 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
|
||||
|
||||
return total_loss / len(dataloader)
|
||||
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):
|
||||
@@ -289,7 +359,8 @@ def train_cloud_removal_model(
|
||||
num_epochs=50,
|
||||
learning_rate=1e-4,
|
||||
device="cuda" if torch.cuda.is_available() else "cpu",
|
||||
save_dir="model_train"
|
||||
save_dir="model_train",
|
||||
status_callback=None
|
||||
):
|
||||
"""
|
||||
Train cloud removal model
|
||||
@@ -302,6 +373,7 @@ def train_cloud_removal_model(
|
||||
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)
|
||||
@@ -375,11 +447,25 @@ def train_cloud_removal_model(
|
||||
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}")
|
||||
@@ -391,18 +477,48 @@ def train_cloud_removal_model(
|
||||
train_losses.append(train_loss)
|
||||
|
||||
# Validate
|
||||
val_loss = validate(model, val_loader, criterion, device)
|
||||
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)
|
||||
|
||||
print(f"\nEpoch {epoch + 1} Summary:")
|
||||
print(f" Train Loss: {train_loss:.6f}")
|
||||
print(f" Val Loss: {val_loss:.6f}")
|
||||
# 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)
|
||||
|
||||
# Save best model
|
||||
if val_loss < best_val_loss:
|
||||
# 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)
|
||||
@@ -413,12 +529,34 @@ def train_cloud_removal_model(
|
||||
'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
|
||||
'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: {save_path}")
|
||||
print(f" 💾 Saved best model (highest PSNR): {save_path}")
|
||||
|
||||
# Visualize every 10 epochs
|
||||
if (epoch + 1) % 10 == 0:
|
||||
@@ -431,17 +569,40 @@ def train_cloud_removal_model(
|
||||
|
||||
print(f" 💾 Saved visualization: {viz_path}")
|
||||
|
||||
# Close Tensorboard writer
|
||||
writer.close()
|
||||
|
||||
# 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)
|
||||
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)
|
||||
@@ -453,10 +614,14 @@ def train_cloud_removal_model(
|
||||
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
|
||||
return model, train_losses, val_losses, val_psnrs, baseline_psnr
|
||||
|
||||
|
||||
# ============ INFERENCE FUNCTION ============
|
||||
|
||||
Reference in New Issue
Block a user