{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ExplicitNamespaces #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE Strict #-}
{-# LANGUAGE TypeApplications #-}

module DataFrame.Operations.Aggregation (
    module DataFrame.Operations.Aggregation,
    groupBy,
    buildRowToGroup,
    changingPoints,
) where

import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

import Control.Exception (throw)
import DataFrame.Errors
import DataFrame.Internal.AggPlan (MomentPlan, planAgg, planMoments)
import DataFrame.Internal.Column (
    Column (..),
    TypedColumn (..),
    atIndicesStable,
 )
import DataFrame.Internal.DataFrame (
    DataFrame (..),
    GroupedDataFrame (..),
    columnNames,
    insertColumn,
 )
import DataFrame.Internal.Expression
import DataFrame.Internal.Grouping (buildRowToGroup, changingPoints, groupBy)
import DataFrame.Internal.Interpreter
import DataFrame.Internal.RowHash (computeRowHashesIO)
import DataFrame.Operations.AggregateScatter (runMomentPlan, runPlan)
import DataFrame.Operations.Core
import DataFrame.Operations.Subset
import System.IO.Unsafe (unsafePerformIO)

{- | Per-row key hash over the selected key columns. Delegates to the shared
'computeRowHashesIO' kernel, which forks over contiguous row ranges for large
frames (the hashing of a wide 1e7-row text/factor join key dominates that join)
and is bit-for-bit identical to a single sequential pass at any capability count.
-}
computeRowHashes :: [Int] -> DataFrame -> VU.Vector Int
computeRowHashes :: [Int] -> DataFrame -> Vector Int
computeRowHashes [Int]
indices DataFrame
df =
    let n :: Int
n = (Int, Int) -> Int
forall a b. (a, b) -> a
fst (DataFrame -> (Int, Int)
dimensions DataFrame
df)
        selectedCols :: [Column]
selectedCols = (Int -> Column) -> [Int] -> [Column]
forall a b. (a -> b) -> [a] -> [b]
map (DataFrame -> Vector Column
columns DataFrame
df Vector Column -> Int -> Column
forall a. Vector a -> Int -> a
V.!) [Int]
indices
     in IO (Vector Int) -> Vector Int
forall a. IO a -> a
unsafePerformIO (Int -> [Column] -> IO (Vector Int)
computeRowHashesIO Int
n [Column]
selectedCols)
{-# NOINLINE computeRowHashes #-}

{- | Aggregate a grouped dataframe using the expressions given.
All ungrouped columns will be dropped.
-}
aggregate :: [NamedExpr] -> GroupedDataFrame -> DataFrame
aggregate :: [NamedExpr] -> GroupedDataFrame -> DataFrame
aggregate [NamedExpr]
aggs gdf :: GroupedDataFrame
gdf@(Grouped DataFrame
df [Text]
groupingColumns Vector Int
valIndices Vector Int
offs Vector Int
rowToGroupV) =
    let
        df' :: DataFrame
df' =
            Vector Int -> DataFrame -> DataFrame
selectIndices
                ((Int -> Int) -> Vector Int -> Vector Int
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (Vector Int
valIndices Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.!) (Vector Int -> Vector Int
forall a. Unbox a => Vector a -> Vector a
VU.init Vector Int
offs))
                ([Text] -> DataFrame -> DataFrame
select [Text]
groupingColumns DataFrame
df)

        !nGroups :: Int
nGroups = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
offs Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1

        -- Fast path: a recognised reduction scatters in one unboxed pass.
        -- Anything 'planAgg' rejects keeps the existing interpreter, so the
        -- general typed + DSL aggregate API stays correct for arbitrary
        -- expressions.
        f :: NamedExpr -> DataFrame -> DataFrame
f ne :: NamedExpr
ne@(Text
name, UExpr
uexpr) DataFrame
d =
            let value :: Column
value = case GroupedDataFrame -> UExpr -> Maybe AggPlan
planAgg GroupedDataFrame
gdf UExpr
uexpr of
                    Just AggPlan
plan -> GroupedDataFrame -> Vector Int -> Int -> AggPlan -> Column
runPlan GroupedDataFrame
gdf Vector Int
rowToGroupV Int
nGroups AggPlan
plan
                    Maybe AggPlan
Nothing -> GroupedDataFrame -> NamedExpr -> Column
interpretNamed GroupedDataFrame
gdf NamedExpr
ne
             in Text -> Column -> DataFrame -> DataFrame
insertColumn Text
name Column
value DataFrame
d

        -- Fused fast path: the Q9 regression family (count + five moment sums
        -- of two base columns) becomes one scatter over the base columns,
        -- dropping the derived product columns and the six separate folds.
        -- 'planMoments' returns 'Nothing' on any other set, falling back below.
        fusedMoments :: Maybe [(Text, Column)]
fusedMoments = do
            MomentPlan
mp <- GroupedDataFrame -> [NamedExpr] -> Maybe MomentPlan
planMoments GroupedDataFrame
gdf [NamedExpr]
aggs :: Maybe MomentPlan
            GroupedDataFrame -> Int -> MomentPlan -> Maybe [(Text, Column)]
runMomentPlan GroupedDataFrame
gdf Int
nGroups MomentPlan
mp
     in
        case Maybe [(Text, Column)]
fusedMoments of
            Just [(Text, Column)]
cols -> ((Text, Column) -> DataFrame -> DataFrame)
-> [(Text, Column)] -> DataFrame -> DataFrame
forall a.
(a -> DataFrame -> DataFrame) -> [a] -> DataFrame -> DataFrame
fold ((Text -> Column -> DataFrame -> DataFrame)
-> (Text, Column) -> DataFrame -> DataFrame
forall a b c. (a -> b -> c) -> (a, b) -> c
uncurry Text -> Column -> DataFrame -> DataFrame
insertColumn) [(Text, Column)]
cols DataFrame
df'
            Maybe [(Text, Column)]
Nothing -> (NamedExpr -> DataFrame -> DataFrame)
-> [NamedExpr] -> DataFrame -> DataFrame
forall a.
(a -> DataFrame -> DataFrame) -> [a] -> DataFrame -> DataFrame
fold NamedExpr -> DataFrame -> DataFrame
f [NamedExpr]
aggs DataFrame
df'

-- | The fall-back path: evaluate one named aggregation via the interpreter.
interpretNamed :: GroupedDataFrame -> NamedExpr -> Column
interpretNamed :: GroupedDataFrame -> NamedExpr -> Column
interpretNamed GroupedDataFrame
gdf (Text
_, UExpr (Expr a
expr :: Expr a)) =
    case forall a.
Columnable a =>
GroupedDataFrame
-> Expr a -> Either DataFrameException (AggregationResult a)
interpretAggregation @a GroupedDataFrame
gdf Expr a
expr of
        Left DataFrameException
e -> DataFrameException -> Column
forall a e. Exception e => e -> a
throw DataFrameException
e
        Right (UnAggregated Column
_) -> DataFrameException -> Column
forall a e. Exception e => e -> a
throw (DataFrameException -> Column) -> DataFrameException -> Column
forall a b. (a -> b) -> a -> b
$ Text -> DataFrameException
UnaggregatedException (String -> Text
T.pack (String -> Text) -> String -> Text
forall a b. (a -> b) -> a -> b
$ Expr a -> String
forall a. Show a => a -> String
show Expr a
expr)
        Right (Aggregated (TColumn Column
col)) -> Column
col

selectIndices :: VU.Vector Int -> DataFrame -> DataFrame
selectIndices :: Vector Int -> DataFrame -> DataFrame
selectIndices Vector Int
xs DataFrame
df =
    DataFrame
df
        { columns = V.map (atIndicesStable xs) (columns df)
        , dataframeDimensions = (VU.length xs, V.length (columns df))
        }

-- | Filter out all non-unique values in a dataframe.
distinct :: DataFrame -> DataFrame
distinct :: DataFrame -> DataFrame
distinct DataFrame
df = Vector Int -> DataFrame -> DataFrame
selectIndices ((Int -> Int) -> Vector Int -> Vector Int
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (Vector Int
indices Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.!) (Vector Int -> Vector Int
forall a. Unbox a => Vector a -> Vector a
VU.init Vector Int
os)) DataFrame
df
  where
    (Grouped DataFrame
_ [Text]
_ Vector Int
indices Vector Int
os Vector Int
_rtg) = [Text] -> DataFrame -> GroupedDataFrame
groupBy (DataFrame -> [Text]
columnNames DataFrame
df) DataFrame
df