Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 15 additions & 1 deletion poetry.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
43 changes: 13 additions & 30 deletions smarttree/_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:

Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -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))
2 changes: 2 additions & 0 deletions smarttree/_criterion.pxd
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
3 changes: 3 additions & 0 deletions smarttree/_criterion.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -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):

Expand Down
44 changes: 27 additions & 17 deletions smarttree/_criterion.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -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):

Expand All @@ -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 = <double>counts[i] / <double>N
if distribution[i] > 0:
p_i = <double>distribution[i] / <double>N
gini -= p_i * p_i

return gini
Expand All @@ -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 = <double>counts[i] / <double>N
if distribution[i] > 0:
p_i = <double>distribution[i] / <double>N
entropy -= p_i * log2(p_i)

return entropy
5 changes: 2 additions & 3 deletions tests/column_splitter/test__base_column_splitter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -16,7 +15,7 @@ def split(
...

return ConcreteColumnSplitter(
dataset=Dataset(X, y),
dataset=dataset,
criterion="gini",
min_samples_split=2,
min_samples_leaf=1,
Expand Down
5 changes: 2 additions & 3 deletions tests/column_splitter/test__cat_column_splitter.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down
5 changes: 2 additions & 3 deletions tests/column_splitter/test__num_column_splitter.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down
5 changes: 2 additions & 3 deletions tests/column_splitter/test__rank_column_splitter.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down
Loading