# -*- coding: utf-8 -*-
# I2 Bab 5: klasifikasi terbimbing (jarak minimum dan kemiripan maksimum, ditulis dengan NumPy). Penulis: Badar Mubarok Yogaswara
# Jalankan di Python Console QGIS. Ubah dua jalur di bawah sesuai komputer Anda.
import os
import numpy as np
import processing
from osgeo import gdal

PAKET_I2 = os.environ.get("I2_PAKET", r"C:/KPH_Contoh/paket-i2")
HASIL = r"C:/temp/hasil_i2"
os.makedirs(HASIL, exist_ok=True)
NAMA = {1: "Hutan", 2: "Kebun", 3: "Sawah", 4: "Lahan Terbuka"}

citra_tif = os.path.join(PAKET_I2, "Citra_Lahan.tif")
ds = gdal.Open(citra_tif)
gt, proj = ds.GetGeoTransform(), ds.GetProjection()
citra = ds.ReadAsArray().astype("float64")
_, tinggi, lebar = citra.shape
piksel = citra.reshape(4, -1).T

# 1) ubah poligon area latih menjadi raster kelas (piksel di dalam poligon bernilai Kode, lainnya 0)
latih_tif = os.path.join(HASIL, "area_latih.tif")
processing.run("gdal:rasterize", {
    "INPUT": os.path.join(PAKET_I2, "Area_Latih.gpkg"), "FIELD": "Kode", "UNITS": 1, "WIDTH": 1.0, "HEIGHT": 1.0,
    "EXTENT": "312000,312300,9996000,9996300 [EPSG:32749]", "NODATA": 0, "DATA_TYPE": 0, "INIT": 0, "OUTPUT": latih_tif})
latih = gdal.Open(latih_tif).ReadAsArray().ravel()

# 2) tanda spektral: rata-rata dan sebaran (kovarians) tiap kelas dari piksel latihnya
rata, kov = {}, {}
print("Kelas          | piksel latih |   R     G     B    NIR (rata-rata)")
for k in NAMA:
    x = piksel[latih == k]
    assert len(x) > 4, "Area latih kelas %d terlalu sedikit" % k
    rata[k], kov[k] = x.mean(axis=0), np.cov(x.T)
    print("%-14s | %6d       | %5.1f %5.1f %5.1f %5.1f" % (NAMA[k], len(x), *rata[k]))

# 3a) jarak minimum: tiap piksel masuk kelas yang rata-ratanya paling dekat
jarak = np.stack([((piksel - rata[k]) ** 2).sum(axis=1) for k in NAMA], axis=1)
jarak_min = jarak.argmin(axis=1) + 1


# 3b) kemiripan maksimum: seperti jarak minimum, tetapi memperhitungkan sebaran tiap kelas (kovarians)
def skor(x, mu, c):
    d = x - mu
    return -0.5 * np.einsum("ij,jk,ik->i", d, np.linalg.inv(c), d) - 0.5 * np.log(np.linalg.det(c))


kemiripan = np.stack([skor(piksel, rata[k], kov[k]) for k in NAMA], axis=1)
maks_lik = kemiripan.argmax(axis=1) + 1


def simpan(arr, nama):
    jalur = os.path.join(HASIL, nama)
    out = gdal.GetDriverByName("GTiff").Create(jalur, lebar, tinggi, 1, gdal.GDT_Byte)
    out.SetGeoTransform(gt)
    out.SetProjection(proj)
    out.GetRasterBand(1).WriteArray(arr.reshape(tinggi, lebar).astype("uint8"))
    out.FlushCache()
    out = None
    print("Tersimpan:", jalur)


simpan(jarak_min, "terbimbing_jarak_min.tif")
simpan(maks_lik, "terbimbing_maks_lik.tif")
for nama, arr in (("jarak minimum", jarak_min), ("kemiripan maksimum", maks_lik)):
    print(nama, "-> piksel per kelas:", {NAMA[k]: int((arr == k).sum()) for k in NAMA})
