# -*- coding: utf-8 -*-
# I2 Bab 4: klasifikasi tak terbimbing dengan k-means (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)
JUMLAH_KLUSTER = [6, 10]    # dua percobaan: 6 kluster, lalu 10 kluster
NAMA = {1: "Hutan", 2: "Kebun", 3: "Sawah", 4: "Lahan Terbuka"}

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

# lokasi yang kelasnya sudah diketahui (area latih) dipakai hanya untuk MEMBERI NAMA kluster
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()


def kmeans(x, k, ulang=30, seed=1):
    """K-means sederhana: pusat awal dipilih acak, lalu dihitung ulang sampai stabil."""
    rng = np.random.default_rng(seed)
    pusat = x[rng.choice(len(x), k, replace=False)]
    for _ in range(ulang):
        jarak = ((x[:, None, :] - pusat[None, :, :]) ** 2).sum(axis=2)     # jarak kuadrat tiap piksel ke tiap pusat
        label = jarak.argmin(axis=1)
        baru = np.array([x[label == i].mean(axis=0) if np.any(label == i) else pusat[i] for i in range(k)])
        if np.allclose(baru, pusat, atol=1e-3):
            break
        pusat = baru
    return label, pusat


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.astype("uint8"))
    out.FlushCache()
    out = None
    return jalur


for k in JUMLAH_KLUSTER:
    label, pusat = kmeans(piksel, k)
    peta = np.zeros(len(piksel), dtype=int)
    print("\n=== %d kluster ===" % k)
    print("Kluster | piksel  persen |   R     G     B    NIR | piksel latih per kelas (H/Kb/Sw/LT) -> diberi nama")
    for i in np.argsort(-pusat[:, 3]):                       # urut menurut NIR menurun
        n = int((label == i).sum())
        hitung = np.bincount(latih[(label == i) & (latih > 0)], minlength=5)[1:]
        kelas = int(hitung.argmax()) + 1 if hitung.sum() else 0     # suara terbanyak; 0 = tidak ada piksel latih
        if kelas == 0:
            print("  PERINGATAN: kluster %d tidak punya piksel latih; kelasnya 0 (belum bernama)." % (i + 1))
        peta[label == i] = kelas
        r, g, b, nir = pusat[i]
        print("  %2d    | %6d  %5.1f%% | %5.1f %5.1f %5.1f %5.1f | %s -> %s" % (
            i + 1, n, 100.0 * n / label.size, r, g, b, nir, "/".join(str(int(v)) for v in hitung), NAMA.get(kelas, "(belum bernama)")))
    print("Tersimpan:", simpan(peta.reshape(tinggi, lebar), "tak_terbimbing_k%d.tif" % k))
