{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeApplications #-}

{- | Oblique split candidates: fit an L1-regularised logistic hyperplane to the
care points (class-balanced) and convert it to a boolean condition, rejecting
all-zero and degenerate (single-side) hyperplanes.
-}
module DataFrame.DecisionTree.Linear (
    bestLinearCandidate,
    fitLinearCandidate,
    careRowsFromFeatures,
    careLabels,
    featName,
    imputeMean,
    materializeFeatureForCare,
) where

import DataFrame.DecisionTree.Numeric (NumExpr (..), numExprCols, numericCols)
import DataFrame.DecisionTree.Types (
    CarePoint (..),
    Direction (..),
    TreeConfig (..),
 )
import DataFrame.Internal.Column (TypedColumn (..), toVector)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr, getColumns)
import DataFrame.Internal.Interpreter (interpret)
import qualified DataFrame.LinearSolver as LS

import Data.Maybe (catMaybes, fromMaybe, mapMaybe)
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

{- | Best oblique candidate, or 'Nothing' when the linear path is disabled or
there are too few care points to fit on. The target column is named so it can
be kept out of the feature set.
-}
bestLinearCandidate ::
    TreeConfig -> T.Text -> DataFrame -> [CarePoint] -> Maybe (Expr Bool)
bestLinearCandidate :: TreeConfig -> Text -> DataFrame -> [CarePoint] -> Maybe (Expr Bool)
bestLinearCandidate TreeConfig
cfg Text
target DataFrame
df [CarePoint]
carePoints
    | Bool -> Bool
not (TreeConfig -> Bool
useLinearSolver TreeConfig
cfg) = Maybe (Expr Bool)
forall a. Maybe a
Nothing
    | [CarePoint] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [CarePoint]
carePoints Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< TreeConfig -> Int
minCarePointsForLinear TreeConfig
cfg = Maybe (Expr Bool)
forall a. Maybe a
Nothing
    | Bool
otherwise = TreeConfig -> Text -> DataFrame -> [CarePoint] -> Maybe (Expr Bool)
fitLinearCandidate TreeConfig
cfg Text
target DataFrame
df [CarePoint]
carePoints

{- | Fit an L1 logistic regression to the care points and convert the resulting
hyperplane to a condition, or 'Nothing' when no numeric features exist or the
fitted model is all-zero or degenerate.
-}
fitLinearCandidate ::
    TreeConfig -> T.Text -> DataFrame -> [CarePoint] -> Maybe (Expr Bool)
fitLinearCandidate :: TreeConfig -> Text -> DataFrame -> [CarePoint] -> Maybe (Expr Bool)
fitLinearCandidate TreeConfig
cfg Text
target DataFrame
df [CarePoint]
carePoints =
    case Text -> DataFrame -> [CarePoint] -> [(Text, Vector Double)]
materializedFeatures Text
target DataFrame
df [CarePoint]
carePoints of
        [] -> Maybe (Expr Bool)
forall a. Maybe a
Nothing
        [(Text, Vector Double)]
mats -> TreeConfig
-> [CarePoint] -> [(Text, Vector Double)] -> Maybe (Expr Bool)
linearFromFeatures TreeConfig
cfg [CarePoint]
carePoints [(Text, Vector Double)]
mats

materializedFeatures ::
    T.Text -> DataFrame -> [CarePoint] -> [(T.Text, VU.Vector Double)]
materializedFeatures :: Text -> DataFrame -> [CarePoint] -> [(Text, Vector Double)]
materializedFeatures Text
target DataFrame
df [CarePoint]
carePoints =
    (NumExpr -> Maybe (Text, Vector Double))
-> [NumExpr] -> [(Text, Vector Double)]
forall a b. (a -> Maybe b) -> [a] -> [b]
mapMaybe (DataFrame -> [CarePoint] -> NumExpr -> Maybe (Text, Vector Double)
materializeFeatureForCare DataFrame
df [CarePoint]
carePoints) (Text -> DataFrame -> [NumExpr]
featureCols Text
target DataFrame
df)

featureCols :: T.Text -> DataFrame -> [NumExpr]
featureCols :: Text -> DataFrame -> [NumExpr]
featureCols Text
target DataFrame
df = (NumExpr -> Bool) -> [NumExpr] -> [NumExpr]
forall a. (a -> Bool) -> [a] -> [a]
filter (Text -> [Text] -> Bool
forall (t :: * -> *) a. (Foldable t, Eq a) => a -> t a -> Bool
notElem Text
target ([Text] -> Bool) -> (NumExpr -> [Text]) -> NumExpr -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. NumExpr -> [Text]
numExprCols) (DataFrame -> [NumExpr]
numericCols DataFrame
df)

linearFromFeatures ::
    TreeConfig -> [CarePoint] -> [(T.Text, VU.Vector Double)] -> Maybe (Expr Bool)
linearFromFeatures :: TreeConfig
-> [CarePoint] -> [(Text, Vector Double)] -> Maybe (Expr Bool)
linearFromFeatures TreeConfig
cfg [CarePoint]
carePoints [(Text, Vector Double)]
mats
    | (Double -> Bool) -> Vector Double -> Bool
forall a. Unbox a => (a -> Bool) -> Vector a -> Bool
VU.all (Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0) Vector Double
weights = Maybe (Expr Bool)
forall a. Maybe a
Nothing
    | Vector (Vector Double) -> Vector Double -> Double -> Bool
degenerateHyperplane Vector (Vector Double)
rows Vector Double
weights (LinearModel -> Double
LS.lmIntercept LinearModel
model) = Maybe (Expr Bool)
forall a. Maybe a
Nothing
    | Bool
otherwise = Expr Bool -> Maybe (Expr Bool)
forall a. a -> Maybe a
Just (LinearModel -> Expr Bool
LS.modelToExpr LinearModel
model)
  where
    rows :: Vector (Vector Double)
rows = Int -> [(Text, Vector Double)] -> Vector (Vector Double)
careRowsFromFeatures ([CarePoint] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [CarePoint]
carePoints) [(Text, Vector Double)]
mats
    labels :: Vector Double
labels = [CarePoint] -> Vector Double
careLabels [CarePoint]
carePoints
    model :: LinearModel
model =
        SolverConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector Text
-> LinearModel
LS.fitL1Logistic
            (TreeConfig -> Vector Double -> SolverConfig
solverConfigFor TreeConfig
cfg Vector Double
labels)
            Vector (Vector Double)
rows
            Vector Double
labels
            ([Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList (((Text, Vector Double) -> Text)
-> [(Text, Vector Double)] -> [Text]
forall a b. (a -> b) -> [a] -> [b]
map (Text, Vector Double) -> Text
forall a b. (a, b) -> a
fst [(Text, Vector Double)]
mats))
    weights :: Vector Double
weights = LinearModel -> Vector Double
LS.lmWeights LinearModel
model

solverConfigFor :: TreeConfig -> VU.Vector Double -> LS.SolverConfig
solverConfigFor :: TreeConfig -> Vector Double -> SolverConfig
solverConfigFor TreeConfig
cfg Vector Double
labels = (TreeConfig -> SolverConfig
linearSolverConfig TreeConfig
cfg){LS.scSampleWeights = classBalancedWeights labels}

{- | Class-balanced sklearn-form weights @w_i = N / (2 · N_class)@ (mean 1), or
'Nothing' in the degenerate one-class case (uniform weighting).
-}
classBalancedWeights :: VU.Vector Double -> Maybe (VU.Vector Double)
classBalancedWeights :: Vector Double -> Maybe (Vector Double)
classBalancedWeights Vector Double
labels
    | Int
nPos Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
0 Bool -> Bool -> Bool
&& Int
nNeg Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
0 = Vector Double -> Maybe (Vector Double)
forall a. a -> Maybe a
Just (Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
nCare Int -> Double
forall {a}. Fractional a => Int -> a
weightAt)
    | Bool
otherwise = Maybe (Vector Double)
forall a. Maybe a
Nothing
  where
    nCare :: Int
nCare = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
labels
    nPos :: Int
nPos = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length ((Double -> Bool) -> Vector Double -> Vector Double
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0) Vector Double
labels)
    nNeg :: Int
nNeg = Int
nCare Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
nPos
    weightAt :: Int -> a
weightAt Int
i
        | Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
labels Int
i Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0 = Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
nCare a -> a -> a
forall a. Fractional a => a -> a -> a
/ (a
2 a -> a -> a
forall a. Num a => a -> a -> a
* Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
nPos)
        | Bool
otherwise = Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
nCare a -> a -> a
forall a. Fractional a => a -> a -> a
/ (a
2 a -> a -> a
forall a. Num a => a -> a -> a
* Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
nNeg)

{- | A hyperplane is degenerate when every care row scores on the same side of
zero (equivalent to an invalid split, caught upstream).
-}
degenerateHyperplane ::
    V.Vector (VU.Vector Double) -> VU.Vector Double -> Double -> Bool
degenerateHyperplane :: Vector (Vector Double) -> Vector Double -> Double -> Bool
degenerateHyperplane Vector (Vector Double)
rows Vector Double
weights Double
bias =
    Int
nCare Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
0 Bool -> Bool -> Bool
&& (Vector Double -> Double
forall a. (Unbox a, Ord a) => Vector a -> a
VU.minimum Vector Double
scores Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0 Bool -> Bool -> Bool
|| Vector Double -> Double
forall a. (Unbox a, Ord a) => Vector a -> a
VU.maximum Vector Double
scores Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
0)
  where
    nCare :: Int
nCare = Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
rows
    scores :: Vector Double
scores =
        Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate
            Int
nCare
            (\Int
i -> Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Vector Double
weights (Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.unsafeIndex Vector (Vector Double)
rows Int
i)) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
bias)

{- | Per-care-point feature rows from materialized columns (each of length
@nCare@, so indexing is in range).
-}
careRowsFromFeatures ::
    Int -> [(T.Text, VU.Vector Double)] -> V.Vector (VU.Vector Double)
careRowsFromFeatures :: Int -> [(Text, Vector Double)] -> Vector (Vector Double)
careRowsFromFeatures Int
nCare [(Text, Vector Double)]
mats =
    Int -> (Int -> Vector Double) -> Vector (Vector Double)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
nCare (\Int
i -> Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
nFeat (\Int
j -> (Text, Vector Double) -> Vector Double
forall a b. (a, b) -> b
snd (Vector (Text, Vector Double)
matsVec Vector (Text, Vector Double) -> Int -> (Text, Vector Double)
forall a. Vector a -> Int -> a
V.! Int
j) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i))
  where
    matsVec :: Vector (Text, Vector Double)
matsVec = [(Text, Vector Double)] -> Vector (Text, Vector Double)
forall a. [a] -> Vector a
V.fromList [(Text, Vector Double)]
mats
    nFeat :: Int
nFeat = Vector (Text, Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Text, Vector Double)
matsVec

-- | Solver labels: @+1@ when 'GoLeft' is correct, @-1@ otherwise.
careLabels :: [CarePoint] -> VU.Vector Double
careLabels :: [CarePoint] -> Vector Double
careLabels [CarePoint]
carePoints =
    [Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList [if CarePoint -> Direction
cpCorrectDir CarePoint
cp Direction -> Direction -> Bool
forall a. Eq a => a -> a -> Bool
== Direction
GoLeft then Double
1.0 else -Double
1.0 | CarePoint
cp <- [CarePoint]
carePoints]

-- | First column referenced by an expression, or a placeholder when none.
featName :: Expr b -> T.Text
featName :: forall b. Expr b -> Text
featName Expr b
expr = case Expr b -> [Text]
forall a. Expr a -> [Text]
getColumns Expr b
expr of
    (Text
c : [Text]
_) -> Text
c
    [] -> Text
"<feat>"

{- | Replace missing values with the mean of present ones; 'Nothing' when
nothing is present so the caller can drop the feature.
-}
imputeMean :: [Maybe Double] -> Maybe (VU.Vector Double)
imputeMean :: [Maybe Double] -> Maybe (Vector Double)
imputeMean [Maybe Double]
careRaw = case [Maybe Double] -> [Double]
forall a. [Maybe a] -> [a]
catMaybes [Maybe Double]
careRaw of
    [] -> Maybe (Vector Double)
forall a. Maybe a
Nothing
    [Double]
present -> Vector Double -> Maybe (Vector Double)
forall a. a -> Maybe a
Just ([Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList [Double -> Maybe Double -> Double
forall a. a -> Maybe a -> a
fromMaybe ([Double] -> Double
forall {a} {t :: * -> *}. (Fractional a, Foldable t) => t a -> a
mean [Double]
present) Maybe Double
mv | Maybe Double
mv <- [Maybe Double]
careRaw])
  where
    mean :: t a -> a
mean t a
xs = t a -> a
forall a. Num a => t a -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum t a
xs a -> a -> a
forall a. Fractional a => a -> a -> a
/ Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral (t a -> Int
forall a. t a -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length t a
xs)

interpretDoubleVals :: DataFrame -> Expr Double -> Maybe (V.Vector Double)
interpretDoubleVals :: DataFrame -> Expr Double -> Maybe (Vector Double)
interpretDoubleVals DataFrame
df Expr Double
expr = case forall a.
Columnable a =>
DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
interpret @Double DataFrame
df Expr Double
expr of
    Right (TColumn Column
column) -> (DataFrameException -> Maybe (Vector Double))
-> (Vector Double -> Maybe (Vector Double))
-> Either DataFrameException (Vector Double)
-> Maybe (Vector Double)
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Maybe (Vector Double)
-> DataFrameException -> Maybe (Vector Double)
forall a b. a -> b -> a
const Maybe (Vector Double)
forall a. Maybe a
Nothing) Vector Double -> Maybe (Vector Double)
forall a. a -> Maybe a
Just (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @Double Column
column)
    Either DataFrameException (TypedColumn Double)
_ -> Maybe (Vector Double)
forall a. Maybe a
Nothing

interpretMaybeDoubleVals ::
    DataFrame -> Expr (Maybe Double) -> Maybe (V.Vector (Maybe Double))
interpretMaybeDoubleVals :: DataFrame -> Expr (Maybe Double) -> Maybe (Vector (Maybe Double))
interpretMaybeDoubleVals DataFrame
df Expr (Maybe Double)
expr = case forall a.
Columnable a =>
DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
interpret @(Maybe Double) DataFrame
df Expr (Maybe Double)
expr of
    Right (TColumn Column
column) -> (DataFrameException -> Maybe (Vector (Maybe Double)))
-> (Vector (Maybe Double) -> Maybe (Vector (Maybe Double)))
-> Either DataFrameException (Vector (Maybe Double))
-> Maybe (Vector (Maybe Double))
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Maybe (Vector (Maybe Double))
-> DataFrameException -> Maybe (Vector (Maybe Double))
forall a b. a -> b -> a
const Maybe (Vector (Maybe Double))
forall a. Maybe a
Nothing) Vector (Maybe Double) -> Maybe (Vector (Maybe Double))
forall a. a -> Maybe a
Just (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @(Maybe Double) Column
column)
    Either DataFrameException (TypedColumn (Maybe Double))
_ -> Maybe (Vector (Maybe Double))
forall a. Maybe a
Nothing

{- | Materialize a 'NumExpr' over the care rows; 'Nothing' on interpret failure
or (nullable) when no care point has a present value, else mean-imputed.
-}
materializeFeatureForCare ::
    DataFrame -> [CarePoint] -> NumExpr -> Maybe (T.Text, VU.Vector Double)
materializeFeatureForCare :: DataFrame -> [CarePoint] -> NumExpr -> Maybe (Text, Vector Double)
materializeFeatureForCare DataFrame
df [CarePoint]
carePoints (NDouble Expr Double
expr) = do
    Vector Double
vals <- DataFrame -> Expr Double -> Maybe (Vector Double)
interpretDoubleVals DataFrame
df Expr Double
expr
    (Text, Vector Double) -> Maybe (Text, Vector Double)
forall a. a -> Maybe a
Just (Expr Double -> Text
forall b. Expr b -> Text
featName Expr Double
expr, [Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList [Vector Double
vals Vector Double -> Int -> Double
forall a. Vector a -> Int -> a
V.! CarePoint -> Int
cpIndex CarePoint
cp | CarePoint
cp <- [CarePoint]
carePoints])
materializeFeatureForCare DataFrame
df [CarePoint]
carePoints (NMaybeDouble Expr (Maybe Double)
expr) = do
    Vector (Maybe Double)
vals <- DataFrame -> Expr (Maybe Double) -> Maybe (Vector (Maybe Double))
interpretMaybeDoubleVals DataFrame
df Expr (Maybe Double)
expr
    Vector Double
imputed <- [Maybe Double] -> Maybe (Vector Double)
imputeMean [Vector (Maybe Double)
vals Vector (Maybe Double) -> Int -> Maybe Double
forall a. Vector a -> Int -> a
V.! CarePoint -> Int
cpIndex CarePoint
cp | CarePoint
cp <- [CarePoint]
carePoints]
    (Text, Vector Double) -> Maybe (Text, Vector Double)
forall a. a -> Maybe a
Just (Expr (Maybe Double) -> Text
forall b. Expr b -> Text
featName Expr (Maybe Double)
expr, Vector Double
imputed)