133 lines
4.1 KiB
Python
133 lines
4.1 KiB
Python
#!/usr/bin/env python
|
|
# coding: utf-8
|
|
|
|
import os
|
|
import json
|
|
import joblib
|
|
import numpy as np
|
|
import importlib
|
|
import new_import_ODC
|
|
importlib.reload(new_import_ODC)
|
|
from new_import_ODC import *
|
|
|
|
import lightgbm as lgb
|
|
from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score
|
|
|
|
print("=" * 70)
|
|
print("🚀 Training LightGBM model for Land Classification (REAL DATA)")
|
|
print("=" * 70)
|
|
|
|
# 1. Cấu hình thời gian và tọa độ
|
|
date_range = ("2022-09-01", "2022-10-01")
|
|
longtitude_range = (105.86, 105.94)
|
|
latitude_range = (9.65, 9.69)
|
|
coordinates = (longtitude_range, latitude_range)
|
|
|
|
# 2. Load Dữ liệu Thật (Từ Cache)
|
|
print("\n[DATA] Đang load dữ liệu Sentinel-2...")
|
|
data = load_data(None, date_range, longtitude_range, latitude_range)
|
|
result = mask_clean(data)
|
|
ds1 = calculate_indices(result, index="NDVI", satellite_mission="s2")
|
|
ndvi = ds1["NDVI"]
|
|
|
|
time_split = [
|
|
slice("2022-09-01", "2023-01-01"),
|
|
slice("2023-01-01", "2023-05-01"),
|
|
slice("2023-05-01", "2023-07-01"),
|
|
slice("2023-07-01", "2022-10-01"),
|
|
]
|
|
fill_nan_ndvi = fill_nan(ndvi, time_split)
|
|
|
|
average_ndvi = fill_nan_ndvi.resample(time="1M").mean().persist()
|
|
average_ndvi = average_ndvi.compute()
|
|
|
|
print("\n[DATA] Đang load dữ liệu Sentinel-1 (VV, VH)...")
|
|
dsvh, dsvv = load_data_sen1(None, date_range, coordinates)
|
|
average_vv = calculate_average(dsvv, time_pattern='1M')
|
|
average_vh = calculate_average(dsvh, time_pattern='1M')
|
|
|
|
# 3. Chuẩn bị tập Train
|
|
print("\n[DATA] Đang trích xuất điểm huấn luyện...")
|
|
train_path = "train/ST_training_data_updated_1130points_new.shp"
|
|
train = load_train_data(train_path)
|
|
|
|
label_mapping = {
|
|
"Lua tom": "0", "Lua": "1", "CHN": "2", "CLN": "3",
|
|
"TS": "4", "Song": "5", "Dat xay dung": "6", "Rung": "7"
|
|
}
|
|
|
|
datasets = get_data_sen1_and_sen2(train, average_ndvi, average_vh, average_vv)
|
|
X_train, X_val, X_test, y_train, y_val, y_test = split_train_data(train, label_mapping, datasets)
|
|
|
|
X_train_np = np.asarray(X_train, dtype=np.float32)
|
|
y_train_np = np.asarray(y_train, dtype=np.int32)
|
|
X_val_np = np.asarray(X_val, dtype=np.float32)
|
|
y_val_np = np.asarray(y_val, dtype=np.int32)
|
|
X_test_np = np.asarray(X_test, dtype=np.float32)
|
|
y_test_np = np.asarray(y_test, dtype=np.int32)
|
|
|
|
# 4. Huấn luyện LightGBM
|
|
print("\n[MODEL] Bắt đầu huấn luyện LightGBM (Cân bằng lớp)...")
|
|
|
|
params = {
|
|
'objective': 'multiclass',
|
|
'num_class': 8,
|
|
'metric': 'multi_error',
|
|
'boosting_type': 'gbdt',
|
|
'learning_rate': 0.05,
|
|
'num_leaves': 31,
|
|
'max_depth': -1,
|
|
'feature_fraction': 0.8,
|
|
'class_weight': 'balanced', # Xử lý mất cân bằng dữ liệu
|
|
'verbose': -1,
|
|
'n_jobs': -1
|
|
}
|
|
|
|
model = lgb.LGBMClassifier(**params, n_estimators=300)
|
|
|
|
model.fit(
|
|
X_train_np, y_train_np,
|
|
eval_set=[(X_val_np, y_val_np)]
|
|
)
|
|
|
|
# 5. Đánh giá mô hình
|
|
print("\n[EVAL] Đang đánh giá trên tập Validation...")
|
|
y_val_pred = model.predict(X_val_np)
|
|
val_accuracy = accuracy_score(y_val_np, y_val_pred)
|
|
print(f"Validation Accuracy: {val_accuracy:.4f}")
|
|
|
|
y_pred_test = model.predict(X_test_np)
|
|
|
|
test_accuracy = accuracy_score(y_test_np, y_pred_test)
|
|
precision = precision_score(y_test_np, y_pred_test, average='weighted', zero_division=0)
|
|
recall = recall_score(y_test_np, y_pred_test, average='weighted', zero_division=0)
|
|
f1 = f1_score(y_test_np, y_pred_test, average='weighted', zero_division=0)
|
|
|
|
# 6. Lưu mô hình và Metadata
|
|
model_dir = "model_train"
|
|
os.makedirs(model_dir, exist_ok=True)
|
|
model_path = os.path.join(model_dir, "model_lightgbm.joblib")
|
|
joblib.dump(model, model_path)
|
|
print(f"\n[SAVE] Model saved to {model_path}")
|
|
|
|
info = {
|
|
"model_type": "LightGBM_Balanced",
|
|
"num_classes": 8,
|
|
"classes": list(label_mapping.keys()),
|
|
"num_features": X_train_np.shape[1],
|
|
"accuracy": float(test_accuracy),
|
|
"precision": float(precision),
|
|
"recall": float(recall),
|
|
"f1_score": float(f1),
|
|
"params": {
|
|
"n_estimators": 300,
|
|
"max_depth": -1
|
|
},
|
|
"description": "LightGBM trained on real Planetary Computer data with balanced class weights"
|
|
}
|
|
|
|
with open(os.path.join(model_dir, "model_lightgbm_info.json"), "w") as f:
|
|
json.dump(info, f, indent=2)
|
|
|
|
print("[SAVE] Model info saved.")
|