# -*- coding: utf-8 -*-
# SKRIP 6.2: Latih hutan acak dari titik latih, uji silang acak vs uji silang blok ruang, lalu petakan seluruh area
# Penulis: Badar Mubarok Yogaswara
# Syarat: Skrip 6.1 sudah dijalankan. Memakai scikit-learn bila terpasang; bila tidak, memakai hutan_mini.py.
import os
import numpy as np
import processing
from qgis.core import QgsVectorLayer
import m3_umum as U

try:
    from sklearn.ensemble import RandomForestClassifier
    print("Memakai scikit-learn")
except ImportError:
    from hutan_mini import RandomForestClassifier
    print("scikit-learn tidak ada; memakai hutan_mini.py")

fitur_tif = os.path.join(U.HASIL, "fitur_deret.tif")
titik = QgsVectorLayer(os.path.join(U.PAKET, "vektor", "Titik_Latih.gpkg"), "latih")
# ambil nilai 12 fitur pada tiap titik latih (Sample raster values)
s = processing.run("native:rastersampling", {"INPUT": titik, "RASTERCOPY": fitur_tif, "COLUMN_PREFIX": "f", "OUTPUT": "TEMPORARY_OUTPUT"})["OUTPUT"]
nama_f = [fl.name() for fl in s.fields() if fl.name().startswith("f") and fl.name() != "fid"]
X = np.array([[f[n] for n in nama_f] for f in s.getFeatures()], dtype=float)
y = np.array([f["KELAS"] for f in s.getFeatures()])
xy = np.array([[f.geometry().asPoint().x(), f.geometry().asPoint().y()] for f in s.getFeatures()])
print("Titik latih: %d, fitur: %d, kelas: %s" % (len(y), X.shape[1], dict(zip(*np.unique(y, return_counts=True)))))

# --- uji silang 1: acak 4 lipatan; uji silang 2: lipatan = blok ruang 200 m x 200 m (4 blok, area 400 m)
rng = np.random.default_rng(5)
lipat_acak = rng.permutation(len(y)) % 4
lipat_blok = (xy[:, 0] >= 312200).astype(int) + 2 * (xy[:, 1] >= 9996200).astype(int)


def uji_silang(lipat):
    benar = []
    for k in range(4):
        tr, te = lipat != k, lipat == k
        m = RandomForestClassifier(n_estimators=100, random_state=1).fit(X[tr], y[tr])
        benar.append((m.predict(X[te]) == y[te]).mean())
    return np.array(benar)


a, b = uji_silang(lipat_acak), uji_silang(lipat_blok)
print("Akurasi uji silang ACAK : %s rata-rata %.3f" % (np.round(a, 2), a.mean()))
print("Akurasi uji silang BLOK : %s rata-rata %.3f" % (np.round(b, 2), b.mean()))

# --- latih akhir dan petakan
model = RandomForestClassifier(n_estimators=200, random_state=1).fit(X, y)
urut = np.argsort(model.feature_importances_)[::-1]
from osgeo import gdal
ds = gdal.Open(fitur_tif)
nama_fitur = [ds.GetRasterBand(i + 1).GetDescription() for i in range(ds.RasterCount)]
print("Kepentingan fitur (5 teratas):", ", ".join("%s %.2f" % (nama_fitur[i], model.feature_importances_[i]) for i in urut[:5]))
st, gt, prj = U.baca_tif(fitur_tif)                  # (12, H, W)
nb, H, W = st.shape
peta = model.predict(st.reshape(nb, -1).T).reshape(H, W).astype("uint8")
U.tulis_tif(os.path.join(U.HASIL, "peta_rf.tif"), peta, gt, prj, nodata=0)
din, _, _ = U.baca_tif(os.path.join(U.PAKET, "acuan", "Acuan_Dinamika.tif"))
nama = {1: "Hutan alam", 2: "Hutan tanaman", 3: "Pertanian", 4: "Air", 5: "Terbuka", 6: "Deforestasi", 7: "Terbakar", 8: "Panen HTI"}
print("Luas hasil peta (ha) / luas acuan (ha):")
for k in range(1, 9):
    print("  %-14s %6.2f / %6.2f" % (nama[k], (peta == k).sum() * 0.01, (din == k).sum() * 0.01))
print("Kesesuaian peta dengan acuan di SEMUA piksel (bukan uji mandiri): %.3f" % (peta == din).mean())
