diff --git a/poetry.lock b/poetry.lock index 64d2f77..13ad053 100644 --- a/poetry.lock +++ b/poetry.lock @@ -2867,6 +2867,20 @@ files = [ {file = "six-1.17.0.tar.gz", hash = "sha256:ff70335d468e7eb6ec65b95b99d3a2836546063f63acc5171de367e834932a81"}, ] +[[package]] +name = "snakeviz" +version = "2.2.2" +description = "A web-based viewer for Python profiler output" +optional = false +python-versions = ">=3.9" +files = [ + {file = "snakeviz-2.2.2-py3-none-any.whl", hash = "sha256:77e7b9c82f6152edc330040319b97612351cd9b48c706434c535c2df31d10ac5"}, + {file = "snakeviz-2.2.2.tar.gz", hash = "sha256:08028c6f8e34a032ff14757a38424770abb8662fb2818985aeea0d9bc13a7d83"}, +] + +[package.dependencies] +tornado = ">=2.0" + [[package]] name = "sniffio" version = "1.3.1" @@ -3121,4 +3135,4 @@ test = ["websockets"] [metadata] lock-version = "2.0" python-versions = "^3.11" -content-hash = "512ec2bb7bbbde49bb32fd398d9e850100eb65765dda8315f526a2844a32fce9" +content-hash = "74a5d61b448b14dda65d214b87971bd4fdbf897042c788dac6b6e1bcb244b56d" diff --git a/pyproject.toml b/pyproject.toml index 5a10658..ab9cc82 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,6 +29,7 @@ graphviz = "^0.20.1" notebook = "^7.0.6" jupyterlab-execute-time = "^3.1.0" cython = "^3.1.4" +snakeviz = "^2.2.2" [tool.poetry.group.test.dependencies] pytest = "^7.4.3" diff --git a/smarttree/_builder.py b/smarttree/_builder.py index 5e099d2..ca37236 100644 --- a/smarttree/_builder.py +++ b/smarttree/_builder.py @@ -2,7 +2,6 @@ import numpy as np import pandas as pd -from numpy.typing import NDArray from ._criterion import ClassificationCriterion, Entropy, Gini from ._dataset import Dataset @@ -23,15 +22,18 @@ def __init__( ) -> None: self.X = X - self.available_features = X.columns.to_list() self.y = y - self.criterion = criterion + self.dataset = Dataset(X, y) + self.available_features = X.columns.to_list() self.splitter = splitter self.max_leaf_nodes = max_leaf_nodes self.hierarchy = hierarchy - if self.criterion in ("gini", "entropy", "log_loss"): - self.class_names = np.sort(self.y.unique()) + self.criterion: ClassificationCriterion + if criterion == "gini": + self.criterion = Gini(self.dataset) + else: # "entropy" | "log_loss" + self.criterion = Entropy(self.dataset) def build(self, tree: Tree) -> None: @@ -43,11 +45,12 @@ def build(self, tree: Tree) -> None: self.available_features.remove(value) mask = self.y.apply(lambda x: True) + mask_np = mask.to_numpy(dtype=np.int8) root = tree.create_node( mask=mask, hierarchy=self.hierarchy, - distribution=self.distribution(mask), - impurity=self.impurity(mask), + distribution=self.criterion.distribution(mask_np), + impurity=self.criterion.impurity(mask_np), label=self.y[mask].mode()[0], available_features=self.available_features, depth=0, @@ -73,11 +76,12 @@ def build(self, tree: Tree) -> None: else: # str node.available_features.append(value) + child_mask_np = child_mask.to_numpy(dtype=np.int8) child_node = tree.create_node( mask=child_mask, hierarchy=node.hierarchy, - distribution=self.distribution(child_mask), - impurity=self.impurity(child_mask), + distribution=self.criterion.distribution(child_mask_np), + impurity=self.criterion.impurity(child_mask_np), label=self.y[child_mask].mode()[0], available_features=node.available_features, depth=node.depth+1, @@ -94,24 +98,3 @@ def build(self, tree: Tree) -> None: node.is_leaf = False tree.leaf_counter -= 1 - - def distribution(self, mask: pd.Series) -> NDArray[np.integer]: - - mask_arr = mask.to_numpy() - y_arr = self.y.to_numpy() - - result = np.zeros(len(self.class_names), dtype=np.int32) - for i, class_name in enumerate(self.class_names): - result[i] = np.sum(mask_arr & (y_arr == class_name)) - - return result - - def impurity(self, mask: pd.Series) -> float: - - criterion: ClassificationCriterion - if self.criterion == "gini": - criterion = Gini(Dataset(self.X, self.y)) - else: # "entropy" | "log_loss" - criterion = Entropy(Dataset(self.X, self.y)) - - return criterion.impurity(mask.to_numpy(dtype=np.int8)) diff --git a/smarttree/_criterion.pxd b/smarttree/_criterion.pxd index 4d6f2ba..138e718 100644 --- a/smarttree/_criterion.pxd +++ b/smarttree/_criterion.pxd @@ -7,6 +7,8 @@ cdef class ClassificationCriterion: cdef Py_ssize_t n_classes cdef Py_ssize_t n_samples + cpdef long[:] distribution(self, int8_t[:] mask) + cdef class Gini(ClassificationCriterion): cpdef double impurity(self, int8_t[:] mask) diff --git a/smarttree/_criterion.pyi b/smarttree/_criterion.pyi index e9a9d12..326571d 100644 --- a/smarttree/_criterion.pyi +++ b/smarttree/_criterion.pyi @@ -14,6 +14,9 @@ class ClassificationCriterion(ABC): def impurity(self, mask: NDArray[np.int8]) -> float: raise NotImplementedError + def distribution(self, mask: NDArray[np.int8]) -> NDArray[np.int32]: + ... + class Gini(ClassificationCriterion): diff --git a/smarttree/_criterion.pyx b/smarttree/_criterion.pyx index 79014ac..6d99a58 100644 --- a/smarttree/_criterion.pyx +++ b/smarttree/_criterion.pyx @@ -14,6 +14,20 @@ cdef class ClassificationCriterion: self.n_classes = len(dataset.classes) self.n_samples = len(dataset.y) + @cython.boundscheck(False) + @cython.wraparound(False) + cpdef long[:] distribution(self, int8_t[:] mask): + + cdef Py_ssize_t i + cdef long[:] result + + result = np.zeros(self.n_classes, dtype=np.int32) + for i in range(self.n_samples): + if mask[i]: + result[self.y[i]] += 1 + + return result + cdef class Gini(ClassificationCriterion): @@ -23,21 +37,19 @@ cdef class Gini(ClassificationCriterion): cpdef double impurity(self, int8_t[:] mask): cdef Py_ssize_t i - cdef long[:] counts + cdef long[:] distribution cdef long N cdef double p_i, gini - counts = np.zeros(self.n_classes, dtype=np.int32) + distribution = self.distribution(mask) N = 0 - for i in range(self.n_samples): - if mask[i]: - N += 1 - counts[self.y[i]] += 1 + for i in range(self.n_classes): + N += distribution[i] gini = 1.0 for i in range(self.n_classes): - if counts[i] > 0: - p_i = counts[i] / N + if distribution[i] > 0: + p_i = distribution[i] / N gini -= p_i * p_i return gini @@ -51,21 +63,19 @@ cdef class Entropy(ClassificationCriterion): cpdef double impurity(self, int8_t[:] mask): cdef Py_ssize_t i - cdef long[:] counts + cdef long[:] distribution cdef long N - cdef double p_i, entropy + cdef double p_i, gini - counts = np.zeros(self.n_classes, dtype=np.int32) + distribution = self.distribution(mask) N = 0 - for i in range(self.n_samples): - if mask[i]: - N += 1 - counts[self.y[i]] += 1 + for i in range(self.n_classes): + N += distribution[i] entropy = 0.0 for i in range(self.n_classes): - if counts[i] > 0: - p_i = counts[i] / N + if distribution[i] > 0: + p_i = distribution[i] / N entropy -= p_i * log2(p_i) return entropy diff --git a/tests/column_splitter/test__base_column_splitter.py b/tests/column_splitter/test__base_column_splitter.py index 7f7ae7e..0e4adf0 100644 --- a/tests/column_splitter/test__base_column_splitter.py +++ b/tests/column_splitter/test__base_column_splitter.py @@ -2,11 +2,10 @@ import pytest from smarttree._column_splitter import BaseColumnSplitter -from smarttree._dataset import Dataset @pytest.fixture(scope="module") -def concrete_column_splitter(X, y, feature_na_mode) -> BaseColumnSplitter: +def concrete_column_splitter(dataset, feature_na_mode) -> BaseColumnSplitter: class ConcreteColumnSplitter(BaseColumnSplitter): def split( self, @@ -16,7 +15,7 @@ def split( ... return ConcreteColumnSplitter( - dataset=Dataset(X, y), + dataset=dataset, criterion="gini", min_samples_split=2, min_samples_leaf=1, diff --git a/tests/column_splitter/test__cat_column_splitter.py b/tests/column_splitter/test__cat_column_splitter.py index 6bb5537..da15fa6 100644 --- a/tests/column_splitter/test__cat_column_splitter.py +++ b/tests/column_splitter/test__cat_column_splitter.py @@ -1,11 +1,10 @@ from smarttree._column_splitter import CatColumnSplitter -from smarttree._dataset import Dataset -def test__split(X, y, root_node, feature_na_mode): +def test__split(dataset, root_node, feature_na_mode): categorical_column_splitter = CatColumnSplitter( - dataset=Dataset(X, y), + dataset=dataset, criterion="gini", min_samples_split=2, min_samples_leaf=1, diff --git a/tests/column_splitter/test__num_column_splitter.py b/tests/column_splitter/test__num_column_splitter.py index 275c230..873b368 100644 --- a/tests/column_splitter/test__num_column_splitter.py +++ b/tests/column_splitter/test__num_column_splitter.py @@ -1,11 +1,10 @@ from smarttree._column_splitter import NumColumnSplitter -from smarttree._dataset import Dataset -def test__split(X, y, root_node, feature_na_mode): +def test__split(dataset, root_node, feature_na_mode): numerical_column_splitter = NumColumnSplitter( - dataset=Dataset(X, y), + dataset=dataset, criterion="gini", min_samples_split=2, min_samples_leaf=1, diff --git a/tests/column_splitter/test__rank_column_splitter.py b/tests/column_splitter/test__rank_column_splitter.py index d95c648..b3b3632 100644 --- a/tests/column_splitter/test__rank_column_splitter.py +++ b/tests/column_splitter/test__rank_column_splitter.py @@ -1,11 +1,10 @@ from smarttree._column_splitter import RankColumnSplitter -from smarttree._dataset import Dataset -def test__split(X, y, root_node, feature_na_mode): +def test__split(dataset, root_node, feature_na_mode): rank_column_splitter = RankColumnSplitter( - dataset=Dataset(X, y), + dataset=dataset, criterion="gini", min_samples_split=2, min_samples_leaf=1,