refactor: reorganize project structure by moving core modules and update import paths in API server
This commit is contained in:
@@ -0,0 +1,122 @@
|
||||
#!/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 *
|
||||
|
||||
from sklearn.ensemble import RandomForestClassifier
|
||||
from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, confusion_matrix
|
||||
|
||||
print("=" * 70)
|
||||
print("🚀 Training Random Forest 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)
|
||||
|
||||
print("[DATA] Đang loại bỏ mây...")
|
||||
result = mask_clean(data)
|
||||
|
||||
print("[DATA] Đang tính toán NDVI...")
|
||||
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"),
|
||||
]
|
||||
|
||||
print("[DATA] Đang nội suy (fill_nan) cho mây...")
|
||||
fill_nan_ndvi = fill_nan(ndvi, time_split)
|
||||
|
||||
print("[DATA] Đang tính trung bình tháng (resample 1M)...")
|
||||
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)
|
||||
|
||||
# 4. Huấn luyện Random Forest
|
||||
print("\n[MODEL] Bắt đầu huấn luyện Random Forest...")
|
||||
model = RandomForestClassifier(
|
||||
n_estimators=100,
|
||||
max_depth=15,
|
||||
random_state=42,
|
||||
n_jobs=-1 # Dùng tất cả nhân CPU
|
||||
)
|
||||
|
||||
model.fit(X_train_np, y_train_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}")
|
||||
|
||||
X_test_np = np.asarray(X_test, dtype=np.float32)
|
||||
y_test_np = np.asarray(y_test, dtype=np.int32)
|
||||
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_randomforest.joblib")
|
||||
joblib.dump(model, model_path)
|
||||
print(f"\n[SAVE] Model saved to {model_path}")
|
||||
|
||||
info = {
|
||||
"model_type": "RandomForest_RealData",
|
||||
"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),
|
||||
"description": "Random Forest trained on real Planetary Computer data"
|
||||
}
|
||||
|
||||
with open(os.path.join(model_dir, "model_randomforest_info.json"), "w") as f:
|
||||
json.dump(info, f, indent=2)
|
||||
|
||||
print("[SAVE] Model info saved.")
|
||||
Reference in New Issue
Block a user