update 01
This commit is contained in:
+558
@@ -0,0 +1,558 @@
|
||||
"""
|
||||
Training module for land classification using Sentinel-2 and Sentinel-1 data
|
||||
from Microsoft Planetary Computer STAC API
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import xarray as xr
|
||||
import geopandas as gpd
|
||||
from sklearn.model_selection import train_test_split
|
||||
from sklearn.preprocessing import LabelEncoder
|
||||
from sklearn.metrics import classification_report, confusion_matrix
|
||||
from sklearn.ensemble import RandomForestClassifier
|
||||
from sklearn.tree import DecisionTreeClassifier
|
||||
from sklearn.svm import SVC
|
||||
from xgboost import XGBClassifier
|
||||
import joblib
|
||||
from datetime import datetime
|
||||
import json
|
||||
import os
|
||||
import warnings
|
||||
import hashlib
|
||||
from pathlib import Path
|
||||
warnings.filterwarnings('ignore')
|
||||
|
||||
# PyTorch for CNN
|
||||
try:
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.optim as optim
|
||||
from torch.utils.data import TensorDataset, DataLoader
|
||||
PYTORCH_AVAILABLE = True
|
||||
except ImportError:
|
||||
PYTORCH_AVAILABLE = False
|
||||
print("Warning: PyTorch not available. CNN model will not work.")
|
||||
|
||||
# Define CNN model class for PyTorch
|
||||
class CNNClassifier(nn.Module):
|
||||
def __init__(self, n_features, n_classes):
|
||||
super(CNNClassifier, self).__init__()
|
||||
self.n_features = n_features
|
||||
self.n_classes = n_classes
|
||||
|
||||
# For small feature sets (like 3 features), use simpler architecture
|
||||
if n_features < 8:
|
||||
# Simple fully connected network for small features
|
||||
self.use_conv = False
|
||||
self.fc1 = nn.Linear(n_features, 64)
|
||||
self.dropout1 = nn.Dropout(0.3)
|
||||
self.fc2 = nn.Linear(64, 128)
|
||||
self.dropout2 = nn.Dropout(0.5)
|
||||
self.fc3 = nn.Linear(128, n_classes)
|
||||
else:
|
||||
# CNN architecture for larger feature sets
|
||||
self.use_conv = True
|
||||
self.conv1 = nn.Conv1d(in_channels=1, out_channels=32, kernel_size=3, padding=1)
|
||||
self.pool1 = nn.MaxPool1d(kernel_size=2)
|
||||
self.conv2 = nn.Conv1d(in_channels=32, out_channels=64, kernel_size=3, padding=1)
|
||||
self.pool2 = nn.MaxPool1d(kernel_size=2)
|
||||
|
||||
# Calculate size after convolutions
|
||||
conv_output_size = (n_features // 2 // 2) * 64
|
||||
|
||||
# Fully connected layers
|
||||
self.fc1 = nn.Linear(conv_output_size, 128)
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
self.fc2 = nn.Linear(128, n_classes)
|
||||
|
||||
def forward(self, x):
|
||||
# x shape: (batch, n_features) or (batch, 1, n_features)
|
||||
if self.use_conv:
|
||||
# CNN path for larger feature sets
|
||||
if len(x.shape) == 2:
|
||||
x = x.unsqueeze(1) # Add channel dimension
|
||||
x = F.relu(self.conv1(x))
|
||||
x = self.pool1(x)
|
||||
x = F.relu(self.conv2(x))
|
||||
x = self.pool2(x)
|
||||
x = x.view(x.size(0), -1) # Flatten
|
||||
x = F.relu(self.fc1(x))
|
||||
x = self.dropout(x)
|
||||
x = self.fc2(x)
|
||||
else:
|
||||
# Fully connected path for small feature sets
|
||||
if len(x.shape) == 3:
|
||||
x = x.squeeze(1) # Remove channel dimension if present
|
||||
x = F.relu(self.fc1(x))
|
||||
x = self.dropout1(x)
|
||||
x = F.relu(self.fc2(x))
|
||||
x = self.dropout2(x)
|
||||
x = self.fc3(x)
|
||||
return x
|
||||
|
||||
def predict(self, X):
|
||||
"""Scikit-learn style predict method"""
|
||||
self.eval()
|
||||
with torch.no_grad():
|
||||
if isinstance(X, np.ndarray):
|
||||
X = torch.FloatTensor(X)
|
||||
# Handle both 2D and 3D inputs
|
||||
if not self.use_conv and len(X.shape) == 3:
|
||||
X = X.squeeze(1)
|
||||
elif self.use_conv and len(X.shape) == 2:
|
||||
X = X.unsqueeze(1)
|
||||
outputs = self(X)
|
||||
_, predicted = torch.max(outputs, 1)
|
||||
return predicted.cpu().numpy()
|
||||
|
||||
def score(self, X, y):
|
||||
"""Scikit-learn style score method"""
|
||||
predictions = self.predict(X)
|
||||
if isinstance(y, torch.Tensor):
|
||||
y = y.cpu().numpy()
|
||||
return np.mean(predictions == y)
|
||||
|
||||
# Microsoft Planetary Computer imports
|
||||
import planetary_computer
|
||||
from pystac_client import Client
|
||||
from odc.stac import load as stac_load
|
||||
|
||||
|
||||
def train_model(
|
||||
bbox=[105.6, 9.3, 106.2, 9.8],
|
||||
time_range='2023-03-01/2023-05-31',
|
||||
max_scenes=12,
|
||||
cloud_cover=30,
|
||||
resolution=20,
|
||||
training_shapefile='train/ST_training data_updated_1130points_new.shp',
|
||||
model_type='xgboost',
|
||||
n_estimators=100,
|
||||
max_depth=20,
|
||||
learning_rate=0.1,
|
||||
use_gpu=True,
|
||||
use_cache=True,
|
||||
test_size=0.2,
|
||||
output_model_path=None,
|
||||
status_callback=None,
|
||||
cancel_check=None
|
||||
):
|
||||
"""
|
||||
Train a land classification model using Sentinel-2 and Sentinel-1 data
|
||||
|
||||
Args:
|
||||
bbox: [min_lon, min_lat, max_lon, max_lat]
|
||||
time_range: "YYYY-MM-DD/YYYY-MM-DD"
|
||||
max_scenes: maximum number of scenes to load
|
||||
cloud_cover: maximum cloud cover percentage
|
||||
resolution: resolution in meters (e.g., 20)
|
||||
training_shapefile: path to training shapefile
|
||||
n_estimators: number of trees for XGBoost
|
||||
max_depth: maximum tree depth
|
||||
learning_rate: learning rate for XGBoost
|
||||
use_gpu: whether to use GPU for training
|
||||
output_model_path: path to save trained model (auto-generated if None)
|
||||
status_callback: Optional callback function to report progress
|
||||
cancel_check: Optional function that returns True if training should be cancelled
|
||||
test_size: Fraction of data to use for test set (0-1)
|
||||
|
||||
Returns:
|
||||
Dictionary containing training results
|
||||
"""
|
||||
|
||||
def update_status(message, progress=None):
|
||||
"""Helper to update status"""
|
||||
if status_callback:
|
||||
# Try calling with both arguments, fallback to just message
|
||||
try:
|
||||
status_callback(message, progress)
|
||||
except TypeError:
|
||||
status_callback(message)
|
||||
print(message)
|
||||
|
||||
def check_cancellation():
|
||||
"""Check if training should be cancelled"""
|
||||
if cancel_check and cancel_check():
|
||||
raise InterruptedError("Training cancelled by user")
|
||||
|
||||
try:
|
||||
# Auto-generate output path if not provided
|
||||
if output_model_path is None:
|
||||
timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
output_model_path = f'model_train/model_{model_type}_{timestamp}.joblib'
|
||||
|
||||
# ============ CACHE SYSTEM ============
|
||||
# Create cache directory
|
||||
cache_dir = Path("dataset_cache")
|
||||
cache_dir.mkdir(exist_ok=True)
|
||||
|
||||
# Generate cache key from parameters
|
||||
cache_params = f"{bbox}_{time_range}_{max_scenes}_{cloud_cover}_{resolution}"
|
||||
cache_key = hashlib.md5(cache_params.encode()).hexdigest()
|
||||
cache_file = cache_dir / f"training_data_{cache_key}.joblib"
|
||||
|
||||
features = None
|
||||
labels = None
|
||||
|
||||
# Try to load from cache
|
||||
if use_cache and cache_file.exists():
|
||||
update_status(f"📦 Loading cached dataset from {cache_file.name}...", 5)
|
||||
try:
|
||||
cached_data = joblib.load(cache_file)
|
||||
features = cached_data['features']
|
||||
labels = cached_data['labels']
|
||||
update_status(f"✅ Loaded {len(features)} samples from cache (skipped satellite download!)", 50)
|
||||
except Exception as e:
|
||||
update_status(f"⚠️ Cache load failed: {str(e)}, downloading fresh data...", 10)
|
||||
features = None
|
||||
|
||||
# If no cache or cache failed, download data
|
||||
if features is None:
|
||||
update_status("📡 Cache not found or disabled, downloading satellite data...", 10)
|
||||
|
||||
# Connect to Microsoft Planetary Computer
|
||||
update_status("Connecting to Microsoft Planetary Computer...", 12)
|
||||
catalog = Client.open("https://planetarycomputer.microsoft.com/api/stac/v1")
|
||||
check_cancellation()
|
||||
|
||||
# Search for Sentinel-2 scenes
|
||||
update_status("Searching for Sentinel-2 scenes...", 10)
|
||||
query_s2 = catalog.search(
|
||||
collections=["sentinel-2-l2a"],
|
||||
bbox=bbox,
|
||||
datetime=time_range,
|
||||
query={"eo:cloud_cover": {"lt": cloud_cover}}
|
||||
)
|
||||
items_s2 = list(query_s2.item_collection())
|
||||
|
||||
check_cancellation()
|
||||
|
||||
# Limit scenes
|
||||
if len(items_s2) > max_scenes:
|
||||
step = len(items_s2) // max_scenes
|
||||
items_s2 = items_s2[::step][:max_scenes]
|
||||
|
||||
update_status(f"Found {len(items_s2)} Sentinel-2 scenes", 20)
|
||||
|
||||
# Sign and load Sentinel-2 data
|
||||
update_status("Loading Sentinel-2 data...", 25)
|
||||
items_s2 = [planetary_computer.sign(item) for item in items_s2]
|
||||
ds_s2 = stac_load(
|
||||
items_s2,
|
||||
bands=["B04", "B08", "SCL"],
|
||||
crs="EPSG:32648",
|
||||
resolution=resolution,
|
||||
bbox=bbox,
|
||||
patch_url=planetary_computer.sign,
|
||||
fail_on_error=False,
|
||||
)
|
||||
ds_s2 = ds_s2.rename({"B04": "red", "B08": "nir", "SCL": "scl"})
|
||||
|
||||
check_cancellation()
|
||||
|
||||
# Search for Sentinel-1 scenes
|
||||
update_status("Searching for Sentinel-1 scenes...", 35)
|
||||
query_s1 = catalog.search(
|
||||
collections=["sentinel-1-rtc"],
|
||||
bbox=bbox,
|
||||
datetime=time_range,
|
||||
)
|
||||
items_s1 = list(query_s1.item_collection())
|
||||
|
||||
# Limit scenes
|
||||
if len(items_s1) > max_scenes:
|
||||
step = len(items_s1) // max_scenes
|
||||
items_s1 = items_s1[::step][:max_scenes]
|
||||
|
||||
update_status(f"Found {len(items_s1)} Sentinel-1 scenes", 40)
|
||||
|
||||
# Sign and load Sentinel-1 data
|
||||
update_status("Loading Sentinel-1 data...", 45)
|
||||
items_s1 = [planetary_computer.sign(item) for item in items_s1]
|
||||
ds_s1 = stac_load(
|
||||
items_s1,
|
||||
bands=["vv", "vh"],
|
||||
crs="EPSG:32648",
|
||||
resolution=resolution,
|
||||
bbox=bbox,
|
||||
patch_url=planetary_computer.sign,
|
||||
fail_on_error=False,
|
||||
)
|
||||
|
||||
# Convert to dB
|
||||
ds_s1['vv_db'] = 10 * np.log10(ds_s1['vv'].where(ds_s1['vv'] > 0))
|
||||
ds_s1['vh_db'] = 10 * np.log10(ds_s1['vh'].where(ds_s1['vh'] > 0))
|
||||
|
||||
check_cancellation()
|
||||
|
||||
# Calculate NDVI
|
||||
update_status("Calculating NDVI...", 50)
|
||||
ndvi = (ds_s2['nir'] - ds_s2['red']) / (ds_s2['nir'] + ds_s2['red'] + 1e-8)
|
||||
|
||||
# Apply cloud mask
|
||||
cloud_mask = ds_s2['scl'].isin([1, 3, 8, 9, 10])
|
||||
ndvi_masked = ndvi.where(~cloud_mask)
|
||||
ndvi_mean = ndvi_masked.mean(dim='time')
|
||||
|
||||
# Load training data
|
||||
update_status("Loading training data...", 55)
|
||||
train_gdf = gpd.read_file(training_shapefile)
|
||||
|
||||
if train_gdf.crs != 'EPSG:32648':
|
||||
train_gdf = train_gdf.to_crs('EPSG:32648')
|
||||
|
||||
# Auto-detect label column
|
||||
label_column = None
|
||||
for col in ['HT_code', 'Ma_LU', 'LU2022', 'Hientrang', 'class', 'Class', 'CLASS']:
|
||||
if col in train_gdf.columns:
|
||||
label_column = col
|
||||
break
|
||||
|
||||
if label_column is None:
|
||||
raise ValueError(f"Cannot find label column in shapefile. Available: {list(train_gdf.columns)}")
|
||||
|
||||
# Extract features
|
||||
update_status("Extracting features from training points...", 60)
|
||||
features = []
|
||||
labels = []
|
||||
|
||||
for idx, row in train_gdf.iterrows():
|
||||
point = row.geometry
|
||||
x_coord = point.x
|
||||
y_coord = point.y
|
||||
label = row[label_column]
|
||||
|
||||
try:
|
||||
ndvi_val = ndvi_mean.sel(x=x_coord, y=y_coord, method='nearest').values
|
||||
vh_val = ds_s1['vh_db'].sel(x=x_coord, y=y_coord, method='nearest').mean(dim='time').values
|
||||
vv_val = ds_s1['vv_db'].sel(x=x_coord, y=y_coord, method='nearest').mean(dim='time').values
|
||||
|
||||
feature_vec = [ndvi_val, vh_val, vv_val]
|
||||
|
||||
if not np.isnan(feature_vec).any():
|
||||
features.append(feature_vec)
|
||||
labels.append(label)
|
||||
except:
|
||||
continue
|
||||
|
||||
features = np.array(features)
|
||||
labels = np.array(labels)
|
||||
|
||||
check_cancellation()
|
||||
|
||||
update_status(f"Extracted {len(features)} valid training samples", 70)
|
||||
|
||||
# ============ SAVE TO CACHE ============
|
||||
if use_cache:
|
||||
update_status(f"💾 Saving dataset to cache for future use...", 72)
|
||||
try:
|
||||
cache_data = {
|
||||
'features': features,
|
||||
'labels': labels,
|
||||
'bbox': bbox,
|
||||
'time_range': time_range,
|
||||
'resolution': resolution,
|
||||
'timestamp': datetime.now().isoformat()
|
||||
}
|
||||
joblib.dump(cache_data, cache_file)
|
||||
update_status(f"✅ Cached to {cache_file.name}", 75)
|
||||
except Exception as e:
|
||||
update_status(f"⚠️ Cache save failed: {str(e)}", 75)
|
||||
|
||||
# Encode labels
|
||||
label_encoder = LabelEncoder()
|
||||
labels_encoded = label_encoder.fit_transform(labels)
|
||||
|
||||
# Split data
|
||||
X_train, X_test, y_train, y_test = train_test_split(
|
||||
features, labels_encoded, test_size=test_size, random_state=42, stratify=labels_encoded
|
||||
)
|
||||
|
||||
# Train model based on selected type
|
||||
update_status(f"Training {model_type.upper()} model...", 75)
|
||||
|
||||
device = 'cuda:0' if use_gpu else 'cpu'
|
||||
|
||||
if model_type == 'xgboost':
|
||||
model = XGBClassifier(
|
||||
n_estimators=n_estimators,
|
||||
max_depth=max_depth,
|
||||
learning_rate=learning_rate,
|
||||
device=device if use_gpu else 'cpu',
|
||||
tree_method='hist',
|
||||
random_state=42,
|
||||
eval_metric='mlogloss',
|
||||
verbosity=0
|
||||
)
|
||||
elif model_type == 'random_forest':
|
||||
model = RandomForestClassifier(
|
||||
n_estimators=n_estimators,
|
||||
max_depth=max_depth,
|
||||
random_state=42,
|
||||
n_jobs=-1, # Use all cores
|
||||
verbose=0
|
||||
)
|
||||
elif model_type == 'decision_tree':
|
||||
model = DecisionTreeClassifier(
|
||||
max_depth=max_depth,
|
||||
random_state=42
|
||||
)
|
||||
elif model_type == 'svm':
|
||||
model = SVC(
|
||||
kernel='rbf',
|
||||
random_state=42,
|
||||
verbose=False
|
||||
)
|
||||
elif model_type == 'cnn':
|
||||
if not PYTORCH_AVAILABLE:
|
||||
raise ImportError("PyTorch is required for CNN. Install: pip install torch")
|
||||
|
||||
# CNN requires reshaping data
|
||||
n_features = X_train.shape[1]
|
||||
n_classes = len(np.unique(y_train))
|
||||
|
||||
# Build PyTorch CNN model
|
||||
device = torch.device('cuda' if torch.cuda.is_available() and use_gpu else 'cpu')
|
||||
update_status(f"Building CNN model on {device}...", 75)
|
||||
|
||||
model = CNNClassifier(n_features, n_classes).to(device)
|
||||
|
||||
# Convert to PyTorch tensors
|
||||
X_train_tensor = torch.FloatTensor(X_train).unsqueeze(1) # Add channel dim: (N, 1, features)
|
||||
y_train_tensor = torch.LongTensor(y_train)
|
||||
X_test_tensor = torch.FloatTensor(X_test).unsqueeze(1)
|
||||
y_test_tensor = torch.LongTensor(y_test)
|
||||
|
||||
# Create data loaders
|
||||
train_dataset = TensorDataset(X_train_tensor, y_train_tensor)
|
||||
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
|
||||
|
||||
# Loss and optimizer
|
||||
criterion = nn.CrossEntropyLoss()
|
||||
optimizer = optim.Adam(model.parameters(), lr=0.001)
|
||||
|
||||
# Train CNN
|
||||
update_status("Training CNN model with PyTorch...", 80)
|
||||
epochs = min(50, n_estimators // 2) # Use n_estimators as epochs
|
||||
|
||||
model.train()
|
||||
for epoch in range(epochs):
|
||||
epoch_loss = 0.0
|
||||
for batch_X, batch_y in train_loader:
|
||||
batch_X, batch_y = batch_X.to(device), batch_y.to(device)
|
||||
|
||||
optimizer.zero_grad()
|
||||
outputs = model(batch_X)
|
||||
loss = criterion(outputs, batch_y)
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
epoch_loss += loss.item()
|
||||
|
||||
if (epoch + 1) % 10 == 0:
|
||||
avg_loss = epoch_loss / len(train_loader)
|
||||
update_status(f"CNN Epoch {epoch+1}/{epochs}, Loss: {avg_loss:.4f}", 80 + (epoch / epochs) * 10)
|
||||
|
||||
# Move model to CPU for saving (compatible with non-GPU systems)
|
||||
model = model.cpu()
|
||||
model.device_used = str(device)
|
||||
else:
|
||||
raise ValueError(f"Unknown model type: {model_type}. Choose: xgboost, random_forest, decision_tree, svm, cnn")
|
||||
|
||||
# Fit non-CNN models
|
||||
if model_type != 'cnn':
|
||||
model.fit(X_train, y_train)
|
||||
|
||||
# Evaluate
|
||||
update_status("Evaluating model...", 90)
|
||||
if model_type == 'cnn':
|
||||
# PyTorch CNN evaluation
|
||||
train_score = model.score(X_train, y_train)
|
||||
test_score = model.score(X_test, y_test)
|
||||
y_pred = model.predict(X_test)
|
||||
else:
|
||||
train_score = model.score(X_train, y_train)
|
||||
test_score = model.score(X_test, y_test)
|
||||
y_pred = model.predict(X_test)
|
||||
|
||||
# Generate classification report and confusion matrix
|
||||
update_status("Generating classification report...", 92)
|
||||
class_names = label_encoder.classes_.tolist()
|
||||
|
||||
# Classification report as dict
|
||||
from sklearn.metrics import classification_report, confusion_matrix
|
||||
cls_report = classification_report(y_test, y_pred, target_names=class_names, output_dict=True, zero_division=0)
|
||||
|
||||
# Confusion matrix
|
||||
conf_matrix = confusion_matrix(y_test, y_pred).tolist()
|
||||
|
||||
# Save model
|
||||
update_status("Saving model...", 95)
|
||||
os.makedirs(os.path.dirname(output_model_path), exist_ok=True)
|
||||
joblib.dump({'model': model, 'label_encoder': label_encoder}, output_model_path)
|
||||
|
||||
# Save model info
|
||||
info = {
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": ["sentinel-2-l2a", "sentinel-1-rtc"],
|
||||
"features": ["NDVI_mean", "VH_dB_mean", "VV_dB_mean"],
|
||||
"training_samples": len(X_train),
|
||||
"testing_samples": len(X_test),
|
||||
"test_size": test_size,
|
||||
"train_accuracy": float(train_score),
|
||||
"test_accuracy": float(test_score),
|
||||
"model_type": model_type,
|
||||
"device": device if model_type == 'xgboost' else 'cpu',
|
||||
"n_estimators": n_estimators if model_type in ['xgboost', 'random_forest', 'cnn'] else None,
|
||||
"max_depth": max_depth if model_type != 'cnn' else None,
|
||||
"learning_rate": learning_rate if model_type == 'xgboost' else None,
|
||||
"cnn_epochs": min(50, n_estimators // 2) if model_type == 'cnn' else None,
|
||||
"n_features": X_train.shape[1],
|
||||
"n_classes": len(np.unique(y_train)),
|
||||
"class_names": class_names,
|
||||
"classification_report": cls_report,
|
||||
"confusion_matrix": conf_matrix,
|
||||
"bbox": bbox,
|
||||
"time_range": time_range,
|
||||
"resolution": resolution
|
||||
}
|
||||
|
||||
info_path = output_model_path.replace('.joblib', '_info.json')
|
||||
with open(info_path, 'w') as f:
|
||||
json.dump(info, f, indent=2)
|
||||
|
||||
update_status("Training complete!", 100)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"model_path": output_model_path,
|
||||
"info_path": info_path,
|
||||
"train_accuracy": train_score,
|
||||
"test_accuracy": test_score,
|
||||
"training_samples": len(X_train),
|
||||
"testing_samples": len(X_test),
|
||||
"test_size": test_size,
|
||||
"classes": class_names,
|
||||
"classification_report": cls_report,
|
||||
"confusion_matrix": conf_matrix,
|
||||
"model_type": model_type,
|
||||
"bbox": bbox,
|
||||
"time_range": time_range,
|
||||
"resolution": resolution
|
||||
}
|
||||
|
||||
except InterruptedError as e:
|
||||
update_status(f"Cancelled: {str(e)}", -1)
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"cancelled": True
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
update_status(f"Error: {str(e)}", -1)
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}
|
||||
Reference in New Issue
Block a user