# -*- coding: utf-8 -*-
# SKRIP 2.3: Isi lubang akibat awan dengan interpolasi linear sepanjang waktu, lalu uji kebenarannya
# Penulis: Badar Mubarok Yogaswara
# Cara uji: sembunyikan 20% pengamatan bersih secara acak, isi dengan interpolasi, bandingkan dengan nilai aslinya.
import os
import numpy as np
import m3_umum as U

tgl, refl, scl, gt, prj = U.baca_deret()
ndvi = U.nilai_ndvi(refl)
baik = U.bersih(scl)
hari = np.array([(t - tgl[0]).days for t in tgl], dtype="float64")     # sumbu waktu dalam hari (selang tidak sama)


def isi(ndvi, baik):
    T, H, W = ndvi.shape
    hasil = np.empty_like(ndvi)
    for i in range(H):
        for j in range(W):
            ok = baik[:, i, j]
            if ok.sum() == 0:
                hasil[:, i, j] = np.nan
            else:
                # np.interp: di luar rentang, nilai ujung dipakai (tidak ekstrapolasi)
                hasil[:, i, j] = np.interp(hari, hari[ok], ndvi[ok, i, j])
    return hasil


lengkap = isi(ndvi, baik)
U.tulis_tif(os.path.join(U.HASIL, "NDVI_isi.tif"), lengkap.astype("float32"), gt, prj,
            deskripsi=[t.isoformat() for t in tgl])
print("Piksel-tanggal awan sebelum diisi: %.1f%%, sesudah: %.1f%%" % (100 * (~baik).mean(), 100 * np.isnan(lengkap).mean()))

# uji sembunyi-isi
rng = np.random.default_rng(1)
sembunyi = baik & (rng.random(baik.shape) < 0.20)
sembunyi[0], sembunyi[-1] = False, False                   # tanggal ujung tidak diuji (tanpa tetangga di satu sisi)
uji = isi(ndvi, baik & ~sembunyi)
galat = (uji - ndvi)[sembunyi]
print("Uji sembunyi-isi: %d pengamatan; RMSE %.3f; galat mutlak median %.3f" % (sembunyi.sum(), np.sqrt(np.nanmean(galat ** 2)), np.nanmedian(np.abs(galat))))
# galat per kelas acuan
dinamika, _, _ = 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"}
for k in range(1, 9):
    m = sembunyi & (dinamika == k)[None]
    e = (uji - ndvi)[m]
    print("  %-14s n=%4d  RMSE %.3f" % (nama[k], m.sum(), np.sqrt(np.nanmean(e ** 2))))
