Files
remote-sensing/01.train_ODC_MobileNet.ipynb
T
2026-03-07 17:14:01 +07:00

370 lines
15 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"id": "17da4353",
"metadata": {},
"outputs": [],
"source": [
"import importlib\n",
"import new_import_ODC as odc_tools\n",
"importlib.reload(odc_tools)\n",
"from new_import_ODC import *\n",
"print(\"✅ Import thành công\")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "9c063be3",
"metadata": {},
"outputs": [],
"source": [
"# Khởi tạo Dask + Datacube + S3\n",
"cluster, client = initialize_dask(use_gateway=True)\n",
"dc = datacube.Datacube()\n",
"configure_s3_access(aws_unsigned=True)\n",
"print(\"✅ Dask + Datacube + S3 sẵn sàng\")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "83784d01",
"metadata": {},
"outputs": [],
"source": [
"# Cấu hình vùng và thời gian\n",
"date_range = (\"2022-09-01\", \"2023-10-01\")\n",
"longtitude_range = (105.5, 106.4)\n",
"latitude_range = (9.2, 10.0)\n",
"\n",
"data_sen2 = load_data(\n",
" dc=dc,\n",
" date_range=date_range,\n",
" longtitude_range=longtitude_range,\n",
" latitude_range=latitude_range,\n",
")\n",
"print(f\"✅ Sentinel-2 raw: {data_sen2.dims}\")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c3faed92",
"metadata": {},
"outputs": [],
"source": [
"# Tiền xử lý Sentinel-2: cloud mask + NDVI + resampling\n",
"data_clean = mask_clean(data_sen2)\n",
"data_ndvi = calculate_indices(data_clean, index=\"NDVI\")\n",
"data_fill = fill_nan(data_ndvi)\n",
"data_sen2_monthly = data_fill.resample(time=\"1MS\").mean().compute()\n",
"print(f\"✅ S2 monthly shape: {data_sen2_monthly.dims}\")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "569bfebb",
"metadata": {},
"outputs": [],
"source": [
"# Tải Sentinel-1 (SAR VV/VH)\n",
"data_sen1 = load_data_sen1(\n",
" dc=dc,\n",
" date_range=date_range,\n",
" longtitude_range=longtitude_range,\n",
" latitude_range=latitude_range,\n",
")\n",
"data_sen1_monthly = calculate_average(data_sen1, [\"VV\", \"VH\"], resample=\"1MS\").compute()\n",
"print(f\"✅ S1 monthly shape: {data_sen1_monthly.dims}\")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "ecc56c2f",
"metadata": {},
"outputs": [],
"source": [
"import numpy as np\n",
"\n",
"# Ánh xạ nhãn lớp đất\n",
"label_mapping = {\n",
" \"Lua tom\": \"0\", \"Lua\": \"1\", \"CHN\": \"2\", \"CLN\": \"3\",\n",
" \"TS\": \"4\", \"Song\": \"5\", \"Dat xay dung\": \"6\", \"Rung\": \"7\",\n",
"}\n",
"\n",
"train_data = load_train_data(label_mapping=label_mapping)\n",
"X, y = get_data_sen1_and_sen2(train_data, data_sen2_monthly, data_sen1_monthly)\n",
"X_train, X_val, X_test, y_train, y_val, y_test = split_train_data(X, y, test_size=0.2, val_size=0.1)\n",
"\n",
"X_train_np = np.array(X_train, dtype=np.float32)\n",
"X_val_np = np.array(X_val, dtype=np.float32)\n",
"X_test_np = np.array(X_test, dtype=np.float32)\n",
"y_train_np = np.array(y_train, dtype=np.int64)\n",
"y_val_np = np.array(y_val, dtype=np.int64)\n",
"y_test_np = np.array(y_test, dtype=np.int64)\n",
"\n",
"n_features = X_train_np.shape[1]\n",
"n_classes = len(np.unique(y_train_np))\n",
"\n",
"print(f\"✅ Train: {X_train_np.shape} Val: {X_val_np.shape} Test: {X_test_np.shape}\")\n",
"print(f\" n_features={n_features} n_classes={n_classes}\")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "5492528b",
"metadata": {},
"outputs": [],
"source": [
"import torch\n",
"import torch.nn as nn\n",
"import torch.optim as optim\n",
"from torch.utils.data import TensorDataset, DataLoader\n",
"\n",
"# ── MobileNetV3 + LR-ASPP classifier ──────────────────────────────────────────\n",
"class MobileNetLRASPPClassifier(nn.Module):\n",
" \"\"\"\n",
" MobileNetV3-inspired backbone with LR-ASPP (Lite Reduced ASPP) head\n",
" for land-use classification on flat feature vectors.\n",
" \"\"\"\n",
" def __init__(self, n_features, n_classes):\n",
" super().__init__()\n",
" # Feature extraction backbone\n",
" self.feature_extractor = nn.Sequential(\n",
" nn.Linear(n_features, 128), nn.BatchNorm1d(128), nn.ReLU(inplace=True), nn.Dropout(0.2),\n",
" nn.Linear(128, 256), nn.BatchNorm1d(256), nn.ReLU(inplace=True), nn.Dropout(0.3),\n",
" nn.Linear(256, 512), nn.BatchNorm1d(512), nn.ReLU(inplace=True), nn.Dropout(0.3),\n",
" )\n",
" # LR-ASPP Branch 1: global pooling → 128\n",
" self.global_pool = nn.AdaptiveAvgPool1d(1)\n",
" self.global_conv = nn.Sequential(nn.Linear(512, 128), nn.ReLU(inplace=True))\n",
" # LR-ASPP Branch 2: direct 1×1 → 128\n",
" self.branch_conv = nn.Sequential(nn.Linear(512, 128), nn.BatchNorm1d(128), nn.ReLU(inplace=True))\n",
" # Fusion → n_classes\n",
" self.classifier = nn.Sequential(\n",
" nn.Linear(256, 128), nn.BatchNorm1d(128), nn.ReLU(inplace=True), nn.Dropout(0.4),\n",
" nn.Linear(128, n_classes),\n",
" )\n",
"\n",
" def forward(self, x):\n",
" feat = self.feature_extractor(x)\n",
" global_feat = self.global_pool(feat.unsqueeze(-1)).squeeze(-1)\n",
" global_feat = self.global_conv(global_feat)\n",
" branch_feat = self.branch_conv(feat)\n",
" fused = torch.cat([global_feat, branch_feat], dim=1)\n",
" return self.classifier(fused)\n",
"\n",
"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",
"print(f\"✅ Device: {device}\")\n",
"model = MobileNetLRASPPClassifier(n_features, n_classes).to(device)\n",
"print(model)\n",
"total_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n",
"print(f\" Trainable params: {total_params:,}\")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "9da40f5a",
"metadata": {},
"outputs": [],
"source": [
"%%time\n",
"import matplotlib.pyplot as plt\n",
"\n",
"# ── Hyper-parameters ───────────────────────────────────────────────────────────\n",
"LEARNING_RATE = 1e-3\n",
"BATCH_SIZE = 64\n",
"N_EPOCHS = 60\n",
"PATIENCE = 10\n",
"\n",
"# ── Tensors & DataLoaders ──────────────────────────────────────────────────────\n",
"X_tr_t = torch.FloatTensor(X_train_np)\n",
"y_tr_t = torch.LongTensor(y_train_np)\n",
"X_va_t = torch.FloatTensor(X_val_np)\n",
"y_va_t = torch.LongTensor(y_val_np)\n",
"\n",
"train_loader = DataLoader(TensorDataset(X_tr_t, y_tr_t), batch_size=BATCH_SIZE, shuffle=True)\n",
"val_loader = DataLoader(TensorDataset(X_va_t, y_va_t), batch_size=BATCH_SIZE, shuffle=False)\n",
"\n",
"# ── Class-weighted loss ────────────────────────────────────────────────────────\n",
"class_counts = np.bincount(y_train_np)\n",
"class_weights = 1.0 / (class_counts + 1e-6)\n",
"class_weights = class_weights / class_weights.sum() * n_classes\n",
"criterion = nn.CrossEntropyLoss(weight=torch.FloatTensor(class_weights).to(device))\n",
"\n",
"optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE, weight_decay=1e-4)\n",
"scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode=\"min\", factor=0.5, patience=5)\n",
"\n",
"# ── Training loop ──────────────────────────────────────────────────────────────\n",
"train_losses, val_losses, val_accs = [], [], []\n",
"best_val_loss = float(\"inf\")\n",
"best_state_dict = None\n",
"patience_counter = 0\n",
"\n",
"print(\"🚀 Training MobileNetV3 + LR-ASPP...\")\n",
"for epoch in range(1, N_EPOCHS + 1):\n",
" # --- Train ---\n",
" model.train()\n",
" epoch_loss = 0.0\n",
" for bx, by in train_loader:\n",
" bx, by = bx.to(device), by.to(device)\n",
" optimizer.zero_grad()\n",
" loss = criterion(model(bx), by)\n",
" loss.backward()\n",
" torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n",
" optimizer.step()\n",
" epoch_loss += loss.item()\n",
" avg_train_loss = epoch_loss / len(train_loader)\n",
"\n",
" # --- Validate ---\n",
" model.eval()\n",
" val_loss = 0.0; correct = 0; total = 0\n",
" with torch.no_grad():\n",
" for bx, by in val_loader:\n",
" bx, by = bx.to(device), by.to(device)\n",
" out = model(bx)\n",
" val_loss += criterion(out, by).item()\n",
" pred = out.argmax(1)\n",
" correct += (pred == by).sum().item()\n",
" total += by.size(0)\n",
" avg_val_loss = val_loss / len(val_loader)\n",
" val_acc = correct / total\n",
"\n",
" scheduler.step(avg_val_loss)\n",
" lr = optimizer.param_groups[0][\"lr\"]\n",
"\n",
" train_losses.append(avg_train_loss)\n",
" val_losses.append(avg_val_loss)\n",
" val_accs.append(val_acc)\n",
"\n",
" if epoch % 5 == 0 or epoch == 1:\n",
" print(f\"Epoch {epoch:3d}/{N_EPOCHS} train_loss={avg_train_loss:.4f} \"\n",
" f\"val_loss={avg_val_loss:.4f} val_acc={val_acc:.4f} lr={lr:.2e}\")\n",
"\n",
" # Early stopping\n",
" if avg_val_loss < best_val_loss:\n",
" best_val_loss = avg_val_loss\n",
" best_state_dict = {k: v.clone() for k, v in model.state_dict().items()}\n",
" patience_counter = 0\n",
" else:\n",
" patience_counter += 1\n",
" if patience_counter >= PATIENCE:\n",
" print(f\"⏹ Early stopping at epoch {epoch}.\")\n",
" break\n",
"\n",
"# Restore best weights\n",
"model.load_state_dict(best_state_dict)\n",
"print(f\"\\n✅ Training hoàn tất! Best val_loss={best_val_loss:.4f}\")\n",
"\n",
"# ── Learning curves ──────────────────────────────────────────────────────────\n",
"fig, axes = plt.subplots(1, 2, figsize=(13, 4))\n",
"axes[0].plot(train_losses, label=\"Train\")\n",
"axes[0].plot(val_losses, label=\"Val\")\n",
"axes[0].set_title(\"Loss\")\n",
"axes[0].legend()\n",
"axes[1].plot(val_accs)\n",
"axes[1].set_title(\"Val Accuracy\")\n",
"plt.tight_layout()\n",
"plt.show()\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "bc690b92",
"metadata": {},
"outputs": [],
"source": [
"from sklearn.metrics import classification_report, confusion_matrix, accuracy_score\n",
"import seaborn as sns\n",
"\n",
"# ── Đánh giá trên tập test ─────────────────────────────────────────────────────\n",
"model.eval()\n",
"with torch.no_grad():\n",
" X_te_t = torch.FloatTensor(X_test_np).to(device)\n",
" logits = model(X_te_t)\n",
" y_pred = logits.argmax(1).cpu().numpy()\n",
"\n",
"acc = accuracy_score(y_test_np, y_pred)\n",
"print(f\"Test Accuracy : {acc:.4f} ({acc*100:.2f}%)\\n\")\n",
"print(classification_report(y_test_np, y_pred, digits=4))\n",
"\n",
"# ── Confusion matrix ────────────────────────────────────────────────────────────\n",
"class_names = list(label_mapping.keys())\n",
"cm = confusion_matrix(y_test_np, y_pred)\n",
"plt.figure(figsize=(9, 7))\n",
"sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Purples\",\n",
" xticklabels=class_names, yticklabels=class_names)\n",
"plt.xlabel(\"Predicted\")\n",
"plt.ylabel(\"Actual\")\n",
"plt.title(\"Confusion Matrix — MobileNetV3 + LR-ASPP\")\n",
"plt.tight_layout()\n",
"plt.show()\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "28f87d0e",
"metadata": {},
"outputs": [],
"source": [
"import json\n",
"from datetime import datetime\n",
"\n",
"# ── Lưu mô hình PyTorch ────────────────────────────────────────────────────────\n",
"model_path = \"model_mobilenet_land_use.pth\"\n",
"torch.save(model.state_dict(), model_path)\n",
"print(f\"✅ Model saved → {model_path}\")\n",
"\n",
"# ── Lưu thông tin mô hình ──────────────────────────────────────────────────────\n",
"info = {\n",
" \"model_type\": \"MobileNetV3_LR-ASPP\",\n",
" \"n_features\": n_features,\n",
" \"n_classes\": n_classes,\n",
" \"learning_rate\": LEARNING_RATE,\n",
" \"batch_size\": BATCH_SIZE,\n",
" \"max_epochs\": N_EPOCHS,\n",
" \"early_stopping_patience\": PATIENCE,\n",
" \"optimizer\": \"Adam\",\n",
" \"scheduler\": \"ReduceLROnPlateau(factor=0.5, patience=5)\",\n",
" \"class_weight\": \"balanced\",\n",
" \"label_mapping\": label_mapping,\n",
" \"test_accuracy\": float(acc),\n",
" \"train_samples\": len(X_train_np),\n",
" \"val_samples\": len(X_val_np),\n",
" \"test_samples\": len(X_test_np),\n",
" \"saved_at\": datetime.now().isoformat(),\n",
"}\n",
"info_path = \"model_mobilenet_land_use_info.json\"\n",
"with open(info_path, \"w\") as f:\n",
" json.dump(info, f, indent=2, ensure_ascii=False)\n",
"print(f\"✅ Info saved → {info_path}\")\n",
"print(json.dumps(info, indent=2, ensure_ascii=False))\n",
"\n",
"# ── Đóng kết nối Dask ──────────────────────────────────────────────────────────\n",
"try:\n",
" client.close()\n",
" cluster.close()\n",
" print(\"✅ Dask cluster closed.\")\n",
"except Exception:\n",
" pass\n"
]
}
],
"metadata": {
"language_info": {
"name": "python"
}
},
"nbformat": 4,
"nbformat_minor": 5
}