Merge branch 'XGBoost'
This commit is contained in:
@@ -0,0 +1,7 @@
|
||||
ThuanHoa/ThuanHoa_VH.tif
|
||||
ThuanHoa/ThuanHoa_VV.tif
|
||||
model_train/model.joblib
|
||||
model_train/model_new.joblib
|
||||
backup_model_train/model.joblib
|
||||
backup_model_train/model_new.joblib
|
||||
dataset_cache/sentinel2_timeseries_40scenes.nc
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:dd3fecbef4c80ff1bd20922e6b2250d94ac05e55c3a2efbff4da20fcea3d739f
|
||||
size 39812
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:350b680a5454d4d080d8d943d723ea01668c9c19219cb8feb396c54462ee89c3
|
||||
size 857629
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:89c347e98288e77ddb7579d4cfa2b49924a993f7173a2d12155ef7a1893082b1
|
||||
size 163627
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:54efd837e5dcbfbd4bd8c2cfc44ff1b5eecbf03de302d41fab0ed47257bfa62c
|
||||
size 4828615
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:72d0f97e1e623ab49fe749054c8b2c3901ff0807b52667289108bbd0b8220e2a
|
||||
size 2240648
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:ff33ab950e3524342a99e70a02598ccb1ca57e60627c887c10b34466eabf3e87
|
||||
size 3536032
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:4a62b68a633d79c53a6fd8893e8ea42dcf2b9a8a3e907b1b9861661f04f21517
|
||||
size 72
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:520c3480efaafdf844c3bba3dd165a8e44f3318ab21209539c78201f18fd6312
|
||||
size 3066028
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7d3bff20682fc510a9c41b5aa2b42e3ebef8098ed81e27b4436f9ae86f47b9a6
|
||||
size 3066028
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:322f48d2d54ee921facfc3c0c44b320a86a90694e1779b683824264f9affa3fc
|
||||
size 3066028
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:2948384b4c43483675d366b75b7bcdceb63fe5a8ad7d9b7a84ee4a1018cfc826
|
||||
size 55333632
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:1d5a53bba7a6db30b7f80bc4250d09edacfd244e97d7993382927a9bde85c6f7
|
||||
size 55738194
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+788
@@ -0,0 +1,788 @@
|
||||
"""
|
||||
API Server for Land Classification Model Training
|
||||
Cho phép chọn dữ liệu và cấu hình training qua giao diện web
|
||||
"""
|
||||
|
||||
from fastapi import FastAPI, BackgroundTasks, HTTPException
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.responses import HTMLResponse, FileResponse
|
||||
from pydantic import BaseModel
|
||||
from typing import Optional, List
|
||||
import uvicorn
|
||||
import joblib
|
||||
import json
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
app = FastAPI(title="Land Classification Training API", version="1.0.0")
|
||||
|
||||
# Enable CORS
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# Global training status
|
||||
training_status = {
|
||||
"is_training": False,
|
||||
"progress": "",
|
||||
"error": None,
|
||||
"result": None,
|
||||
"start_time": None,
|
||||
"end_time": None,
|
||||
"cancel_requested": False
|
||||
}
|
||||
|
||||
# Global prediction status
|
||||
prediction_status = {
|
||||
"is_predicting": False,
|
||||
"progress": "",
|
||||
"error": None,
|
||||
"result": None,
|
||||
"output_file": None,
|
||||
"start_time": None,
|
||||
"end_time": None
|
||||
}
|
||||
|
||||
|
||||
class TrainingConfig(BaseModel):
|
||||
"""Cấu hình training"""
|
||||
# Khu vực (bbox)
|
||||
min_lon: float = 105.6
|
||||
min_lat: float = 9.3
|
||||
max_lon: float = 106.2
|
||||
max_lat: float = 9.8
|
||||
|
||||
# Thời gian
|
||||
start_date: str = "2023-03-01"
|
||||
end_date: str = "2023-05-31"
|
||||
|
||||
# Dữ liệu
|
||||
max_scenes: int = 12
|
||||
cloud_cover: int = 30
|
||||
resolution: int = 20 # 10m hoặc 20m
|
||||
|
||||
# Model parameters
|
||||
model_type: str = "xgboost" # xgboost, random_forest, decision_tree, svm, cnn
|
||||
n_estimators: int = 100
|
||||
max_depth: int = 20
|
||||
learning_rate: float = 0.1
|
||||
use_gpu: bool = True
|
||||
|
||||
# Cache
|
||||
use_cache: bool = True # Cache dataset để test nhanh hơn
|
||||
|
||||
# Training data
|
||||
training_shapefile: str = "train/ST_training data_updated_1130points_new.shp"
|
||||
|
||||
|
||||
class PredictionConfig(BaseModel):
|
||||
"""Cấu hình dự đoán"""
|
||||
# Model to use
|
||||
model_filename: str
|
||||
|
||||
# Khu vực (bbox)
|
||||
min_lon: float = 105.6
|
||||
min_lat: float = 9.3
|
||||
max_lon: float = 106.2
|
||||
max_lat: float = 9.8
|
||||
|
||||
# Thời gian
|
||||
start_date: str = "2023-03-01"
|
||||
end_date: str = "2023-05-31"
|
||||
|
||||
# Dữ liệu
|
||||
max_scenes: int = 12
|
||||
cloud_cover: int = 30
|
||||
resolution: int = 20
|
||||
|
||||
|
||||
class TrainingStatus(BaseModel):
|
||||
"""Trạng thái training"""
|
||||
is_training: bool
|
||||
progress: str
|
||||
error: Optional[str]
|
||||
result: Optional[dict]
|
||||
start_time: Optional[str]
|
||||
end_time: Optional[str]
|
||||
|
||||
|
||||
@app.get("/", response_class=HTMLResponse)
|
||||
async def root():
|
||||
"""Serve giao diện web"""
|
||||
html_file = Path(__file__).parent / "training_interface.html"
|
||||
if html_file.exists():
|
||||
return FileResponse(html_file)
|
||||
else:
|
||||
return HTMLResponse("""
|
||||
<html>
|
||||
<head><title>Training Interface</title></head>
|
||||
<body>
|
||||
<h1>Land Classification Training API</h1>
|
||||
<p>API Documentation: <a href="/docs">/docs</a></p>
|
||||
<p>Training Interface: Tạo file training_interface.html</p>
|
||||
</body>
|
||||
</html>
|
||||
""")
|
||||
|
||||
|
||||
@app.get("/api/config/presets")
|
||||
async def get_presets():
|
||||
"""Lấy các preset cấu hình sẵn"""
|
||||
return {
|
||||
"presets": [
|
||||
{
|
||||
"name": "PC - Nhỏ (3 tháng, 20m, 12 scenes)",
|
||||
"config": {
|
||||
"min_lon": 105.6, "min_lat": 9.3, "max_lon": 106.2, "max_lat": 9.8,
|
||||
"start_date": "2023-03-01", "end_date": "2023-05-31",
|
||||
"max_scenes": 12, "cloud_cover": 30, "resolution": 20
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "Server - Trung bình (6 tháng, 10m, 30 scenes)",
|
||||
"config": {
|
||||
"min_lon": 105.5, "min_lat": 9.2, "max_lon": 106.4, "max_lat": 10.0,
|
||||
"start_date": "2023-01-01", "end_date": "2023-06-30",
|
||||
"max_scenes": 30, "cloud_cover": 30, "resolution": 10
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "Full - Lớn (1 năm, 10m, 60 scenes)",
|
||||
"config": {
|
||||
"min_lon": 105.5, "min_lat": 9.2, "max_lon": 106.4, "max_lat": 10.0,
|
||||
"start_date": "2022-09-01", "end_date": "2023-10-01",
|
||||
"max_scenes": 60, "cloud_cover": 50, "resolution": 10
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@app.get("/api/training/status", response_model=TrainingStatus)
|
||||
async def get_training_status():
|
||||
"""Kiểm tra trạng thái training"""
|
||||
return training_status
|
||||
|
||||
|
||||
@app.post("/api/training/start")
|
||||
async def start_training(config: TrainingConfig, background_tasks: BackgroundTasks):
|
||||
"""Bắt đầu training với config đã chọn"""
|
||||
global training_status
|
||||
|
||||
if training_status["is_training"]:
|
||||
raise HTTPException(status_code=400, detail="Training đang chạy, vui lòng đợi")
|
||||
|
||||
# Reset status
|
||||
training_status = {
|
||||
"is_training": True,
|
||||
"progress": "Đang khởi tạo...",
|
||||
"error": None,
|
||||
"result": None,
|
||||
"start_time": datetime.now().isoformat(),
|
||||
"end_time": None
|
||||
}
|
||||
|
||||
# Run training in background
|
||||
background_tasks.add_task(run_training, config)
|
||||
|
||||
return {"message": "Training đã bắt đầu", "status": training_status}
|
||||
|
||||
|
||||
@app.post("/api/training/stop")
|
||||
async def stop_training():
|
||||
"""Dừng training (nếu đang chạy)"""
|
||||
global training_status
|
||||
|
||||
if not training_status["is_training"]:
|
||||
return {"message": "Không có training nào đang chạy"}
|
||||
|
||||
# Set cancel flag - the training will check this and stop
|
||||
training_status["cancel_requested"] = True
|
||||
training_status["progress"] = "Đang hủy training..."
|
||||
|
||||
return {"message": "Đang dừng training..."}
|
||||
|
||||
|
||||
@app.post("/api/cache/clear")
|
||||
async def clear_cache():
|
||||
"""Xóa cache dataset"""
|
||||
import shutil
|
||||
cache_dir = Path("dataset_cache")
|
||||
|
||||
if not cache_dir.exists():
|
||||
return {"message": "Không có cache để xóa", "deleted": 0}
|
||||
|
||||
# Count files
|
||||
cache_files = list(cache_dir.glob("*.joblib"))
|
||||
count = len(cache_files)
|
||||
|
||||
# Delete all cache files
|
||||
for cache_file in cache_files:
|
||||
try:
|
||||
cache_file.unlink()
|
||||
except:
|
||||
pass
|
||||
|
||||
return {"message": f"Đã xóa {count} file cache", "deleted": count}
|
||||
|
||||
|
||||
@app.get("/api/cache/info")
|
||||
async def get_cache_info():
|
||||
"""Lấy thông tin về cache với metadata đầy đủ"""
|
||||
cache_dir = Path("dataset_cache")
|
||||
|
||||
if not cache_dir.exists():
|
||||
return {"exists": False, "files": [], "total_size_mb": 0}
|
||||
|
||||
cache_files = []
|
||||
total_size = 0
|
||||
|
||||
for cache_file in cache_dir.glob("*.joblib"):
|
||||
size = cache_file.stat().st_size
|
||||
total_size += size
|
||||
|
||||
# Try to load metadata from cache
|
||||
metadata = {}
|
||||
try:
|
||||
cached_data = joblib.load(cache_file)
|
||||
if isinstance(cached_data, dict):
|
||||
metadata = {
|
||||
"bbox": cached_data.get("bbox", []),
|
||||
"time_range": cached_data.get("time_range", ""),
|
||||
"resolution": cached_data.get("resolution", 20),
|
||||
"n_samples": len(cached_data.get("features", [])),
|
||||
"created": cached_data.get("timestamp", "")
|
||||
}
|
||||
# Parse time_range to get start/end dates
|
||||
if metadata["time_range"]:
|
||||
time_parts = metadata["time_range"].split("/")
|
||||
if len(time_parts) == 2:
|
||||
metadata["start_date"] = time_parts[0]
|
||||
metadata["end_date"] = time_parts[1]
|
||||
# Parse bbox to get min/max lon/lat
|
||||
if metadata["bbox"] and len(metadata["bbox"]) == 4:
|
||||
metadata["min_lon"] = metadata["bbox"][0]
|
||||
metadata["min_lat"] = metadata["bbox"][1]
|
||||
metadata["max_lon"] = metadata["bbox"][2]
|
||||
metadata["max_lat"] = metadata["bbox"][3]
|
||||
except Exception as e:
|
||||
print(f"Error loading cache metadata: {e}")
|
||||
|
||||
cache_files.append({
|
||||
"filename": cache_file.name,
|
||||
"size_mb": round(size / 1024 / 1024, 2),
|
||||
"modified": datetime.fromtimestamp(cache_file.stat().st_mtime).isoformat(),
|
||||
"metadata": metadata
|
||||
})
|
||||
|
||||
# Sort by modified time (newest first)
|
||||
cache_files.sort(key=lambda x: x["modified"], reverse=True)
|
||||
|
||||
return {
|
||||
"exists": True,
|
||||
"files": cache_files,
|
||||
"count": len(cache_files),
|
||||
"total_size_mb": round(total_size / 1024 / 1024, 2)
|
||||
}
|
||||
|
||||
|
||||
|
||||
@app.get("/api/models/list")
|
||||
async def list_models():
|
||||
"""Liệt kê các model đã train"""
|
||||
model_dir = Path("model_train")
|
||||
if not model_dir.exists():
|
||||
return {"models": []}
|
||||
|
||||
models = []
|
||||
for model_file in model_dir.glob("*.joblib"):
|
||||
info_file = model_file.with_suffix('.json')
|
||||
info = {}
|
||||
if info_file.exists():
|
||||
with open(info_file) as f:
|
||||
info = json.load(f)
|
||||
|
||||
models.append({
|
||||
"filename": model_file.name,
|
||||
"created": datetime.fromtimestamp(model_file.stat().st_mtime).isoformat(),
|
||||
"size_mb": round(model_file.stat().st_size / 1024 / 1024, 2),
|
||||
"info": info
|
||||
})
|
||||
|
||||
# Sort by creation time (newest first)
|
||||
models.sort(key=lambda x: x["created"], reverse=True)
|
||||
return {"models": models}
|
||||
|
||||
|
||||
@app.post("/api/prediction/start")
|
||||
async def start_prediction(config: PredictionConfig, background_tasks: BackgroundTasks):
|
||||
"""Bắt đầu dự đoán"""
|
||||
global prediction_status
|
||||
|
||||
if prediction_status["is_predicting"]:
|
||||
raise HTTPException(status_code=400, detail="Đang có dự đoán khác đang chạy")
|
||||
|
||||
# Reset status
|
||||
prediction_status = {
|
||||
"is_predicting": True,
|
||||
"progress": "Đang khởi động...",
|
||||
"error": None,
|
||||
"result": None,
|
||||
"output_file": None,
|
||||
"start_time": datetime.now().isoformat(),
|
||||
"end_time": None
|
||||
}
|
||||
|
||||
# Run prediction in background
|
||||
background_tasks.add_task(run_prediction, config)
|
||||
|
||||
return {"message": "Đã bắt đầu dự đoán", "status": prediction_status}
|
||||
|
||||
|
||||
@app.get("/api/prediction/status")
|
||||
async def get_prediction_status():
|
||||
"""Kiểm tra trạng thái dự đoán"""
|
||||
return prediction_status
|
||||
|
||||
|
||||
async def run_training(config: TrainingConfig):
|
||||
"""Chạy training process"""
|
||||
global training_status
|
||||
|
||||
try:
|
||||
training_status["cancel_requested"] = False
|
||||
training_status["progress"] = "Đang import thư viện..."
|
||||
|
||||
# Import training module
|
||||
from train_module import train_model
|
||||
|
||||
training_status["progress"] = "Đang load dữ liệu Sentinel-2..."
|
||||
|
||||
# Function to check if training should be cancelled
|
||||
def should_cancel():
|
||||
return training_status.get("cancel_requested", False)
|
||||
|
||||
# Run training
|
||||
result = train_model(
|
||||
bbox=[config.min_lon, config.min_lat, config.max_lon, config.max_lat],
|
||||
time_range=f"{config.start_date}/{config.end_date}",
|
||||
max_scenes=config.max_scenes,
|
||||
cloud_cover=config.cloud_cover,
|
||||
resolution=config.resolution,
|
||||
training_shapefile=config.training_shapefile,
|
||||
model_type=config.model_type,
|
||||
n_estimators=config.n_estimators,
|
||||
max_depth=config.max_depth,
|
||||
learning_rate=config.learning_rate,
|
||||
use_gpu=config.use_gpu,
|
||||
use_cache=config.use_cache,
|
||||
status_callback=lambda msg: update_progress(msg),
|
||||
cancel_check=should_cancel
|
||||
)
|
||||
|
||||
if training_status.get("cancel_requested", False):
|
||||
training_status["is_training"] = False
|
||||
training_status["progress"] = "Đã hủy training"
|
||||
training_status["error"] = "Training cancelled by user"
|
||||
else:
|
||||
training_status["is_training"] = False
|
||||
training_status["progress"] = "Hoàn thành!"
|
||||
training_status["result"] = result
|
||||
training_status["end_time"] = datetime.now().isoformat()
|
||||
|
||||
except Exception as e:
|
||||
training_status["is_training"] = False
|
||||
training_status["error"] = str(e)
|
||||
training_status["progress"] = f"Lỗi: {str(e)}"
|
||||
training_status["end_time"] = datetime.now().isoformat()
|
||||
import traceback
|
||||
print(traceback.format_exc())
|
||||
|
||||
|
||||
def update_progress(message: str):
|
||||
"""Cập nhật progress message"""
|
||||
global training_status
|
||||
training_status["progress"] = message
|
||||
print(f"[PROGRESS] {message}")
|
||||
|
||||
|
||||
def update_prediction_progress(message: str):
|
||||
"""Cập nhật prediction progress message"""
|
||||
global prediction_status
|
||||
prediction_status["progress"] = message
|
||||
print(f"[PREDICTION PROGRESS] {message}")
|
||||
|
||||
|
||||
async def run_prediction(config: PredictionConfig):
|
||||
"""Chạy prediction process - Áp dụng phương pháp từ 02.predict_ODC.ipynb"""
|
||||
global prediction_status
|
||||
|
||||
try:
|
||||
prediction_status["progress"] = "Đang import thư viện..."
|
||||
|
||||
# Import required libraries
|
||||
import xarray as xr
|
||||
import numpy as np
|
||||
from datetime import datetime as dt
|
||||
import rioxarray
|
||||
import dask.array as da
|
||||
|
||||
prediction_status["progress"] = "Đang load model..."
|
||||
|
||||
# Load model
|
||||
model_path = Path("model_train") / config.model_filename
|
||||
if not model_path.exists():
|
||||
raise FileNotFoundError(f"Model không tồn tại: {config.model_filename}")
|
||||
|
||||
model_data = joblib.load(model_path)
|
||||
|
||||
# Extract model from dict (models are saved as {'model': xgb_model, 'label_encoder': encoder})
|
||||
if isinstance(model_data, dict):
|
||||
model = model_data.get('model')
|
||||
label_encoder = model_data.get('label_encoder')
|
||||
else:
|
||||
model = model_data
|
||||
label_encoder = None
|
||||
|
||||
# Check if it's a CNN model (PyTorch)
|
||||
is_cnn_model = hasattr(model, '__class__') and 'CNN' in model.__class__.__name__
|
||||
if is_cnn_model:
|
||||
prediction_status["progress"] = "Phát hiện PyTorch CNN model..."
|
||||
# Import PyTorch if needed
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
raise ImportError("PyTorch is required for CNN prediction. Install: pip install torch")
|
||||
|
||||
prediction_status["progress"] = "Đang kết nối Microsoft Planetary Computer..."
|
||||
|
||||
# Import and use Microsoft Planetary Computer STAC API
|
||||
import pystac_client
|
||||
import planetary_computer
|
||||
from odc.stac import load
|
||||
|
||||
catalog = pystac_client.Client.open(
|
||||
"https://planetarycomputer.microsoft.com/api/stac/v1",
|
||||
modifier=planetary_computer.sign_inplace,
|
||||
)
|
||||
|
||||
bbox = [config.min_lon, config.min_lat, config.max_lon, config.max_lat]
|
||||
time_range = f"{config.start_date}/{config.end_date}"
|
||||
|
||||
# ============ BƯỚC 1: TẢI DỮ LIỆU SENTINEL-2 ============
|
||||
prediction_status["progress"] = "Đang tải dữ liệu Sentinel-2..."
|
||||
|
||||
# Search Sentinel-2 data
|
||||
s2_search = catalog.search(
|
||||
collections=["sentinel-2-l2a"],
|
||||
bbox=bbox,
|
||||
datetime=time_range,
|
||||
query={"eo:cloud_cover": {"lt": config.cloud_cover}}
|
||||
)
|
||||
|
||||
s2_items = list(s2_search.items())
|
||||
if not s2_items:
|
||||
raise ValueError("Không tìm thấy dữ liệu Sentinel-2 cho khu vực và thời gian này")
|
||||
|
||||
s2_items = s2_items[:config.max_scenes]
|
||||
prediction_status["progress"] = f"Đang xử lý {len(s2_items)} scenes Sentinel-2..."
|
||||
|
||||
# Load Sentinel-2 data
|
||||
s2_data = load(
|
||||
s2_items,
|
||||
bbox=bbox,
|
||||
chunks={"time": 1, "x": 2048, "y": 2048},
|
||||
groupby="solar_day",
|
||||
resolution=config.resolution
|
||||
)
|
||||
|
||||
# ============ BƯỚC 2: TÍNH NDVI VÀ XỬ LÝ MÂY ============
|
||||
prediction_status["progress"] = "Đang tính toán NDVI và xử lý mây..."
|
||||
|
||||
# Calculate NDVI using Sentinel-2 band names (B08 = NIR, B04 = Red)
|
||||
nir = s2_data["B08"].astype('float32')
|
||||
red = s2_data["B04"].astype('float32')
|
||||
ndvi = (nir - red) / (nir + red + 1e-8)
|
||||
|
||||
# Mask clouds using SCL band if available
|
||||
if "SCL" in s2_data:
|
||||
scl = s2_data["SCL"]
|
||||
# SCL values: 4=vegetation, 5=bare soil, 6=water - these are clear
|
||||
# 3=cloud shadow, 8=cloud medium, 9=cloud high, 10=cirrus - mask these
|
||||
cloud_mask = (scl == 3) | (scl == 8) | (scl == 9) | (scl == 10)
|
||||
ndvi = ndvi.where(~cloud_mask)
|
||||
|
||||
# ============ BƯỚC 3: ĐIỀN GIÁ TRỊ NAN (FILL NAN) ============
|
||||
prediction_status["progress"] = "Đang điền giá trị bị che mây..."
|
||||
|
||||
# Fill NaN using forward fill and backward fill
|
||||
ndvi_filled = ndvi.ffill(dim='time').bfill(dim='time')
|
||||
|
||||
# Resample to monthly average
|
||||
prediction_status["progress"] = "Đang tính trung bình NDVI theo tháng..."
|
||||
ndvi_monthly = ndvi_filled.resample(time="1ME").mean()
|
||||
|
||||
# Compute NDVI (convert from dask to numpy)
|
||||
ndvi_monthly = ndvi_monthly.compute()
|
||||
|
||||
# ============ BƯỚC 4: TẢI DỮ LIỆU SENTINEL-1 (VH, VV) ============
|
||||
prediction_status["progress"] = "Đang tải dữ liệu Sentinel-1 (Radar)..."
|
||||
|
||||
# Search Sentinel-1 data
|
||||
s1_search = catalog.search(
|
||||
collections=["sentinel-1-rtc"],
|
||||
bbox=bbox,
|
||||
datetime=time_range,
|
||||
)
|
||||
|
||||
s1_items = list(s1_search.items())
|
||||
|
||||
if s1_items:
|
||||
s1_items = s1_items[:config.max_scenes]
|
||||
prediction_status["progress"] = f"Đang xử lý {len(s1_items)} scenes Sentinel-1..."
|
||||
|
||||
# Load Sentinel-1 data (without like= to avoid conflict with bbox/resolution)
|
||||
s1_data = load(
|
||||
s1_items,
|
||||
bbox=bbox,
|
||||
chunks={"time": 1, "x": 2048, "y": 2048},
|
||||
groupby="sat:absolute_orbit",
|
||||
resolution=config.resolution
|
||||
)
|
||||
|
||||
# Extract VH and VV bands
|
||||
if "vh" in s1_data and "vv" in s1_data:
|
||||
vh = s1_data["vh"].astype('float32')
|
||||
vv = s1_data["vv"].astype('float32')
|
||||
|
||||
# Resample to monthly average
|
||||
prediction_status["progress"] = "Đang tính trung bình VH/VV theo tháng..."
|
||||
vh_monthly = vh.resample(time="1ME").mean().compute()
|
||||
vv_monthly = vv.resample(time="1ME").mean().compute()
|
||||
|
||||
use_radar = True
|
||||
else:
|
||||
prediction_status["progress"] = "Không tìm thấy bands VH/VV, tiếp tục với NDVI..."
|
||||
use_radar = False
|
||||
else:
|
||||
prediction_status["progress"] = "Không có dữ liệu Sentinel-1, tiếp tục với NDVI..."
|
||||
use_radar = False
|
||||
|
||||
# ============ BƯỚC 5: CHUẨN BỊ FEATURES CHO DỰ ĐOÁN ============
|
||||
prediction_status["progress"] = "Đang chuẩn bị features cho dự đoán..."
|
||||
|
||||
# Get shape information
|
||||
n_times_ndvi = len(ndvi_monthly.time)
|
||||
y_size = len(ndvi_monthly.y)
|
||||
x_size = len(ndvi_monthly.x)
|
||||
n_pixels = y_size * x_size
|
||||
|
||||
# Prepare NDVI features (flatten each time step)
|
||||
ndvi_features = []
|
||||
for t in range(n_times_ndvi):
|
||||
ndvi_t = ndvi_monthly.isel(time=t).values.flatten()
|
||||
ndvi_features.append(ndvi_t)
|
||||
|
||||
# Stack NDVI features
|
||||
features = np.column_stack(ndvi_features)
|
||||
|
||||
# Add radar features if available
|
||||
if use_radar:
|
||||
n_times_vh = len(vh_monthly.time)
|
||||
n_times_vv = len(vv_monthly.time)
|
||||
|
||||
# Add VH features
|
||||
for t in range(min(n_times_vh, n_times_ndvi)):
|
||||
vh_t = vh_monthly.isel(time=t).values.flatten()
|
||||
# Resize if needed
|
||||
if len(vh_t) != n_pixels:
|
||||
vh_t = np.resize(vh_t, n_pixels)
|
||||
features = np.column_stack([features, vh_t])
|
||||
|
||||
# Add VV features
|
||||
for t in range(min(n_times_vv, n_times_ndvi)):
|
||||
vv_t = vv_monthly.isel(time=t).values.flatten()
|
||||
# Resize if needed
|
||||
if len(vv_t) != n_pixels:
|
||||
vv_t = np.resize(vv_t, n_pixels)
|
||||
features = np.column_stack([features, vv_t])
|
||||
|
||||
# Handle NaN values in features✓ CNN PyTorch: Mạnh nhất với ảnh vệ tinh, tự học features, tương thích GPU tốt, cần pip install torch
|
||||
features = np.nan_to_num(features, nan=0.0)
|
||||
|
||||
# ============ BƯỚC 6: DỰ ĐOÁN ============
|
||||
# Check model's expected feature count and adjust
|
||||
try:
|
||||
# Get expected number of features from model
|
||||
if is_cnn_model:
|
||||
# For PyTorch CNN, get n_features from model
|
||||
expected_features = model.n_features
|
||||
elif hasattr(model, 'n_features_in_'):
|
||||
expected_features = model.n_features_in_
|
||||
elif hasattr(model, 'feature_names_in_'):
|
||||
expected_features = len(model.feature_names_in_)
|
||||
else:
|
||||
# Try to get from booster for XGBoost
|
||||
try:
|
||||
expected_features = model.get_booster().num_features()
|
||||
except:
|
||||
expected_features = features.shape[1]
|
||||
|
||||
prediction_status["progress"] = f"Model cần {expected_features} features, đang có {features.shape[1]} features..."
|
||||
|
||||
# Adjust features to match model
|
||||
if features.shape[1] > expected_features:
|
||||
# Trim to expected number (use only first N features - NDVI only)
|
||||
prediction_status["progress"] = f"Cắt bớt features từ {features.shape[1]} xuống {expected_features}..."
|
||||
features = features[:, :expected_features]
|
||||
elif features.shape[1] < expected_features:
|
||||
# Pad with zeros or repeat last features
|
||||
prediction_status["progress"] = f"Thêm features từ {features.shape[1]} lên {expected_features}..."
|
||||
n_missing = expected_features - features.shape[1]
|
||||
# Repeat last feature column to fill
|
||||
padding = np.tile(features[:, -1:], (1, n_missing))
|
||||
features = np.column_stack([features, padding])
|
||||
except Exception as e:
|
||||
prediction_status["progress"] = f"Không thể xác định số features của model, tiếp tục với {features.shape[1]} features..."
|
||||
|
||||
prediction_status["progress"] = f"Đang dự đoán với {features.shape[1]} features..."
|
||||
|
||||
# Make prediction
|
||||
if is_cnn_model:
|
||||
# PyTorch CNN prediction
|
||||
predictions = model.predict(features)
|
||||
else:
|
||||
predictions = model.predict(features)
|
||||
|
||||
# Decode labels if label_encoder exists
|
||||
if label_encoder is not None:
|
||||
try:
|
||||
predictions = label_encoder.inverse_transform(predictions)
|
||||
except:
|
||||
pass # Keep numeric predictions if inverse_transform fails
|
||||
|
||||
# Reshape to original shape
|
||||
pred_shape = (y_size, x_size)
|
||||
predictions_2d = predictions.reshape(pred_shape)
|
||||
|
||||
# ============ BƯỚC 7: TẠO OUTPUT VÀ LƯU KẾT QUẢ ============
|
||||
prediction_status["progress"] = "Đang tạo bản đồ phân loại..."
|
||||
|
||||
# Create output xarray
|
||||
prediction_da = xr.DataArray(
|
||||
predictions_2d,
|
||||
coords={
|
||||
"y": ndvi_monthly.y,
|
||||
"x": ndvi_monthly.x
|
||||
},
|
||||
dims=["y", "x"],
|
||||
name="classification"
|
||||
)
|
||||
|
||||
# Save output
|
||||
output_dir = Path("predictions")
|
||||
output_dir.mkdir(exist_ok=True)
|
||||
|
||||
timestamp = dt.now().strftime("%Y%m%d_%H%M%S")
|
||||
output_file = output_dir / f"prediction_{timestamp}.tif"
|
||||
|
||||
prediction_status["progress"] = "Đang lưu kết quả GeoTIFF..."
|
||||
|
||||
# Set CRS and save as GeoTIFF
|
||||
if hasattr(s2_data, 'rio') and s2_data.rio.crs is not None:
|
||||
prediction_da.rio.write_crs(s2_data.rio.crs, inplace=True)
|
||||
else:
|
||||
prediction_da.rio.write_crs("EPSG:4326", inplace=True)
|
||||
|
||||
prediction_da.rio.to_raster(str(output_file), driver="GTiff")
|
||||
|
||||
# Get unique classes for result
|
||||
unique_classes = np.unique(predictions_2d)
|
||||
unique_classes = unique_classes[~np.isnan(unique_classes)].tolist()
|
||||
|
||||
prediction_status["is_predicting"] = False
|
||||
prediction_status["progress"] = "Hoàn thành!"
|
||||
prediction_status["output_file"] = str(output_file)
|
||||
prediction_status["result"] = {
|
||||
"output_file": str(output_file),
|
||||
"shape": list(pred_shape),
|
||||
"unique_classes": unique_classes,
|
||||
"bbox": bbox,
|
||||
"time_range": time_range,
|
||||
"n_features": features.shape[1],
|
||||
"n_times_ndvi": n_times_ndvi,
|
||||
"used_radar": use_radar,
|
||||
"model_used": config.model_filename
|
||||
}
|
||||
prediction_status["end_time"] = dt.now().isoformat()
|
||||
|
||||
except Exception as e:
|
||||
prediction_status["is_predicting"] = False
|
||||
prediction_status["error"] = str(e)
|
||||
prediction_status["progress"] = f"Lỗi: {str(e)}"
|
||||
prediction_status["end_time"] = dt.now().isoformat()
|
||||
import traceback
|
||||
print(traceback.format_exc())
|
||||
|
||||
|
||||
@app.get("/api/predictions/list")
|
||||
async def list_predictions():
|
||||
"""Lấy danh sách các file prediction đã tạo"""
|
||||
predictions_dir = Path("predictions")
|
||||
predictions_dir.mkdir(exist_ok=True)
|
||||
|
||||
predictions = []
|
||||
for pred_file in predictions_dir.glob("*.tif"):
|
||||
predictions.append({
|
||||
"filename": pred_file.name,
|
||||
"created": datetime.fromtimestamp(pred_file.stat().st_mtime).isoformat(),
|
||||
"size_mb": round(pred_file.stat().st_size / 1024 / 1024, 2),
|
||||
"download_url": f"/api/predictions/download/{pred_file.name}"
|
||||
})
|
||||
|
||||
# Sort by creation time (newest first)
|
||||
predictions.sort(key=lambda x: x["created"], reverse=True)
|
||||
return {"predictions": predictions}
|
||||
|
||||
|
||||
@app.get("/api/predictions/download/{filename}")
|
||||
async def download_prediction(filename: str):
|
||||
"""Download file prediction GeoTIFF"""
|
||||
predictions_dir = Path("predictions")
|
||||
file_path = predictions_dir / filename
|
||||
|
||||
# Security check: ensure filename doesn't contain path traversal
|
||||
if ".." in filename or "/" in filename or "\\" in filename:
|
||||
raise HTTPException(status_code=400, detail="Invalid filename")
|
||||
|
||||
if not file_path.exists():
|
||||
raise HTTPException(status_code=404, detail=f"File không tồn tại: {filename}")
|
||||
|
||||
return FileResponse(
|
||||
path=str(file_path),
|
||||
filename=filename,
|
||||
media_type="image/tiff",
|
||||
headers={
|
||||
"Content-Disposition": f"attachment; filename={filename}"
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("=" * 70)
|
||||
print("🚀 LAND CLASSIFICATION TRAINING API SERVER")
|
||||
print("=" * 70)
|
||||
print("\n📍 Endpoints:")
|
||||
print(" - Web Interface: http://localhost:8000")
|
||||
print(" - API Docs: http://localhost:8000/docs")
|
||||
print(" - Start Training: POST http://localhost:8000/api/training/start")
|
||||
print(" - Check Status: GET http://localhost:8000/api/training/status")
|
||||
print("\n" + "=" * 70)
|
||||
|
||||
uvicorn.run(app, host="0.0.0.0", port=8000, log_level="info")
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:37686f6bbdd4590ce2dc5474eef8268f680498c2b78881a486625e264d2a1015
|
||||
size 841901
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:78d28d677d39cc67e607bb418948b9e3e921351762d63b4737036fc02d1d361a
|
||||
size 3375913
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:892cabc8da5cfcc0f9295ddb3ad32fda077027efa53e7b48e2f81965ad6ca202
|
||||
size 3361173
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:fbab8f823ca6965dcbe43d9ce914f8912e3976833de1b205032eec48bfa83970
|
||||
size 1670930
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7c33115ec7d6896758abf1eb8d07eba94ae15444c63f318e8adbd5d032df9a4f
|
||||
size 1671094
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:e71c696efbcb9e766eb1a666cc85b4f604992a65efc7258a46a46c7f81c596db
|
||||
size 3536032
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:520c3480efaafdf844c3bba3dd165a8e44f3318ab21209539c78201f18fd6312
|
||||
size 3066028
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7d3bff20682fc510a9c41b5aa2b42e3ebef8098ed81e27b4436f9ae86f47b9a6
|
||||
size 3066028
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:322f48d2d54ee921facfc3c0c44b320a86a90694e1779b683824264f9affa3fc
|
||||
size 3066028
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:a446d333bf7f6d0cb7df014f12b0da3f7298f85bdfb4de06893173e90fbd5ccb
|
||||
size 14112695
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:3eec576dffe1cc1393fdb584fb99d28db62abdd0977d41030aee5c4fa5377180
|
||||
size 11474743
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:cc125de4c172c98d9e50123bfcc2f8dd856e4ccd6bb04edd639badf090af2eef
|
||||
size 5077735
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:8004f197bab9800b257d62064bf1220b908f2b17eadce826b87f0284241a1842
|
||||
size 1420071
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:4e989c28641beb9af91d2b12a9a0b957b2b8c6510f437a68b69271ea136cce65
|
||||
size 1604860
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7115a28f7fb41599b59f9337d7457ed4578ba989878a4ccd907c60e69be2311a
|
||||
size 973302
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:da63ad454fd0e6de316e17df9345abe1b193998d1eb97c545b6a87a66058e884
|
||||
size 766972
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:cfb8412f0ee12abb7f1bb067522b3b800102ffdfa18b84027603c5d1870c968e
|
||||
size 518845
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7318566231680cb0c97883b7a5e4177aa1bfeec462a0f52ed34aa103ddec210b
|
||||
size 20903
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:9cabf5e0241dcc3a73133ac8ae11171f34042c491f895b268c016407619bdfe1
|
||||
size 20903
|
||||
@@ -0,0 +1,164 @@
|
||||
# 🌾 Giải Thích Quy Trình Phân Loại Đất Trồng Cây
|
||||
|
||||
File notebook `02.predict_ODC.ipynb` sử dụng **Machine Learning** kết hợp với **dữ liệu vệ tinh** để phân loại các loại đất/cây trồng. Dưới đây là quy trình chi tiết:
|
||||
|
||||
---
|
||||
|
||||
## **Bước 1: Thu thập dữ liệu vệ tinh** (Cell 3-4)
|
||||
|
||||
```python
|
||||
date_range = ('2022-09-01', '2023-10-01')
|
||||
longtitude_range = (105.86575, 105.94120)
|
||||
latitude_range = (9.65070, 9.69850)
|
||||
data = load_data(dc, date_range, longtitude_range, latitude_range)
|
||||
```
|
||||
|
||||
- Lấy ảnh **Sentinel-2** (ảnh quang học) từ kho dữ liệu trong khoảng thời gian và vị trí cụ thể
|
||||
|
||||
---
|
||||
|
||||
## **Bước 2: Xử lý mây** (Cell 5)
|
||||
|
||||
```python
|
||||
result = mask_clean(data)
|
||||
```
|
||||
|
||||
- Loại bỏ các pixel bị mây che phủ để đảm bảo dữ liệu chính xác
|
||||
|
||||
---
|
||||
|
||||
## **Bước 3: Tính chỉ số NDVI** (Cell 6-10)
|
||||
|
||||
```python
|
||||
ndvi = calculate_indices(result, index='NDVI', satellite_mission='s2')
|
||||
fill_nan_ndvi = fill_nan(ndvi, time_split)
|
||||
average_ndvi = fill_nan_ndvi.resample(time='1M').mean()
|
||||
```
|
||||
|
||||
- **NDVI** (Normalized Difference Vegetation Index) = (NIR - Red) / (NIR + Red)
|
||||
- Giá trị từ **-1 đến 1**: cao = thực vật xanh tốt, thấp = đất trống/nước
|
||||
- Điền giá trị nan (mây) và tính trung bình theo tháng
|
||||
|
||||
---
|
||||
|
||||
## **Bước 4: Lấy dữ liệu Radar Sentinel-1** (Cell 11)
|
||||
|
||||
```python
|
||||
dsvh, dsvv = load_data_sen1(dc, date_range, coordinates)
|
||||
average_vv = calculate_average(dsvv, time_pattern='1M')
|
||||
average_vh = calculate_average(dsvh, time_pattern='1M')
|
||||
```
|
||||
|
||||
- **VH, VV**: Dữ liệu radar (xuyên mây), cho biết cấu trúc bề mặt
|
||||
- Giúp phân biệt lúa ngập nước, cây trồng cạn, mặt nước...
|
||||
|
||||
---
|
||||
|
||||
## **Bước 5: Dự đoán bằng Model ML** (Cell 12) ⭐ **QUAN TRỌNG NHẤT**
|
||||
|
||||
```python
|
||||
loaded_model = joblib.load("model_train/model_odc.joblib")
|
||||
data_array = predict(loaded_model, data.rio.crs, average_ndvi, average_vh, average_vv)
|
||||
```
|
||||
|
||||
**Model đã được train trước** với dữ liệu mẫu (training data) gồm:
|
||||
- **Đầu vào (Features)**: NDVI theo tháng + VH + VV (chuỗi thời gian)
|
||||
- **Đầu ra (Labels)**: Loại đất đã được gắn nhãn thủ công
|
||||
|
||||
### Cách model phân loại:
|
||||
|
||||
| Đặc điểm | Loại đất |
|
||||
|----------|----------|
|
||||
| NDVI cao đều, VV thấp | Rừng |
|
||||
| NDVI biến đổi theo mùa vụ, VH cao (nước) | Lúa |
|
||||
| NDVI thấp, VV rất thấp | Sông/nước |
|
||||
| NDVI trung bình ổn định | Cây lâu năm (CLN) |
|
||||
|
||||
---
|
||||
|
||||
## **Bước 6: Hiển thị kết quả** (Cell 13-15)
|
||||
|
||||
```python
|
||||
colors = ["#abcee9", "#ffef44", "#c4ff9e", "#ffd6a8", "#93ddda", "#1aeef7", "#ffa7f2", "#33ee33"]
|
||||
labels = ["Lúa tôm", "Lúa", "CHN", "CLN", "TS", "Sông", "Đất xây dựng", "Rừng"]
|
||||
```
|
||||
|
||||
### 8 lớp phân loại:
|
||||
|
||||
| Mã | Tên | Màu | Ý nghĩa |
|
||||
|----|-----|-----|---------|
|
||||
| 0 | Lúa tôm | 🔵 Xanh nhạt | Luân canh lúa-tôm |
|
||||
| 1 | Lúa | 🟡 Vàng | Đất trồng lúa |
|
||||
| 2 | CHN | 🟢 Xanh lá nhạt | Cây hàng năm |
|
||||
| 3 | CLN | 🟠 Cam nhạt | Cây lâu năm (cây ăn trái) |
|
||||
| 4 | TS | 🩵 Xanh ngọc | Thủy sản |
|
||||
| 5 | Sông | 🔷 Cyan | Mặt nước sông |
|
||||
| 6 | Đất XD | 💗 Hồng | Đất xây dựng |
|
||||
| 7 | Rừng | 💚 Xanh đậm | Rừng |
|
||||
|
||||
---
|
||||
|
||||
## **Bước 7: Lưu kết quả** (Cell 16)
|
||||
|
||||
```python
|
||||
region_result.rio.to_raster("KetQuaPhanLoaiDatODC.tif")
|
||||
```
|
||||
|
||||
- Xuất file GeoTIFF chứa mã phân loại (0-7) cho từng pixel
|
||||
|
||||
---
|
||||
|
||||
## 📊 **Tóm tắt quy trình:**
|
||||
|
||||
```
|
||||
Ảnh vệ tinh (Sentinel-1 + Sentinel-2)
|
||||
↓
|
||||
Xử lý (loại mây, tính NDVI, VH, VV)
|
||||
↓
|
||||
Kết hợp features theo thời gian (13 tháng)
|
||||
↓
|
||||
Model ML (Random Forest/XGBoost) dự đoán
|
||||
↓
|
||||
Bản đồ phân loại 8 lớp đất
|
||||
↓
|
||||
File .tif (mỗi pixel = 1 mã loại đất)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📁 Cấu trúc dữ liệu đầu vào cho Model
|
||||
|
||||
### Features (Đặc trưng):
|
||||
- **NDVI theo 13 tháng**: 13 bands
|
||||
- **VH (radar) theo 13 tháng**: 13 bands
|
||||
- **VV (radar) theo 13 tháng**: 13 bands
|
||||
- **Tổng cộng**: ~39 features cho mỗi pixel
|
||||
|
||||
### Labels (Nhãn):
|
||||
- Được lấy từ shapefile training: `train/ST_training data_updated_1130points_new.shp`
|
||||
- 1130 điểm mẫu đã được gắn nhãn thủ công bởi chuyên gia
|
||||
|
||||
---
|
||||
|
||||
## 🔧 Các thư viện sử dụng
|
||||
|
||||
| Thư viện | Mục đích |
|
||||
|----------|----------|
|
||||
| `datacube` | Truy vấn dữ liệu vệ tinh |
|
||||
| `xarray` | Xử lý dữ liệu đa chiều |
|
||||
| `rioxarray` | Đọc/ghi GeoTIFF |
|
||||
| `joblib` | Load/save model ML |
|
||||
| `sklearn` / `xgboost` | Training model |
|
||||
| `matplotlib` / `hvplot` | Trực quan hóa |
|
||||
|
||||
---
|
||||
|
||||
## 📝 Ghi chú
|
||||
|
||||
- **Độ phân giải**: 10-20m (tùy cấu hình)
|
||||
- **Thời gian xử lý**: Phụ thuộc vào kích thước vùng và số scenes
|
||||
- **Yêu cầu**: Cần kết nối internet để tải dữ liệu vệ tinh từ Planetary Computer hoặc ODC
|
||||
|
||||
---
|
||||
|
||||
*Tài liệu được tạo ngày 14/12/2025*
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:a446d333bf7f6d0cb7df014f12b0da3f7298f85bdfb4de06893173e90fbd5ccb
|
||||
size 14112695
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:ae2f82f6c837396729cc63efa41ee3048d9a7de3197e28318dc846be830239b9
|
||||
size 41536
|
||||
@@ -0,0 +1,33 @@
|
||||
{
|
||||
"timestamp": "2025-12-14T18:04:27.406540",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 510,
|
||||
"testing_samples": 128,
|
||||
"train_accuracy": 0.515686274509804,
|
||||
"test_accuracy": 0.5,
|
||||
"model_type": "cnn",
|
||||
"device": "cpu",
|
||||
"n_estimators": 50,
|
||||
"max_depth": null,
|
||||
"learning_rate": null,
|
||||
"cnn_epochs": 25,
|
||||
"n_features": 3,
|
||||
"n_classes": 7,
|
||||
"bbox": [
|
||||
105.6,
|
||||
9.3,
|
||||
106.2,
|
||||
9.8
|
||||
],
|
||||
"time_range": "2023-03-01/2023-05-31",
|
||||
"resolution": 20
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:abc03d62c7f620b88150a6481143026d516fd3a34bd5ab67c734adac4a8900f9
|
||||
size 41536
|
||||
@@ -0,0 +1,33 @@
|
||||
{
|
||||
"timestamp": "2025-12-14T18:17:48.399298",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 510,
|
||||
"testing_samples": 128,
|
||||
"train_accuracy": 0.4803921568627451,
|
||||
"test_accuracy": 0.484375,
|
||||
"model_type": "cnn",
|
||||
"device": "cpu",
|
||||
"n_estimators": 50,
|
||||
"max_depth": null,
|
||||
"learning_rate": null,
|
||||
"cnn_epochs": 25,
|
||||
"n_features": 3,
|
||||
"n_classes": 7,
|
||||
"bbox": [
|
||||
105.6,
|
||||
9.3,
|
||||
106.2,
|
||||
9.8
|
||||
],
|
||||
"time_range": "2023-03-01/2023-05-25",
|
||||
"resolution": 20
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:21f906aa2a61d2e793a3463df95fe3e364134934af326771efae90d80caca419
|
||||
size 41536
|
||||
@@ -0,0 +1,33 @@
|
||||
{
|
||||
"timestamp": "2025-12-14T18:23:10.713912",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 510,
|
||||
"testing_samples": 128,
|
||||
"train_accuracy": 0.46862745098039216,
|
||||
"test_accuracy": 0.4609375,
|
||||
"model_type": "cnn",
|
||||
"device": "cpu",
|
||||
"n_estimators": 50,
|
||||
"max_depth": null,
|
||||
"learning_rate": null,
|
||||
"cnn_epochs": 25,
|
||||
"n_features": 3,
|
||||
"n_classes": 7,
|
||||
"bbox": [
|
||||
105.6,
|
||||
9.3,
|
||||
106.2,
|
||||
9.8
|
||||
],
|
||||
"time_range": "2023-03-01/2023-05-31",
|
||||
"resolution": 20
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:3eec576dffe1cc1393fdb584fb99d28db62abdd0977d41030aee5c4fa5377180
|
||||
size 11474743
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:0fe99f96ad3d3ba7aaacc0e742572a8f5b22947a328c74b245e0aa5f2913c757
|
||||
size 1347520
|
||||
@@ -0,0 +1,24 @@
|
||||
{
|
||||
"timestamp": "2025-12-12T12:57:54.509336",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 510,
|
||||
"testing_samples": 128,
|
||||
"train_accuracy": 1.0,
|
||||
"test_accuracy": 0.578125,
|
||||
"model_type": "XGBClassifier",
|
||||
"device": "cuda:0",
|
||||
"gpu_device": "RTX 4060",
|
||||
"tree_method": "hist",
|
||||
"n_estimators": 100,
|
||||
"max_depth": 20,
|
||||
"learning_rate": 0.1
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:b717f564f9413a6e5c9cd3f7011cbc18be02479691d01c554986defb400f0490
|
||||
size 1347520
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"timestamp": "2025-12-12T22:15:25.794614",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 510,
|
||||
"testing_samples": 128,
|
||||
"train_accuracy": 1.0,
|
||||
"test_accuracy": 0.578125,
|
||||
"model_type": "XGBClassifier",
|
||||
"device": "cuda:0",
|
||||
"tree_method": "hist",
|
||||
"n_estimators": 100,
|
||||
"max_depth": 20,
|
||||
"learning_rate": 0.1,
|
||||
"bbox": [
|
||||
105.6,
|
||||
9.3,
|
||||
106.2,
|
||||
9.8
|
||||
],
|
||||
"time_range": "2023-03-01/2023-05-31",
|
||||
"resolution": 20
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:b9c6cabb59d9cdbba22a438ae935911f1711d7434dcf3728db4e518ec1b90190
|
||||
size 556184
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"timestamp": "2025-12-12T22:33:42.950629",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 904,
|
||||
"testing_samples": 226,
|
||||
"train_accuracy": 0.19911504424778761,
|
||||
"test_accuracy": 0.19911504424778761,
|
||||
"model_type": "XGBClassifier",
|
||||
"device": "cuda:0",
|
||||
"tree_method": "hist",
|
||||
"n_estimators": 100,
|
||||
"max_depth": 20,
|
||||
"learning_rate": 0.1,
|
||||
"bbox": [
|
||||
104.89032,
|
||||
10.944563,
|
||||
104.972717,
|
||||
11.016689
|
||||
],
|
||||
"time_range": "2023-03-01/2023-05-31",
|
||||
"resolution": 20
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:b717f564f9413a6e5c9cd3f7011cbc18be02479691d01c554986defb400f0490
|
||||
size 1347520
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"timestamp": "2025-12-14T13:40:03.325930",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 510,
|
||||
"testing_samples": 128,
|
||||
"train_accuracy": 1.0,
|
||||
"test_accuracy": 0.578125,
|
||||
"model_type": "XGBClassifier",
|
||||
"device": "cuda:0",
|
||||
"tree_method": "hist",
|
||||
"n_estimators": 100,
|
||||
"max_depth": 20,
|
||||
"learning_rate": 0.1,
|
||||
"bbox": [
|
||||
105.6,
|
||||
9.3,
|
||||
106.2,
|
||||
9.8
|
||||
],
|
||||
"time_range": "2023-03-01/2023-05-31",
|
||||
"resolution": 20
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:a44efc173619782c8add224e9024f33ef0150ab192b5435d009cb119c379f03c
|
||||
size 2088632
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"timestamp": "2025-12-14T16:52:19.862770",
|
||||
"data_source": "Microsoft Planetary Computer STAC",
|
||||
"collections": [
|
||||
"sentinel-2-l2a",
|
||||
"sentinel-1-rtc"
|
||||
],
|
||||
"features": [
|
||||
"NDVI_mean",
|
||||
"VH_dB_mean",
|
||||
"VV_dB_mean"
|
||||
],
|
||||
"training_samples": 859,
|
||||
"testing_samples": 215,
|
||||
"train_accuracy": 0.9976717112922002,
|
||||
"test_accuracy": 0.6837209302325581,
|
||||
"model_type": "XGBClassifier",
|
||||
"device": "cuda:0",
|
||||
"tree_method": "hist",
|
||||
"n_estimators": 100,
|
||||
"max_depth": 20,
|
||||
"learning_rate": 0.1,
|
||||
"bbox": [
|
||||
105.6,
|
||||
9.3,
|
||||
106.2,
|
||||
9.8
|
||||
],
|
||||
"time_range": "2023-03-01/2023-12-31",
|
||||
"resolution": 20
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:eb89e36fa4f4c740d5a079baf53b02bec2ef1120ac770bed6b6be4aa9fc99a1b
|
||||
size 208470
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:179d3a61420c1552e9653de5c7657665c9880ad29d751e56bae646fb3b634687
|
||||
size 73272920
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:179d3a61420c1552e9653de5c7657665c9880ad29d751e56bae646fb3b634687
|
||||
size 73272920
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:179d3a61420c1552e9653de5c7657665c9880ad29d751e56bae646fb3b634687
|
||||
size 73272920
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:d8fb13c2e466b9811104cbd7747ea6ca8c14ddb44802d934aadfe52a4b5dd916
|
||||
size 73272920
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:a2fe092e1fd96e56da519acffe5ca5c9a8246e796bc5700cffb9694dc99f3aec
|
||||
size 73272920
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:809b41520bb499e1032042e806ffb9f3c798609a88edc771efb75e39dd80f20d
|
||||
size 1335865
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:6c282f1a104a47a2837f1882da8011293051f507b0bb1d7e6ab27ba68b24dc19
|
||||
size 787409
|
||||
+309
@@ -0,0 +1,309 @@
|
||||
affine @ file:///home/conda/feedstock_root/build_artifacts/affine_1733762038348/work
|
||||
aiobotocore==2.25.0
|
||||
aiohappyeyeballs==2.6.1
|
||||
aiohttp==3.12.15
|
||||
aioitertools==0.12.0
|
||||
aiosignal==1.4.0
|
||||
alembic==1.16.5
|
||||
annotated-doc==0.0.4
|
||||
annotated-types==0.7.0
|
||||
antimeridian @ file:///home/conda/feedstock_root/build_artifacts/antimeridian_1753706324394/work
|
||||
anyio @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_anyio_1758634638/work
|
||||
argon2-cffi @ file:///home/conda/feedstock_root/build_artifacts/argon2-cffi_1749017159514/work
|
||||
argon2-cffi-bindings @ file:///home/conda/feedstock_root/build_artifacts/argon2-cffi-bindings_1649500328244/work
|
||||
arrow @ file:///home/conda/feedstock_root/build_artifacts/arrow_1733584251875/work
|
||||
asciitree==0.3.3
|
||||
asttokens @ file:///home/conda/feedstock_root/build_artifacts/asttokens_1733250440834/work
|
||||
async-lru @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_async-lru_1742153708/work
|
||||
async-timeout==3.0.1
|
||||
attrs @ file:///home/conda/feedstock_root/build_artifacts/attrs_1741918516150/work
|
||||
babel @ file:///home/conda/feedstock_root/build_artifacts/babel_1738490167835/work
|
||||
beautifulsoup4 @ file:///home/conda/feedstock_root/build_artifacts/beautifulsoup4_1759146011391/work
|
||||
bleach @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_bleach_1737382993/work
|
||||
blinker==1.9.0
|
||||
bokeh==3.7.3
|
||||
boto3==1.40.18
|
||||
botocore==1.40.49
|
||||
Bottleneck @ file:///croot/bottleneck_1731058641041/work
|
||||
branca @ file:///croot/branca_1675157607453/work
|
||||
Brotli @ file:///croot/brotli-split_1736182456865/work
|
||||
brotlicffi @ file:///croot/brotlicffi_1736182461069/work
|
||||
cached-property @ file:///home/conda/feedstock_root/build_artifacts/cached_property_1615209429212/work
|
||||
cachetools==6.2.0
|
||||
Cartopy==0.25.0
|
||||
certifi @ file:///home/conda/feedstock_root/build_artifacts/certifi_1759648874697/work/certifi
|
||||
cffi @ file:///croot/cffi_1736182485317/work
|
||||
cftime @ file:///home/conda/feedstock_root/build_artifacts/cftime_1649636873066/work
|
||||
chardet @ file:///home/conda/feedstock_root/build_artifacts/chardet_1649184137891/work
|
||||
charset-normalizer @ file:///croot/charset-normalizer_1721748349566/work
|
||||
ciso8601==2.3.3
|
||||
click @ file:///home/conda/feedstock_root/build_artifacts/click_1747811314515/work
|
||||
click-plugins @ file:///home/conda/feedstock_root/build_artifacts/click-plugins_1750848229740/work
|
||||
cligj @ file:///home/conda/feedstock_root/build_artifacts/cligj_1733749956636/work
|
||||
cloudpickle @ file:///home/conda/feedstock_root/build_artifacts/cloudpickle_1736947526808/work
|
||||
colorama==0.4.6
|
||||
colorcet==3.1.0
|
||||
comm @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_comm_1753453984/work
|
||||
contourpy @ file:///croot/contourpy_1732540045555/work
|
||||
cycler @ file:///tmp/build/80754af9/cycler_1637851556182/work
|
||||
cytoolz==0.11.2
|
||||
dask @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_dask-core_1760473436/work
|
||||
dask-gateway @ file:///Users/runner/miniforge3/conda-bld/bld/rattler-build_dask-gateway_1744370153/work/dask-gateway
|
||||
dask-glm @ file:///home/conda/feedstock_root/build_artifacts/dask-glm_1701346265909/work
|
||||
dask-image==2024.5.3
|
||||
dask-ml @ file:///home/conda/feedstock_root/build_artifacts/dask-ml_1679705292494/work
|
||||
datacube==1.8.15
|
||||
datacube_ows==1.9.4
|
||||
datashader==0.18.2
|
||||
dea-tools==0.3.0
|
||||
debugpy @ file:///home/task_175706711740264/conda-bld/debugpy_1757067131873/work
|
||||
decorator @ file:///home/conda/feedstock_root/build_artifacts/decorator_1740384970518/work
|
||||
deepdiff==8.6.1
|
||||
defusedxml @ file:///home/conda/feedstock_root/build_artifacts/defusedxml_1615232257335/work
|
||||
deprecat @ file:///home/conda/feedstock_root/build_artifacts/deprecat_1734684036993/work
|
||||
distributed @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_distributed_1760476147/work
|
||||
eo-tides==0.8.2
|
||||
exceptiongroup @ file:///home/conda/feedstock_root/build_artifacts/exceptiongroup_1746947292760/work
|
||||
executing @ file:///home/conda/feedstock_root/build_artifacts/executing_1756729339227/work
|
||||
fastapi==0.124.3
|
||||
fasteners @ file:///home/conda/feedstock_root/build_artifacts/fasteners_1734943108928/work
|
||||
fastjsonschema @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_python-fastjsonschema_1755304154/work/dist
|
||||
filelock==3.19.1
|
||||
fiona==1.10.1
|
||||
Flask==3.1.2
|
||||
flask-babel==4.0.0
|
||||
flatbuffers==25.2.10
|
||||
folium==0.20.0
|
||||
fonttools @ file:///croot/fonttools_1737039080035/work
|
||||
fqdn @ file:///home/conda/feedstock_root/build_artifacts/fqdn_1733327382592/work/dist
|
||||
frozenlist==1.7.0
|
||||
fsspec @ file:///home/conda/feedstock_root/build_artifacts/fsspec_1756908513222/work
|
||||
GDAL @ file:///croot/gdal-split_1734448174900/work/build/swig/python
|
||||
GeoAlchemy2 @ file:///home/conda/feedstock_root/build_artifacts/geoalchemy2_1753372953474/work
|
||||
geographiclib==2.1
|
||||
geojson==3.2.0
|
||||
geomad==1.0.0
|
||||
geopandas @ file:///croot/geopandas-split_1755761494241/work
|
||||
geopy==2.4.1
|
||||
greenlet @ file:///home/conda/feedstock_root/build_artifacts/greenlet_1648882383677/work
|
||||
h11 @ file:///home/conda/feedstock_root/build_artifacts/h11_1745526374115/work
|
||||
h2 @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_h2_1756364871/work
|
||||
h3==4.3.1
|
||||
hdstats==0.2.1
|
||||
holoviews==1.21.0
|
||||
hpack @ file:///home/conda/feedstock_root/build_artifacts/hpack_1737618293087/work
|
||||
httpcore @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_httpcore_1745602916/work
|
||||
httpx @ file:///home/conda/feedstock_root/build_artifacts/httpx_1733663348460/work
|
||||
hvplot==0.12.1
|
||||
hyperframe @ file:///home/conda/feedstock_root/build_artifacts/hyperframe_1737618333194/work
|
||||
idna==3.10
|
||||
imagecodecs==2025.3.30
|
||||
imageio==2.37.0
|
||||
importlib_metadata @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_importlib-metadata_1747934053/work
|
||||
ipykernel @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_ipykernel_1760459840/work
|
||||
ipyleaflet==0.20.0
|
||||
ipython @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_ipython_1748711175/work
|
||||
ipywidgets==8.1.7
|
||||
iso8601==2.1.0
|
||||
isoduration @ file:///home/conda/feedstock_root/build_artifacts/isoduration_1733493628631/work/dist
|
||||
itsdangerous==2.2.0
|
||||
jedi @ file:///home/conda/feedstock_root/build_artifacts/jedi_1733300866624/work
|
||||
Jinja2 @ file:///croot/jinja2_1741710844255/work
|
||||
jmespath @ file:///home/conda/feedstock_root/build_artifacts/jmespath_1733229141657/work
|
||||
joblib @ file:///home/conda/feedstock_root/build_artifacts/joblib_1756321760188/work
|
||||
json5 @ file:///home/conda/feedstock_root/build_artifacts/json5_1755034879854/work
|
||||
jsonpointer @ file:///home/conda/feedstock_root/build_artifacts/jsonpointer_1756754132747/work
|
||||
jsonschema @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_jsonschema_1755595646/work
|
||||
jsonschema-specifications==2025.4.1
|
||||
jupyter-events @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_jupyter_events_1738765986/work
|
||||
jupyter-leaflet==0.20.0
|
||||
jupyter-lsp @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_jupyter-lsp_1756388269/work/jupyter-lsp
|
||||
jupyter-ui-poll==1.0.0
|
||||
jupyter_client @ file:///home/conda/feedstock_root/build_artifacts/jupyter_client_1733440914442/work
|
||||
jupyter_core @ file:///home/conda/feedstock_root/build_artifacts/jupyter_core_1748333051527/work
|
||||
jupyter_server @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_jupyter_server_1755870522/work
|
||||
jupyter_server_terminals @ file:///home/conda/feedstock_root/build_artifacts/jupyter_server_terminals_1733427956852/work
|
||||
jupyterlab @ file:///home/conda/feedstock_root/build_artifacts/jupyterlab_1758913905644/work
|
||||
jupyterlab_pygments @ file:///home/conda/feedstock_root/build_artifacts/jupyterlab_pygments_1733328101776/work
|
||||
jupyterlab_server @ file:///home/conda/feedstock_root/build_artifacts/jupyterlab_server_1733599573484/work
|
||||
jupyterlab_widgets==3.0.15
|
||||
kiwisolver @ file:///croot/kiwisolver_1737039087198/work
|
||||
lark==1.2.2
|
||||
lark-parser==0.12.0
|
||||
lazy_loader==0.4
|
||||
linkify-it-py==2.0.3
|
||||
llvmlite @ file:///croot/llvmlite_1741209858218/work
|
||||
locket @ file:///home/conda/feedstock_root/build_artifacts/locket_1650660393415/work
|
||||
lxml==5.4.0
|
||||
lz4 @ file:///croot/lz4_1736366683208/work
|
||||
Mako @ file:///home/conda/feedstock_root/build_artifacts/mako_1744317760971/work
|
||||
mapclassify @ file:///croot/mapclassify_1675157730177/work
|
||||
Markdown==3.9
|
||||
markdown-it-py==4.0.0
|
||||
MarkupSafe @ file:///croot/markupsafe_1738584038848/work
|
||||
matplotlib==3.10.5
|
||||
matplotlib-inline @ file:///home/conda/feedstock_root/build_artifacts/matplotlib-inline_1733416936468/work
|
||||
mdit-py-plugins==0.5.0
|
||||
mdurl==0.1.2
|
||||
mistune @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_mistune_1756495311/work
|
||||
mpmath==1.3.0
|
||||
msgpack @ file:///home/conda/feedstock_root/build_artifacts/msgpack-python_1648745999384/work
|
||||
multidict @ file:///home/conda/feedstock_root/build_artifacts/multidict_1648882415384/work
|
||||
multipledispatch @ file:///home/conda/feedstock_root/build_artifacts/multipledispatch_1721907546485/work
|
||||
narwhals==2.3.0
|
||||
nbclient @ file:///home/conda/feedstock_root/build_artifacts/nbclient_1734628800805/work
|
||||
nbconvert @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_nbconvert-core_1738067871/work
|
||||
nbformat @ file:///home/conda/feedstock_root/build_artifacts/nbformat_1733402752141/work
|
||||
nest_asyncio @ file:///home/conda/feedstock_root/build_artifacts/nest-asyncio_1733325553580/work
|
||||
netCDF4 @ file:///croot/netcdf4_1743512888672/work
|
||||
networkx @ file:///croot/networkx_1737039604450/work
|
||||
notebook @ file:///home/conda/feedstock_root/build_artifacts/notebook_1759152069573/work
|
||||
notebook_shim @ file:///home/conda/feedstock_root/build_artifacts/notebook-shim_1733408315203/work
|
||||
numba @ file:///croot/numba_1750798165355/work
|
||||
numcodecs @ file:///croot/numcodecs_1707513121886/work
|
||||
numexpr @ file:///croot/numexpr_1755766469354/work
|
||||
numpy @ file:///croot/numpy_and_numpy_base_1755590845055/work/dist/numpy-1.26.4-cp310-cp310-linux_x86_64.whl#sha256=1096d33ad9a9757a1b4b46634d809e894263fc8b78780bff36801684b6e8cc88
|
||||
nvidia-cublas-cu12==12.8.4.1
|
||||
nvidia-cuda-cupti-cu12==12.8.90
|
||||
nvidia-cuda-nvrtc-cu12==12.8.93
|
||||
nvidia-cuda-runtime-cu12==12.8.90
|
||||
nvidia-cudnn-cu12==9.10.2.21
|
||||
nvidia-cufft-cu12==11.3.3.83
|
||||
nvidia-cufile-cu12==1.13.1.3
|
||||
nvidia-curand-cu12==10.3.9.90
|
||||
nvidia-cusolver-cu12==11.7.3.90
|
||||
nvidia-cusparse-cu12==12.5.8.93
|
||||
nvidia-cusparselt-cu12==0.7.1
|
||||
nvidia-nccl-cu12==2.27.3
|
||||
nvidia-nvjitlink-cu12==12.8.93
|
||||
nvidia-nvtx-cu12==12.8.90
|
||||
odc-algo==0.2.3
|
||||
odc-geo==0.4.10
|
||||
odc-io==0.2.2
|
||||
odc-loader @ file:///home/conda/feedstock_root/build_artifacts/odc-loader_1743656085024/work
|
||||
odc-stac @ file:///home/conda/feedstock_root/build_artifacts/odc-stac_1746136311934/work
|
||||
odc-ui==0.2.1
|
||||
orderly-set==5.5.0
|
||||
overrides @ file:///home/conda/feedstock_root/build_artifacts/overrides_1734587627321/work
|
||||
OWSLib==0.34.1
|
||||
packaging @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_packaging_1745345660/work
|
||||
pandas @ file:///home/task_175982153789305/conda-bld/pandas_1759822248912/work/dist/pandas-2.3.3-cp310-cp310-linux_x86_64.whl#sha256=0de7c83109c411cc2a74419a396c92f65e3d1e457fb4d835e5f100cfb04393a7
|
||||
pandocfilters @ file:///home/conda/feedstock_root/build_artifacts/pandocfilters_1631603243851/work
|
||||
panel==1.7.5
|
||||
param==2.2.1
|
||||
parso @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_parso_1755974222/work
|
||||
partd @ file:///home/conda/feedstock_root/build_artifacts/partd_1715026491486/work
|
||||
pexpect @ file:///home/conda/feedstock_root/build_artifacts/pexpect_1733301927746/work
|
||||
pickleshare @ file:///home/conda/feedstock_root/build_artifacts/pickleshare_1733327343728/work
|
||||
pillow @ file:///croot/pillow_1738010226202/work
|
||||
PIMS==0.7
|
||||
planetary-computer==1.0.0
|
||||
platformdirs @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_platformdirs_1756227402/work
|
||||
prometheus_client==0.22.1
|
||||
prometheus_flask_exporter==0.23.2
|
||||
prompt_toolkit @ file:///home/conda/feedstock_root/build_artifacts/prompt-toolkit_1756321756983/work
|
||||
propcache==0.3.2
|
||||
psutil @ file:///home/conda/feedstock_root/build_artifacts/psutil_1653089181607/work
|
||||
psycopg2 @ file:///croot/psycopg2_1744919787325/work
|
||||
ptyprocess @ file:///home/conda/feedstock_root/build_artifacts/ptyprocess_1733302279685/work/dist/ptyprocess-0.7.0-py2.py3-none-any.whl#sha256=92c32ff62b5fd8cf325bec5ab90d7be3d2a8ca8c8a3813ff487a8d2002630d1f
|
||||
pure_eval @ file:///home/conda/feedstock_root/build_artifacts/pure_eval_1733569405015/work
|
||||
pyarrow @ file:///home/task_175983338836370/conda-bld/pyarrow_1759833584228/work/python
|
||||
pycparser @ file:///tmp/build/80754af9/pycparser_1636541352034/work
|
||||
pyct==0.5.0
|
||||
pydantic==2.11.7
|
||||
pydantic_core==2.33.2
|
||||
Pygments @ file:///home/conda/feedstock_root/build_artifacts/pygments_1750615794071/work
|
||||
pyogrio @ file:///croot/pyogrio_1741107161422/work
|
||||
pyows==0.3.1
|
||||
pyparsing @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_pyparsing_1753873557/work
|
||||
pyproj @ file:///croot/pyproj_1739284761968/work
|
||||
PyQt6==6.7.1
|
||||
PyQt6_sip @ file:///croot/pyqt-split_1753427276959/work/pyqt_sip
|
||||
pyshp==2.3.1
|
||||
PySocks @ file:///home/builder/ci_310/pysocks_1640793678128/work
|
||||
pystac @ file:///home/conda/feedstock_root/build_artifacts/pystac_1758218055393/work
|
||||
pystac-client==0.9.0
|
||||
python-dateutil==2.9.0.post0
|
||||
python-dotenv==1.1.1
|
||||
python-json-logger @ file:///home/conda/feedstock_root/build_artifacts/python-json-logger_1677079630776/work
|
||||
python-slugify==8.0.4
|
||||
pyTMD==2.2.8
|
||||
pytz @ file:///home/conda/feedstock_root/build_artifacts/pytz_1742920838005/work
|
||||
pyviz_comms==3.0.6
|
||||
PyYAML==6.0.2
|
||||
pyzmq @ file:///croot/pyzmq_1734687138743/work
|
||||
rasterio @ file:///croot/rasterio_1740069178893/work
|
||||
rasterstats==0.20.0
|
||||
referencing==0.36.2
|
||||
regex==2025.9.1
|
||||
requests @ file:///croot/requests_1756709366904/work
|
||||
rfc3339_validator @ file:///home/conda/feedstock_root/build_artifacts/rfc3339-validator_1733599910982/work
|
||||
rfc3986-validator @ file:///home/conda/feedstock_root/build_artifacts/rfc3986-validator_1598024191506/work
|
||||
rfc3987==1.3.8
|
||||
rfc3987-syntax @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_rfc3987-syntax_1752876729/work
|
||||
rioxarray @ file:///home/conda/feedstock_root/build_artifacts/rioxarray_1737140588464/work
|
||||
rpds-py @ file:///croot/rpds-py_1736541261634/work
|
||||
ruamel.yaml @ file:///home/conda/feedstock_root/build_artifacts/ruamel.yaml_1649033201098/work
|
||||
ruamel.yaml.clib==0.2.12
|
||||
s3fs==2025.9.0
|
||||
s3transfer==0.13.1
|
||||
scikit-image==0.25.2
|
||||
scikit-learn==1.7.1
|
||||
scipy @ file:///croot/scipy_1747238027288/work/dist/scipy-1.15.3-cp310-cp310-linux_x86_64.whl#sha256=2a791554880ad4f358fcc4cd2a982ffe1e9d472e9241011216b2be797457f1f9
|
||||
seaborn==0.13.2
|
||||
Send2Trash @ file:///home/conda/feedstock_root/build_artifacts/send2trash_1733322040660/work
|
||||
setuptools-scm==9.2.0
|
||||
shapely @ file:///croot/shapely_1754380812723/work
|
||||
simplejson==3.20.1
|
||||
sip @ file:///croot/sip_1738856193618/work
|
||||
six==1.17.0
|
||||
slicerator==1.1.0
|
||||
sniffio @ file:///home/conda/feedstock_root/build_artifacts/sniffio_1733244044561/work
|
||||
snuggs @ file:///home/conda/feedstock_root/build_artifacts/snuggs_1733818638588/work
|
||||
sortedcontainers @ file:///home/conda/feedstock_root/build_artifacts/sortedcontainers_1738440353519/work
|
||||
soupsieve @ file:///home/conda/feedstock_root/build_artifacts/soupsieve_1756330469801/work
|
||||
sparse @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_sparse_1747799051/work
|
||||
SQLAlchemy==1.4.54
|
||||
stack_data @ file:///home/conda/feedstock_root/build_artifacts/stack_data_1733569443808/work
|
||||
starlette==0.50.0
|
||||
sympy==1.14.0
|
||||
tblib @ file:///home/conda/feedstock_root/build_artifacts/tblib_1743515515538/work
|
||||
terminado @ file:///home/conda/feedstock_root/build_artifacts/terminado_1710262609923/work
|
||||
text-unidecode==1.3
|
||||
threadpoolctl @ file:///home/conda/feedstock_root/build_artifacts/threadpoolctl_1741878222898/work
|
||||
tifffile==2025.5.10
|
||||
timescale==0.0.9
|
||||
timezonefinder==8.0.0
|
||||
tinycss2 @ file:///home/conda/feedstock_root/build_artifacts/tinycss2_1729802851396/work
|
||||
tomli @ file:///croot/tomli_1753774587605/work
|
||||
toolz @ file:///home/conda/feedstock_root/build_artifacts/toolz_1733736030883/work
|
||||
torch==2.8.0
|
||||
tornado @ file:///croot/tornado_1748956929273/work
|
||||
tqdm==4.67.1
|
||||
traitlets @ file:///home/conda/feedstock_root/build_artifacts/traitlets_1733367359838/work
|
||||
traittypes==0.2.1
|
||||
triton==3.4.0
|
||||
types-python-dateutil @ file:///home/conda/feedstock_root/build_artifacts/types-python-dateutil_1759899809376/work
|
||||
typing-inspection==0.4.1
|
||||
typing_extensions @ file:///home/conda/feedstock_root/build_artifacts/bld/rattler-build_typing_extensions_1756220668/work
|
||||
typing_utils @ file:///home/conda/feedstock_root/build_artifacts/typing_utils_1733331286120/work
|
||||
tzdata @ file:///croot/python-tzdata_1746123641790/work
|
||||
uc-micro-py==1.0.3
|
||||
unicodedata2 @ file:///croot/unicodedata2_1736541023050/work
|
||||
uri-template @ file:///home/conda/feedstock_root/build_artifacts/uri-template_1733323593477/work/dist
|
||||
urllib3 @ file:///croot/urllib3_1750775463400/work
|
||||
uvicorn==0.38.0
|
||||
wcwidth @ file:///home/conda/feedstock_root/build_artifacts/wcwidth_1733231326287/work
|
||||
webcolors @ file:///home/conda/feedstock_root/build_artifacts/webcolors_1733359735138/work
|
||||
webencodings @ file:///home/conda/feedstock_root/build_artifacts/webencodings_1733236011802/work
|
||||
websocket-client @ file:///home/conda/feedstock_root/build_artifacts/websocket-client_1759928050786/work
|
||||
Werkzeug==3.1.3
|
||||
widgetsnbextension==4.0.14
|
||||
wrapt @ file:///home/conda/feedstock_root/build_artifacts/wrapt_1651495243689/work
|
||||
xarray @ file:///home/conda/feedstock_root/build_artifacts/xarray_1749743207754/work
|
||||
xgboost==3.1.2
|
||||
xyzservices @ file:///croot/xyzservices_1675159059961/work
|
||||
yarl==1.20.1
|
||||
zarr @ file:///home/conda/feedstock_root/build_artifacts/zarr_1733237197728/work
|
||||
zict @ file:///home/conda/feedstock_root/build_artifacts/zict_1733261551178/work
|
||||
zipp @ file:///home/conda/feedstock_root/build_artifacts/zipp_1749421620841/work
|
||||
@@ -0,0 +1,3 @@
|
||||
fastapi
|
||||
uvicorn
|
||||
pydantic
|
||||
@@ -0,0 +1 @@
|
||||
uvicorn api_server:app --reload --host 0.0.0.0 --port 8000
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:c8fa50e2debed3c65b599f956095550a72844082245ae2500d1ab196ae23641e
|
||||
size 393276
|
||||
+532
@@ -0,0 +1,532 @@
|
||||
"""
|
||||
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,
|
||||
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
|
||||
|
||||
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=0.2, 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)
|
||||
else:
|
||||
train_score = model.score(X_train, y_train)
|
||||
test_score = model.score(X_test, y_test)
|
||||
|
||||
# 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),
|
||||
"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)),
|
||||
"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),
|
||||
"classes": label_encoder.classes_.tolist()
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user