{-# LANGUAGE DataKinds #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE TypeOperators #-}
{-# LANGUAGE UndecidableInstances #-}

-- | Affine indexing morphisms between harpie-style shapes.
--
-- An @Affine p q@ is an affine map @q = Λ·p + v@ where @Λ@ is a
-- @len q × len p@ matrix of natural numbers and @v@ is a @len q@ offset
-- vector.  In-bounds is checked at construction time.
--
-- These are the change-of-basis joints that let @Mat@ constructions flow
-- between carriers: flatten, shapen, axis swap, and arbitrary permutations
-- are all affine maps, and harpie consumes them as stride/index maps.
module Circuit.Mat.Affine
  ( -- * Type
    Affine (..),

    -- * Smart constructor
    affine,

    -- * Category structure
    identityAffine,
    composeAffine,
    applyAffine,

    -- * Canonical maps
    flattenAffine,
    swapAxesAffine,
    permuteAxes,

    -- * Harpie seam
    toIndexMap,
    toHarpieBackpermute,
  )
where

import Data.Bool (bool)
import Data.List (lookup, sort)
import GHC.TypeNats (KnownNat, Nat)
import GHC.TypeNats qualified as TN
import Harpie.Fixed qualified as F
import Harpie.Shape (Fins (..), KnownNats, valuesOf)
import Prelude

-- $setup
-- >>> :set -XDataKinds
-- >>> :set -XTypeApplications
-- >>> import Circuit.Mat.Affine

-- | Type-level product of a shape list.
type family SizeOf (p :: [Nat]) :: Nat where
  SizeOf '[] = 1
  SizeOf (n ': ns) = n TN.* SizeOf ns

-- | An affine map from shape @p@ to shape @q@.
--
-- The phantom types carry the shape; the fields carry the concrete matrix and
-- offset.  Use 'affine' to construct a validated value.
data Affine (p :: [Nat]) (q :: [Nat]) = Affine
  { forall (p :: [Nat]) (q :: [Nat]). Affine p q -> [[Int]]
lambda :: [[Int]],
    forall (p :: [Nat]) (q :: [Nat]). Affine p q -> [Int]
offset :: [Int]
  }
  deriving (Affine p q -> Affine p q -> Bool
(Affine p q -> Affine p q -> Bool)
-> (Affine p q -> Affine p q -> Bool) -> Eq (Affine p q)
forall (p :: [Nat]) (q :: [Nat]). Affine p q -> Affine p q -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall (p :: [Nat]) (q :: [Nat]). Affine p q -> Affine p q -> Bool
== :: Affine p q -> Affine p q -> Bool
$c/= :: forall (p :: [Nat]) (q :: [Nat]). Affine p q -> Affine p q -> Bool
/= :: Affine p q -> Affine p q -> Bool
Eq, Int -> Affine p q -> ShowS
[Affine p q] -> ShowS
Affine p q -> String
(Int -> Affine p q -> ShowS)
-> (Affine p q -> String)
-> ([Affine p q] -> ShowS)
-> Show (Affine p q)
forall (p :: [Nat]) (q :: [Nat]). Int -> Affine p q -> ShowS
forall (p :: [Nat]) (q :: [Nat]). [Affine p q] -> ShowS
forall (p :: [Nat]) (q :: [Nat]). Affine p q -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall (p :: [Nat]) (q :: [Nat]). Int -> Affine p q -> ShowS
showsPrec :: Int -> Affine p q -> ShowS
$cshow :: forall (p :: [Nat]) (q :: [Nat]). Affine p q -> String
show :: Affine p q -> String
$cshowList :: forall (p :: [Nat]) (q :: [Nat]). [Affine p q] -> ShowS
showList :: [Affine p q] -> ShowS
Show)

-- | Smart constructor.  Returns 'Nothing' if the dimensions do not match the
-- shape, if entries are negative, or if offsets are out of bounds.
affine ::
  forall p q.
  (KnownNats p, KnownNats q) =>
  [[Int]] ->
  [Int] ->
  Maybe (Affine p q)
affine :: forall (p :: [Nat]) (q :: [Nat]).
(KnownNats p, KnownNats q) =>
[[Int]] -> [Int] -> Maybe (Affine p q)
affine [[Int]]
lam [Int]
off
  | Bool -> Bool
not Bool
shapeOK = Maybe (Affine p q)
forall a. Maybe a
Nothing
  | ([Int] -> Bool) -> [[Int]] -> Bool
forall (t :: * -> *) a. Foldable t => (a -> Bool) -> t a -> Bool
any ((Int -> Bool) -> [Int] -> Bool
forall (t :: * -> *) a. Foldable t => (a -> Bool) -> t a -> Bool
any (Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0)) [[Int]]
lam = Maybe (Affine p q)
forall a. Maybe a
Nothing
  | (Int -> Bool) -> [Int] -> Bool
forall (t :: * -> *) a. Foldable t => (a -> Bool) -> t a -> Bool
any (Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0) [Int]
off = Maybe (Affine p q)
forall a. Maybe a
Nothing
  | Bool -> Bool
not Bool
offsetsOK = Maybe (Affine p q)
forall a. Maybe a
Nothing
  | Bool
otherwise = Affine p q -> Maybe (Affine p q)
forall a. a -> Maybe a
Just ([[Int]] -> [Int] -> Affine p q
forall (p :: [Nat]) (q :: [Nat]). [[Int]] -> [Int] -> Affine p q
Affine [[Int]]
lam [Int]
off)
  where
    sp :: [Int]
sp = forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @p
    sq :: [Int]
sq = forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @q
    shapeOK :: Bool
shapeOK =
      [[Int]] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [[Int]]
lam Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Int]
sq
        Bool -> Bool -> Bool
&& ([Int] -> Bool) -> [[Int]] -> Bool
forall (t :: * -> *) a. Foldable t => (a -> Bool) -> t a -> Bool
all ((Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Int]
sp) (Int -> Bool) -> ([Int] -> Int) -> [Int] -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length) [[Int]]
lam
    offsetsOK :: Bool
offsetsOK = [Bool] -> Bool
forall (t :: * -> *). Foldable t => t Bool -> Bool
and ((Int -> Int -> Bool) -> [Int] -> [Int] -> [Bool]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
(<) [Int]
off [Int]
sq)

-- | Identity affine map.
identityAffine :: forall p. (KnownNats p) => Affine p p
identityAffine :: forall (p :: [Nat]). KnownNats p => Affine p p
identityAffine =
  let n :: Int
n = [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length (forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @p)
   in Affine
        { lambda :: [[Int]]
lambda = [[Int -> Int -> Bool -> Int
forall a. a -> a -> Bool -> a
bool Int
0 Int
1 (Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j) | Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]],
          offset :: [Int]
offset = Int -> Int -> [Int]
forall a. Int -> a -> [a]
replicate Int
n Int
0
        }

-- | Composition of affine maps.
--
-- >>> let Just f = affine @'[2,3] @'[6] [[3,1]] [0]
-- >>> let Just g = affine @'[6] @'[2,3] [[1],[2]] [0,0]
-- >>> applyAffine (composeAffine g f) [1,2]
-- [5,10]
composeAffine ::
  forall p q r.
  (KnownNats p, KnownNats q, KnownNats r) =>
  Affine q r ->
  Affine p q ->
  Affine p r
composeAffine :: forall (p :: [Nat]) (q :: [Nat]) (r :: [Nat]).
(KnownNats p, KnownNats q, KnownNats r) =>
Affine q r -> Affine p q -> Affine p r
composeAffine (Affine [[Int]]
lam2 [Int]
off2) (Affine [[Int]]
lam1 [Int]
off1) =
  Affine
    { lambda :: [[Int]]
lambda = [[[Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [[[Int]]
lam2 [[Int]] -> Int -> [Int]
forall a. HasCallStack => [a] -> Int -> a
!! Int
i [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!! Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* [[Int]]
lam1 [[Int]] -> Int -> [Int]
forall a. HasCallStack => [a] -> Int -> a
!! Int
k [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!! Int
j | Int
k <- [Int
0 .. Int
q Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
j <- [Int
0 .. Int
p Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
i <- [Int
0 .. Int
r Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]],
      offset :: [Int]
offset = (Int -> Int -> Int) -> [Int] -> [Int] -> [Int]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Int -> Int -> Int
forall a. Num a => a -> a -> a
(+) [Int]
off2 [[Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [[[Int]]
lam2 [[Int]] -> Int -> [Int]
forall a. HasCallStack => [a] -> Int -> a
!! Int
i [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!! Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* [Int]
off1 [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!! Int
k | Int
k <- [Int
0 .. Int
q Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
i <- [Int
0 .. Int
r Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
    }
  where
    p :: Int
p = [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length (forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @p)
    q :: Int
q = [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length (forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @q)
    r :: Int
r = [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length (forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @r)

-- | Apply an affine map to an index vector.  Assumes the input is in bounds.
applyAffine :: Affine p q -> [Int] -> [Int]
applyAffine :: forall (p :: [Nat]) (q :: [Nat]). Affine p q -> [Int] -> [Int]
applyAffine (Affine [[Int]]
lam [Int]
off) [Int]
x = (Int -> Int -> Int) -> [Int] -> [Int] -> [Int]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Int -> Int -> Int
forall a. Num a => a -> a -> a
(+) [Int]
off [[Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum ((Int -> Int -> Int) -> [Int] -> [Int] -> [Int]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Int -> Int -> Int
forall a. Num a => a -> a -> a
(*) [Int]
row [Int]
x) | [Int]
row <- [[Int]]
lam]

-- | Flatten a shape to its total size.
flattenAffine ::
  forall p.
  (KnownNats p) =>
  Affine p '[SizeOf p]
flattenAffine :: forall (p :: [Nat]). KnownNats p => Affine p '[SizeOf p]
flattenAffine =
  let sp :: [Int]
sp = forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @p
      strides :: [Int]
strides = Int -> [Int] -> [Int]
forall a. Int -> [a] -> [a]
drop Int
1 ((Int -> Int -> Int) -> Int -> [Int] -> [Int]
forall a b. (a -> b -> b) -> b -> [a] -> [b]
scanr Int -> Int -> Int
forall a. Num a => a -> a -> a
(*) Int
1 [Int]
sp)
   in Affine
        { lambda :: [[Int]]
lambda = [[Int]
strides],
          offset :: [Int]
offset = [Int
0]
        }

-- | Swap the two axes of a rank-2 shape.
swapAxesAffine :: Affine '[m, n] '[n, m]
swapAxesAffine :: forall (m :: Nat) (n :: Nat). Affine '[m, n] '[n, m]
swapAxesAffine =
  Affine
    { lambda :: [[Int]]
lambda = [[Int
0, Int
1], [Int
1, Int
0]],
      offset :: [Int]
offset = [Int
0, Int
0]
    }

-- | Permute axes by a value-level permutation list.
--
-- The permutation must be a reordering of @[0 .. length p - 1]@.  No
-- type-level tracking of the output shape is attempted; the caller supplies
-- the desired shape phantom.
permuteAxes ::
  forall p q.
  (KnownNats p, KnownNats q) =>
  [Int] ->
  Maybe (Affine p q)
permuteAxes :: forall (p :: [Nat]) (q :: [Nat]).
(KnownNats p, KnownNats q) =>
[Int] -> Maybe (Affine p q)
permuteAxes [Int]
perm
  | Bool -> Bool
not Bool
ok = Maybe (Affine p q)
forall a. Maybe a
Nothing
  | Bool
otherwise = [[Int]] -> [Int] -> Maybe (Affine p q)
forall (p :: [Nat]) (q :: [Nat]).
(KnownNats p, KnownNats q) =>
[[Int]] -> [Int] -> Maybe (Affine p q)
affine [[Int]]
lam (Int -> Int -> [Int]
forall a. Int -> a -> [a]
replicate ([Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Int]
sq) Int
0)
  where
    sp :: [Int]
sp = forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @p
    sq :: [Int]
sq = forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @q
    n :: Int
n = [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Int]
sp
    m :: Int
m = [Int] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Int]
sq
    ok :: Bool
ok =
      Int
m Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
n
        Bool -> Bool -> Bool
&& [Int] -> [Int]
forall a. Ord a => [a] -> [a]
sort [Int]
perm [Int] -> [Int] -> Bool
forall a. Eq a => a -> a -> Bool
== [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        Bool -> Bool -> Bool
&& [Int]
sq [Int] -> [Int] -> Bool
forall a. Eq a => a -> a -> Bool
== (Int -> Int) -> [Int] -> [Int]
forall a b. (a -> b) -> [a] -> [b]
map ([Int]
sp [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!!) [Int]
perm
    lam :: [[Int]]
lam = [[Int -> Int -> Bool -> Int
forall a. a -> a -> Bool -> a
bool Int
0 Int
1 ([Int]
perm [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!! Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j) | Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
i <- [Int
0 .. Int
m Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]

-- | Apply the affine map to harpie 'Fins'.
toIndexMap :: Affine p q -> Fins p -> Fins q
toIndexMap :: forall (p :: [Nat]) (q :: [Nat]). Affine p q -> Fins p -> Fins q
toIndexMap Affine p q
a (UnsafeFins [Int]
x) = [Int] -> Fins q
forall {k} (s :: k). [Int] -> Fins s
UnsafeFins (Affine p q -> [Int] -> [Int]
forall (p :: [Nat]) (q :: [Nat]). Affine p q -> [Int] -> [Int]
applyAffine Affine p q
a [Int]
x)

-- | Build a harpie 'F.backpermute' from an affine map.
--
-- This is the code generator: an affine reindexing becomes a harpie
-- backpermute with an index-map derived from the affine inverse.  The map is
-- precomputed, so the result fuses with surrounding harpie operations.
--
-- Precondition: the affine map must be a bijection between the index sets
-- (true for flatten, permutation, and axis-swap).  If it is not, the lookup
-- will fail at runtime.
toHarpieBackpermute ::
  forall p q s.
  (KnownNats p, KnownNats q) =>
  Affine p q ->
  F.Array p s ->
  F.Array q s
toHarpieBackpermute :: forall (p :: [Nat]) (q :: [Nat]) s.
(KnownNats p, KnownNats q) =>
Affine p q -> Array p s -> Array q s
toHarpieBackpermute Affine p q
a Array p s
arr = (Fins q -> Fins p) -> Array p s -> Array q s
forall (s' :: [Nat]) (s :: [Nat]) a.
(KnownNats s, KnownNats s') =>
(Fins s' -> Fins s) -> Array s a -> Array s' a
F.backpermute Fins q -> Fins p
inverseFins Array p s
arr
  where
    inverseFins :: Fins q -> Fins p
inverseFins (UnsafeFins [Int]
q) =
      case [Int] -> [([Int], [Int])] -> Maybe [Int]
forall a b. Eq a => a -> [(a, b)] -> Maybe b
lookup [Int]
q [([Int], [Int])]
inverseMap of
        Just [Int]
p -> [Int] -> Fins p
forall {k} (s :: k). [Int] -> Fins s
UnsafeFins [Int]
p
        Maybe [Int]
Nothing -> String -> Fins p
forall a. HasCallStack => String -> a
error String
"toHarpieBackpermute: affine map is not a bijection on indices"
    inverseMap :: [([Int], [Int])]
inverseMap =
      [ (Affine p q -> [Int] -> [Int]
forall (p :: [Nat]) (q :: [Nat]). Affine p q -> [Int] -> [Int]
applyAffine Affine p q
a [Int]
ps, [Int]
ps)
      | [Int]
ps <- [Int] -> [[Int]]
shapeIndices (forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @p)
      ]

-- | All index vectors for a shape.
shapeIndices :: [Int] -> [[Int]]
shapeIndices :: [Int] -> [[Int]]
shapeIndices = (Int -> [Int]) -> [Int] -> [[Int]]
forall (t :: * -> *) (m :: * -> *) a b.
(Traversable t, Monad m) =>
(a -> m b) -> t a -> m (t b)
forall (m :: * -> *) a b. Monad m => (a -> m b) -> [a] -> m [b]
mapM (\Int
d -> [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1])