#!/usr/bin/python3
"""Compare exported inference with sklearn's native tree evaluator.

Rebuild Tree objects from the original saved nodes using the installed
sklearn's node dtype. These trees independently evaluate the same parameters;
the comparison covers both models, float32 split boundaries and random inputs.
"""

import importlib.util
import os
from pathlib import Path

os.environ['OMP_NUM_THREADS'] = '1'
os.environ['OPENBLAS_NUM_THREADS'] = '1'

import numpy as np
from sklearn.tree import _tree
from scipy.special import expit

from pepsickle.model_functions import initialize_digestion_gb_model

root = Path(__file__).resolve().parents[2]
spec = importlib.util.spec_from_file_location('converter', root / 'debian/convert-models.py')
converter = importlib.util.module_from_spec(spec)
spec.loader.exec_module(converter)

for human_only in (False, True):
    name = 'in-vitro_human' if human_only else 'in-vitro_mammal'
    path = root / 'pepsickle/in-vitro-models' / name / 'model.joblib'
    with path.open('rb') as stream:
        original = converter.ModelUnpickler(
            str(path), stream, ensure_native_byte_order=True
        ).load()
    native_trees = []
    features = [np.random.default_rng(2026).uniform(
        0, 1, (256, original.n_features_in_)).astype(np.float32)]
    for index, estimator in enumerate(original.estimators_[:, 0]):
        saved = estimator.tree_
        nodes = np.zeros(saved.node_count, dtype=_tree.NODE_DTYPE)
        for field in saved.nodes.dtype.names:
            nodes[field] = saved.nodes[field]
        tree = _tree.Tree(original.n_features_in_, np.array([1], dtype=np.intp), 1)
        tree.__setstate__({
            'max_depth': saved.max_depth, 'node_count': saved.node_count,
            'nodes': nodes, 'values': saved.values,
        })
        native_trees.append(tree)
        if index < 10:
            for node in saved.nodes:
                if node['left_child'] == -1:
                    continue
                threshold = np.float32(node['threshold'])
                sample = np.full((3, original.n_features_in_), 0.5, dtype=np.float32)
                sample[:, node['feature']] = [
                    np.nextafter(threshold, np.float32(-np.inf)), threshold,
                    np.nextafter(threshold, np.float32(np.inf)),
                ]
                features.append(sample)
    features = np.concatenate(features)
    prior = original.init_.class_prior_[1]
    scores = np.full(features.shape[0], np.log(prior / (1 - prior)))
    for tree in native_trees:
        scores += original.learning_rate * tree.predict(features)[:, 0]
    actual = initialize_digestion_gb_model(human_only).predict_proba(features)
    np.testing.assert_allclose(actual[:, 1], expit(scores), rtol=1e-13, atol=1e-15)
    np.testing.assert_allclose(actual.sum(axis=1), 1)
    print(f'{name}: {features.shape[0]} predictions match native sklearn trees')
