From fa877c6d9f92c70afe8af46c11a8c405c36aa064 Mon Sep 17 00:00:00 2001 From: Mikhail Martin Date: Sat, 27 Sep 2025 23:48:50 +0300 Subject: [PATCH 1/5] refactoring tests --- tests/column_splitter/test__base_column_splitter.py | 5 ++--- tests/column_splitter/test__cat_column_splitter.py | 5 ++--- tests/column_splitter/test__num_column_splitter.py | 5 ++--- tests/column_splitter/test__rank_column_splitter.py | 5 ++--- 4 files changed, 8 insertions(+), 12 deletions(-) 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, From 0bbd4a7debb60af37337cfc602d6387f78cee2db Mon Sep 17 00:00:00 2001 From: Mikhail Martin Date: Sat, 27 Sep 2025 23:56:50 +0300 Subject: [PATCH 2/5] up dataset --- smarttree/_builder.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/smarttree/_builder.py b/smarttree/_builder.py index 5e099d2..9401768 100644 --- a/smarttree/_builder.py +++ b/smarttree/_builder.py @@ -23,8 +23,9 @@ def __init__( ) -> None: self.X = X - self.available_features = X.columns.to_list() self.y = y + self.dataset = Dataset(X, y) + self.available_features = X.columns.to_list() self.criterion = criterion self.splitter = splitter self.max_leaf_nodes = max_leaf_nodes @@ -110,8 +111,8 @@ def impurity(self, mask: pd.Series) -> float: criterion: ClassificationCriterion if self.criterion == "gini": - criterion = Gini(Dataset(self.X, self.y)) + criterion = Gini(self.dataset) else: # "entropy" | "log_loss" - criterion = Entropy(Dataset(self.X, self.y)) + criterion = Entropy(self.dataset) return criterion.impurity(mask.to_numpy(dtype=np.int8)) From 53f4c51bd9483590522d8351cc201b2adefedc0a Mon Sep 17 00:00:00 2001 From: Mikhail Martin Date: Sun, 28 Sep 2025 00:16:59 +0300 Subject: [PATCH 3/5] refactoring --- smarttree/_builder.py | 26 +++++++++----------------- 1 file changed, 9 insertions(+), 17 deletions(-) diff --git a/smarttree/_builder.py b/smarttree/_builder.py index 9401768..a011087 100644 --- a/smarttree/_builder.py +++ b/smarttree/_builder.py @@ -26,13 +26,15 @@ def __init__( self.y = y self.dataset = Dataset(X, y) self.available_features = X.columns.to_list() - self.criterion = criterion 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: @@ -48,7 +50,7 @@ def build(self, tree: Tree) -> None: mask=mask, hierarchy=self.hierarchy, distribution=self.distribution(mask), - impurity=self.impurity(mask), + impurity=self.criterion.impurity(mask.to_numpy(dtype=np.int8)), label=self.y[mask].mode()[0], available_features=self.available_features, depth=0, @@ -78,7 +80,7 @@ def build(self, tree: Tree) -> None: mask=child_mask, hierarchy=node.hierarchy, distribution=self.distribution(child_mask), - impurity=self.impurity(child_mask), + impurity=self.criterion.impurity(mask.to_numpy(dtype=np.int8)), label=self.y[child_mask].mode()[0], available_features=node.available_features, depth=node.depth+1, @@ -101,18 +103,8 @@ 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 = np.zeros(len(self.dataset.classes), dtype=np.int32) + for i, class_name in enumerate(self.dataset.classes): 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(self.dataset) - else: # "entropy" | "log_loss" - criterion = Entropy(self.dataset) - - return criterion.impurity(mask.to_numpy(dtype=np.int8)) From 7bc486d6537edf349e7c9bf7b77e82fbcf4d9e90 Mon Sep 17 00:00:00 2001 From: Mikhail Martin Date: Sun, 28 Sep 2025 00:52:43 +0300 Subject: [PATCH 4/5] extract ClassificationCriterion.distribution() --- smarttree/_builder.py | 22 ++++++-------------- smarttree/_criterion.pxd | 2 ++ smarttree/_criterion.pyi | 3 +++ smarttree/_criterion.pyx | 44 ++++++++++++++++++++++++---------------- 4 files changed, 38 insertions(+), 33 deletions(-) diff --git a/smarttree/_builder.py b/smarttree/_builder.py index a011087..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 @@ -46,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.criterion.impurity(mask.to_numpy(dtype=np.int8)), + 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, @@ -76,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.criterion.impurity(mask.to_numpy(dtype=np.int8)), + 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, @@ -97,14 +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.dataset.classes), dtype=np.int32) - for i, class_name in enumerate(self.dataset.classes): - result[i] = np.sum(mask_arr & (y_arr == class_name)) - - return result 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 From 89b1330bfe682f16c6e8b46459997d51c1c92d41 Mon Sep 17 00:00:00 2001 From: Mikhail Martin Date: Sun, 28 Sep 2025 02:05:52 +0300 Subject: [PATCH 5/5] added snakeviz --- poetry.lock | 16 +++++++++++++++- pyproject.toml | 1 + 2 files changed, 16 insertions(+), 1 deletion(-) 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"