# -*- coding: utf-8 -*-
# SKRIP 7.1: Validasi peta: sampel acak berstrata, matriks kebingungan, akurasi, dan estimasi luas terkoreksi
# Penulis: Badar Mubarok Yogaswara
# Dua peta diuji: (A) hutan acak dengan 12 fitur (hasil Skrip 6.2) dan (B) hutan acak yang hanya memakai fitur NDVI.
# Acuan: Acuan_Dinamika.tif. Dalam praktik, acuan berasal dari tafsir citra resolusi tinggi atau survei lapangan.
import os
import numpy as np
from qgis.core import QgsVectorLayer
import m3_umum as U

try:
    from sklearn.ensemble import RandomForestClassifier
except ImportError:
    from hutan_mini import RandomForestClassifier

NAMA = {1: "Hutan alam", 2: "Hutan tanaman", 3: "Pertanian", 4: "Air", 5: "Terbuka", 6: "Deforestasi", 7: "Terbakar", 8: "Panen HTI"}
N_STRATUM = 20            # sampel per kelas peta
LUAS_PIKSEL = 0.01        # ha

st, gt, prj = U.baca_tif(os.path.join(U.HASIL, "fitur_deret.tif"))
acuan, _, _ = U.baca_tif(os.path.join(U.PAKET, "acuan", "Acuan_Dinamika.tif"))
latih = QgsVectorLayer(os.path.join(U.PAKET, "vektor", "Titik_Latih.gpkg"), "latih")
pl = [(U.xy_ke_piksel(gt, f.geometry().asPoint().x(), f.geometry().asPoint().y()), f["KELAS"]) for f in latih.getFeatures()]
tr = np.array([p[0] for p in pl])
yl = np.array([p[1] for p in pl])
pakai_latih = np.zeros(acuan.shape, bool)
pakai_latih[tr[:, 0], tr[:, 1]] = True

# peta B: hanya 7 fitur NDVI (band 1 sampai 7 pada tumpukan fitur)
idx = list(range(7))
mB = RandomForestClassifier(n_estimators=100, random_state=1).fit(st[idx][:, tr[:, 0], tr[:, 1]].T, yl)
petaB = mB.predict(st[idx].reshape(len(idx), -1).T).reshape(acuan.shape).astype("uint8")
jB = os.path.join(U.HASIL, "peta_ndvi_saja.tif")
U.tulis_tif(jB, petaB, gt, prj, nodata=0)
jA = os.path.join(U.HASIL, "peta_rf.tif")


def sampel_berstrata(peta, seed):
    """Sampel acak berstrata: N_STRATUM piksel acak dari tiap kelas peta, di luar titik latih.
    Padanan di antarmuka QGIS: poligonkan peta, lebur per kelas, lalu Random points in polygons (hasil tiap jalan tidak identik)."""
    rng = np.random.default_rng(seed)
    hasil = {}
    for k in np.unique(peta):
        if k == 0:
            continue
        kand = np.argwhere((peta == k) & ~pakai_latih)
        pilih = rng.choice(len(kand), size=min(N_STRATUM, len(kand)), replace=False)
        hasil[int(k)] = [((int(kand[p][0]), int(kand[p][1])), int(acuan[kand[p][0], kand[p][1]])) for p in pilih]
    return hasil

def analisis(peta, sampel, nama, cetak=True):
    kelas = sorted(sampel)
    pop = peta[~pakai_latih]                          # populasi sampel = piksel di luar titik latih
    A = pop.size * LUAS_PIKSEL
    W = {k: (pop == k).sum() / pop.size for k in kelas}
    n = {k: len(sampel[k]) for k in kelas}
    nij = np.zeros((len(kelas), len(kelas)))
    for a, k in enumerate(kelas):
        for _, r in sampel[k]:
            if r in kelas:
                nij[a, kelas.index(r)] += 1
    p = np.array([[W[k] * nij[a, b] / n[k] for b in range(len(kelas))] for a, k in enumerate(kelas)])
    oa = np.trace(p)
    ua = np.diag(p) / p.sum(1)
    pa = np.diag(p) / np.maximum(p.sum(0), 1e-12)
    # piksel titik latih sudah punya label pasti: luasnya ditambahkan apa adanya
    luas_peta = np.array([(peta == k).sum() * LUAS_PIKSEL for k in kelas])
    luas_adj = p.sum(0) * A + np.array([(acuan[pakai_latih] == k).sum() * LUAS_PIKSEL for k in kelas])
    se = np.array([np.sqrt(sum(W[ki] ** 2 * (nij[a, b] / n[ki]) * (1 - nij[a, b] / n[ki]) / (n[ki] - 1) for a, ki in enumerate(kelas))) for b in range(len(kelas))]) * A
    if cetak:
        print("\n=== %s: %d sampel; baris = kelas peta, kolom = kelas acuan ===" % (nama, nij.sum()))
        print("          " + " ".join("%5d" % k for k in kelas))
        for a, k in enumerate(kelas):
            print("kelas %d   " % k + " ".join("%5d" % v for v in nij[a]) + "   | n=%d" % n[k])
        print("Akurasi keseluruhan (berbobot luas): %.3f" % oa)
        print("%-14s %8s %8s %10s %12s %10s" % ("Kelas", "UA", "PA", "luas peta", "luas terkoreksi", "+/- 95%"))
        for a, k in enumerate(kelas):
            print("%-14s %8.2f %8.2f %9.2f ha %11.2f ha %8.2f ha" % (NAMA[k], ua[a], pa[a], luas_peta[a], luas_adj[a], 1.96 * se[a]))
    return oa, kelas, luas_adj, se


for nama, jalur in (("Peta A (12 fitur)", jA), ("Peta B (NDVI saja)", jB)):
    peta, _, _ = U.baca_tif(jalur)
    oa, kelas, luas_adj, se = analisis(peta, sampel_berstrata(peta, 11), nama)
    print("Luas acuan sebenarnya: " + ", ".join("%s %.2f" % (NAMA[k], (acuan == k).sum() * LUAS_PIKSEL) for k in (6, 7)))
    # kestabilan terhadap pengundian sampel
    oas = [analisis(peta, sampel_berstrata(peta, s), nama, cetak=False)[0] for s in range(21, 26)]
    print("Akurasi keseluruhan pada 5 undian sampel lain: " + ", ".join("%.3f" % v for v in oas))


