# -*- coding: utf-8 -*-
"""M2 Bab 5: interpolasi curah hujan tahunan dari 24 stasiun latih, dinilai dengan 6 stasiun uji.
Metode: IDW (p = 1, 2, 3), TIN linear, TIN Clough-Tocher, B-spline (GRASS), RST (GRASS), kriging biasa (NumPy), dan regresi + residu IDW (NumPy).
Penulis: Badar Mubarok Yogaswara. Pemakaian: python-qgis.bat m2_05a_interpolasi.py <folder paket-m2>
Keluaran: <paket-m2>/hasil/CH_tahunan.tif (metode terbaik) dan tabel galat (RMSE) di layar."""
import os
import sys
import numpy as np
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import _inisialisasi  # noqa: F401
import processing
from osgeo import gdal, ogr
from scipy.optimize import curve_fit

gdal.UseExceptions()
ogr.UseExceptions()
paket = sys.argv[1]
out = os.path.join(paket, "hasil")
os.makedirs(out, exist_ok=True)
P = lambda n: os.path.join(out, n)
X0, Y0, L, RES = 312000.0, 9996000.0, 2000.0, 10.0
EXT = "%f,%f,%f,%f [EPSG:32749]" % (X0, X0 + L, Y0, Y0 + L)

# ---- data stasiun: latih dan uji
ds = ogr.Open(os.path.join(paket, "Stasiun_Hujan.gpkg"))
lyr = ds.GetLayer(0)
st = [(f.GetGeometryRef().GetX(), f.GetGeometryRef().GetY(), f["elevasi_m"], f["ch_tahunan_mm"], f["peran"]) for f in lyr]
latih = np.array([s[:4] for s in st if s[4] == "latih"])
uji = np.array([s[:4] for s in st if s[4] == "uji"])
print("stasiun latih: %d, uji: %d" % (len(latih), len(uji)))
processing.run("native:extractbyexpression", {"INPUT": os.path.join(paket, "Stasiun_Hujan.gpkg|layername=Stasiun_Hujan"),
                                              "EXPRESSION": "\"peran\" = 'latih'", "OUTPUT": P("stasiun_latih.gpkg")})
from qgis.core import QgsVectorLayer
fi = QgsVectorLayer(P("stasiun_latih.gpkg"), "latih", "ogr").fields().indexOf("ch_tahunan_mm")      # indeks kolom menurut QGIS (kolom fid ikut dihitung)
data_in = "%s::~::0::~::%d::~::0" % (P("stasiun_latih.gpkg"), fi)


def contoh(raster_path, titik):
    r = gdal.Open(raster_path)
    gt = r.GetGeoTransform()
    a = r.GetRasterBand(1).ReadAsArray()
    nd = r.GetRasterBand(1).GetNoDataValue()
    h = np.array([a[int((gt[3] - y) / -gt[5]), int((x - gt[0]) / gt[1])] for x, y in titik[:, :2]], dtype="float64")
    if nd is not None:
        h[h == nd] = np.nan                           # di luar segitiga TIN: tidak ada nilai
    return h


hasil = {}
for p in (1, 2, 3):
    processing.run("qgis:idwinterpolation", {"INTERPOLATION_DATA": data_in, "DISTANCE_COEFFICIENT": p, "EXTENT": EXT, "PIXEL_SIZE": RES,
                                             "OUTPUT": P("CH_idw_p%d.tif" % p)})
    hasil["IDW p=%d" % p] = P("CH_idw_p%d.tif" % p)
for nama, m in (("TIN linear", 0), ("TIN Clough-Tocher", 1)):
    fn = P("CH_tin_%d.tif" % m)
    processing.run("qgis:tininterpolation", {"INTERPOLATION_DATA": data_in, "METHOD": m, "EXTENT": EXT, "PIXEL_SIZE": RES, "OUTPUT": fn})
    hasil[nama] = fn
BAKU = {"GRASS_REGION_PARAMETER": EXT, "GRASS_REGION_CELLSIZE_PARAMETER": RES}
processing.run("grass:v.surf.bspline", dict(BAKU, input=P("stasiun_latih.gpkg"), column="ch_tahunan_mm", ew_step=500, ns_step=500, method=1,
               lambda_i=0.01, solver=0, maxit=10000, error=1e-6, memory=300, raster_output=P("CH_bspline.tif")))
hasil["B-spline bikubik"] = P("CH_bspline.tif")
processing.run("grass:v.surf.rst", dict(BAKU, input=P("stasiun_latih.gpkg"), zcolumn="ch_tahunan_mm", tension=40, smooth=0.5, segmax=40, npmin=150,
               elevation=P("CH_rst.tif")))
hasil["RST (spline tegangan)"] = P("CH_rst.tif")

# ---- baris dan kolom sel pusat grid, untuk metode NumPy
n = int(L / RES)
gx = X0 + (np.arange(n) + 0.5) * RES
gy = Y0 + L - (np.arange(n) + 0.5) * RES
GX, GY = np.meshgrid(gx, gy)
dem = gdal.Open(os.path.join(paket, "DEM_10m.tif")).ReadAsArray().astype("float64")


def idw_np(xy, v, tx, ty, p=2):
    d = np.hypot(tx[..., None] - xy[:, 0], ty[..., None] - xy[:, 1])
    d = np.maximum(d, 1e-9)
    w = 1.0 / d ** p
    return (w * v).sum(-1) / w.sum(-1)


# kriging biasa dengan semivariogram eksponensial yang disesuaikan ke data latih
xy = latih[:, :2]
v = latih[:, 3]
ii, jj = np.triu_indices(len(xy), 1)
h = np.hypot(*(xy[ii] - xy[jj]).T)
gam = 0.5 * (v[ii] - v[jj]) ** 2
tepi = np.linspace(0, h.max() * 0.6, 8)
hc, gc = [], []
for a, b in zip(tepi[:-1], tepi[1:]):
    m = (h >= a) & (h < b)
    if m.sum() >= 5:
        hc.append(h[m].mean())
        gc.append(gam[m].mean())
model = lambda hh, c0, c, a: c0 + c * (1 - np.exp(-3 * hh / a))
(c0, c, a), _ = curve_fit(model, hc, gc, p0=[gc[0], max(gc) - gc[0], hc[-1]], bounds=([0, 1, 100], [1e7, 1e7, 5000]))
print("semivariogram eksponensial: nugget=%.0f, sill parsial=%.0f, jangkauan=%.0f m" % (c0, c, a))


def krig(tx, ty):
    N = len(xy)
    D = np.hypot(xy[:, None, 0] - xy[None, :, 0], xy[:, None, 1] - xy[None, :, 1])
    A = np.ones((N + 1, N + 1))
    A[:N, :N] = model(D, 0, c, a) + np.where(D == 0, 0, c0)
    A[N, N] = 0
    Ainv = np.linalg.inv(A)
    shp = tx.shape
    d0 = np.hypot(tx.ravel()[:, None] - xy[:, 0], ty.ravel()[:, None] - xy[:, 1])
    B = np.ones((d0.shape[0], N + 1))
    B[:, :N] = model(d0, 0, c, a) + c0
    lam = B @ Ainv
    return (lam[:, :N] @ v).reshape(shp)


# regresi pada elevasi (dan koordinat X), lalu residunya dengan IDW p=2
def elev_di(x, y):
    return dem[np.clip(((Y0 + L - y) / RES).astype(int), 0, n - 1), np.clip(((x - X0) / RES).astype(int), 0, n - 1)]


Xr = np.column_stack([np.ones(len(latih)), latih[:, 2], latih[:, 0] - X0])
beta, *_ = np.linalg.lstsq(Xr, v, rcond=None)
res = v - Xr @ beta
print("regresi CH = %.1f + %.2f * elevasi + %.4f * (X - X0)   (R2 = %.3f)" % (beta[0], beta[1], beta[2], 1 - res.var() / v.var()))


def regresi(tx, ty):
    return beta[0] + beta[1] * elev_di(tx, ty) + beta[2] * (tx - X0) + idw_np(xy, res, tx, ty, 2)


np_metode = {"Kriging biasa (NumPy)": krig, "Regresi elevasi + residu IDW (NumPy)": regresi}
tabel = []
for nama, fn in hasil.items():
    pred = contoh(fn, uji)
    tabel.append((nama, pred))
for nama, fn in np_metode.items():
    pred = fn(uji[:, 0], uji[:, 1])
    tabel.append((nama, pred))
print("\nGalat pada 6 stasiun uji (mm/tahun), dihitung dari prediksi - pengamatan")
print("%-40s %8s %8s %8s %4s" % ("metode", "RMSE", "MAE", "bias", "n"))
skor = {}
for nama, pred in tabel:
    e = (pred - uji[:, 3])
    e = e[~np.isnan(e)]
    skor[nama] = float(np.sqrt((e ** 2).mean())) if len(e) == len(uji) else 1e9      # metode yang tak mencakup semua titik uji tidak dipilih
    print("%-40s %8.1f %8.1f %8.1f %4d" % (nama, np.sqrt((e ** 2).mean()), np.abs(e).mean(), e.mean(), len(e)))
print("terbaik menurut RMSE:", min(skor, key=skor.get))

# simpan peta terbaik dari metode NumPy bila memang terbaik; jika tidak, salin raster QGIS terbaik
terbaik = min(skor, key=skor.get)
if terbaik in np_metode:
    arr = np_metode[terbaik](GX, GY).astype("float32")
    d = gdal.GetDriverByName("GTiff").Create(P("CH_tahunan.tif"), n, n, 1, gdal.GDT_Float32, ["COMPRESS=DEFLATE"])
    d.SetGeoTransform((X0, RES, 0, Y0 + L, 0, -RES))
    d.SetProjection(gdal.Open(os.path.join(paket, "DEM_10m.tif")).GetProjection())
    d.GetRasterBand(1).WriteArray(arr)
    d = None
else:
    gdal.Translate(P("CH_tahunan.tif"), hasil[terbaik], creationOptions=["COMPRESS=DEFLATE"])
print("CH_tahunan.tif dibuat dari:", terbaik)

# validasi silang tinggalkan-satu (LOO) untuk pangkat IDW pada 24 stasiun latih
print("\nTinggalkan-satu (LOO) pada 24 stasiun latih, RMSE IDW menurut pangkat:")
for p in (0.5, 1, 2, 3, 4):
    e = []
    for k in range(len(xy)):
        m = np.arange(len(xy)) != k
        e.append(idw_np(xy[m], v[m], np.array([xy[k, 0]]), np.array([xy[k, 1]]), p)[0] - v[k])
    print("  p=%.1f  RMSE=%.1f" % (p, np.sqrt(np.mean(np.square(e)))))
