horde-ad-0.3.0.0: src/HordeAd/Core/AstTools.hs
{-# LANGUAGE LambdaCase, ViewPatterns #-}
{-# OPTIONS_GHC -fplugin GHC.TypeLits.KnownNat.Solver #-}
-- | An assortment of operations working on AST of the code to be differentiated
-- or the code resulting from differentiation.
module HordeAd.Core.AstTools
( -- * Full tensor kind derivation
ftkAst, stkAstX
-- * Variable occurrence detection
, varInAst, varInIxS, varNameInAst, varNameInIxS
-- * Tools related to sharing
, astIsSmall, ixIsSmall, astLetDown, astVar, astShare
-- * Odds and ends
, bounds, intBounds
, pattern AstConvUpSFromK, pattern AstConvUp, AstConvUpMaybe(..)
, convDownMaybe, convUpMaybe
, setTotalSharing
) where
import Prelude
import Control.Exception.Assert.Sugar
import Data.Bifunctor (second)
import Data.IORef
import Data.Maybe (fromMaybe)
import Data.Proxy (Proxy (Proxy))
import Data.Type.Equality (gcastWith, testEquality, (:~:) (Refl))
import Data.Vector.Generic qualified as V
import System.IO.Unsafe (unsafePerformIO)
import Type.Reflection (Typeable, typeRep)
import Data.Array.Nested (type (++))
import Data.Array.Nested qualified as Nested
import Data.Array.Nested.Lemmas
import Data.Array.Nested.Mixed.Shape
import Data.Array.Nested.Ranked.Shape
import Data.Array.Nested.Shaped.Shape
import Data.Array.Nested.Types (fromSNat', snatPlus, unsafeCoerceRefl)
import HordeAd.Core.Ast
import HordeAd.Core.Conversion
import HordeAd.Core.TensorKind
import HordeAd.Core.Types
-- * Full tensor kind derivation
-- | This is cheap and dirty. We don't shape-check the terms and we don't
-- unify or produce (partial) results with unknowns. Instead, we investigate
-- only one path and fail if it doesn't contain enough information
-- to determine shape (which rarely happens in our AST that is based
-- on shaped tensors and always indicates an ill-typed term).
ftkAst :: forall s y ms. AstTensor ms s y -> FullShapeTK y
ftkAst t = case t of
AstPair t1 t2 -> FTKProduct (ftkAst t1) (ftkAst t2)
AstProject1 v -> case ftkAst v of
FTKProduct ftk1 _ -> ftk1
AstProject2 v -> case ftkAst v of
FTKProduct _ ftk2 -> ftk2
AstMapAccumLDer k bftk _eftk _f _df _rf acc0 _es ->
FTKProduct (ftkAst acc0) (buildFTK k bftk)
AstApply (AstLambda !_ !u) _ -> ftkAst u
AstVar var -> varNameToFTK var
AstBuild1 snat _ (_var, v) -> buildFTK snat (ftkAst v)
AstLet _ _ v -> ftkAst v
AstShare var _ -> varNameToFTK var
AstToShare v -> ftkAst v
AstPrimalPart a -> ftkAst a
AstDualPart a -> ftkAst a
AstPlainPart a -> ftkAst a
AstFromPrimal a -> ftkAst a
AstFromDual a -> ftkAst a
AstFromPlain a -> ftkAst a
AstPlusK{} -> FTKScalar
AstTimesK{} -> FTKScalar
AstN1K{} -> FTKScalar
AstR1K{} -> FTKScalar
AstR2K{} -> FTKScalar
AstI2K{} -> FTKScalar
AstConcreteK _ -> FTKScalar
AstFloorK{} -> FTKScalar
AstFromIntegralK{} -> FTKScalar
AstCastK{} -> FTKScalar
AstArgMinK{} -> FTKScalar
AstArgMaxK{} -> FTKScalar
AstIndexK{} -> FTKScalar
AstPlusS v _ -> ftkAst v
AstTimesS v _ -> ftkAst v
AstN1S _ v -> ftkAst v
AstR1S _ v -> ftkAst v
AstR2S _ v _ -> ftkAst v
AstI2S _ v _ -> ftkAst v
AstConcreteS a -> FTKS (Nested.sshape a) FTKScalar
AstFloorS v -> case ftkAst v of
FTKS sh FTKScalar -> FTKS sh FTKScalar
AstFromIntegralS v -> case ftkAst v of
FTKS sh FTKScalar -> FTKS sh FTKScalar
AstCastS v -> case ftkAst v of
FTKS sh FTKScalar -> FTKS sh FTKScalar
AstArgMinS v -> case ftkAst v of
FTKS sh FTKScalar -> FTKS (shsInit sh) FTKScalar
AstArgMaxS v -> case ftkAst v of
FTKS sh FTKScalar -> FTKS (shsInit sh) FTKScalar
AstIndexS shn v _ix -> case ftkAst v of
FTKS _ x -> FTKS shn x
AstCondK _b v _w -> ftkAst v
AstCondS _b v _w -> ftkAst v
AstFromVectorK shm _ -> FTKS shm FTKScalar
AstFromVectorS shm l -> case V.uncons l of
Just (v, _) | FTKS shn x <- ftkAst v -> FTKS (shm `shsAppend` shn) x
Nothing -> error "ftkAst: empty vector in AstFromVectorS"
AstSumK{} -> FTKScalar
AstSumS @shm @shn shm v -> case ftkAst v of
FTKS shmshn x | SNat <- shsRank shm ->
gcastWith (unsafeCoerceRefl :: Drop (Rank shm) (shm ++ shn) :~: shn) $
FTKS (shsDrop @(Rank shm) shmshn) x
AstScatterS _ shn shp v _ -> case ftkAst v of
FTKS _ x -> FTKS (shp `shsAppend` shn) x
AstReplicateK shm _v -> FTKS shm FTKScalar
AstReplicateS shm v -> case ftkAst v of
FTKS shn x -> FTKS (shm `shsAppend` shn) x
AstGatherS shm shn _shp v _ -> case ftkAst v of
FTKS _ x -> FTKS (shm `shsAppend` shn) x
AstIotaS n@SNat -> FTKS (n :$$ ZSS) FTKScalar
AstAppendS a b -> case (ftkAst a, ftkAst b) of
(FTKS (m :$$ sh) x, FTKS (n :$$ _) _) -> FTKS (snatPlus m n :$$ sh) x
AstSliceS _ n@SNat _ a -> case ftkAst a of
FTKS (_ :$$ sh) x -> FTKS (n :$$ sh) x
AstReverseS v -> ftkAst v
AstTransposeS perm v -> case ftkAst v of
FTKS sh x -> FTKS (shsPermutePrefix perm sh) x
AstReshapeS sh2 v -> case ftkAst v of
FTKS _ x -> FTKS sh2 x
AstConvert c u -> convertFTK c $ ftkAst u
AstDot0{} -> FTKScalar
AstDot1InS sh _ _u _v -> FTKS sh FTKScalar
AstMatmul2S m@SNat _ p@SNat _u _v -> FTKS (m :$$ p :$$ ZSS) FTKScalar
AstBoolNotK{} -> FTKScalar
AstBoolNotS a -> ftkAst a
AstBoolAndK{} -> FTKScalar
AstBoolAndS a _ -> ftkAst a
AstLeqK{} -> FTKScalar
AstLeq{} -> FTKScalar
AstLeqS shb _ _ _ -> FTKS shb FTKScalar
stkAstX :: forall s x ms sh. AstTensor ms s (TKS2 sh x) -> SingletonTK x
{-# INLINE stkAstX #-}
stkAstX t = case ftkAst t of
FTKS _ x -> ftkToSTK x
-- * Variable occurrence detection
-- | We assume no variable is shared between a binding and its nested binding
-- and nobody asks about occurrences of variables that are bound.
-- This keeps the occurrence checking code simple, because we never need
-- to compare variables to any variable in the bindings.
varInAst :: AstVarId -> AstTensor ms s y -> Bool
varInAst var = \case
AstPair t1 t2 -> varInAst var t1 || varInAst var t2
AstProject1 t -> varInAst var t
AstProject2 t -> varInAst var t
AstMapAccumLDer _k _bftk _eftk _f _df _rf acc0 es ->
varInAst var acc0 || varInAst var es
AstApply t ll -> varInAstHFun var t || varInAst var ll
AstVar var2 -> var == varNameToAstVarId var2
AstBuild1 _ _ (var2, v) ->
assert (varNameToAstVarId var2 /= var) $
varInAst var v
AstLet var2 u v ->
assert (varNameToAstVarId var2 /= var) $
varInAst var u || varInAst var v
AstShare _ v -> varInAst var v
AstToShare v -> varInAst var v
AstPrimalPart a -> varInAst var a
AstDualPart a -> varInAst var a
AstPlainPart a -> varInAst var a
AstFromPrimal v -> varInAst var v
AstFromDual v -> varInAst var v
AstFromPlain v -> varInAst var v
AstPlusK t u -> varInAst var t || varInAst var u
AstTimesK t u -> varInAst var t || varInAst var u
AstN1K _ t -> varInAst var t
AstR1K _ t -> varInAst var t
AstR2K _ t u -> varInAst var t || varInAst var u
AstI2K _ t u -> varInAst var t || varInAst var u
AstConcreteK{} -> False
AstFloorK a -> varInAst var a
AstFromIntegralK t -> varInAst var t
AstCastK t -> varInAst var t
AstArgMinK t -> varInAst var t
AstArgMaxK t -> varInAst var t
AstIndexK v ix -> varInAst var v || varInIxS var ix
AstPlusS t u -> varInAst var t || varInAst var u
AstTimesS t u -> varInAst var t || varInAst var u
AstN1S _ t -> varInAst var t
AstR1S _ t -> varInAst var t
AstR2S _ t u -> varInAst var t || varInAst var u
AstI2S _ t u -> varInAst var t || varInAst var u
AstConcreteS{} -> False
AstFloorS a -> varInAst var a
AstFromIntegralS a -> varInAst var a
AstCastS t -> varInAst var t
AstArgMinS a -> varInAst var a
AstArgMaxS a -> varInAst var a
AstIndexS _ v ix -> varInAst var v || varInIxS var ix
AstCondK b v w -> varInAst var b || varInAst var v || varInAst var w
AstCondS b v w -> varInAst var b || varInAst var v || varInAst var w
AstFromVectorK _ l -> any (varInAst var) l
AstFromVectorS _ l -> any (varInAst var) l
AstSumK v -> varInAst var v
AstSumS _ v -> varInAst var v
AstScatterS _ _ _ t (vars, ix) ->
assert (all (\v -> var /= varNameToAstVarId v) vars) $
varInIxS var ix || varInAst var t
AstReplicateK _ v -> varInAst var v
AstReplicateS _ v -> varInAst var v
AstGatherS _ _ _ t (vars, ix) ->
assert (all (\v -> var /= varNameToAstVarId v) vars) $
varInIxS var ix || varInAst var t
AstIotaS{} -> False
AstAppendS v u -> varInAst var v || varInAst var u
AstSliceS _ _ _ v -> varInAst var v
AstReverseS v -> varInAst var v
AstTransposeS _perm v -> varInAst var v
AstReshapeS _ v -> varInAst var v
AstConvert _ v -> varInAst var v
AstDot0 u v -> varInAst var u || varInAst var v
AstDot1InS _ _ u v -> varInAst var u || varInAst var v
AstMatmul2S _ _ _ u v -> varInAst var u || varInAst var v
AstBoolNotK b -> varInAst var b
AstBoolNotS b -> varInAst var b
AstBoolAndK arg1 arg2 -> varInAst var arg1 || varInAst var arg2
AstBoolAndS arg1 arg2 -> varInAst var arg1 || varInAst var arg2
AstLeqK arg1 arg2 -> varInAst var arg1 || varInAst var arg2
AstLeq arg1 arg2 -> varInAst var arg1 || varInAst var arg2
AstLeqS _ _ arg1 arg2 -> varInAst var arg1 || varInAst var arg2
varInIxS :: AstVarId -> AstIxS ms sh -> Bool
varInIxS var = any (varInAst var)
varInAstHFun :: AstVarId -> AstHFun s x y -> Bool
varInAstHFun var (AstLambda var2 _) =
assert (varNameToAstVarId var2 /= var)
False -- we take advantage of the term being closed
varNameInAst :: AstVarName '(s, y) -> AstTensor ms s2 y2 -> Bool
varNameInAst var = varInAst (varNameToAstVarId var)
varNameInIxS :: AstVarName '(s, y) -> AstIxS ms sh -> Bool
varNameInIxS var = varInIxS (varNameToAstVarId var)
-- * Tools related to sharing
-- Turns off all but the most trivial cases of astIsSmall.
-- For tests only. Affects all simplification and inlining taking place
-- in parallel in the program at the time it's changed.
unsafeTotalSharingRef :: IORef Bool
{-# NOINLINE unsafeTotalSharingRef #-}
unsafeTotalSharingRef = unsafePerformIO $ newIORef False
setTotalSharing :: Bool -> IO ()
setTotalSharing = atomicWriteIORef unsafeTotalSharingRef
-- | A term requires sharing if it's too large as a term and so duplicating
-- it could affect the performance of simplification
-- or if it's too expensive when interpreted and so duplicating it
-- would increase the work done at runtime.
astIsSmall :: Bool -> AstTensor ms s y -> Bool
astIsSmall _ AstVar{} = True
astIsSmall _ AstShare{} = True
astIsSmall _ AstConcreteK{} = True
astIsSmall _ (AstConcreteS a) | fromSNat' (Nested.srank a) == 0 = True
astIsSmall lax t = unsafePerformIO $ do
unsafeTotalSharing <- readIORef unsafeTotalSharingRef
return $! if | unsafeTotalSharing -> False
| lax -> astIsSmallN 50 t > 0
| otherwise -> astIsSmallN 20 t > 0
-- The cases with n <= 20 are usually good redex candidates,
-- so we expose them, but only if they are not burried too deeply.
-- Some of these constructors change tensor metadata into
-- a non-canonical form, which sometimes incurs the cost of converting
-- the vector to canonical form. The cost can be shared only
-- when the constructor is not the root of the shared term,
-- so when inlining (the False argument) we share them
-- unless they are at the root of the term tree.
astIsSmallN :: Int -> AstTensor ms s y -> Int
astIsSmallN n _ | n <= 0 = 0
astIsSmallN n t0 = case t0 of
AstPair t1 t2 -> astIsSmallN (astIsSmallN (n - 1) t1) t2
AstProject1 t -> astIsSmallN (n - 1) t
AstProject2 t -> astIsSmallN (n - 1) t
AstVar{} -> n
AstShare{} -> n
AstPrimalPart v -> astIsSmallN (n - 1) v
AstDualPart v -> astIsSmallN (n - 1) v
AstPlainPart v -> astIsSmallN (n - 1) v
AstFromPrimal v -> astIsSmallN (n - 1) v
AstFromDual v -> astIsSmallN (n - 1) v
AstFromPlain v -> astIsSmallN (n - 1) v
AstPlusK AstConcreteK{} AstVar{} -> n - 1 -- likely index offset
AstTimesK AstConcreteK{} AstVar{} -> n - 1 -- likely index manipulation
AstConcreteK{} -> n -- small terms with zero interpretation cost;
AstConcreteS{} -> n -- the physical arrays are shared on GHC heap
-- This often appears from user writing (-1), often reduces away
-- and it has only one argument.
AstN1K NegateOp v -> astIsSmallN (n - 1) v
AstCondK b u v -> astIsSmallN (astIsSmallN (astIsSmallN (n - 1) b) u) v
AstCondS b u v -> astIsSmallN (astIsSmallN (astIsSmallN (n - 1) b) u) v
-- This is a really good redex, often nested, executed as a metadata change,
-- but not completely free, hence non-zero cost.
AstReplicateK _ v -> astIsSmallN (n - 1) v
AstReplicateS _ v -> astIsSmallN (n - 1) v
AstIotaS{} -> n
AstSliceS _ _ _ v ->
if n <= 20 then 0 else astIsSmallN (n - 1) v -- executed as metadata change
AstReverseS v ->
astIsSmallN (n - 1) v -- executed as a cheap metadata change
AstTransposeS _perm v ->
if n <= 20 then 0 else astIsSmallN (n - 1) v -- executed as metadata change
AstConvert _ v -> astIsSmallN (n - 1) v
AstBoolNotK v -> astIsSmallN (n - 1) v
AstBoolAndK u v -> astIsSmallN (astIsSmallN (n - 1) u) v
AstLeqK u v -> astIsSmallN (astIsSmallN (n - 1) u) v
AstLeq u v -> astIsSmallN (astIsSmallN (n - 1) u) v
_ -> 0
ixIsSmall :: AstIxS ms sh -> Bool
ixIsSmall = all (astIsSmall True)
-- | Try to limit the scope of the let cheaply.
--
-- Note that this doesn't inline lets into indexes, so it can be safely
-- performed on user code without the risk of removing the workaround lets for
-- big non-constant values in indexes that prevent the loss of sharing occurring
-- when differentiating indexing. gathers and scatters.
astLetDown :: forall y z s s2. KnownSpan s2
=> AstVarName '(s, y) -> AstTensor AstMethodLet s y
-> AstTensor AstMethodLet s2 z
-> AstTensor AstMethodLet s2 z
astLetDown var u v@(AstVar var2) =
if varNameToAstVarId var2 == varNameToAstVarId var
then case testEquality var var2 of
Just Refl -> u
_ -> error "astLetDown: wrong variable types at AstVar"
else v
astLetDown var u v = case v of
-- Normaly the type bounds pair nesting, so the check is cheap.
AstPair t1 t2 ->
if | not (varNameInAst var t1) -> AstPair t1 (astLetDown var u t2)
| not (varNameInAst var t2) -> AstPair (astLetDown var u t1) t2
| otherwise -> AstLet var u v
AstProject1 v2 -> AstProject1 (astLetDown var u v2)
AstProject2 v2 -> AstProject2 (astLetDown var u v2)
-- Plausibly, accumulators are small and mapAccums are rare,
-- so the check is cheap.
AstMapAccumLDer k bftk eftk f df rf acc0 es ->
if varNameInAst var acc0
then AstLet var u v
else AstMapAccumLDer k bftk eftk f df rf acc0 (astLetDown var u es)
AstApply f t -> AstApply f (astLetDown var u t)
-- handled above: AstVar
AstBuild1 k stk (var2, v2) ->
let !v3 = astLetDown var u v2
in AstBuild1 k stk (var2, v3)
AstLet{} -> AstLet var u v
AstPrimalPart v2 -> AstPrimalPart (astLetDown var u v2)
AstDualPart v2 -> AstDualPart (astLetDown var u v2)
AstPlainPart v2 -> AstPlainPart (astLetDown var u v2)
AstFromPrimal v2 -> fromPrimal (astLetDown var u v2)
AstFromDual v2 -> fromDual (astLetDown var u v2)
AstFromPlain v2 -> fromPlain (astLetDown var u v2)
AstPlusK{} -> AstLet var u v
AstTimesK{} -> AstLet var u v
AstN1K op u2 -> AstN1K op (astLetDown var u u2)
AstR1K op u2 -> AstR1K op (astLetDown var u u2)
AstR2K{} -> AstLet var u v
AstI2K{} -> AstLet var u v
AstConcreteK{} -> v
AstFloorK a -> AstFloorK (astLetDown var u a)
AstFromIntegralK v2 -> AstFromIntegralK (astLetDown var u v2)
AstCastK v2 -> AstCastK (astLetDown var u v2)
AstArgMinK v2 -> AstArgMinK (astLetDown var u v2)
AstArgMaxK v2 -> AstArgMaxK (astLetDown var u v2)
AstIndexK v2 ix ->
if varNameInIxS var ix
then AstLet var u v
else AstIndexK (astLetDown var u v2) ix
AstPlusS{} -> AstLet var u v
AstTimesS{} -> AstLet var u v
AstN1S op u2 -> AstN1S op (astLetDown var u u2)
AstR1S op u2 -> AstR1S op (astLetDown var u u2)
AstR2S{} -> AstLet var u v
AstI2S{} -> AstLet var u v
AstConcreteS{} -> v
AstFloorS a -> AstFloorS (astLetDown var u a)
AstFromIntegralS v2 -> AstFromIntegralS (astLetDown var u v2)
AstCastS v2 -> AstCastS (astLetDown var u v2)
AstArgMinS a -> AstArgMinS (astLetDown var u a)
AstArgMaxS a -> AstArgMaxS (astLetDown var u a)
-- In these three, index terms are usually small, so the check is cheap.
-- Also, this undoes precisely the pushing of the lets up that rules
-- for these three perform when simplifying. Note that we never push
-- lets down indexes, which is important for user workarounds
-- and is a legal inlining only when the iteration is trivial and so
-- the number of dynamic occurrences of the inlined variable is trivial.
AstIndexS shn v2 ix ->
if varNameInIxS var ix
then AstLet var u v
else AstIndexS shn (astLetDown var u v2) ix
AstCondK{} -> AstLet var u v
AstCondS{} -> AstLet var u v
AstFromVectorK{} -> AstLet var u v
AstFromVectorS{} -> AstLet var u v
AstSumK v2 -> AstSumK (astLetDown var u v2)
AstSumS shm v2 -> AstSumS shm (astLetDown var u v2)
AstScatterS shm shn shp v2 (vars, ix) ->
if varNameInIxS var ix
then AstLet var u v
else AstScatterS shm shn shp (astLetDown var u v2) (vars, ix)
AstReplicateK shm v2 -> AstReplicateK shm (astLetDown var u v2)
AstReplicateS shm v2 -> AstReplicateS shm (astLetDown var u v2)
AstGatherS shm shn shp v2 (vars, ix) ->
if varNameInIxS var ix
then AstLet var u v
else AstGatherS shm shn shp (astLetDown var u v2) (vars, ix)
AstIotaS{} -> v
AstAppendS{} -> AstLet var u v
AstSliceS i n k v2 -> AstSliceS i n k (astLetDown var u v2)
AstReverseS v2 -> AstReverseS (astLetDown var u v2)
AstTransposeS perm v2 -> AstTransposeS perm (astLetDown var u v2)
AstReshapeS sh v2 -> AstReshapeS sh (astLetDown var u v2)
AstConvert c v2 -> AstConvert c (astLetDown var u v2)
AstDot0{} -> AstLet var u v
AstDot1InS{} -> AstLet var u v
AstMatmul2S{} -> AstLet var u v
AstBoolNotK arg -> AstBoolNotK (astLetDown var u arg)
AstBoolNotS arg -> AstBoolNotS (astLetDown var u arg)
AstBoolAndK{} -> AstLet var u v
AstBoolAndS{} -> AstLet var u v
AstLeqK{} -> AstLet var u v
AstLeq{} -> AstLet var u v
AstLeqS{} -> AstLet var u v
astVar :: AstVarName '(s, y) -> AstTensor ms s y
astVar (AstVarName _ (FtkAndBoundsBounds lb ub)) | lb == ub =
AstConcreteK lb
astVar var = AstVar var
astShare :: AstVarName '(s, y) -> AstTensor AstMethodShare s y
-> AstTensor AstMethodShare s y
astShare (AstVarName _ (FtkAndBoundsBounds lb ub)) _ | lb == ub =
AstConcreteK lb
astShare var t = AstShare var t
-- * Odds and ends
bounds :: forall r s ms. (Typeable r, KnownSpan s)
=> AstTensor ms s (TKScalar r) -> Maybe (r, r)
bounds t | SPlainSpan <- knownSpan @s
, Just Refl <- testEquality (typeRep @r) (typeRep @Int) = intBounds t
bounds _ = Nothing
-- An approximation: lower and upper bound.
-- TODO: extend, e.g., to general quot and rem.
intBounds :: AstTensor ms PlainSpan (TKScalar Int) -> Maybe (Int, Int)
intBounds (AstConcreteK u) = Just (u, u)
intBounds (AstApply (AstLambda _ u) _) = intBounds u
intBounds (AstVar var) = varNameToBounds var
intBounds (AstLet _ _ u) = intBounds u -- TODO: substitute?
intBounds (AstShare var _) = varNameToBounds var
intBounds (AstToShare u) = intBounds u
intBounds (AstPlusK u v) = do
(u1, u2) <- intBounds u
(v1, v2) <- intBounds v
pure (u1 + v1, u2 + v2)
intBounds (AstN1K NegateOp u) = do
(u1, u2) <- intBounds u
pure (- u2, - u1)
intBounds (AstTimesK u v) = case (intBounds u, intBounds v) of
(Nothing, Nothing) -> Nothing
(mu, mv) -> -- multiplication by zero makes even one operand enough;
-- TODO: but this is, in theory, unsound, because the below
-- are not real infinities
let (u1, u2) = fromMaybe (-1000000000, 1000000000) mu
(v1, v2) = fromMaybe (-1000000000, 1000000000) mv
l = [u1 * v1, u1 * v2, u2 * v1, u2 * v2]
(lb, ub) = (minimum l, maximum l)
in if lb == -1000000000 && ub == 1000000000
then Nothing
else Just (lb, ub)
intBounds (AstI2K QuotOp u (AstConcreteK v)) | v > 0 = do -- a common case
(u1, u2) <- intBounds u
pure (u1 `quotH` v, u2 `quotH` v)
intBounds (AstI2K RemOp u (AstConcreteK v)) | v > 0 = do
(u1, u2) <- intBounds u
pure $ if | u1 >= 0 -> (0, min u2 (v - 1)) -- very crude
| u2 <= 0 -> (max u1 (- v + 1), 0)
| otherwise -> (- v + 1, v - 1)
intBounds (AstCondK _b u v) = do
(u1, u2) <- intBounds u
(v1, v2) <- intBounds v
pure (min u1 v1, max u2 v2)
intBounds _ = Nothing
pattern AstConvUpSFromK :: forall r sh s ms. () => sh ~ '[]
=> AstTensor ms s (TKScalar r)
-> AstTensor ms s (TKS sh r)
pattern AstConvUpSFromK t <- (matchAstConvUpSFromK -> Just (Refl, t))
matchAstConvUpSFromK :: AstTensor ms s (TKS sh r)
-> Maybe ( sh :~: '[]
, AstTensor ms s (TKScalar r) )
matchAstConvUpSFromK = \case
AstConvert c t
| FTKScalar @ry <- ftkAst t
, FTKS ZSS (FTKScalar @r) <- convertFTK c (ftkAst t)
, Just Refl <- testEquality (typeRep @ry) (typeRep @r) ->
Just (Refl, t)
AstConvert c t
| FTKR ZSR (FTKScalar @ry) <- ftkAst t
, FTKS ZSS (FTKScalar @r) <- convertFTK c (ftkAst t)
, Just Refl <- testEquality (typeRep @ry) (typeRep @r) ->
Just (Refl, AstConvert (ConvCmp ConvX0 ConvRX) t)
AstConvert c t
| FTKX ZSX (FTKScalar @ry) <- ftkAst t
, FTKS ZSS (FTKScalar @r) <- convertFTK c (ftkAst t)
, Just Refl <- testEquality (typeRep @ry) (typeRep @r) ->
Just (Refl, AstConvert ConvX0 t)
AstConcreteS a | ZSS <- Nested.sshape a ->
Just (Refl, AstConcreteK $ Nested.sunScalar a)
AstFromPrimal t -> second AstFromPrimal <$> matchAstConvUpSFromK t
AstFromDual t -> second AstFromDual <$> matchAstConvUpSFromK t
AstFromPlain t -> second AstFromPlain <$> matchAstConvUpSFromK t
_ -> Nothing
-- TODO: simplify this monstrosity, if possible
pattern AstConvUp :: forall {z1} {ms1} {s1}.
forall y z ms s. (z ~ z1, ms ~ ms1, s ~ s1)
=> TKConversion y z -> FullShapeTK z -> AstTensor ms s y
-> AstTensor ms1 s1 z1
pattern AstConvUp c zftk a <- (matchAstConvUp -> AstConvUpJust c zftk a)
type role AstConvUpMaybe nominal nominal nominal
data AstConvUpMaybe z ms s =
forall y.
AstConvUpJust (TKConversion y z) (FullShapeTK z) (AstTensor ms s y)
| AstConvUpNothing
matchAstConvUp :: AstTensor ms s z -> AstConvUpMaybe z ms s
matchAstConvUp = \case
AstConvert c t
| FTKR ZSR (FTKScalar @ry) <- ftkAst t
, let zftk = convertFTK c (ftkAst t)
, FTKS ZSS (FTKScalar @r) <- zftk
, Just Refl <- testEquality (typeRep @ry) (typeRep @r) ->
AstConvUpJust (ConvCmp ConvXS (Conv0X STKScalar))
zftk
(AstConvert (ConvCmp ConvX0 ConvRX) t)
AstConvert c t
| FTKX ZSX (FTKScalar @ry) <- ftkAst t
, let zftk = convertFTK c (ftkAst t)
, FTKS ZSS (FTKScalar @r) <- zftk
, Just Refl <- testEquality (typeRep @ry) (typeRep @r) ->
AstConvUpJust (ConvCmp ConvXS (Conv0X STKScalar))
zftk
(AstConvert ConvX0 t)
AstConvert c t
| FTKS ZSS (FTKScalar @r) <- ftkAst t
, let zftk = convertFTK c (ftkAst t)
, FTKR ZSR (FTKScalar @ry) <- zftk
, Just Refl <- testEquality (typeRep @ry) (typeRep @r) ->
AstConvUpJust (ConvCmp (ConvXR STKScalar) (Conv0X STKScalar))
zftk
(AstConvert (ConvCmp ConvX0 ConvSX) t)
AstConvert c t
| FTKS ZSS (FTKScalar @r) <- ftkAst t
, let zftk = convertFTK c (ftkAst t)
, FTKX ZSX (FTKScalar @ry) <- zftk
, Just Refl <- testEquality (typeRep @ry) (typeRep @r) ->
AstConvUpJust (Conv0X STKScalar)
zftk
(AstConvert (ConvCmp ConvX0 ConvSX) t)
AstConvert c t ->
let yftk = ftkAst t
zftk = convertFTK c yftk
in case convUpMaybe (ftkAst t) zftk of
Just c2 -> AstConvUpJust c2 zftk t
Nothing -> AstConvUpNothing
AstConcreteS a | ZSS <- Nested.sshape a ->
AstConvUpJust (ConvCmp ConvXS (Conv0X STKScalar))
(FTKS ZSS FTKScalar)
(AstConcreteK $ Nested.sunScalar a)
AstFromPrimal t -> case matchAstConvUp t of
AstConvUpJust c zftk u -> AstConvUpJust c zftk (AstFromPrimal u)
AstConvUpNothing -> AstConvUpNothing
AstFromDual t -> case matchAstConvUp t of
AstConvUpJust c zftk u -> AstConvUpJust c zftk (AstFromDual u)
AstConvUpNothing -> AstConvUpNothing
AstFromPlain t -> case matchAstConvUp t of
AstConvUpJust c zftk u -> AstConvUpJust c zftk (AstFromPlain u)
AstConvUpNothing -> AstConvUpNothing
_ -> AstConvUpNothing
convDownMaybe :: FullShapeTK y0 -> SingletonTK z0 -> Maybe (TKConversion y0 z0)
convDownMaybe = \cases
yftk0 zstk0 | Just Refl <- sameSTK (ftkToSTK yftk0) zstk0 -> Just ConvId
(FTKS ZSS (FTKScalar @ry)) (STKScalar @rz)
| Just Refl <- testEquality (typeRep @ry) (typeRep @rz) ->
Just $ convCmp ConvX0 ConvSX
(FTKR ZSR (FTKScalar @ry)) (STKScalar @rz)
| Just Refl <- testEquality (typeRep @ry) (typeRep @rz) ->
Just $ convCmp ConvX0 ConvRX
(FTKX ZSX (FTKScalar @ry)) (STKScalar @rz)
| Just Refl <- testEquality (typeRep @ry) (typeRep @rz) ->
Just ConvX0
(FTKR rsh rx) (STKS @sh sh@(_ :$$ _) x)
| Just Refl <- sameSTK x (ftkToSTK rx)
, Just Refl <- testEquality (shsRank sh) (shrRank rsh)
, Refl <- lemRankReplicate (Proxy @(Rank sh)) ->
Just $ convCmp (ConvXS' (FTKS sh rx)) ConvRX
(FTKX xsh xx) (STKS sh@(_ :$$ _) x)
| Just Refl <- sameSTK x (ftkToSTK xx)
, Just Refl <- testEquality (shsRank sh) (shxRank xsh)
, Refl <- lemRankMapJust sh ->
Just $ ConvXS' (FTKS sh xx)
(FTKProduct yftk1 yftk2) (STKProduct zstk1 zstk2) -> do
c1 <- convDownMaybe yftk1 zstk1
c2 <- convDownMaybe yftk2 zstk2
Just $ ConvT2 c1 c2
(FTKProduct (FTKS sh' yftk1) (FTKS sh'' yftk2)) (STKS sh (STKProduct ystk1
ystk2))
| Just Refl <- testEquality sh sh'
, Just Refl <- testEquality sh sh''
, Just Refl <- sameSTK ystk1 (ftkToSTK yftk1)
, Just Refl <- sameSTK ystk2 (ftkToSTK yftk2) ->
Just
$ convCmp
ConvXS
(convCmp
(ConvZip ystk1 ystk2)
(ConvT2 ConvSX ConvSX))
_ _ -> Nothing
convUpMaybe :: FullShapeTK y0 -> FullShapeTK z0 -> Maybe (TKConversion y0 z0)
convUpMaybe = \cases
yftk0 zftk0 | Just Refl <- matchingFTK yftk0 zftk0 -> Just ConvId
(FTKScalar @rz) (FTKS ZSS (FTKScalar @ry))
| Just Refl <- testEquality (typeRep @ry) (typeRep @rz) ->
Just $ convCmp ConvXS (Conv0X STKScalar)
(FTKScalar @rz) (FTKR ZSR (FTKScalar @ry))
| Just Refl <- testEquality (typeRep @ry) (typeRep @rz) ->
Just $ convCmp (ConvXR STKScalar) (Conv0X STKScalar)
(FTKScalar @rz) (FTKX ZSX (FTKScalar @ry))
| Just Refl <- testEquality (typeRep @ry) (typeRep @rz) ->
Just $ Conv0X STKScalar
(FTKS sh@(_ :$$ _) x) (FTKR rsh rx)
| Just Refl <- matchingFTK x rx
, Just Refl <- testEquality (shsRank sh) (shrRank rsh)
, Refl <- lemRankMapJust sh ->
Just $ convCmp (ConvXR (ftkToSTK x)) ConvSX
(FTKS sh@(_ :$$ _) x) zftk0@(FTKX xsh xx)
| Just Refl <- matchingFTK x xx
, Just Refl <- testEquality (shsRank sh) (shxRank xsh)
, Refl <- lemRankMapJust sh ->
Just $ convCmp (ConvXX' zftk0) ConvSX
(FTKProduct yftk1 yftk2) (FTKProduct zftk1 zftk2) -> do
c1 <- convUpMaybe yftk1 zftk1
c2 <- convUpMaybe yftk2 zftk2
Just $ ConvT2 c1 c2
(FTKS sh (FTKProduct yftk1 yftk2)) (FTKProduct (FTKS sh' yftk1')
(FTKS sh'' yftk2'))
| Just Refl <- testEquality sh sh'
, Just Refl <- testEquality sh sh''
, Just Refl <- matchingFTK yftk1 yftk1'
, Just Refl <- matchingFTK yftk2 yftk2' ->
Just
$ convCmp
(ConvT2 ConvXS ConvXS)
(convCmp
(ConvUnzip (ftkToSTK yftk1) (ftkToSTK yftk2))
ConvSX)
_ _ -> Nothing