hoàn thành model swing-unet
This commit is contained in:
@@ -82,3 +82,4 @@ model_train/*.log
|
|||||||
|
|
||||||
# Jupyter checkpoints
|
# Jupyter checkpoints
|
||||||
.ipynb_checkpoints/
|
.ipynb_checkpoints/
|
||||||
|
reports/
|
||||||
|
|||||||
+226
-83
@@ -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 = {
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user