{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ExplicitNamespaces #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}

{- | The aggregation fast-path planner and the two-column moment scatter.
'planAgg' recognises a supported aggregate shape over a clean unboxed Int/Double
column and returns an 'AggPlan'; 'momentScatter' fuses the six regression sums.
-}
module DataFrame.Internal.AggPlan (
    AggPlan (..),
    planAgg,
    Moments (..),
    momentScatter,
    MomentPlan (..),
    planMoments,
) where

import qualified Data.Map.Strict as M
import qualified Data.Text as T
import Data.Type.Equality (TestEquality (..), type (:~:) (Refl))
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM

import Control.Monad.ST (runST)
import DataFrame.Internal.AggKernel (Reduction (..), scatterColumnToDouble)
import DataFrame.Internal.Column (Column (..), fromUnboxedVector)
import DataFrame.Internal.DataFrame (
    DataFrame (derivingExpressions),
    GroupedDataFrame (..),
    getColumn,
 )
import DataFrame.Internal.Expression (
    AggStrategy (..),
    BinaryOp (binaryCommutative, binaryName),
    Expr (..),
    UExpr (..),
 )
import Type.Reflection (Typeable, typeRep)

{- | The plan 'planAgg' produces for a recognised output expression. The median
plan carries only the column name (the holistic grouped sort lives in the
operations layer, where @vector-algorithms@ is available).
-}
data AggPlan
    = -- | A single scatter reduction over one named column.
      PlanScatter Reduction T.Text
    | -- | @max a - min b@ (Q7): two scatters then a vectorized combine.
      PlanMaxMinusMin T.Text T.Text
    | -- | Holistic median over one named column.
      PlanMedian T.Text

{- | Inspect a named output expression; return @Just plan@ on a recognised shape
over a present clean column, else 'Nothing'. Nullable or non-Int/Double columns
are rejected here so the scatter only sees a clean unboxed vector.
-}
planAgg :: GroupedDataFrame -> UExpr -> Maybe AggPlan
planAgg :: GroupedDataFrame -> UExpr -> Maybe AggPlan
planAgg GroupedDataFrame
gdf (UExpr (Expr a
expr :: Expr a)) = case Expr a
expr of
    Agg (FoldAgg Text
tag Maybe a
_ a -> b -> a
_) (Col Text
name) -> Text -> Text -> Maybe AggPlan
foldPlan Text
tag Text
name
    Agg (MergeAgg Text
tag acc
_ acc -> b -> acc
_ acc -> acc -> acc
_ acc -> a
_) (Col Text
name) -> Text -> Text -> Maybe AggPlan
mergePlan Text
tag Text
name
    Agg (CollectAgg Text
tag v b -> a
_) (Col Text
name) -> Text -> Text -> Maybe AggPlan
collectPlan Text
tag Text
name
    Binary
        op c b a
op
        (Agg (FoldAgg Text
lt Maybe c
Nothing c -> b -> c
_) (Col Text
a))
        (Agg (FoldAgg Text
rt Maybe b
Nothing b -> b -> b
_) (Col Text
b)) ->
            if op c b a -> Text
forall a b c. op a b c -> Text
forall (op :: * -> * -> * -> *) a b c.
BinaryOp op =>
op a b c -> Text
binaryName op c b a
op Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Text
"sub" Bool -> Bool -> Bool
&& Text
lt Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Text
"maximum" Bool -> Bool -> Bool
&& Text
rt Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Text
"minimum"
                then Text -> Text -> AggPlan -> Maybe AggPlan
requireBoth Text
a Text
b (Text -> Text -> AggPlan
PlanMaxMinusMin Text
a Text
b)
                else Maybe AggPlan
forall a. Maybe a
Nothing
    Expr a
_ -> Maybe AggPlan
forall a. Maybe a
Nothing
  where
    foldPlan :: Text -> Text -> Maybe AggPlan
foldPlan Text
tag Text
name = case Text
tag of
        Text
"sum" -> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RSum Text
name)
        Text
"minimum" -> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RMin Text
name)
        Text
"maximum" -> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RMax Text
name)
        Text
_ -> Maybe AggPlan
forall a. Maybe a
Nothing
    mergePlan :: Text -> Text -> Maybe AggPlan
mergePlan Text
tag Text
name = case Text
tag of
        Text
"mean" -> forall t. Typeable t => Maybe ()
outputType @Double Maybe () -> Maybe AggPlan -> Maybe AggPlan
forall a b. Maybe a -> Maybe b -> Maybe b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RMean Text
name)
        Text
"count" -> forall t. Typeable t => Maybe ()
outputType @Int Maybe () -> Maybe AggPlan -> Maybe AggPlan
forall a b. Maybe a -> Maybe b -> Maybe b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RCount Text
name)
        Text
_ -> Maybe AggPlan
forall a. Maybe a
Nothing
    outputType :: forall t. (Typeable t) => Maybe ()
    outputType :: forall t. Typeable t => Maybe ()
outputType = case TypeRep a -> TypeRep t -> Maybe (a :~: t)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @t) of
        Just a :~: t
Refl -> () -> Maybe ()
forall a. a -> Maybe a
Just ()
        Maybe (a :~: t)
Nothing -> Maybe ()
forall a. Maybe a
Nothing
    collectPlan :: Text -> Text -> Maybe AggPlan
collectPlan Text
tag Text
name = case Text
tag of
        Text
"stddev" -> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RStd Text
name)
        Text
"variance" -> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RVar Text
name)
        Text
"top2Sum" -> Text -> AggPlan -> Maybe AggPlan
require Text
name (Reduction -> Text -> AggPlan
PlanScatter Reduction
RTop2Sum Text
name)
        Text
"median" -> Text -> AggPlan -> Maybe AggPlan
require Text
name (Text -> AggPlan
PlanMedian Text
name)
        Text
_ -> Maybe AggPlan
forall a. Maybe a
Nothing
    require :: Text -> AggPlan -> Maybe AggPlan
require Text
name AggPlan
plan = Text -> Maybe ()
colUnboxedNumeric Text
name Maybe () -> Maybe AggPlan -> Maybe AggPlan
forall a b. Maybe a -> Maybe b -> Maybe b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> AggPlan -> Maybe AggPlan
forall a. a -> Maybe a
Just AggPlan
plan
    requireBoth :: Text -> Text -> AggPlan -> Maybe AggPlan
requireBoth Text
a Text
b AggPlan
plan = Text -> Maybe ()
colUnboxedNumeric Text
a Maybe () -> Maybe () -> Maybe ()
forall a b. Maybe a -> Maybe b -> Maybe b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Text -> Maybe ()
colUnboxedNumeric Text
b Maybe () -> Maybe AggPlan -> Maybe AggPlan
forall a b. Maybe a -> Maybe b -> Maybe b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> AggPlan -> Maybe AggPlan
forall a. a -> Maybe a
Just AggPlan
plan
    colUnboxedNumeric :: Text -> Maybe ()
colUnboxedNumeric Text
name = case Text -> DataFrame -> Maybe Column
getColumn Text
name (GroupedDataFrame -> DataFrame
fullDataframe GroupedDataFrame
gdf) of
        Just Column
c | Column -> Bool
isUnboxedNumeric Column
c -> () -> Maybe ()
forall a. a -> Maybe a
Just ()
        Maybe Column
_ -> Maybe ()
forall a. Maybe a
Nothing

-- | The matcher only fires on non-null unboxed Int/Double columns.
isUnboxedNumeric :: Column -> Bool
isUnboxedNumeric :: Column -> Bool
isUnboxedNumeric = \case
    UnboxedColumn Maybe Bitmap
Nothing (Vector a
_ :: VU.Vector a) ->
        case TypeRep a -> TypeRep Int -> Maybe (a :~: Int)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Int) of
            Just a :~: Int
Refl -> Bool
True
            Maybe (a :~: Int)
Nothing -> case TypeRep a -> TypeRep Double -> Maybe (a :~: Double)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Double) of
                Just a :~: Double
Refl -> Bool
True
                Maybe (a :~: Double)
Nothing -> Bool
False
    Column
_ -> Bool
False

{- | A recognised moment (Q9 regression) aggregate group: six output columns that
form the sufficient statistics of two base columns @x@ and @y@. The caller runs
'momentScatter' once and binds each output name to a field of the result.
-}
data MomentPlan = MomentPlan
    { MomentPlan -> Text
mpColX :: T.Text
    , MomentPlan -> Text
mpColY :: T.Text
    , MomentPlan -> Text
mpNName :: T.Text
    , MomentPlan -> Text
mpSxName :: T.Text
    , MomentPlan -> Text
mpSyName :: T.Text
    , MomentPlan -> Text
mpSxxName :: T.Text
    , MomentPlan -> Text
mpSyyName :: T.Text
    , MomentPlan -> Text
mpSxyName :: T.Text
    }

{- | The shape of a sum's argument once unary coercions are peeled and derived
columns are resolved through @derivingExpressions@: either linear in one base
column or the product of two base columns (sorted).
-}
data Term
    = Lin T.Text
    | Prod T.Text T.Text
    deriving (Term -> Term -> Bool
(Term -> Term -> Bool) -> (Term -> Term -> Bool) -> Eq Term
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Term -> Term -> Bool
== :: Term -> Term -> Bool
$c/= :: Term -> Term -> Bool
/= :: Term -> Term -> Bool
Eq, Eq Term
Eq Term =>
(Term -> Term -> Ordering)
-> (Term -> Term -> Bool)
-> (Term -> Term -> Bool)
-> (Term -> Term -> Bool)
-> (Term -> Term -> Bool)
-> (Term -> Term -> Term)
-> (Term -> Term -> Term)
-> Ord Term
Term -> Term -> Bool
Term -> Term -> Ordering
Term -> Term -> Term
forall a.
Eq a =>
(a -> a -> Ordering)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> a)
-> (a -> a -> a)
-> Ord a
$ccompare :: Term -> Term -> Ordering
compare :: Term -> Term -> Ordering
$c< :: Term -> Term -> Bool
< :: Term -> Term -> Bool
$c<= :: Term -> Term -> Bool
<= :: Term -> Term -> Bool
$c> :: Term -> Term -> Bool
> :: Term -> Term -> Bool
$c>= :: Term -> Term -> Bool
>= :: Term -> Term -> Bool
$cmax :: Term -> Term -> Term
max :: Term -> Term -> Term
$cmin :: Term -> Term -> Term
min :: Term -> Term -> Term
Ord, Int -> Term -> ShowS
[Term] -> ShowS
Term -> String
(Int -> Term -> ShowS)
-> (Term -> String) -> ([Term] -> ShowS) -> Show Term
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Term -> ShowS
showsPrec :: Int -> Term -> ShowS
$cshow :: Term -> String
show :: Term -> String
$cshowList :: [Term] -> ShowS
showList :: [Term] -> ShowS
Show)

{- | Recognise the moment shape across a whole @aggregate@ list: exactly
@count@, @sum(x)@, @sum(y)@, @sum(x*x)@, @sum(y*y)@, @sum(x*y)@ over two distinct
clean unboxed base columns. 'Nothing' on any other set.
-}
planMoments :: GroupedDataFrame -> [(T.Text, UExpr)] -> Maybe MomentPlan
planMoments :: GroupedDataFrame -> [(Text, UExpr)] -> Maybe MomentPlan
planMoments GroupedDataFrame
gdf [(Text, UExpr)]
aggs
    | [(Text, UExpr)] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [(Text, UExpr)]
aggs Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
/= Int
6 = Maybe MomentPlan
forall a. Maybe a
Nothing
    | Bool
otherwise = do
        let exprs :: Map Text UExpr
exprs = DataFrame -> Map Text UExpr
derivingExpressions (GroupedDataFrame -> DataFrame
fullDataframe GroupedDataFrame
gdf)
        [(Text, Role)]
roles <- ((Text, UExpr) -> Maybe (Text, Role))
-> [(Text, UExpr)] -> Maybe [(Text, Role)]
forall (t :: * -> *) (f :: * -> *) a b.
(Traversable t, Applicative f) =>
(a -> f b) -> t a -> f (t b)
forall (f :: * -> *) a b.
Applicative f =>
(a -> f b) -> [a] -> f [b]
traverse (Map Text UExpr -> (Text, UExpr) -> Maybe (Text, Role)
classify Map Text UExpr
exprs) [(Text, UExpr)]
aggs
        let names :: Map Role Text
names = [(Role, Text)] -> Map Role Text
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList [(Role
r, Text
nm) | (Text
nm, Role
r) <- [(Text, Role)]
roles]
        Text
nName <- Role -> Map Role Text -> Maybe Text
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup Role
RoleN Map Role Text
names
        (Text
x, Text
y) <- [(Text, Role)] -> Maybe (Text, Text)
pickBaseColumns [(Text, Role)]
roles
        Text
sxName <- Role -> Map Role Text -> Maybe Text
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup (Text -> Role
RoleLin Text
x) Map Role Text
names
        Text
syName <- Role -> Map Role Text -> Maybe Text
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup (Text -> Role
RoleLin Text
y) Map Role Text
names
        Text
sxxName <- Role -> Map Role Text -> Maybe Text
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup (Text -> Text -> Role
RoleProd Text
x Text
x) Map Role Text
names
        Text
syyName <- Role -> Map Role Text -> Maybe Text
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup (Text -> Text -> Role
RoleProd Text
y Text
y) Map Role Text
names
        Text
sxyName <- Role -> Map Role Text -> Maybe Text
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup (Text -> Text -> Role
RoleProd Text
x Text
y) Map Role Text
names
        ()
_ <- if Text
x Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
/= Text
y then () -> Maybe ()
forall a. a -> Maybe a
Just () else Maybe ()
forall a. Maybe a
Nothing
        ()
_ <- Text -> Maybe ()
colUnboxedNumeric Text
x
        ()
_ <- Text -> Maybe ()
colUnboxedNumeric Text
y
        MomentPlan -> Maybe MomentPlan
forall a. a -> Maybe a
forall (f :: * -> *) a. Applicative f => a -> f a
pure
            MomentPlan
                { mpColX :: Text
mpColX = Text
x
                , mpColY :: Text
mpColY = Text
y
                , mpNName :: Text
mpNName = Text
nName
                , mpSxName :: Text
mpSxName = Text
sxName
                , mpSyName :: Text
mpSyName = Text
syName
                , mpSxxName :: Text
mpSxxName = Text
sxxName
                , mpSyyName :: Text
mpSyyName = Text
syyName
                , mpSxyName :: Text
mpSxyName = Text
sxyName
                }
  where
    colUnboxedNumeric :: Text -> Maybe ()
colUnboxedNumeric Text
name = case Text -> DataFrame -> Maybe Column
getColumn Text
name (GroupedDataFrame -> DataFrame
fullDataframe GroupedDataFrame
gdf) of
        Just Column
c | Column -> Bool
isUnboxedNumeric Column
c -> () -> Maybe ()
forall a. a -> Maybe a
Just ()
        Maybe Column
_ -> Maybe ()
forall a. Maybe a
Nothing

-- | The output role each named aggregation plays in the moment shape.
data Role
    = RoleN
    | RoleLin T.Text
    | RoleProd T.Text T.Text
    deriving (Role -> Role -> Bool
(Role -> Role -> Bool) -> (Role -> Role -> Bool) -> Eq Role
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Role -> Role -> Bool
== :: Role -> Role -> Bool
$c/= :: Role -> Role -> Bool
/= :: Role -> Role -> Bool
Eq, Eq Role
Eq Role =>
(Role -> Role -> Ordering)
-> (Role -> Role -> Bool)
-> (Role -> Role -> Bool)
-> (Role -> Role -> Bool)
-> (Role -> Role -> Bool)
-> (Role -> Role -> Role)
-> (Role -> Role -> Role)
-> Ord Role
Role -> Role -> Bool
Role -> Role -> Ordering
Role -> Role -> Role
forall a.
Eq a =>
(a -> a -> Ordering)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> a)
-> (a -> a -> a)
-> Ord a
$ccompare :: Role -> Role -> Ordering
compare :: Role -> Role -> Ordering
$c< :: Role -> Role -> Bool
< :: Role -> Role -> Bool
$c<= :: Role -> Role -> Bool
<= :: Role -> Role -> Bool
$c> :: Role -> Role -> Bool
> :: Role -> Role -> Bool
$c>= :: Role -> Role -> Bool
>= :: Role -> Role -> Bool
$cmax :: Role -> Role -> Role
max :: Role -> Role -> Role
$cmin :: Role -> Role -> Role
min :: Role -> Role -> Role
Ord, Int -> Role -> ShowS
[Role] -> ShowS
Role -> String
(Int -> Role -> ShowS)
-> (Role -> String) -> ([Role] -> ShowS) -> Show Role
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Role -> ShowS
showsPrec :: Int -> Role -> ShowS
$cshow :: Role -> String
show :: Role -> String
$cshowList :: [Role] -> ShowS
showList :: [Role] -> ShowS
Show)

-- | Tag a single named aggregation with its moment role, or reject the group.
classify :: M.Map T.Text UExpr -> (T.Text, UExpr) -> Maybe (T.Text, Role)
classify :: Map Text UExpr -> (Text, UExpr) -> Maybe (Text, Role)
classify Map Text UExpr
exprs (Text
name, UExpr Expr a
expr) = case Expr a
expr of
    Agg (MergeAgg Text
"count" acc
_ acc -> b -> acc
_ acc -> acc -> acc
_ acc -> a
_) Expr b
_ -> (Text, Role) -> Maybe (Text, Role)
forall a. a -> Maybe a
Just (Text
name, Role
RoleN)
    Agg (FoldAgg Text
"sum" Maybe a
_ a -> b -> a
_) Expr b
arg -> (\Term
t -> (Text
name, Term -> Role
termRole Term
t)) (Term -> (Text, Role)) -> Maybe Term -> Maybe (Text, Role)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Map Text UExpr -> UExpr -> Maybe Term
resolveTerm Map Text UExpr
exprs (Expr b -> UExpr
forall a. Columnable a => Expr a -> UExpr
UExpr Expr b
arg)
    Expr a
_ -> Maybe (Text, Role)
forall a. Maybe a
Nothing

termRole :: Term -> Role
termRole :: Term -> Role
termRole (Lin Text
a) = Text -> Role
RoleLin Text
a
termRole (Prod Text
a Text
b) = Text -> Text -> Role
RoleProd Text
a Text
b

{- | Resolve a (sum-argument) expression to its 'Term'. Peels @toDouble@-style
unary coercions, follows a derived column to its stored expression, and
recognises a commutative product of two linear terms.
-}
resolveTerm :: M.Map T.Text UExpr -> UExpr -> Maybe Term
resolveTerm :: Map Text UExpr -> UExpr -> Maybe Term
resolveTerm Map Text UExpr
exprs = Int -> UExpr -> Maybe Term
go (Int
8 :: Int)
  where
    go :: Int -> UExpr -> Maybe Term
go Int
0 UExpr
_ = Maybe Term
forall a. Maybe a
Nothing
    go Int
fuel (UExpr Expr a
e) = case Expr a
e of
        Col Text
nm -> case Text -> Map Text UExpr -> Maybe UExpr
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup Text
nm Map Text UExpr
exprs of
            Just UExpr
ue -> Int -> UExpr -> Maybe Term
go (Int
fuel Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) UExpr
ue
            Maybe UExpr
Nothing -> Term -> Maybe Term
forall a. a -> Maybe a
Just (Text -> Term
Lin Text
nm)
        Unary op b a
_ Expr b
inner -> Int -> UExpr -> Maybe Term
go (Int
fuel Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) (Expr b -> UExpr
forall a. Columnable a => Expr a -> UExpr
UExpr Expr b
inner)
        Binary op c b a
op Expr c
l Expr b
r
            | op c b a -> Text
forall a b c. op a b c -> Text
forall (op :: * -> * -> * -> *) a b c.
BinaryOp op =>
op a b c -> Text
binaryName op c b a
op Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Text
"mult" Bool -> Bool -> Bool
&& op c b a -> Bool
forall a b c. op a b c -> Bool
forall (op :: * -> * -> * -> *) a b c.
BinaryOp op =>
op a b c -> Bool
binaryCommutative op c b a
op -> do
                Lin Text
a <- Int -> UExpr -> Maybe Term
go (Int
fuel Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) (Expr c -> UExpr
forall a. Columnable a => Expr a -> UExpr
UExpr Expr c
l)
                Lin Text
b <- Int -> UExpr -> Maybe Term
go (Int
fuel Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) (Expr b -> UExpr
forall a. Columnable a => Expr a -> UExpr
UExpr Expr b
r)
                Term -> Maybe Term
forall a. a -> Maybe a
Just (Text -> Text -> Term
sortProd Text
a Text
b)
        Expr a
_ -> Maybe Term
forall a. Maybe a
Nothing

-- | Products are unordered: store the pair sorted so @x*y@ and @y*x@ unify.
sortProd :: T.Text -> T.Text -> Term
sortProd :: Text -> Text -> Term
sortProd Text
a Text
b
    | Text
a Text -> Text -> Bool
forall a. Ord a => a -> a -> Bool
<= Text
b = Text -> Text -> Term
Prod Text
a Text
b
    | Bool
otherwise = Text -> Text -> Term
Prod Text
b Text
a

{- | From the classified roles, find the unordered pair of base columns that the
linear sums name. There must be exactly two distinct linear-sum columns.
-}
pickBaseColumns :: [(T.Text, Role)] -> Maybe (T.Text, T.Text)
pickBaseColumns :: [(Text, Role)] -> Maybe (Text, Text)
pickBaseColumns [(Text, Role)]
roles =
    case [Text]
lins of
        [Text
a, Text
b] | Text
a Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
/= Text
b -> (Text, Text) -> Maybe (Text, Text)
forall a. a -> Maybe a
Just (Text
a, Text
b)
        [Text]
_ -> Maybe (Text, Text)
forall a. Maybe a
Nothing
  where
    lins :: [Text]
lins = Map Text () -> [Text]
forall k a. Map k a -> [k]
M.keys ([(Text, ())] -> Map Text ()
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList [(Text
c, ()) | (Text
_, RoleLin Text
c) <- [(Text, Role)]
roles])

{- | The additive moment sums of two columns, each an @nGroups@-length column:
@(n, Sx, Sy, Sxx, Syy, Sxy)@.
-}
data Moments = Moments
    { Moments -> Column
mN :: Column
    , Moments -> Column
mSx :: Column
    , Moments -> Column
mSy :: Column
    , Moments -> Column
mSxx :: Column
    , Moments -> Column
mSyy :: Column
    , Moments -> Column
mSxy :: Column
    }

{- | One pass over two Double-coercible columns @x@ and @y@ filling the count and
five sums, collapsing the Q9 regression family's six folds into a single pass.
'Nothing' unless both columns are non-null unboxed Int/Double.
-}
momentScatter :: VU.Vector Int -> Int -> Column -> Column -> Maybe Moments
momentScatter :: Vector Int -> Int -> Column -> Column -> Maybe Moments
momentScatter Vector Int
g Int
nGroups Column
colX Column
colY = do
    Vector Double
xs <- Column -> Maybe (Vector Double)
scatterColumnToDouble Column
colX
    Vector Double
ys <- Column -> Maybe (Vector Double)
scatterColumnToDouble Column
colY
    let (Vector Int
cnt, Vector Double
sx, Vector Double
sy, Vector Double
sxx, Vector Double
syy, Vector Double
sxy) = Vector Int
-> Int
-> Vector Double
-> Vector Double
-> (Vector Int, Vector Double, Vector Double, Vector Double,
    Vector Double, Vector Double)
momentPass Vector Int
g Int
nGroups Vector Double
xs Vector Double
ys
    Moments -> Maybe Moments
forall a. a -> Maybe a
forall (f :: * -> *) a. Applicative f => a -> f a
pure
        Moments
            { mN :: Column
mN = Vector Int -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector Vector Int
cnt
            , mSx :: Column
mSx = Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector Vector Double
sx
            , mSy :: Column
mSy = Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector Vector Double
sy
            , mSxx :: Column
mSxx = Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector Vector Double
sxx
            , mSyy :: Column
mSyy = Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector Vector Double
syy
            , mSxy :: Column
mSxy = Vector Double -> Column
forall a. (Columnable a, Unbox a) => Vector a -> Column
fromUnboxedVector Vector Double
sxy
            }

momentPass ::
    VU.Vector Int ->
    Int ->
    VU.Vector Double ->
    VU.Vector Double ->
    ( VU.Vector Int
    , VU.Vector Double
    , VU.Vector Double
    , VU.Vector Double
    , VU.Vector Double
    , VU.Vector Double
    )
momentPass :: Vector Int
-> Int
-> Vector Double
-> Vector Double
-> (Vector Int, Vector Double, Vector Double, Vector Double,
    Vector Double, Vector Double)
momentPass Vector Int
g Int
nGroups Vector Double
xs Vector Double
ys = (forall s.
 ST
   s
   (Vector Int, Vector Double, Vector Double, Vector Double,
    Vector Double, Vector Double))
-> (Vector Int, Vector Double, Vector Double, Vector Double,
    Vector Double, Vector Double)
forall a. (forall s. ST s a) -> a
runST ((forall s.
  ST
    s
    (Vector Int, Vector Double, Vector Double, Vector Double,
     Vector Double, Vector Double))
 -> (Vector Int, Vector Double, Vector Double, Vector Double,
     Vector Double, Vector Double))
-> (forall s.
    ST
      s
      (Vector Int, Vector Double, Vector Double, Vector Double,
       Vector Double, Vector Double))
-> (Vector Int, Vector Double, Vector Double, Vector Double,
    Vector Double, Vector Double)
forall a b. (a -> b) -> a -> b
$ do
    MVector s Int
cnt <- Int -> Int -> ST s (MVector (PrimState (ST s)) Int)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Int
0 :: Int)
    MVector s Double
sx <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    MVector s Double
sy <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    MVector s Double
sxx <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    MVector s Double
syy <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    MVector s Double
sxy <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
nGroups (Double
0 :: Double)
    let n :: Int
n = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
xs
        bump :: MVector (PrimState m) a -> Int -> a -> m ()
bump MVector (PrimState m) a
arr Int
k a
d = MVector (PrimState m) a -> Int -> m a
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector (PrimState m) a
arr Int
k m a -> (a -> m ()) -> m ()
forall a b. m a -> (a -> m b) -> m b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= \a
c -> MVector (PrimState m) a -> Int -> a -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector (PrimState m) a
arr Int
k (a
c a -> a -> a
forall a. Num a => a -> a -> a
+ a
d)
        go :: Int -> ST s ()
go !Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = () -> ST s ()
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
            | Bool
otherwise = do
                let !k :: Int
k = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
g Int
i
                    !x :: Double
x = Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
xs Int
i
                    !y :: Double
y = Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
ys Int
i
                MVector (PrimState (ST s)) Int -> Int -> ST s Int
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.unsafeRead MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k ST s Int -> (Int -> ST s ()) -> ST s ()
forall a b. ST s a -> (a -> ST s b) -> ST s b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= \Int
c -> MVector (PrimState (ST s)) Int -> Int -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Int
MVector (PrimState (ST s)) Int
cnt Int
k (Int
c Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall {m :: * -> *} {a}.
(PrimMonad m, Unbox a, Num a) =>
MVector (PrimState m) a -> Int -> a -> m ()
bump MVector s Double
MVector (PrimState (ST s)) Double
sx Int
k Double
x
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall {m :: * -> *} {a}.
(PrimMonad m, Unbox a, Num a) =>
MVector (PrimState m) a -> Int -> a -> m ()
bump MVector s Double
MVector (PrimState (ST s)) Double
sy Int
k Double
y
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall {m :: * -> *} {a}.
(PrimMonad m, Unbox a, Num a) =>
MVector (PrimState m) a -> Int -> a -> m ()
bump MVector s Double
MVector (PrimState (ST s)) Double
sxx Int
k (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x)
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall {m :: * -> *} {a}.
(PrimMonad m, Unbox a, Num a) =>
MVector (PrimState m) a -> Int -> a -> m ()
bump MVector s Double
MVector (PrimState (ST s)) Double
syy Int
k (Double
y Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
y)
                MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall {m :: * -> *} {a}.
(PrimMonad m, Unbox a, Num a) =>
MVector (PrimState m) a -> Int -> a -> m ()
bump MVector s Double
MVector (PrimState (ST s)) Double
sxy Int
k (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
y)
                Int -> ST s ()
go (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    Int -> ST s ()
go Int
0
    (,,,,,)
        (Vector Int
 -> Vector Double
 -> Vector Double
 -> Vector Double
 -> Vector Double
 -> Vector Double
 -> (Vector Int, Vector Double, Vector Double, Vector Double,
     Vector Double, Vector Double))
-> ST s (Vector Int)
-> ST
     s
     (Vector Double
      -> Vector Double
      -> Vector Double
      -> Vector Double
      -> Vector Double
      -> (Vector Int, Vector Double, Vector Double, Vector Double,
          Vector Double, Vector Double))
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> MVector (PrimState (ST s)) Int -> ST s (Vector Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Int
MVector (PrimState (ST s)) Int
cnt
        ST
  s
  (Vector Double
   -> Vector Double
   -> Vector Double
   -> Vector Double
   -> Vector Double
   -> (Vector Int, Vector Double, Vector Double, Vector Double,
       Vector Double, Vector Double))
-> ST s (Vector Double)
-> ST
     s
     (Vector Double
      -> Vector Double
      -> Vector Double
      -> Vector Double
      -> (Vector Int, Vector Double, Vector Double, Vector Double,
          Vector Double, Vector Double))
forall a b. ST s (a -> b) -> ST s a -> ST s b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
sx
        ST
  s
  (Vector Double
   -> Vector Double
   -> Vector Double
   -> Vector Double
   -> (Vector Int, Vector Double, Vector Double, Vector Double,
       Vector Double, Vector Double))
-> ST s (Vector Double)
-> ST
     s
     (Vector Double
      -> Vector Double
      -> Vector Double
      -> (Vector Int, Vector Double, Vector Double, Vector Double,
          Vector Double, Vector Double))
forall a b. ST s (a -> b) -> ST s a -> ST s b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
sy
        ST
  s
  (Vector Double
   -> Vector Double
   -> Vector Double
   -> (Vector Int, Vector Double, Vector Double, Vector Double,
       Vector Double, Vector Double))
-> ST s (Vector Double)
-> ST
     s
     (Vector Double
      -> Vector Double
      -> (Vector Int, Vector Double, Vector Double, Vector Double,
          Vector Double, Vector Double))
forall a b. ST s (a -> b) -> ST s a -> ST s b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
sxx
        ST
  s
  (Vector Double
   -> Vector Double
   -> (Vector Int, Vector Double, Vector Double, Vector Double,
       Vector Double, Vector Double))
-> ST s (Vector Double)
-> ST
     s
     (Vector Double
      -> (Vector Int, Vector Double, Vector Double, Vector Double,
          Vector Double, Vector Double))
forall a b. ST s (a -> b) -> ST s a -> ST s b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
syy
        ST
  s
  (Vector Double
   -> (Vector Int, Vector Double, Vector Double, Vector Double,
       Vector Double, Vector Double))
-> ST s (Vector Double)
-> ST
     s
     (Vector Int, Vector Double, Vector Double, Vector Double,
      Vector Double, Vector Double)
forall a b. ST s (a -> b) -> ST s a -> ST s b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
sxy