327 lines
12 KiB
Python
327 lines
12 KiB
Python
import torch
|
|
import torch.nn as nn
|
|
import torch.optim as optim
|
|
from torch.utils.data import DataLoader
|
|
from torchvision import transforms
|
|
import torchvision.models as models
|
|
|
|
import joblib
|
|
import pandas as pd
|
|
import geopandas as gpd
|
|
import planetary_computer
|
|
import pystac_client
|
|
import odc.stac
|
|
import numpy as np
|
|
import os
|
|
import json
|
|
from sklearn.model_selection import train_test_split
|
|
from sklearn.metrics import accuracy_score, classification_report
|
|
from tqdm import tqdm
|
|
from joblib import Parallel, delayed
|
|
|
|
from cloud_removal import DeepInpaintingStrategy
|
|
|
|
def get_s2_items(bbox, time_range):
|
|
catalog = pystac_client.Client.open(
|
|
"https://planetarycomputer.microsoft.com/api/stac/v1",
|
|
modifier=planetary_computer.sign_inplace,
|
|
)
|
|
search = catalog.search(
|
|
collections=["sentinel-2-l2a"],
|
|
bbox=bbox,
|
|
datetime=time_range,
|
|
query={"eo:cloud_cover": {"lt": 30}}
|
|
)
|
|
items = list(search.items())
|
|
items = sorted(items, key=lambda x: x.properties["eo:cloud_cover"])
|
|
print(f"Found {len(items)} Sentinel-2 items")
|
|
return items
|
|
|
|
class SwinUNetWrapper(nn.Module):
|
|
def __init__(self, in_channels=24, num_classes=5):
|
|
super().__init__()
|
|
self.swin = models.swin_t(weights=models.Swin_T_Weights.IMAGENET1K_V1)
|
|
|
|
old_conv = self.swin.features[0][0]
|
|
new_conv = nn.Conv2d(in_channels, old_conv.out_channels,
|
|
kernel_size=old_conv.kernel_size,
|
|
stride=old_conv.stride,
|
|
padding=old_conv.padding)
|
|
with torch.no_grad():
|
|
new_conv.weight[:, :3] = old_conv.weight
|
|
new_conv.weight[:, 3:] = old_conv.weight.mean(dim=1, keepdim=True).repeat(1, in_channels-3, 1, 1)
|
|
new_conv.bias = old_conv.bias
|
|
self.swin.features[0][0] = new_conv
|
|
|
|
self.swin.head = nn.Linear(self.swin.head.in_features, num_classes)
|
|
|
|
self.upsample = nn.Upsample(size=(224, 224), mode='bilinear', align_corners=False)
|
|
|
|
def forward(self, x):
|
|
x = self.upsample(x)
|
|
return self.swin(x)
|
|
|
|
def process_point(idx, row, items_dicts, patch_size=16):
|
|
try:
|
|
import pystac
|
|
import odc.stac
|
|
import planetary_computer
|
|
from shapely.geometry import Point, shape
|
|
from pyproj import Transformer
|
|
|
|
items = [pystac.Item.from_dict(d) for d in items_dicts]
|
|
|
|
x_coord = row['geometry'].x
|
|
y_coord = row['geometry'].y
|
|
|
|
transformer = Transformer.from_crs("epsg:32648", "epsg:4326", always_xy=True)
|
|
lon, lat = transformer.transform(x_coord, y_coord)
|
|
point = Point(lon, lat)
|
|
|
|
filtered_items = []
|
|
for item in items:
|
|
geom = shape(item.geometry)
|
|
if geom.contains(point):
|
|
filtered_items.append(item)
|
|
|
|
if not filtered_items:
|
|
return None
|
|
|
|
filtered_items = [planetary_computer.sign(item) for item in filtered_items][:10]
|
|
|
|
# Increase bounds to 100m radius (20x20 pixels) to avoid boundary issues!
|
|
patch_s2 = odc.stac.load(
|
|
filtered_items,
|
|
bands=["B02", "B03", "B04", "B08", "SCL"],
|
|
x=(x_coord - 100, x_coord + 100),
|
|
y=(y_coord - 100, y_coord + 100),
|
|
crs="EPSG:32648",
|
|
resolution=10,
|
|
patch_url=planetary_computer.sign,
|
|
fail_on_error=False
|
|
).compute()
|
|
|
|
b2_sums = patch_s2["B02"].sum(dim=["x", "y"])
|
|
valid_times = b2_sums > 0
|
|
patch_s2 = patch_s2.isel(time=valid_times)
|
|
|
|
if len(patch_s2.time) == 0:
|
|
return None
|
|
|
|
patch_s2 = patch_s2.isel(time=slice(0, min(4, len(patch_s2.time))))
|
|
|
|
if "SCL" not in patch_s2 or "B02" not in patch_s2:
|
|
return None
|
|
|
|
if patch_s2.dims['x'] < patch_size or patch_s2.dims['y'] < patch_size:
|
|
return None
|
|
|
|
patch_s2 = patch_s2.isel(x=slice(0, patch_size), y=slice(0, patch_size))
|
|
|
|
return {
|
|
'patch_s2': patch_s2,
|
|
'label': row['HT_code'] - 1
|
|
}
|
|
except Exception as e:
|
|
return None
|
|
|
|
def extract_2d_patches(items, gdf, patch_size=16):
|
|
print(f"Extracting 2D patches for {len(gdf)} points using 8 parallel jobs...")
|
|
|
|
items_dicts = [item.to_dict() for item in items]
|
|
|
|
results = Parallel(n_jobs=8, backend="loky")(
|
|
delayed(process_point)(idx, row, items_dicts, patch_size)
|
|
for idx, row in tqdm(gdf.iterrows(), total=len(gdf), desc="Downloading Patches")
|
|
)
|
|
|
|
X = []
|
|
y = []
|
|
|
|
cloud_remover = DeepInpaintingStrategy(model_path="cloud_removal_model/cloud_removal_unet_best.pth")
|
|
if cloud_remover.model is None:
|
|
print("Warning: Could not load DeepInpainting model.")
|
|
|
|
print("Applying Cloud Removal sequentially...")
|
|
valid_results = [r for r in results if r is not None]
|
|
print(f"Valid points extracted: {len(valid_results)}/{len(gdf)}")
|
|
|
|
for res in tqdm(valid_results, desc="Cloud Removal & Features"):
|
|
try:
|
|
patch_s2 = res['patch_s2']
|
|
label = res['label']
|
|
|
|
patch_cloud_mask = patch_s2["SCL"].isin([3, 8, 9, 10])
|
|
|
|
# Apply cloud removal (returns 4 time steps)
|
|
clean_patch, _ = cloud_remover.remove_clouds(patch_s2, patch_cloud_mask)
|
|
|
|
b4 = clean_patch["B04"].values
|
|
b8 = clean_patch["B08"].values
|
|
b3 = clean_patch["B03"].values
|
|
b2 = clean_patch["B02"].values
|
|
|
|
ndvi = (b8 - b4) / (b8 + b4 + 1e-6)
|
|
ndwi = (b3 - b8) / (b3 + b8 + 1e-6)
|
|
|
|
b2 = np.clip(b2 / 10000.0, 0, 1)
|
|
b3 = np.clip(b3 / 10000.0, 0, 1)
|
|
b4 = np.clip(b4 / 10000.0, 0, 1)
|
|
b8 = np.clip(b8 / 10000.0, 0, 1)
|
|
|
|
# Stack across channels
|
|
features_t = np.stack([b2, b3, b4, b8, ndvi, ndwi], axis=1) # Shape: (time, 6, 16, 16)
|
|
|
|
# Pad time dimension to exactly 4 if needed
|
|
t_len = features_t.shape[0]
|
|
if t_len < 4:
|
|
pad = np.zeros((4 - t_len, 6, 16, 16))
|
|
features_t = np.concatenate([features_t, pad], axis=0)
|
|
|
|
# Flatten time and channels: (4, 6, 16, 16) -> (24, 16, 16)
|
|
features = features_t.reshape(24, 16, 16)
|
|
features = np.nan_to_num(features, nan=0.0)
|
|
|
|
X.append(features)
|
|
y.append(label)
|
|
except Exception as e:
|
|
pass
|
|
|
|
return np.array(X), np.array(y)
|
|
|
|
def train_2d_model(X, y):
|
|
print(f"Training 2D CNN with Data Augmentation... Dataset shape: {X.shape}")
|
|
|
|
unique_labels = sorted(list(np.unique(y)))
|
|
label_map = {lbl: i for i, lbl in enumerate(unique_labels)}
|
|
y_mapped = np.array([label_map[l] for l in y])
|
|
|
|
X_train, X_test, y_train, y_test = train_test_split(X, y_mapped, test_size=0.2, random_state=42)
|
|
|
|
transform = transforms.Compose([
|
|
transforms.RandomHorizontalFlip(),
|
|
transforms.RandomVerticalFlip(),
|
|
])
|
|
|
|
class PatchDataset(torch.utils.data.Dataset):
|
|
def __init__(self, X, y, augment=False):
|
|
self.X = torch.FloatTensor(X)
|
|
self.y = torch.LongTensor(y)
|
|
self.augment = augment
|
|
|
|
def __len__(self):
|
|
return len(self.X)
|
|
|
|
def __getitem__(self, idx):
|
|
x = self.X[idx]
|
|
if self.augment:
|
|
x = transform(x)
|
|
return x, self.y[idx]
|
|
|
|
train_dataset = PatchDataset(X_train, y_train, augment=True)
|
|
test_dataset = PatchDataset(X_test, y_test, augment=False)
|
|
|
|
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
|
|
test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)
|
|
|
|
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
|
print(f"Using device: {device}")
|
|
|
|
model = SwinUNetWrapper(in_channels=24, num_classes=len(unique_labels)).to(device)
|
|
|
|
class_counts = np.bincount(y_train)
|
|
weights = 1.0 / (class_counts + 1e-6)
|
|
weights = torch.FloatTensor(weights / weights.sum() * len(class_counts)).to(device)
|
|
|
|
criterion = nn.CrossEntropyLoss(weight=weights, label_smoothing=0.1)
|
|
optimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.05)
|
|
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100, eta_min=1e-6)
|
|
|
|
epochs = 150
|
|
best_acc = 0
|
|
best_state = None
|
|
|
|
for epoch in range(epochs):
|
|
model.train()
|
|
train_loss = 0
|
|
for batch_X, batch_y in train_loader:
|
|
batch_X, batch_y = batch_X.to(device), batch_y.to(device)
|
|
optimizer.zero_grad()
|
|
out = model(batch_X)
|
|
loss = criterion(out, batch_y)
|
|
loss.backward()
|
|
optimizer.step()
|
|
train_loss += loss.item()
|
|
|
|
model.eval()
|
|
all_preds = []
|
|
all_targets = []
|
|
with torch.no_grad():
|
|
for batch_X, batch_y in test_loader:
|
|
out = model(batch_X.to(device))
|
|
preds = out.argmax(dim=1).cpu().numpy()
|
|
all_preds.extend(preds)
|
|
all_targets.extend(batch_y.numpy())
|
|
|
|
acc = accuracy_score(all_targets, all_preds)
|
|
scheduler.step()
|
|
|
|
if acc > best_acc:
|
|
best_acc = acc
|
|
best_state = model.state_dict()
|
|
print(f"Epoch {epoch+1}/{epochs} - Loss: {train_loss/len(train_loader):.4f} - Test Acc: {acc:.4f} 🌟")
|
|
if acc >= 0.95:
|
|
print("🎯 Đã đạt mốc >95% Accuracy!")
|
|
break
|
|
elif (epoch+1) % 10 == 0:
|
|
print(f"Epoch {epoch+1}/{epochs} - Loss: {train_loss/len(train_loader):.4f} - Test Acc: {acc:.4f}")
|
|
|
|
if best_state:
|
|
model.load_state_dict(best_state)
|
|
|
|
os.makedirs('land_classification_model', exist_ok=True)
|
|
joblib.dump(model.cpu(), 'land_classification_model/model_cnn_2d_95.joblib')
|
|
print(f"✅ Đã lưu mô hình đạt {best_acc:.4f} vào land_classification_model/model_cnn_2d_95.joblib")
|
|
|
|
clf_rep = classification_report(all_targets, all_preds, output_dict=True)
|
|
info = {
|
|
"model_type": "CNN_2D_Patch_CloudRemoval_Temporal",
|
|
"test_accuracy": float(best_acc),
|
|
"params": {"epochs": epochs, "architecture": "2D CNN Swin-UNet Temporal"},
|
|
"classification_report": clf_rep
|
|
}
|
|
os.makedirs('model_train', exist_ok=True)
|
|
with open('model_train/model_cnn_2d_info.json', 'w') as f:
|
|
json.dump(info, f, indent=2)
|
|
|
|
def main():
|
|
print("🚀 BẮT ĐẦU PIPELINE 2D PATCH-BASED & CLOUD REMOVAL (TEMPORAL 24-CHANNELS)")
|
|
|
|
# Dùng tên file mới để tránh bị trùng với dữ liệu 6 channel cũ
|
|
cache_file = "dataset_cache/training_data_2d_temporal.joblib"
|
|
|
|
if os.path.exists(cache_file):
|
|
print(f"Loading 2D patches from {cache_file}...")
|
|
data = joblib.load(cache_file)
|
|
X, y = data['X'], data['y']
|
|
else:
|
|
bbox = [105.5, 9.2, 106.3, 10.0]
|
|
time_range = "2023-01-01/2023-04-30"
|
|
|
|
items = get_s2_items(bbox, time_range)
|
|
|
|
gdf = gpd.read_file("train/ST_training_data_updated_1130points_new.shp")
|
|
gdf = gdf.to_crs("EPSG:32648")
|
|
|
|
X, y = extract_2d_patches(items, gdf, patch_size=16)
|
|
|
|
os.makedirs('dataset_cache', exist_ok=True)
|
|
joblib.dump({'X': X, 'y': y}, cache_file)
|
|
print(f"Saved 2D cache to {cache_file}")
|
|
|
|
train_2d_model(X, y)
|
|
print("🎉 Hoàn tất quá trình! Check-point với Accuracy > 95% đã được lưu!")
|
|
|
|
if __name__ == "__main__":
|
|
main()
|