refactor: reorganize project structure by moving core modules and update import paths in API server

This commit is contained in:
2026-07-18 01:24:30 +07:00
parent abab846884
commit a82b2f6fa5
155 changed files with 25 additions and 370 deletions
+122
View File
@@ -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.")