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
6 changes: 4 additions & 2 deletions dataframe-learn/src-internal/DataFrame/DecisionTree/Cart.hs
Original file line number Diff line number Diff line change
Expand Up @@ -344,15 +344,17 @@ nullableLeq c
| otherwise = Nothing

oneHotFeatures ::
forall b. (Columnable b) => T.Text -> Maybe Bitmap -> V.Vector b -> [CartFeature]
forall b.
(Columnable b) => T.Text -> Maybe Bitmap -> V.Vector b -> [CartFeature]
oneHotFeatures c bm v = case testEquality (typeRep @b) (typeRep @T.Text) of
Just Refl -> [oneHot c nulls v cat | cat <- Set.toList (Set.fromList present)]
Nothing -> []
where
nulls = fmap (nullFlags (V.length v)) bm
present = [x | (i, x) <- zip [0 ..] (V.toList v), maybe True (not . (VU.! i)) nulls]

oneHot :: T.Text -> Maybe (VU.Vector Bool) -> V.Vector T.Text -> T.Text -> CartFeature
oneHot ::
T.Text -> Maybe (VU.Vector Bool) -> V.Vector T.Text -> T.Text -> CartFeature
oneHot c Nothing v cat =
CartFeature
(VU.generate (V.length v) (\i -> if v V.! i == cat then 1 else 0))
Expand Down
26 changes: 19 additions & 7 deletions dataframe-learn/src-internal/DataFrame/DecisionTree/Histogram.hs
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,8 @@
runs = runLengths sorted
distinct = map fst runs
edges
| length runs <= maxBins = VU.fromList (zipWith splitMidpoint distinct (drop 1 distinct))
| length runs <= maxBins =
VU.fromList (zipWith splitMidpoint distinct (drop 1 distinct))
| otherwise = VU.fromList (cutEdges maxBins (VU.length sorted) runs)
nb = VU.length edges + 1
-- The last bin's threshold is the largest value, so "every non-null row
Expand Down Expand Up @@ -112,7 +113,8 @@
-}
type Hist = VU.Vector Double

buildHist :: VU.Vector Double -> VU.Vector Double -> VU.Vector Int -> Binned -> Hist
buildHist ::
VU.Vector Double -> VU.Vector Double -> VU.Vector Int -> Binned -> Hist
buildHist w y idxs bn = runST $ do
m <- VUM.replicate (4 * (bnNBins bn + 1)) 0
VU.forM_ idxs $ \i -> do
Expand Down Expand Up @@ -141,11 +143,11 @@
fitBinnedTree lim binned y mw = (toTree root, inSample)
where
n = VU.length y
w = maybe (VU.replicate n 1) id mw
w = Data.Maybe.fromMaybe (VU.replicate n 1) mw
allIdx = VU.enumFromN 0 n
root = node 0 allIdx (hists allIdx)
inSample = VU.update (VU.replicate n 0) (VU.concat (leafRows root))
leafRows (NLeaf v idxs) = [VU.map (\i -> (i, v)) idxs]

Check warning on line 150 in dataframe-learn/src-internal/DataFrame/DecisionTree/Histogram.hs

View workflow job for this annotation

GitHub Actions / Lint (hlint)

Suggestion in fitBinnedTree in module DataFrame.DecisionTree.Histogram: Use tuple-section ▫︎ Found: "\\ i -> (i, v)" ▫︎ Perhaps: "(, v)" ▫︎ Note: may require `{-# LANGUAGE TupleSections #-}` adding to the top of the file
leafRows (NBranch _ l r) = leafRows l ++ leafRows r

-- Feature-parallel only where a node is big enough to pay for the sparks.
Expand All @@ -155,7 +157,8 @@
strategy = if VU.length idxs >= 20000 then parList rseq else evalList rseq

node depth idxs hs
| depth >= tlMaxDepth lim || VU.length idxs < tlMinSamplesSplit lim || V.null hs = leaf
| depth >= tlMaxDepth lim || VU.length idxs < tlMinSamplesSplit lim || V.null hs =
leaf
| otherwise = maybe leaf split (bestSplit lim binned hs)
where
Totals tw tsy _ _ = histTotals (V.head hs)
Expand Down Expand Up @@ -199,7 +202,8 @@
scored with the node's null rows on the right and then on the left; ties keep
the earliest candidate, as in the exact sweep.
-}
bestSplit :: TreeLimits -> V.Vector Binned -> V.Vector Hist -> Maybe (Int, Int, Bool)
bestSplit ::
TreeLimits -> V.Vector Binned -> V.Vector Hist -> Maybe (Int, Int, Bool)
bestSplit lim binned hs
| null candidates = Nothing
| red > 0 && red >= tlMinImpurityDecrease lim = Just sp
Expand All @@ -225,7 +229,14 @@
| otherwise = Nothing
where
wr = totW - wl
scan :: Int -> Double -> Double -> Double -> Int -> Maybe (Int, Bool, Double) -> Maybe (Int, Bool, Double)
scan ::
Int ->
Double ->
Double ->
Double ->
Int ->
Maybe (Int, Bool, Double) ->
Maybe (Int, Bool, Double)
scan !b !wl !syl !syl2 !cl best
| b > lastB = best
| otherwise = scan (b + 1) wl' syl' syl2' cl' best'
Expand All @@ -236,7 +247,8 @@
cl' = cl + round (h VU.! (4 * b + 3))
right = (,) False <$> score cl' wl' syl' syl2'
left
| hasNulls = (,) True <$> score (cl' + nC) (wl' + nW) (syl' + nSY) (syl2' + nSY2)
| hasNulls =
(,) True <$> score (cl' + nC) (wl' + nW) (syl' + nSY) (syl2' + nSY2)
| otherwise = Nothing
consider bst (dir, r)
| maybe True (\(_, _, rb) -> r > rb) bst = Just (b, dir, r)
Expand Down
4 changes: 2 additions & 2 deletions dataframe-learn/src/DataFrame/Boosting/AdaBoost.hs
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,10 @@ import DataFrame.Errors (DataFrameException (..))

import DataFrame.DecisionTree.Cart (
CartFeature (..),
cfPred,
splitMidpoint,
cartFeatures,
cfPred,
sortIndicesByValue,
splitMidpoint,
)
import DataFrame.DecisionTree.Fit (treeToExpr)
import DataFrame.DecisionTree.Types (Tree (..))
Expand Down
13 changes: 9 additions & 4 deletions dataframe-learn/src/DataFrame/Boosting/GBM.hs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}

{-# LANGUAGE TypeFamilies #-}

{- | Gradient boosting of regression trees (Friedman). Trees are fitted to the
Expand Down Expand Up @@ -34,7 +34,11 @@ import DataFrame.Errors (DataFrameException (..))

import DataFrame.DecisionTree.Cart (cartFeatures)
import DataFrame.DecisionTree.Fit (treeToExpr)
import DataFrame.DecisionTree.Histogram (TreeLimits (..), binFeatures, fitBinnedTree)
import DataFrame.DecisionTree.Histogram (
TreeLimits (..),
binFeatures,
fitBinnedTree,
)
import DataFrame.DecisionTree.Types (Tree)
import DataFrame.Expression.Operators ((.*.), (.+.), (.>.))
import DataFrame.Featurize.Internal (targetDoubles)
Expand All @@ -53,8 +57,9 @@ data GBConfig = GBConfig
, gbLearningRate :: !Double
, gbMaxDepth :: !Int
, gbMaxBins :: !Int
-- ^ Bins per feature for split finding. Features with at most this many
-- distinct values split exactly as on the raw values.
{- ^ Bins per feature for split finding. Features with at most this many
distinct values split exactly as on the raw values.
-}
, gbSeed :: !Int
}
deriving (Eq, Show)
Expand Down
6 changes: 5 additions & 1 deletion dataframe-learn/src/DataFrame/DecisionTree/Model.hs
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,11 @@ import qualified Data.Vector as V

import DataFrame.DecisionTree.Cart (cartFeatures)
import DataFrame.DecisionTree.Fit (fitDecisionTree, treeToExpr)
import DataFrame.DecisionTree.Regression (RegFit (..), RegTreeConfig, fitRegTree)
import DataFrame.DecisionTree.Regression (
RegFit (..),
RegTreeConfig,
fitRegTree,
)
import DataFrame.DecisionTree.Types (Tree (..), TreeConfig)
import DataFrame.Featurize.Internal (targetDoubles)
import DataFrame.Internal.Column (Columnable)
Expand Down
7 changes: 4 additions & 3 deletions dataframe-learn/src/DataFrame/DecisionTree/Regression.hs
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@
where
root = buildNode 0 (VU.enumFromN 0 n) featSorted
inSample = VU.update (VU.replicate n 0) (VU.concat (leafRows root))
leafRows (FLeaf v idxs) = [VU.map (\i -> (i, v)) idxs]

Check warning on line 66 in dataframe-learn/src/DataFrame/DecisionTree/Regression.hs

View workflow job for this annotation

GitHub Actions / Lint (hlint)

Suggestion in fitRegTree in module DataFrame.DecisionTree.Regression: Use tuple-section ▫︎ Found: "\\ i -> (i, v)" ▫︎ Perhaps: "(, v)" ▫︎ Note: may require `{-# LANGUAGE TupleSections #-}` adding to the top of the file
leafRows (FBranch _ l r) = leafRows l ++ leafRows r
n = VU.length y
weightAt i = maybe 1 (VU.! i) mw
Expand All @@ -83,8 +83,8 @@
splitNode depth idxs sortedByFeat (fj, thr, side)
| VU.null lefts || VU.null rights = FLeaf (weightedMean idxs) idxs
| otherwise =
forceTree l
`par` (forceTree r `pseq` FBranch (cfSplit feat thr side) l r)
forceTree l `par`
(forceTree r `pseq` FBranch (cfSplit feat thr side) l r)
where
feat = feats V.! fj
vals = cfValues feat
Expand Down Expand Up @@ -134,7 +134,8 @@
lastK = if hasNulls then m - 1 else m - 2
score nl wl syl syl2 =
let wr = totW - wl
ok = nl >= rtMinLeafSize cfg && nNode - nl >= rtMinLeafSize cfg && wl > 0 && wr > 0
ok =
nl >= rtMinLeafSize cfg && nNode - nl >= rtMinLeafSize cfg && wl > 0 && wr > 0
in if ok
then Just (nodeSSE - (sse syl syl2 wl + sse (totSY - syl) (totSY2 - syl2) wr))
else Nothing
Expand Down
57 changes: 46 additions & 11 deletions dataframe-learn/tests-internal/HistogramTrees.hs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,11 @@ import Test.HUnit

import DataFrame.DecisionTree.Cart (cartFeatures)
import DataFrame.DecisionTree.Fit (treeToExpr)
import DataFrame.DecisionTree.Histogram (TreeLimits (..), binFeatures, fitBinnedTree)
import DataFrame.DecisionTree.Histogram (
TreeLimits (..),
binFeatures,
fitBinnedTree,
)
import DataFrame.DecisionTree.Regression (
RegFit (..),
RegTreeConfig (..),
Expand All @@ -30,17 +34,26 @@ import qualified DataFrameApi as D

tests :: [Test]
tests =
[ TestLabel "histogram tree: in-sample == interpreted (mixed nullable frame)" inSampleMatchesInterpret
, TestLabel "histogram tree: in-sample == interpreted (more values than bins)" coarseBinsMatchInterpret
[ TestLabel
"histogram tree: in-sample == interpreted (mixed nullable frame)"
inSampleMatchesInterpret
, TestLabel
"histogram tree: in-sample == interpreted (more values than bins)"
coarseBinsMatchInterpret
, TestLabel "histogram tree: same partition as the exact tree" matchesExactTree
]

limits :: Int -> TreeLimits
limits depth = TreeLimits depth 2 1 0

fitBinned :: Int -> Int -> D.DataFrame -> VU.Vector Double -> (Tree Double, VU.Vector Double)
fitBinned ::
Int -> Int -> D.DataFrame -> VU.Vector Double -> (Tree Double, VU.Vector Double)
fitBinned bins depth df y =
fitBinnedTree (limits depth) (binFeatures bins (V.fromList (cartFeatures "y" df))) y Nothing
fitBinnedTree
(limits depth)
(binFeatures bins (V.fromList (cartFeatures "y" df)))
y
Nothing

interpreted :: D.DataFrame -> Tree Double -> VU.Vector Double
interpreted df t = case interpret @Double df (treeToExpr t) of
Expand All @@ -55,26 +68,48 @@ mixedFrame :: Int -> ([Double], D.DataFrame)
mixedFrame n = (y, D.fromColumns (("y", DI.fromList y) : cols))
where
rows = [0 .. n - 1]
md = [if i `mod` 7 == 0 then Nothing else Just (fromIntegral ((i * 37) `mod` 23) :: Double) | i <- rows]
mi = [if i `mod` 5 == 2 then Nothing else Just ((i * 11) `mod` 9 :: Int) | i <- rows]
mt = [if i `mod` 6 == 4 then Nothing else Just (["a", "b", "c"] !! (i `mod` 3) :: T.Text) | i <- rows]
md =
[ if i `mod` 7 == 0
then Nothing
else Just (fromIntegral ((i * 37) `mod` 23) :: Double)
| i <- rows
]
mi =
[if i `mod` 5 == 2 then Nothing else Just ((i * 11) `mod` 9 :: Int) | i <- rows]
mt =
[ if i `mod` 6 == 4
then Nothing
else Just (["a", "b", "c"] !! (i `mod` 3) :: T.Text)
| i <- rows
]
z = [fromIntegral ((i * 13) `mod` 17) :: Double | i <- rows]
y = [fromIntegral ((i * 29) `mod` 31) / 7 :: Double | i <- rows]
cols = [("md", maybeCol md), ("mi", maybeCol mi), ("mt", maybeCol mt), ("z", DI.fromList z)]
cols =
[ ("md", maybeCol md)
, ("mi", maybeCol mi)
, ("mt", maybeCol mt)
, ("z", DI.fromList z)
]

inSampleMatchesInterpret :: Test
inSampleMatchesInterpret = TestCase $ do
let (y, df) = mixedFrame 60
(t, inSample) = fitBinned 1024 6 df (VU.fromList y)
assertEqual "in-sample predictions equal the interpreted tree" (interpreted df t) inSample
assertEqual
"in-sample predictions equal the interpreted tree"
(interpreted df t)
inSample

-- 23 distinct values of md squeezed into 4 bins: thresholds come from the
-- equal-count cuts, and routing must still agree.
coarseBinsMatchInterpret :: Test
coarseBinsMatchInterpret = TestCase $ do
let (y, df) = mixedFrame 200
(t, inSample) = fitBinned 4 5 df (VU.fromList y)
assertEqual "in-sample predictions equal the interpreted tree" (interpreted df t) inSample
assertEqual
"in-sample predictions equal the interpreted tree"
(interpreted df t)
inSample

matchesExactTree :: Test
matchesExactTree = TestCase $ do
Expand Down
47 changes: 37 additions & 10 deletions dataframe-learn/tests-internal/NullSplits.hs
Original file line number Diff line number Diff line change
Expand Up @@ -23,23 +23,26 @@ import DataFrame.DecisionTree.Regression (
fitRegTree,
)
import DataFrame.DecisionTree.Types (Tree, TreeConfig (..), defaultTreeConfig)
import DataFrame.Internal.Expression (Expr (..))
import qualified DataFrame.Internal.Column as DI
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.Internal.Interpreter (interpret)
import DataFrame.Model (Fit (..), Predict (..))
import qualified DataFrameApi as D

tests :: [Test]
tests =
[ TestLabel "null splits: in-sample == interpreted (mixed nullable frame)" inSampleMatchesInterpret
[ TestLabel
"null splits: in-sample == interpreted (mixed nullable frame)"
inSampleMatchesInterpret
, TestLabel "null splits: nulls join the side they fit" learnedDirection
, TestLabel "null splits: null-indicator split" nullIndicator
, TestLabel "null splits: unseen nulls still route" unseenNullsRoute
, TestLabel "null splits: CART null indicator" cartNullIndicator
, TestLabel "null splits: AdaBoost null indicator" adaBoostNullIndicator
]

fitAll :: Int -> D.DataFrame -> VU.Vector Double -> (Tree Double, VU.Vector Double)
fitAll ::
Int -> D.DataFrame -> VU.Vector Double -> (Tree Double, VU.Vector Double)
fitAll depth df y = (rfTree fitted, rfFitted fitted)
where
fitted =
Expand Down Expand Up @@ -67,17 +70,35 @@ inSampleMatchesInterpret :: Test
inSampleMatchesInterpret = TestCase $ do
let n = 60 :: Int
rows = [0 .. n - 1]
md = [if i `mod` 7 == 0 then Nothing else Just (fromIntegral ((i * 37) `mod` 23) :: Double) | i <- rows]
mi = [if i `mod` 5 == 2 then Nothing else Just ((i * 11) `mod` 9 :: Int) | i <- rows]
mt = [if i `mod` 6 == 4 then Nothing else Just (["a", "b", "c"] !! (i `mod` 3) :: T.Text) | i <- rows]
md =
[ if i `mod` 7 == 0
then Nothing
else Just (fromIntegral ((i * 37) `mod` 23) :: Double)
| i <- rows
]
mi =
[if i `mod` 5 == 2 then Nothing else Just ((i * 11) `mod` 9 :: Int) | i <- rows]
mt =
[ if i `mod` 6 == 4
then Nothing
else Just (["a", "b", "c"] !! (i `mod` 3) :: T.Text)
| i <- rows
]
z = [fromIntegral ((i * 13) `mod` 17) :: Double | i <- rows]
y = [fromIntegral ((i * 29) `mod` 31) / 7 :: Double | i <- rows]
df =
withTarget
y
[("md", maybeCol md), ("mi", maybeCol mi), ("mt", maybeCol mt), ("z", DI.fromList z)]
[ ("md", maybeCol md)
, ("mi", maybeCol mi)
, ("mt", maybeCol mt)
, ("z", DI.fromList z)
]
(t, inSample) = fitAll 6 df (VU.fromList y)
assertEqual "in-sample predictions equal the interpreted tree" (interpreted df t) inSample
assertEqual
"in-sample predictions equal the interpreted tree"
(interpreted df t)
inSample

-- The null rows have the same target as the low values of x, so one split
-- fits exactly only if it sends the nulls left.
Expand Down Expand Up @@ -108,7 +129,9 @@ unseenNullsRoute = TestCase $ do
scored = interpreted testDf t
assertEqual "fits exactly" (VU.fromList y) inSample
assertEqual "present rows unchanged" (VU.take 4 inSample) (VU.take 4 scored)
assertBool "null rows land in a leaf" (VU.all (`elem` [0, 1]) (VU.drop 4 scored))
assertBool
"null rows land in a leaf"
(VU.all (`elem` [0, 1]) (VU.drop 4 scored))

-- The target depends only on whether x is null. If the split's threshold were
-- +Infinity instead of 3, CART and AdaBoost would send the null rows left
Expand All @@ -122,7 +145,11 @@ indicatorFrame = (y, withTarget y [("x", maybeCol x)])
cartNullIndicator :: Test
cartNullIndicator = TestCase $ do
let (y, df) = indicatorFrame
t = buildCartTree @Double defaultTreeConfig{maxTreeDepth = 1, minLeafSize = 1} "y" df
t =
buildCartTree @Double
defaultTreeConfig{maxTreeDepth = 1, minLeafSize = 1}
"y"
df
assertEqual "fits exactly" (VU.fromList y) (interpreted df t)

adaBoostNullIndicator :: Test
Expand Down
Loading
Loading