# -*- coding: utf-8 -*-
"""hutan_mini.py: hutan acak (random forest) ringkas dengan NumPy saja. Penulis: Badar Mubarok Yogaswara
Tujuan: (1) membuka "kotak hitam" supaya Anda paham cara kerjanya, (2) tetap bisa jalan di QGIS yang belum memasang scikit-learn.
Antarmukanya sengaja dibuat mirip scikit-learn: RandomForestClassifier(...).fit(X, y), .predict(X), .feature_importances_.
Ini alat belajar; untuk pekerjaan besar pakai scikit-learn."""
import numpy as np


class _Pohon:
    def __init__(self, maks_fitur, kedalaman, daun_min, rng):
        self.mf, self.kd, self.dm, self.rng = maks_fitur, kedalaman, daun_min, rng
        self.fitur, self.ambang, self.kiri, self.kanan, self.nilai = [], [], [], [], []
        self.penting = None

    def _gini(self, hitung):
        n = hitung.sum()
        return 1.0 - ((hitung / n) ** 2).sum() if n else 0.0

    def _simpul(self, X, y, K, kdl):
        hitung = np.bincount(y, minlength=K).astype(float)
        idx = len(self.fitur)
        self.fitur.append(-1); self.ambang.append(0.0); self.kiri.append(-1); self.kanan.append(-1); self.nilai.append(hitung / hitung.sum())
        if kdl == self.kd or len(y) < 2 * self.dm or (hitung > 0).sum() == 1:
            return idx
        terbaik = (0.0, None, None)
        g0 = self._gini(hitung)
        n = len(y)
        for f in self.rng.choice(X.shape[1], size=self.mf, replace=False):
            urut = np.argsort(X[:, f], kind="stable")
            xs, ys = X[urut, f], y[urut]
            kiri = np.cumsum(np.eye(K)[ys], axis=0)             # jumlah kelas di sisi kiri untuk tiap titik potong
            kanan = hitung - kiri
            nk = np.arange(1, n + 1)[:, None]
            nr = n - nk
            gk = 1 - ((kiri / nk) ** 2).sum(1)
            gr = np.where(nr[:, 0] > 0, 1 - ((kanan / np.maximum(nr, 1)) ** 2).sum(1), 0)
            skor = g0 - (nk[:, 0] * gk + nr[:, 0] * gr) / n      # penurunan ketidakmurnian
            sah = (xs[:-1] < xs[1:]) & (nk[:-1, 0] >= self.dm) & (nr[:-1, 0] >= self.dm)
            if not sah.any():
                continue
            s = np.where(sah, skor[:-1], -1)
            k = int(s.argmax())
            if s[k] > terbaik[0]:
                terbaik = (s[k], f, (xs[k] + xs[k + 1]) / 2)
        if terbaik[1] is None:
            return idx
        gain, f, t = terbaik
        self.penting[f] += gain * n
        m = X[:, f] <= t
        self.fitur[idx], self.ambang[idx] = f, t
        self.kiri[idx] = self._simpul(X[m], y[m], K, kdl + 1)
        self.kanan[idx] = self._simpul(X[~m], y[~m], K, kdl + 1)
        return idx

    def latih(self, X, y, K):
        self.penting = np.zeros(X.shape[1])
        self._simpul(X, y, K, 0)
        self.fitur, self.ambang = np.array(self.fitur), np.array(self.ambang)
        self.kiri, self.kanan, self.nilai = np.array(self.kiri), np.array(self.kanan), np.array(self.nilai)

    def proba(self, X):
        simpul = np.zeros(len(X), dtype=int)
        while True:
            daun = self.fitur[simpul] < 0
            if daun.all():
                break
            f = self.fitur[simpul]
            ke_kiri = X[np.arange(len(X)), np.maximum(f, 0)] <= self.ambang[simpul]
            baru = np.where(ke_kiri, self.kiri[simpul], self.kanan[simpul])
            simpul = np.where(daun, simpul, baru)
        return self.nilai[simpul]


class RandomForestClassifier:
    def __init__(self, n_estimators=100, max_features="sqrt", max_depth=None, min_samples_leaf=1, random_state=0):
        self.n_estimators, self.max_features, self.max_depth = n_estimators, max_features, max_depth
        self.min_samples_leaf, self.random_state = min_samples_leaf, random_state

    def fit(self, X, y):
        X = np.asarray(X, dtype=float)
        self.classes_, yk = np.unique(y, return_inverse=True)
        K, (n, d) = len(self.classes_), X.shape
        mf = max(1, int(np.sqrt(d))) if self.max_features == "sqrt" else (d if self.max_features is None else int(self.max_features))
        rng = np.random.default_rng(self.random_state)
        self.pohon_, imp = [], np.zeros(d)
        for _ in range(self.n_estimators):
            b = rng.integers(0, n, n)                            # contoh bootstrap: ambil n baris dengan pengembalian
            p = _Pohon(mf, self.max_depth or 10 ** 6, self.min_samples_leaf, rng)
            p.latih(X[b], yk[b], K)
            self.pohon_.append(p)
            imp += p.penting
        self.feature_importances_ = imp / imp.sum() if imp.sum() else imp
        return self

    def predict_proba(self, X):
        X = np.asarray(X, dtype=float)
        return np.mean([p.proba(X) for p in self.pohon_], axis=0)

    def predict(self, X):
        return self.classes_[self.predict_proba(X).argmax(1)]
