hoàn thành model swing-unet

This commit is contained in:
Victor Phan
2026-01-05 16:20:58 +07:00
parent f401765996
commit ae3e5fffdd
4 changed files with 509 additions and 95 deletions
+1
View File
@@ -82,3 +82,4 @@ model_train/*.log
# Jupyter checkpoints # Jupyter checkpoints
.ipynb_checkpoints/ .ipynb_checkpoints/
reports/
+226 -83
View File
@@ -108,61 +108,61 @@ DEFAULT_LABEL_NAMES = {
class TrainingConfig(BaseModel): class TrainingConfig(BaseModel):
"""Cấu hình training""" """Cấu hình training - Tất cả bắt buộc nhập từ giao diện"""
# Khu vực (bbox) - từ 01.train_ODC.ipynb # Khu vực (bbox)
min_lon: float = 105.5 min_lon: float
min_lat: float = 9.2 min_lat: float
max_lon: float = 106.4 max_lon: float
max_lat: float = 10.0 max_lat: float
# Thời gian - từ 01.train_ODC.ipynb # Thời gian
start_date: str = "2023-03-01" start_date: str
end_date: str = "2023-12-31" end_date: str
# Dữ liệu # Dữ liệu
max_scenes: int = 12 max_scenes: int
cloud_cover: int = 30 cloud_cover: int
resolution: int = 20 # 10m hoặc 20m resolution: int # 10m hoặc 20m
# Model parameters # Model parameters
model_type: str = "xgboost" # xgboost, random_forest, decision_tree, svm, cnn, swin-unet model_type: str # xgboost, random_forest, decision_tree, svm, cnn, swin-unet
n_estimators: int = 100 n_estimators: int
max_depth: int = 20 max_depth: int
learning_rate: float = 0.1 learning_rate: float
use_gpu: bool = True use_gpu: bool
# Train/test split # Train/test split
test_size: float = 0.2 # Tỷ lệ dữ liệu dùng làm test (0-1) test_size: float # Tỷ lệ dữ liệu dùng làm test (0-1)
# Cache # Cache
use_cache: bool = True # Cache dataset để test nhanh hơn use_cache: bool # Cache dataset để test nhanh hơn
# Training data # Training data
training_shapefile: str = "train/ST_training data_updated_1130points_new.shp" training_shapefile: str
class PredictionConfig(BaseModel): class PredictionConfig(BaseModel):
"""Cấu hình dự đoán""" """Cấu hình dự đoán - Tất cả bắt buộc nhập từ giao diện"""
# Model to use # Model to use
model_filename: str model_filename: str
# Khu vực (bbox) - từ 01.train_ODC.ipynb # Khu vực (bbox)
min_lon: float = 105.5 min_lon: float
min_lat: float = 9.2 min_lat: float
max_lon: float = 106.4 max_lon: float
max_lat: float = 10.0 max_lat: float
# Thời gian - từ 01.train_ODC.ipynb # Thời gian
start_date: str = "2023-03-01" start_date: str
end_date: str = "2023-12-31" end_date: str
# Dữ liệu # Dữ liệu
max_scenes: int = 12 max_scenes: int
cloud_cover: int = 30 cloud_cover: int
resolution: int = 20 resolution: int
# GPU support for deep learning models # GPU support for deep learning models
use_gpu: bool = True use_gpu: bool
class TrainingStatus(BaseModel): class TrainingStatus(BaseModel):
@@ -1258,6 +1258,9 @@ async def run_prediction(config: PredictionConfig):
# ============ EXTRACT FEATURES ============ # ============ EXTRACT FEATURES ============
prediction_status["progress"] = f"Đang trích xuất features (mode={feature_mode})..." prediction_status["progress"] = f"Đang trích xuất features (mode={feature_mode})..."
print(f"[PREDICTION DEBUG] Feature mode: {feature_mode}")
print(f"[PREDICTION DEBUG] S2 bands available: {list(s2_data.data_vars)}")
print(f"[PREDICTION DEBUG] S2 dimensions: {dict(s2_data.dims)}")
# Always fill NaN for all bands in s2_data if present # Always fill NaN for all bands in s2_data if present
for band in ["B02", "B03", "B04", "B08", "B11"]: for band in ["B02", "B03", "B04", "B08", "B11"]:
@@ -1282,6 +1285,12 @@ async def run_prediction(config: PredictionConfig):
# Handle NaN values # Handle NaN values
features = np.nan_to_num(features, nan=0.0) features = np.nan_to_num(features, nan=0.0)
print(f"[PREDICTION DEBUG] Features extracted: shape={features.shape}")
print(f"[PREDICTION DEBUG] Features range: [{features.min():.3f}, {features.max():.3f}]")
print(f"[PREDICTION DEBUG] Features mean: {features.mean():.3f}, std: {features.std():.3f}")
print(f"[PREDICTION DEBUG] NaN count: {np.isnan(features).sum()}")
print(f"[PREDICTION DEBUG] First pixel features: {features[0][:min(8, features.shape[1])]}")
# Ensure features shape matches model expectation # Ensure features shape matches model expectation
if features.shape[1] != n_features_expected: if features.shape[1] != n_features_expected:
@@ -1292,9 +1301,18 @@ async def run_prediction(config: PredictionConfig):
# ============ PREDICT ============ # ============ PREDICT ============
prediction_status["progress"] = "Đang dự đoán..." prediction_status["progress"] = "Đang dự đoán..."
print(f"[PREDICTION DEBUG] Starting prediction with {features.shape[0]} pixels, {features.shape[1]} features")
# Make prediction (all PyTorch models have the same predict interface) # Make prediction (all PyTorch models have the same predict interface)
predictions = model.predict(features) predictions = model.predict(features)
print(f"[PREDICTION DEBUG] Predictions shape: {predictions.shape}")
print(f"[PREDICTION DEBUG] Unique predicted classes: {np.unique(predictions)}")
print(f"[PREDICTION DEBUG] Class distribution:")
unique, counts = np.unique(predictions, return_counts=True)
for cls, cnt in zip(unique, counts):
print(f" Class {cls}: {cnt} pixels ({cnt/len(predictions)*100:.1f}%)")
# Decode labels if label_encoder exists # Decode labels if label_encoder exists
if label_encoder is not None: if label_encoder is not None:
try: try:
@@ -2976,72 +2994,111 @@ async def predict_with_ndvi(config: PredictionWithNDVIConfig, background_tasks:
print(f"[PREDICT+NDVI] Data has {n_times} time steps, spatial size: {height}x{width}") print(f"[PREDICT+NDVI] Data has {n_times} time steps, spatial size: {height}x{width}")
# Build features based on what model expects # Build features using FeatureExtractor to match training
# Model metadata should tell us what features were used feature_mode = model_metadata.get("feature_mode", "simple")
model_features = model_metadata.get("features", ["NDVI_mean", "VH_dB_mean", "VV_dB_mean"]) expected_n_features = model_metadata.get("n_features", 3)
# If model was trained with temporal features (multiple time steps) print(f"[PREDICT+NDVI] Model feature_mode: {feature_mode}")
if expected_n_features > 10: # Likely temporal features print(f"[PREDICT+NDVI] Model n_features: {expected_n_features}")
print(f"[PREDICT+NDVI] Building temporal features (all time steps)")
# Use all time steps for each index # COMPATIBILITY FIX: If model was trained with buggy code (n_features=3 but feature_mode='odc'),
feature_list = [] # fallback to simple mode to match what model actually expects
if feature_mode == 'odc' and expected_n_features == 3:
print(f"[PREDICT+NDVI] ⚠️ WARNING: Model metadata shows odc mode but only 3 features")
print(f"[PREDICT+NDVI] This model was trained with old buggy code - using simple mode for compatibility")
feature_mode = 'simple'
# Use FeatureExtractor for consistent feature building
from feature_extractor import FeatureExtractor
extractor = FeatureExtractor(mode=feature_mode)
print(f"[PREDICT+NDVI] Using FeatureExtractor with mode='{feature_mode}'")
print(f"[PREDICT+NDVI] Expected features: {extractor.get_info()}")
# Extract features from data
# data is already an xr.Dataset with B02, B03, B04, B08
# Need to add B11 for ODC mode (NDBI calculation uses SWIR)
if feature_mode == 'odc' and 'B11' not in data:
# Load B11 if needed for ODC mode
print(f"[PREDICT+NDVI] ODC mode requires B11 (SWIR), loading...")
# Add NDVI for each time step # Fetch B11 band (whether from cache or fresh fetch)
for t in range(n_times): try:
feature_list.append(ndvi[t].flatten()) # If we haven't fetched items yet (cache scenario), do it now
if 'signed_items' not in locals():
catalog = Client.open(
"https://planetarycomputer.microsoft.com/api/stac/v1",
modifier=planetary_computer.sign_inplace
)
time_range = f"{config.start_date}/{config.end_date}"
search = catalog.search(
collections=["sentinel-2-l2a"],
bbox=bbox,
datetime=time_range,
query={"eo:cloud_cover": {"lt": config.cloud_cover}}
)
signed_items = [planetary_computer.sign(item) for item in list(search.items())[:config.max_scenes]]
print(f"[PREDICT+NDVI] Fetched {len(signed_items)} scenes for B11")
b11_data = odc.stac.load(
signed_items,
bbox=bbox,
bands=["B11"],
resolution=config.resolution,
chunks={"x": 2048, "y": 2048}
).compute()
# Merge B11 into existing data
data = xr.merge([data, b11_data])
print(f"[PREDICT+NDVI] Added B11 to data")
except Exception as e:
print(f"[PREDICT+NDVI] Warning: Failed to load B11: {e}")
print(f"[PREDICT+NDVI] Will proceed without B11 (may affect NDBI accuracy)")
# Extract features using FeatureExtractor
if feature_mode == 'simple':
# Simple mode needs pre-calculated NDVI
print(f"[PREDICT+NDVI] Calculating NDVI for simple mode...")
red_band = data["B04"].values
nir_band = data["B08"].values
ndvi_array = (nir_band - red_band) / (nir_band + red_band + 1e-8)
# If model has more features, add NDWI and NDBI time series # Convert to xarray DataArray with proper dims
if expected_n_features >= n_times * 2: ndvi_data = xr.DataArray(
for t in range(n_times): ndvi_array,
feature_list.append(ndwi[t].flatten()) dims=data["B04"].dims,
coords=data["B04"].coords
)
if expected_n_features >= n_times * 3: # Simple mode also needs VH/VV radar data, but we don't have it for this endpoint
for t in range(n_times): # Pass None and let extractor handle it
feature_list.append(ndbi[t].flatten()) features = extractor.extract(s2_data=None, ndvi_data=ndvi_data, vh_data=None, vv_data=None)
features = np.stack(feature_list, axis=1)
# Adjust to match expected features
if features.shape[1] < expected_n_features:
# Pad with mean values
n_missing = expected_n_features - features.shape[1]
padding = np.tile(features[:, -1:], (1, n_missing))
features = np.column_stack([features, padding])
elif features.shape[1] > expected_n_features:
# Trim to expected
features = features[:, :expected_n_features]
else: else:
# Use mean values (aggregate features) # ODC/extended modes use s2_data directly
print(f"[PREDICT+NDVI] Building aggregate features (mean values)") features = extractor.extract(s2_data=data, vh_data=None, vv_data=None)
# Average over time dimension
ndvi_mean = np.nanmean(ndvi, axis=0)
ndwi_mean = np.nanmean(ndwi, axis=0)
ndbi_mean = np.nanmean(ndbi, axis=0)
# Reshape for prediction
features = np.stack([ndvi_mean.flatten(), ndwi_mean.flatten(), ndbi_mean.flatten()], axis=1)
# Adjust to match expected features if needed
if features.shape[1] < expected_n_features:
n_missing = expected_n_features - features.shape[1]
padding = np.tile(features[:, -1:], (1, n_missing))
features = np.column_stack([features, padding])
elif features.shape[1] > expected_n_features:
features = features[:, :expected_n_features]
print(f"[PREDICT+NDVI] Built features shape: {features.shape}") print(f"[PREDICT+NDVI] Built features shape: {features.shape}")
print(f"[PREDICT+NDVI] Features per pixel: {features.shape[1] if len(features.shape) > 1 else 1}")
# Handle NaN values # Handle NaN, inf, and extreme values
valid_mask = ~np.isnan(features).any(axis=1) # Replace inf with 0
features = np.nan_to_num(features, nan=0.0, posinf=0.0, neginf=0.0)
# Clip extreme values to reasonable range
features = np.clip(features, -1e6, 1e6)
# Double-check no inf/nan remain
valid_mask = np.isfinite(features).all(axis=1)
features_clean = features[valid_mask] features_clean = features[valid_mask]
print(f"[PREDICT+NDVI] Predicting {features_clean.shape[0]} valid pixels...") print(f"[PREDICT+NDVI] Predicting {features_clean.shape[0]} valid pixels...")
print(f"[PREDICT+NDVI] Features range: [{features_clean.min():.3f}, {features_clean.max():.3f}]")
# Check if model is PyTorch/deep learning model and use GPU if available # Check if model is PyTorch/deep learning model and use GPU if available
is_pytorch_model = hasattr(model, '__class__') and ('CNN' in model.__class__.__name__ or 'Swin' in model.__class__.__name__ or 'UNet' in model.__class__.__name__) is_pytorch_model = hasattr(model, '__class__') and ('CNN' in model.__class__.__name__ or 'Swin' in model.__class__.__name__ or 'UNet' in model.__class__.__name__)
if is_pytorch_model and config.use_gpu: if is_pytorch_model and config.use_gpu:
try: try:
import torch import torch
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
@@ -3077,6 +3134,12 @@ async def predict_with_ndvi(config: PredictionWithNDVIConfig, background_tasks:
predictions = model.predict(features_clean) predictions = model.predict(features_clean)
except Exception as gpu_error: except Exception as gpu_error:
print(f"[PREDICT+NDVI] GPU prediction failed: {gpu_error}, falling back to CPU") print(f"[PREDICT+NDVI] GPU prediction failed: {gpu_error}, falling back to CPU")
# Move model back to CPU before retrying
try:
model = model.cpu()
print(f"[PREDICT+NDVI] Moved model to CPU")
except:
pass
predictions = model.predict(features_clean) predictions = model.predict(features_clean)
else: else:
# Use CPU for traditional ML models # Use CPU for traditional ML models
@@ -3142,6 +3205,13 @@ async def predict_with_ndvi(config: PredictionWithNDVIConfig, background_tasks:
# Export NDVI if requested # Export NDVI if requested
if config.export_ndvi: if config.export_ndvi:
# Calculate NDVI from data for export (data contains B04=red, B08=nir)
red_band = data["B04"].values
nir_band = data["B08"].values
ndvi_array = (nir_band - red_band) / (nir_band + red_band + 1e-8)
# Average over time dimension to get mean NDVI
ndvi_mean = np.nanmean(ndvi_array, axis=0)
ndvi_file = output_dir / f"ndvi_{timestamp}.tif" ndvi_file = output_dir / f"ndvi_{timestamp}.tif"
transform = from_bounds(bbox[0], bbox[1], bbox[2], bbox[3], width, height) transform = from_bounds(bbox[0], bbox[1], bbox[2], bbox[3], width, height)
@@ -3159,6 +3229,32 @@ async def predict_with_ndvi(config: PredictionWithNDVIConfig, background_tasks:
output_files.append({"type": "ndvi", "path": str(ndvi_file)}) output_files.append({"type": "ndvi", "path": str(ndvi_file)})
print(f"[PREDICT+NDVI] Saved NDVI to {ndvi_file}") print(f"[PREDICT+NDVI] Saved NDVI to {ndvi_file}")
# Create PNG preview for NDVI
ndvi_png = output_dir / f"ndvi_{timestamp}.png"
try:
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
fig, ax = plt.subplots(figsize=(12, 10), dpi=150)
im = ax.imshow(ndvi_mean, cmap='RdYlGn', vmin=-1, vmax=1, interpolation='nearest')
ax.set_title(f'NDVI - {timestamp}', fontsize=14, fontweight='bold')
ax.set_xlabel('X (pixels)', fontsize=10)
ax.set_ylabel('Y (pixels)', fontsize=10)
cbar = plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
cbar.set_label('NDVI', rotation=270, labelpad=15)
ax.grid(True, alpha=0.3, linestyle='--', linewidth=0.5)
plt.tight_layout()
plt.savefig(str(ndvi_png), dpi=150, bbox_inches='tight')
plt.close(fig)
output_files.append({"type": "ndvi_png", "path": str(ndvi_png)})
print(f"[PREDICT+NDVI] Created PNG: {ndvi_png}")
except Exception as e:
print(f"[PREDICT+NDVI] PNG creation failed: {e}")
# Export classification if requested # Export classification if requested
if config.export_classification: if config.export_classification:
@@ -3179,6 +3275,53 @@ async def predict_with_ndvi(config: PredictionWithNDVIConfig, background_tasks:
output_files.append({"type": "classification", "path": str(class_file)}) output_files.append({"type": "classification", "path": str(class_file)})
print(f"[PREDICT+NDVI] Saved classification to {class_file}") print(f"[PREDICT+NDVI] Saved classification to {class_file}")
# Create PNG preview for classification
class_png = output_dir / f"classification_{timestamp}.png"
try:
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
fig, ax = plt.subplots(figsize=(12, 10), dpi=150)
im = ax.imshow(prediction_raster, cmap='tab20', interpolation='nearest')
ax.set_title(f'Land Classification - {timestamp}', fontsize=14, fontweight='bold')
ax.set_xlabel('X (pixels)', fontsize=10)
ax.set_ylabel('Y (pixels)', fontsize=10)
# Try to get class labels
class_map = None
if label_encoder is not None:
try:
le_classes = list(label_encoder.classes_)
if all(isinstance(x, str) for x in le_classes):
class_map = {i: name for i, name in enumerate(le_classes)}
else:
class_map = {int(v): str(v) for v in le_classes}
except:
pass
cbar = plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
cbar.set_label('Class', rotation=270, labelpad=15)
if class_map:
try:
vals = np.array(sorted(class_map.keys()))
cbar.set_ticks(vals)
cbar.set_ticklabels([class_map[int(v)] for v in vals])
except:
pass
ax.grid(True, alpha=0.3, linestyle='--', linewidth=0.5)
plt.tight_layout()
plt.savefig(str(class_png), dpi=150, bbox_inches='tight')
plt.close(fig)
output_files.append({"type": "classification_png", "path": str(class_png)})
print(f"[PREDICT+NDVI] Created PNG: {class_png}")
except Exception as e:
print(f"[PREDICT+NDVI] PNG creation failed: {e}")
# Calculate statistics # Calculate statistics
ndvi_stats = { ndvi_stats = {
+132
View File
@@ -0,0 +1,132 @@
#!/usr/bin/env python3
"""
Generate PNG previews for existing GeoTIFF prediction files
"""
import numpy as np
import rasterio
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from pathlib import Path
import sys
def generate_png_preview(tif_file, output_png=None):
"""Generate PNG preview from GeoTIFF file"""
tif_path = Path(tif_file)
if not tif_path.exists():
print(f"❌ File not found: {tif_file}")
return False
# Determine output PNG path
if output_png is None:
output_png = tif_path.with_suffix('.png')
else:
output_png = Path(output_png)
try:
# Read GeoTIFF
with rasterio.open(tif_path) as src:
data = src.read(1)
print(f"📊 Data shape: {data.shape}, range: [{np.nanmin(data):.3f}, {np.nanmax(data):.3f}]")
# Determine if it's classification or NDVI based on filename
is_classification = 'classification' in tif_path.name.lower() or 'prediction' in tif_path.name.lower()
is_ndvi = 'ndvi' in tif_path.name.lower()
# Create figure
fig, ax = plt.subplots(figsize=(12, 10), dpi=150)
if is_ndvi:
# NDVI: use RdYlGn colormap, range -1 to 1
im = ax.imshow(data, cmap='RdYlGn', vmin=-1, vmax=1, interpolation='nearest')
ax.set_title(f'NDVI - {tif_path.stem}', fontsize=14, fontweight='bold')
cbar_label = 'NDVI'
elif is_classification:
# Classification: use tab20 colormap
im = ax.imshow(data, cmap='tab20', interpolation='nearest')
ax.set_title(f'Land Classification - {tif_path.stem}', fontsize=14, fontweight='bold')
cbar_label = 'Class'
else:
# Generic: use viridis
im = ax.imshow(data, cmap='viridis', interpolation='nearest')
ax.set_title(f'{tif_path.stem}', fontsize=14, fontweight='bold')
cbar_label = 'Value'
ax.set_xlabel('X (pixels)', fontsize=10)
ax.set_ylabel('Y (pixels)', fontsize=10)
# Add colorbar
cbar = plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
cbar.set_label(cbar_label, rotation=270, labelpad=15)
# For classification, try to set integer ticks
if is_classification:
try:
unique_vals = np.unique(data[~np.isnan(data)])
if len(unique_vals) < 20: # Only if not too many classes
cbar.set_ticks(unique_vals)
cbar.set_ticklabels([str(int(v)) for v in unique_vals])
except:
pass
# Add grid
ax.grid(True, alpha=0.3, linestyle='--', linewidth=0.5)
# Save PNG
plt.tight_layout()
plt.savefig(str(output_png), dpi=150, bbox_inches='tight')
plt.close(fig)
print(f"✅ Created PNG: {output_png}")
return True
except Exception as e:
print(f"❌ Error creating PNG: {e}")
import traceback
traceback.print_exc()
return False
def generate_all_previews(predictions_dir="predictions"):
"""Generate PNG previews for all GeoTIFF files without PNGs"""
pred_path = Path(predictions_dir)
if not pred_path.exists():
print(f"❌ Directory not found: {predictions_dir}")
return
tif_files = list(pred_path.glob("*.tif"))
print(f"🔍 Found {len(tif_files)} GeoTIFF files")
generated = 0
skipped = 0
for tif_file in tif_files:
png_file = tif_file.with_suffix('.png')
if png_file.exists():
print(f"⏭️ Skipping {tif_file.name} (PNG already exists)")
skipped += 1
continue
print(f"\n🎨 Processing {tif_file.name}...")
if generate_png_preview(tif_file):
generated += 1
print(f"\n{'='*60}")
print(f"✅ Generated {generated} new PNG previews")
print(f"⏭️ Skipped {skipped} files (already have PNGs)")
print(f"{'='*60}")
if __name__ == "__main__":
if len(sys.argv) > 1:
# Process specific file
tif_file = sys.argv[1]
generate_png_preview(tif_file)
else:
# Process all files in predictions directory
generate_all_previews()
+150 -12
View File
@@ -277,7 +277,7 @@ def train_model(
use_gpu=True, use_gpu=True,
use_cache=True, use_cache=True,
test_size=0.2, test_size=0.2,
feature_mode='simple', # Changed from 'odc' - simple mode works with B04, B08, SCL only feature_mode='odc', # ODC mode: 8 features (NDVI stats + NDWI/NDBI/EVI) for better accuracy
output_model_path=None, output_model_path=None,
status_callback=None, status_callback=None,
cancel_check=None cancel_check=None
@@ -406,9 +406,11 @@ def train_model(
# Load different bands based on feature mode # Load different bands based on feature mode
if feature_mode == 'simple': if feature_mode == 'simple':
bands_to_load = ["B04", "B08", "SCL"] bands_to_load = ["B04", "B08", "SCL"]
else: # temporal or extended else: # odc, temporal, or extended - all need full spectral bands
bands_to_load = ["B02", "B03", "B04", "B08", "B11", "SCL"] bands_to_load = ["B02", "B03", "B04", "B08", "B11", "SCL"]
update_status(f"Loading bands: {bands_to_load} for mode={feature_mode}", 26)
ds_s2 = stac_load( ds_s2 = stac_load(
items_s2, items_s2,
bands=bands_to_load, bands=bands_to_load,
@@ -428,9 +430,11 @@ def train_model(
if 'time' in ds_s2.dims: if 'time' in ds_s2.dims:
print(f"[DEBUG S2] Time range: {ds_s2.time.min().values} to {ds_s2.time.max().values}") print(f"[DEBUG S2] Time range: {ds_s2.time.min().values} to {ds_s2.time.max().values}")
# Rename for compatibility (simple mode) # Rename bands ONLY for simple mode (simple mode uses 'red', 'nir', 'scl' names)
if "B04" in ds_s2 and "red" not in ds_s2: # Other modes (odc, extended, temporal) use original band names (B02, B03, B04, B08, B11, SCL)
if feature_mode == 'simple' and "B04" in ds_s2 and "red" not in ds_s2:
ds_s2 = ds_s2.rename({"B04": "red", "B08": "nir", "SCL": "scl"}) ds_s2 = ds_s2.rename({"B04": "red", "B08": "nir", "SCL": "scl"})
print(f"[DEBUG S2] Renamed bands for simple mode: B04→red, B08→nir, SCL→scl")
check_cancellation() check_cancellation()
@@ -575,6 +579,7 @@ def train_model(
update_status("Extracting features from satellite data...", 60) update_status("Extracting features from satellite data...", 60)
print(f"[DEBUG] Starting feature extraction...") print(f"[DEBUG] Starting feature extraction...")
print(f"[DEBUG] Feature mode: {feature_mode}")
print(f"[DEBUG] Training GDF has {len(train_gdf)} points") print(f"[DEBUG] Training GDF has {len(train_gdf)} points")
print(f"[DEBUG] Training GDF CRS: {train_gdf.crs}") print(f"[DEBUG] Training GDF CRS: {train_gdf.crs}")
print(f"[DEBUG] Training GDF bounds (UTM): {train_gdf.total_bounds}") print(f"[DEBUG] Training GDF bounds (UTM): {train_gdf.total_bounds}")
@@ -637,7 +642,85 @@ def train_model(
features = np.array(features) features = np.array(features)
labels = np.array(labels) labels = np.array(labels)
else: # temporal or extended mode elif feature_mode in ['odc', 'extended']:
# For odc/extended: Extract features for full raster first, then sample at points
update_status(f"Extracting {feature_mode} features from full raster...", 62)
# Apply cloud mask first
if 'SCL' in ds_s2:
scl_band = ds_s2['SCL']
cloud_mask = scl_band.isin([1, 3, 8, 9, 10])
for band in ds_s2.data_vars:
if band != 'SCL':
ds_s2[band] = ds_s2[band].where(~cloud_mask)
# Extract features using FeatureExtractor for entire raster
raster_features = extractor.extract(
s2_data=ds_s2,
vh_data=None, # ODC/extended don't use radar in aggregate
vv_data=None
)
print(f"[DEBUG] Extracted raster features: shape={raster_features.shape}")
print(f"[DEBUG] Feature range: [{raster_features.min()}, {raster_features.max()}]")
# Now sample at each training point
features = []
labels = []
failed_extractions = 0
# Get spatial dimensions
y_coords = ds_s2.y.values
x_coords = ds_s2.x.values
print(f"[DEBUG] S2 spatial grid: x=[{x_coords.min()}, {x_coords.max()}], y=[{y_coords.min()}, {y_coords.max()}]")
for idx, row in train_gdf.iterrows():
point = row.geometry
x_coord = point.x
y_coord = point.y
label = row[label_column]
try:
# Find nearest pixel indices
x_idx = np.argmin(np.abs(x_coords - x_coord))
y_idx = np.argmin(np.abs(y_coords - y_coord))
# Get features at this pixel
# raster_features shape: (n_pixels, n_features)
# Need to convert 2D (y, x) index to 1D pixel index
pixel_idx = y_idx * len(x_coords) + x_idx
if pixel_idx < len(raster_features):
feature_vec = raster_features[pixel_idx]
if idx < 3:
print(f"[DEBUG] Point {idx}: coords=({x_coord:.2f}, {y_coord:.2f}) -> pixel[{y_idx},{x_idx}] -> idx={pixel_idx}, features={feature_vec[:3]}...")
if not np.isnan(feature_vec).any():
features.append(feature_vec)
labels.append(label)
else:
failed_extractions += 1
if idx < 3:
print(f"[DEBUG] Point {idx} has NaN features")
else:
failed_extractions += 1
if idx < 3:
print(f"[DEBUG] Point {idx} pixel_idx {pixel_idx} out of range (max={len(raster_features)})")
except Exception as e:
failed_extractions += 1
if idx < 3:
print(f"[DEBUG] Point {idx} extraction failed: {e}")
continue
if failed_extractions > 0:
update_status(f"⚠️ {failed_extractions}/{len(train_gdf)} points had NaN/missing data", 65)
features = np.array(features)
labels = np.array(labels)
else: # temporal mode
# Apply cloud mask for temporal/extended modes # Apply cloud mask for temporal/extended modes
if 'scl' in ds_s2 or 'SCL' in ds_s2: if 'scl' in ds_s2 or 'SCL' in ds_s2:
scl_band = ds_s2['scl'] if 'scl' in ds_s2 else ds_s2['SCL'] scl_band = ds_s2['scl'] if 'scl' in ds_s2 else ds_s2['SCL']
@@ -882,19 +965,39 @@ def train_model(
train_dataset = TensorDataset(X_train_tensor, y_train_tensor) train_dataset = TensorDataset(X_train_tensor, y_train_tensor)
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
# Loss and optimizer with weight decay # Calculate class weights for imbalanced data
criterion = nn.CrossEntropyLoss() class_counts = np.bincount(y_train)
class_weights = 1.0 / (class_counts + 1e-6) # Avoid division by zero
class_weights = class_weights / class_weights.sum() * len(class_counts) # Normalize
class_weights_tensor = torch.FloatTensor(class_weights).to(device)
print(f"[SWIN-UNET] Class distribution: {class_counts}")
print(f"[SWIN-UNET] Class weights: {class_weights}")
# Loss with class weights and optimizer with weight decay
criterion = nn.CrossEntropyLoss(weight=class_weights_tensor)
optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=0.01) optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=0.01)
# LR scheduler for better convergence # LR scheduler for better convergence
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50) scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50)
# Early stopping to prevent overfitting
best_val_loss = float('inf')
patience = 10
patience_counter = 0
# Train Swin-UNet # Train Swin-UNet
update_status("Training Swin-UNet model with PyTorch...", 80) update_status("Training Swin-UNet model with PyTorch (with class weights)...", 80)
epochs = min(60, n_estimators // 2) # Swin-UNet benefits from more epochs epochs = min(60, n_estimators // 2) # Swin-UNet benefits from more epochs
# Validation dataset
val_dataset = TensorDataset(X_test_tensor, y_test_tensor)
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)
model.train() model.train()
for epoch in range(epochs): for epoch in range(epochs):
# Training phase
model.train()
epoch_loss = 0.0 epoch_loss = 0.0
for batch_X, batch_y in train_loader: for batch_X, batch_y in train_loader:
batch_X, batch_y = batch_X.to(device), batch_y.to(device) batch_X, batch_y = batch_X.to(device), batch_y.to(device)
@@ -903,16 +1006,51 @@ def train_model(
outputs = model(batch_X) outputs = model(batch_X)
loss = criterion(outputs, batch_y) loss = criterion(outputs, batch_y)
loss.backward() loss.backward()
# Gradient clipping to prevent exploding gradients
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step() optimizer.step()
epoch_loss += loss.item() epoch_loss += loss.item()
scheduler.step() scheduler.step()
if (epoch + 1) % 10 == 0: # Validation phase
avg_loss = epoch_loss / len(train_loader) model.eval()
lr = optimizer.param_groups[0]['lr'] val_loss = 0.0
update_status(f"Swin-UNet Epoch {epoch+1}/{epochs}, Loss: {avg_loss:.4f}, LR: {lr:.6f}", 80 + (epoch / epochs) * 10) correct = 0
total = 0
with torch.no_grad():
for batch_X, batch_y in val_loader:
batch_X, batch_y = batch_X.to(device), batch_y.to(device)
outputs = model(batch_X)
loss = criterion(outputs, batch_y)
val_loss += loss.item()
_, predicted = torch.max(outputs, 1)
total += batch_y.size(0)
correct += (predicted == batch_y).sum().item()
avg_train_loss = epoch_loss / len(train_loader)
avg_val_loss = val_loss / len(val_loader)
val_acc = 100 * correct / total
lr = optimizer.param_groups[0]['lr']
if (epoch + 1) % 5 == 0:
update_status(f"Swin-UNet Epoch {epoch+1}/{epochs}, Train Loss: {avg_train_loss:.4f}, Val Loss: {avg_val_loss:.4f}, Val Acc: {val_acc:.2f}%, LR: {lr:.6f}", 80 + (epoch / epochs) * 10)
print(f"[SWIN-UNET] Epoch {epoch+1}/{epochs} - Train Loss: {avg_train_loss:.4f}, Val Loss: {avg_val_loss:.4f}, Val Acc: {val_acc:.2f}%")
# Early stopping check
if avg_val_loss < best_val_loss:
best_val_loss = avg_val_loss
patience_counter = 0
else:
patience_counter += 1
if patience_counter >= patience:
print(f"[SWIN-UNET] Early stopping at epoch {epoch+1} (best val loss: {best_val_loss:.4f})")
update_status(f"Swin-UNet early stopped at epoch {epoch+1}", 90)
break
model = model.cpu() model = model.cpu()
model.device_used = str(device) model.device_used = str(device)