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
12 changes: 5 additions & 7 deletions smarttree/_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,13 +43,12 @@ def build(self, tree: Tree) -> None:
else: # str
self.available_features.remove(value)

mask = self.y.apply(lambda x: True)
mask_np = mask.to_numpy()
mask = self.y.apply(lambda x: True).to_numpy()
root = tree.create_node(
mask=mask,
hierarchy=self.hierarchy,
distribution=self.criterion.distribution(mask_np),
impurity=self.criterion.impurity(mask_np),
distribution=self.criterion.distribution(mask),
impurity=self.criterion.impurity(mask),
label=self.y[mask].mode()[0],
available_features=self.available_features,
depth=0,
Expand All @@ -75,12 +74,11 @@ def build(self, tree: Tree) -> None:
else: # str
node.available_features.append(value)

child_mask_np = child_mask.to_numpy()
child_node = tree.create_node(
mask=child_mask,
hierarchy=node.hierarchy,
distribution=self.criterion.distribution(child_mask_np),
impurity=self.criterion.impurity(child_mask_np),
distribution=self.criterion.distribution(child_mask),
impurity=self.criterion.impurity(child_mask),
label=self.y[child_mask].mode()[0],
available_features=node.available_features,
depth=node.depth+1,
Expand Down
82 changes: 31 additions & 51 deletions smarttree/_column_splitter.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
from typing import NamedTuple

import numpy as np
import pandas as pd
from numpy.typing import NDArray

from ._cy_column_splitter import CyBaseColumnSplitter
Expand All @@ -22,7 +21,7 @@ class ColumnSplitResult(NamedTuple):

information_gain: float
feature_values: list[list]
child_masks: list[pd.Series]
child_masks: list[NDArray[np.bool_]]
child_na_index: int = -1

@classmethod
Expand Down Expand Up @@ -61,21 +60,12 @@ def __init__(
def split(self, *args, **kwargs) -> ColumnSplitResult:
raise NotImplementedError

def pre_information_gain(
self,
parent_mask: pd.Series,
child_masks: list[pd.Series],
) -> tuple[NDArray[np.bool_], list[NDArray[np.bool_]]]:
parent_mask_np = parent_mask.to_numpy()
child_masks_np = [child_mask.to_numpy() for child_mask in child_masks]
return parent_mask_np, child_masks_np

def foo(
self,
parent_mask: pd.Series,
parent_mask: NDArray[np.bool_],
split_feature: str,
child_masks: list[pd.Series],
) -> tuple[float, list[pd.Series], int]:
child_masks: list[NDArray[np.bool_]],
) -> tuple[float, list[NDArray[np.bool_]], int]:

if self.dataset.has_na[split_feature]:
mask_na = parent_mask & self.dataset.mask_na[split_feature]
Expand All @@ -87,33 +77,31 @@ def foo(
else:
assert False
else:
parent_mask_np, child_masks_np = self.pre_information_gain(parent_mask, child_masks)
information_gain = self.information_gain(parent_mask_np, child_masks_np)
information_gain = self.information_gain(parent_mask, child_masks)
return information_gain, child_masks, -1

def include_all_split(
self,
parent_mask: pd.Series,
mask_na: pd.Series,
child_masks: list[pd.Series],
) -> tuple[float, list[pd.Series], int]:
parent_mask: NDArray[np.bool_],
mask_na: NDArray[np.bool_],
child_masks: list[NDArray[np.bool_]],
) -> tuple[float, list[NDArray[np.bool_]], int]:

for i, child_mask in enumerate(child_masks):
child_masks[i] = child_mask | (parent_mask & mask_na)
if child_masks[i].sum() < self.min_samples_leaf:
return NO_INFORMATION_GAIN, [], -1

parent_mask_np, child_masks_np = self.pre_information_gain(parent_mask, child_masks)
information_gain = self.information_gain(parent_mask_np, child_masks_np, normalize=True)
information_gain = self.information_gain(parent_mask, child_masks, normalize=True)

return information_gain, child_masks, -1

def include_best_split(
self,
parent_mask: pd.Series,
mask_na: pd.Series,
child_masks: list[pd.Series],
) -> tuple[float, list[pd.Series], int]:
parent_mask: NDArray[np.bool_],
mask_na: NDArray[np.bool_],
child_masks: list[NDArray[np.bool_]],
) -> tuple[float, list[NDArray[np.bool_]], int]:

candidates = []
origin_child_masks = child_masks
Expand All @@ -130,8 +118,7 @@ def include_best_split(
best_child_masks = []
best_child_na_index = -1
for child_na_index, child_masks in enumerate(candidates):
parent_mask_np, child_masks_np = self.pre_information_gain(parent_mask, child_masks)
information_gain = self.information_gain(parent_mask_np, child_masks_np)
information_gain = self.information_gain(parent_mask, child_masks)
if best_information_gain < information_gain:
best_information_gain = information_gain
best_child_masks = child_masks
Expand Down Expand Up @@ -204,9 +191,9 @@ def __init__(

def split(self, node: TreeNode, split_feature: str) -> ColumnSplitResult:

numerical_column = self.dataset.X.loc[node.mask, split_feature]
points = numerical_column.dropna().to_numpy()
thresholds = self.__get_thresholds(points)
num_column = self.dataset[split_feature][node.mask]
points = np.sort(np.unique(num_column[~np.isnan(num_column)]))
thresholds = np.array([]) if len(points) <= 1 else self.__moving_average(points)

best_split_result = ColumnSplitResult.no_split()
for threshold in thresholds:
Expand All @@ -221,26 +208,19 @@ def split(self, node: TreeNode, split_feature: str) -> ColumnSplitResult:

return best_split_result

def __get_thresholds(self, array: NDArray) -> NDArray:

array = np.sort(np.unique(array))
thresholds = np.array([]) if len(array) <= 1 else self.__moving_average(array)

return thresholds

@staticmethod
def __moving_average(array: NDArray, window: int = 2) -> NDArray:
return np.convolve(array, np.ones(window), mode="valid") / window

def __num_split(
self,
parent_mask: pd.Series,
parent_mask: NDArray[np.bool_],
split_feature: str,
threshold: float,
) -> tuple[float, list[pd.Series], int]:
) -> tuple[float, list[NDArray[np.bool_]], int]:

mask_less = parent_mask & (self.dataset.X[split_feature] <= threshold)
mask_more = parent_mask & (self.dataset.X[split_feature] > threshold)
mask_less = parent_mask & (self.dataset[split_feature] <= threshold)
mask_more = parent_mask & (self.dataset[split_feature] > threshold)
child_masks = [mask_less, mask_more]

return self.foo(parent_mask, split_feature, child_masks)
Expand Down Expand Up @@ -276,8 +256,8 @@ def split(
leaf_counter: int,
) -> ColumnSplitResult:

category_column: pd.Series = self.dataset.X.loc[node.mask, split_feature] # type: ignore
categories = category_column.dropna().unique().tolist()
cat_column = self.dataset[split_feature][node.mask & ~self.dataset.mask_na[split_feature]]
categories = np.unique(cat_column).tolist()

if len(categories) <= 1:
return ColumnSplitResult.no_split()
Expand Down Expand Up @@ -306,14 +286,14 @@ def split(

def __cat_split(
self,
parent_mask: pd.Series,
parent_mask: NDArray[np.bool_],
split_feature: str,
feature_values: list[list],
) -> tuple[float, list[pd.Series], int]:
) -> tuple[float, list[NDArray[np.bool_]], int]:

child_masks = []
for partition in feature_values:
partition_mask = self.dataset.X[split_feature].isin(partition)
partition_mask = np.isin(self.dataset[split_feature], partition)
child_mask = parent_mask & partition_mask
child_masks.append(child_mask)

Expand Down Expand Up @@ -376,15 +356,15 @@ def split(self, node: TreeNode, split_feature: str) -> ColumnSplitResult:

def __rank_split(
self,
parent_mask: pd.Series,
parent_mask: NDArray[np.bool_],
split_feature: str,
feature_values: tuple[list, list],
) -> tuple[float, list[pd.Series], int]:
) -> tuple[float, list[NDArray[np.bool_]], int]:

feature_values_left, feature_values_right = feature_values

mask_left = parent_mask & self.dataset.X[split_feature].isin(feature_values_left)
mask_right = parent_mask & self.dataset.X[split_feature].isin(feature_values_right)
mask_left = parent_mask & np.isin(self.dataset[split_feature], feature_values_left)
mask_right = parent_mask & np.isin(self.dataset[split_feature], feature_values_right)
child_masks = [mask_left, mask_right]

return self.foo(parent_mask, split_feature, child_masks)
Expand Down
20 changes: 11 additions & 9 deletions smarttree/_dataset.py
Original file line number Diff line number Diff line change
@@ -1,22 +1,24 @@
import numpy as np
import pandas as pd
from numpy.typing import NDArray


class Dataset:

def __init__(self, X: pd.DataFrame, y: pd.Series) -> None:
self.X = X

self.classes = np.sort(y.unique())
self.y = np.searchsorted(self.classes, y.to_numpy()).astype(np.int64)
self.has_na: dict[str, bool] = dict()
self.mask_na: dict[str, pd.Series] = dict()
for column in self.X.columns:
mask_na = self.X[column].isna()
self.mask_na: dict[str, NDArray] = dict()
self.columns: dict[str, NDArray] = dict()
for column in X.columns:
mask_na = X[column].isna()
has_na = mask_na.any()
self.has_na[column] = has_na
if has_na:
self.mask_na[column] = mask_na
self.mask_na[column] = mask_na.to_numpy()
self.columns[column] = X[column].to_numpy()
self.n_samples, self.n_columns = X.shape

@property
def size(self) -> int:
return self.X.shape[0]
def __getitem__(self, item) -> NDArray:
return self.columns[item]
4 changes: 3 additions & 1 deletion smarttree/_node_splitter.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
from typing import NamedTuple

import numpy as np
import pandas as pd
from numpy.typing import NDArray

from ._column_splitter import CatColumnSplitter, NumColumnSplitter, RankColumnSplitter
from ._dataset import Dataset
Expand All @@ -14,7 +16,7 @@ class NodeSplitResult(NamedTuple):
split_type: str
split_feature: str
feature_values: list[list[str]]
child_masks: list[pd.Series]
child_masks: list[NDArray[np.bool_]]
child_na_index: int = -1

@classmethod
Expand Down
9 changes: 4 additions & 5 deletions smarttree/_tree.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
from typing import Self

import numpy as np
import pandas as pd
from numpy.typing import NDArray


Expand All @@ -13,7 +12,7 @@ class TreeNode:
number: int
num_samples: int
depth: int = field(repr=False)
mask: pd.Series = field(repr=False)
mask: NDArray[np.bool_] = field(repr=False)
hierarchy: dict[str, str | list[str]] = field(repr=False)
available_features: list[str] = field(repr=False)

Expand Down Expand Up @@ -54,7 +53,7 @@ def dummy(cls):
number=-1,
num_samples=-1,
depth=-1,
mask=pd.Series(),
mask=np.array([]),
hierarchy={},
available_features=[],
distribution=np.array([]),
Expand All @@ -73,8 +72,8 @@ def __init__(self) -> None:

def create_node(
self,
mask: pd.Series,
distribution: NDArray[np.integer],
mask: NDArray[np.bool_],
distribution: NDArray[np.int64],
impurity: float,
label: str,
hierarchy: dict[str, str | list[str]],
Expand Down
2 changes: 1 addition & 1 deletion tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,7 +175,7 @@ def root_node(X, y):
number=0,
num_samples=y.apply(lambda x: True).sum(),
depth=0,
mask=y.apply(lambda x: True),
mask=y.apply(lambda x: True).to_numpy(),
hierarchy=dict(),
available_features=X.columns.to_list(),
impurity=0.67,
Expand Down
Loading