refactor: reorganize project structure by moving core modules and update import paths in API server
This commit is contained in:
@@ -0,0 +1,306 @@
|
||||
"""
|
||||
Updated run_prediction function for api_server.py
|
||||
Uses FeatureExtractor for consistent feature extraction
|
||||
"""
|
||||
|
||||
async def run_prediction(config: PredictionConfig):
|
||||
"""Chạy prediction process - Sử dụng FeatureExtractor để đồng bộ với training"""
|
||||
global prediction_status
|
||||
|
||||
try:
|
||||
prediction_status["progress"] = "Đang import thư viện..."
|
||||
|
||||
# Import required libraries
|
||||
import numpy as np
|
||||
import xarray as xr
|
||||
from datetime import datetime as dt
|
||||
import hashlib
|
||||
from feature_extractor import get_feature_extractor
|
||||
|
||||
# Validate bbox
|
||||
if (config.min_lon < -180 or config.max_lon > 180 or
|
||||
config.min_lat < -90 or config.max_lat > 90):
|
||||
raise ValueError(f"Bbox không hợp lệ: ({config.min_lon}, {config.min_lat}, {config.max_lon}, {config.max_lat}). "
|
||||
f"Phải trong phạm vi (-180, -90, 180, 90)")
|
||||
|
||||
prediction_status["progress"] = "Đang load model..."
|
||||
|
||||
# Load model using ModelManager
|
||||
model_manager = get_model_manager()
|
||||
model, label_encoder, model_metadata = model_manager.load_model(config.model_filename)
|
||||
|
||||
# Get feature_mode from metadata (default to 'simple' if not specified)
|
||||
feature_mode = model_metadata.get("feature_mode", "simple")
|
||||
required_features = model_metadata.get("features", [])
|
||||
n_features_expected = model_metadata.get("n_features", len(required_features))
|
||||
|
||||
prediction_status["progress"] = f"Model: {model_metadata.get('model_type', 'unknown')}, mode={feature_mode}, features={n_features_expected}"
|
||||
|
||||
# Initialize FeatureExtractor with same mode as training
|
||||
extractor = get_feature_extractor(mode=feature_mode)
|
||||
|
||||
# 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..."
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
raise ImportError("PyTorch required for CNN models. Install: pip install torch")
|
||||
|
||||
# Initialize common variables
|
||||
bbox = [config.min_lon, config.min_lat, config.max_lon, config.max_lat]
|
||||
time_range = f"{config.start_date}/{config.end_date}"
|
||||
|
||||
# ============ LOAD SENTINEL-2 DATA ============
|
||||
prediction_status["progress"] = "Đang kết nối Microsoft Planetary Computer..."
|
||||
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,
|
||||
)
|
||||
|
||||
prediction_status["progress"] = "Đang tải dữ liệu Sentinel-2..."
|
||||
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 different bands based on feature mode
|
||||
if feature_mode == 'simple':
|
||||
bands_to_load = ["B04", "B08", "SCL"]
|
||||
else: # temporal or extended
|
||||
bands_to_load = ["B02", "B03", "B04", "B08", "B11", "SCL"]
|
||||
|
||||
s2_data = load(
|
||||
s2_items,
|
||||
bbox=bbox,
|
||||
bands=bands_to_load,
|
||||
chunks={"time": 1, "x": 2048, "y": 2048},
|
||||
groupby="solar_day",
|
||||
resolution=config.resolution
|
||||
).compute()
|
||||
|
||||
prediction_status["progress"] = "Đã load Sentinel-2 data"
|
||||
|
||||
# ============ LOAD SENTINEL-1 DATA (RADAR) ============
|
||||
prediction_status["progress"] = "Đang tải dữ liệu Sentinel-1 (Radar)..."
|
||||
use_radar = False
|
||||
vh_data = None
|
||||
vv_data = None
|
||||
|
||||
try:
|
||||
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]
|
||||
s1_data = load(
|
||||
s1_items,
|
||||
bbox=bbox,
|
||||
bands=["vh", "vv"],
|
||||
chunks={"time": 1, "x": 2048, "y": 2048},
|
||||
groupby="solar_day",
|
||||
resolution=config.resolution
|
||||
).compute()
|
||||
|
||||
# Convert to dB
|
||||
vh_data = 10 * np.log10(s1_data['vh'].where(s1_data['vh'] > 0))
|
||||
vv_data = 10 * np.log10(s1_data['vv'].where(s1_data['vv'] > 0))
|
||||
use_radar = True
|
||||
prediction_status["progress"] = f"Đã load Sentinel-1 data ({len(s1_items)} scenes)"
|
||||
else:
|
||||
prediction_status["progress"] = "Không có dữ liệu Sentinel-1, bỏ qua radar features"
|
||||
except Exception as e:
|
||||
prediction_status["progress"] = f"Lỗi load Sentinel-1: {str(e)}, bỏ qua radar features"
|
||||
|
||||
# ============ APPLY CLOUD MASK ============
|
||||
prediction_status["progress"] = "Đang xử lý mây..."
|
||||
if "SCL" in s2_data:
|
||||
scl = s2_data["SCL"]
|
||||
# SCL values: 3=cloud shadow, 8=cloud medium, 9=cloud high, 10=cirrus
|
||||
cloud_mask = (scl == 3) | (scl == 8) | (scl == 9) | (scl == 10)
|
||||
for band in s2_data.data_vars:
|
||||
if band != "SCL":
|
||||
s2_data[band] = s2_data[band].where(~cloud_mask)
|
||||
|
||||
# ============ EXTRACT FEATURES ============
|
||||
prediction_status["progress"] = f"Đang trích xuất features (mode={feature_mode})..."
|
||||
|
||||
if feature_mode == 'simple':
|
||||
# Calculate NDVI for simple mode
|
||||
nir = s2_data["B08"].astype('float32')
|
||||
red = s2_data["B04"].astype('float32')
|
||||
ndvi = (nir - red) / (nir + red + 1e-8)
|
||||
|
||||
# Fill NaN
|
||||
ndvi_filled = ndvi.ffill(dim='time').bfill(dim='time')
|
||||
|
||||
# Extract features using FeatureExtractor
|
||||
features = extractor.extract(
|
||||
ndvi_data=ndvi_filled,
|
||||
vh_data=vh_data,
|
||||
vv_data=vv_data
|
||||
)
|
||||
else:
|
||||
# temporal or extended mode
|
||||
# Fill NaN values in spectral bands
|
||||
for band in ["B02", "B03", "B04", "B08", "B11"]:
|
||||
if band in s2_data:
|
||||
s2_data[band] = s2_data[band].ffill(dim='time').bfill(dim='time')
|
||||
|
||||
# Extract features using FeatureExtractor
|
||||
features = extractor.extract(
|
||||
s2_data=s2_data,
|
||||
vh_data=vh_data,
|
||||
vv_data=vv_data
|
||||
)
|
||||
|
||||
# Handle NaN values
|
||||
features = np.nan_to_num(features, nan=0.0)
|
||||
|
||||
prediction_status["progress"] = f"Đã extract {features.shape[1]} features cho {features.shape[0]} pixels"
|
||||
|
||||
# ============ PREDICT ============
|
||||
prediction_status["progress"] = "Đang dự đoán..."
|
||||
|
||||
# Make prediction
|
||||
if is_cnn_model:
|
||||
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.astype(int))
|
||||
except:
|
||||
pass
|
||||
|
||||
# Reshape to original shape
|
||||
if feature_mode == 'simple' and 'B08' in s2_data:
|
||||
# Use B08 to get shape
|
||||
y_size = len(s2_data.y)
|
||||
x_size = len(s2_data.x)
|
||||
else:
|
||||
y_size = len(s2_data.y)
|
||||
x_size = len(s2_data.x)
|
||||
|
||||
pred_shape = (y_size, x_size)
|
||||
predictions_2d = predictions.reshape(pred_shape)
|
||||
|
||||
# ============ CREATE OUTPUT ============
|
||||
prediction_status["progress"] = "Đang tạo bản đồ phân loại..."
|
||||
|
||||
# Create output xarray
|
||||
prediction_da = xr.DataArray(
|
||||
predictions_2d,
|
||||
coords={
|
||||
"y": s2_data.y,
|
||||
"x": s2_data.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")
|
||||
|
||||
# Generate PNG preview
|
||||
prediction_status["progress"] = "Đang tạo PNG preview..."
|
||||
png_file = output_dir / f"prediction_{timestamp}.png"
|
||||
try:
|
||||
import matplotlib
|
||||
matplotlib.use('Agg')
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
fig, ax = plt.subplots(figsize=(12, 10), dpi=150)
|
||||
im = ax.imshow(predictions_2d, cmap='tab20', interpolation='nearest')
|
||||
ax.set_title(f'Prediction Result - {timestamp}', fontsize=14, fontweight='bold')
|
||||
ax.set_xlabel('X (pixels)', fontsize=10)
|
||||
ax.set_ylabel('Y (pixels)', fontsize=10)
|
||||
|
||||
cbar = plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
|
||||
cbar.set_label('Class', rotation=270, labelpad=15)
|
||||
ax.grid(True, alpha=0.3, linestyle='--', linewidth=0.5)
|
||||
|
||||
plt.tight_layout()
|
||||
plt.savefig(str(png_file), dpi=150, bbox_inches='tight')
|
||||
plt.close(fig)
|
||||
|
||||
print(f"[PNG PREVIEW] Created: {png_file}")
|
||||
except Exception as e:
|
||||
print(f"[PNG PREVIEW ERROR] Failed to create PNG: {e}")
|
||||
png_file = None
|
||||
|
||||
# Get unique classes
|
||||
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! Đang tạo báo cáo..."
|
||||
prediction_status["output_file"] = str(output_file)
|
||||
prediction_status["result"] = {
|
||||
"output_file": str(output_file),
|
||||
"png_file": str(png_file) if png_file else None,
|
||||
"shape": list(pred_shape),
|
||||
"unique_classes": unique_classes,
|
||||
"bbox": bbox,
|
||||
"time_range": time_range,
|
||||
"n_features": features.shape[1],
|
||||
"feature_mode": feature_mode,
|
||||
"used_radar": use_radar,
|
||||
"model_used": config.model_filename
|
||||
}
|
||||
|
||||
# Auto generate prediction report
|
||||
try:
|
||||
report_path, _ = generate_prediction_report(prediction_status["result"])
|
||||
prediction_status["result"]["report_path"] = report_path
|
||||
prediction_status["result"]["report_filename"] = Path(report_path).name
|
||||
prediction_status["progress"] = "Hoàn thành! Báo cáo đã được tạo."
|
||||
print(f"[PREDICTION REPORT] Generated: {report_path}")
|
||||
except Exception as e:
|
||||
print(f"[PREDICTION REPORT ERROR] Failed to generate report: {e}")
|
||||
prediction_status["progress"] = "Hoàn thành! (Không thể tạo báo cáo)"
|
||||
|
||||
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(f"[PREDICTION ERROR] {str(e)}")
|
||||
print(traceback.format_exc())
|
||||
Reference in New Issue
Block a user