packages feed

mini-1.6.4.0: src/Mini/Linear/Matrix.hs

{-# LANGUAGE ImpredicativeTypes #-}

-- | Column-major matrix operations
module Mini.Linear.Matrix (
  -- * Class
  Square (
    adj,
    det,
    diagonal
  ),

  -- * Construction
  identity,
  zero,

  -- * Operations
  (#*#),
  (#*^),
  (#+#),
  (#-#),
  (^*#),
  (~*#),
  inverse,
  trace,
  transpose,
) where

import Control.Applicative (
  liftA2,
 )
import Mini.Linear.Space (
  V0 (V0),
  V1,
  V2 (V2),
  V3 (V3),
  V4 (V4),
  w,
  x,
  xy,
  xyw,
  xyz,
  xz,
  xzw,
  y,
  yz,
  yzw,
  z,
 )
import Mini.Linear.Vector (
  Vector,
  axes,
  dot,
  (^+^),
  (^-^),
  (~*^),
 )
import qualified Mini.Linear.Vector as V (
  zero,
 )
import Mini.Optics.Lens (
  Lens,
  set,
  view,
 )
import Prelude (
  Fractional,
  Num,
  fmap,
  foldr,
  id,
  negate,
  pure,
  recip,
  sum,
  ($),
  (*),
  (+),
  (-),
  (.),
  (<$),
  (<$>),
 )

-- Class

-- | The class of square matrices
class (Vector v) => Square v where
  -- | Adjoint of a matrix
  adj :: (Num a) => v (v a) -> v (v a)

  -- | Determinant of a matrix
  det :: (Num a) => v (v a) -> a

  -- | Diagonal lens
  diagonal :: Lens (v (v a)) (v (v a)) (v a) (v a)

instance Square V0 where
  adj = id
  det _ = 1
  diagonal f _ = V0 <$ f V0

instance Square V1 where
  adj = set (x . x) 1
  det = view (x . x)
  diagonal = x

instance Square V2 where
  adj m =
    V2
      (V2 (view (y . y) m) (negate $ view (x . y) m))
      (V2 (negate $ view (y . x) m) (view (x . x) m))
  det m =
    view (x . x) m * view (y . y) m
      - view (y . x) m * view (x . y) m
  diagonal f m =
    ( \v ->
        V2
          (V2 (view x v) (view (x . y) m))
          (V2 (view (y . x) m) (view y v))
    )
      <$> f (V2 (view (x . x) m) (view (y . y) m))

instance Square V3 where
  adj m =
    let d00 = det $ view yz <$> view yz m
        d01 = det $ view yz <$> view xz m
        d02 = det $ view yz <$> view xy m
        d10 = det $ view xz <$> view yz m
        d11 = det $ view xz <$> view xz m
        d12 = det $ view xz <$> view xy m
        d20 = det $ view xy <$> view yz m
        d21 = det $ view xy <$> view xz m
        d22 = det $ view xy <$> view xy m
     in V3
          (V3 d00 (negate d01) d02)
          (V3 (negate d10) d11 (negate d12))
          (V3 d20 (negate d21) d22)
  det m =
    view (x . x) m * det (view yz <$> view yz m)
      - view (y . x) m * det (view yz <$> view xz m)
      + view (z . x) m * det (view yz <$> view xy m)
  diagonal f m =
    ( \v ->
        V3
          (V3 (view x v) (view (x . y) m) (view (x . z) m))
          (V3 (view (y . x) m) (view y v) (view (y . z) m))
          (V3 (view (z . x) m) (view (z . y) m) (view z v))
    )
      <$> f (V3 (view (x . x) m) (view (y . y) m) (view (z . z) m))

instance Square V4 where
  adj m =
    let d00 = det $ view yzw <$> view yzw m
        d01 = det $ view yzw <$> view xzw m
        d02 = det $ view yzw <$> view xyw m
        d03 = det $ view yzw <$> view xyz m
        d10 = det $ view xzw <$> view yzw m
        d11 = det $ view xzw <$> view xzw m
        d12 = det $ view xzw <$> view xyw m
        d13 = det $ view xzw <$> view xyz m
        d20 = det $ view xyw <$> view yzw m
        d21 = det $ view xyw <$> view xzw m
        d22 = det $ view xyw <$> view xyw m
        d23 = det $ view xyw <$> view xyz m
        d30 = det $ view xyz <$> view yzw m
        d31 = det $ view xyz <$> view xzw m
        d32 = det $ view xyz <$> view xyw m
        d33 = det $ view xyz <$> view xyz m
     in V4
          (V4 d00 (negate d01) d02 (negate d03))
          (V4 (negate d10) d11 (negate d12) d13)
          (V4 d20 (negate d21) d22 (negate d23))
          (V4 (negate d30) d31 (negate d32) d33)
  det m =
    view (x . x) m * det (view yzw <$> view yzw m)
      - view (y . x) m * det (view yzw <$> view xzw m)
      + view (z . x) m * det (view yzw <$> view xyw m)
      - view (w . x) m * det (view yzw <$> view xyz m)
  diagonal f m =
    ( \v ->
        V4
          (V4 (view x v) (view (x . y) m) (view (x . z) m) (view (x . w) m))
          (V4 (view (y . x) m) (view y v) (view (y . z) m) (view (y . w) m))
          (V4 (view (z . x) m) (view (z . y) m) (view z v) (view (z . w) m))
          (V4 (view (w . x) m) (view (w . y) m) (view (w . z) m) (view w v))
    )
      <$> f
        (V4 (view (x . x) m) (view (y . y) m) (view (z . z) m) (view (w . w) m))

-- Construction

-- | Multiplicative identity matrix
identity :: (Square v, Num a) => v (v a)
identity = set diagonal (pure 1) zero

-- | Additive identity matrix
zero :: (Vector p, Vector q, Num a) => q (p a)
zero = pure V.zero

-- Operations

infixr 7 #*#

-- | Matrix-matrix multiplication
(#*#) :: (Vector p, Vector q, Vector r, Num a) => q (p a) -> r (q a) -> r (p a)
qp #*# rq = fmap (\q -> foldr (^+^) V.zero $ liftA2 (~*^) q qp) rq

infix 7 #*^

-- | Matrix-vector multiplication
(#*^) :: (Vector p, Vector q, Num a) => q (p a) -> q a -> p a
m #*^ v = foldr (^+^) V.zero $ liftA2 (~*^) v m

infixl 6 #+#

-- | Matrix-matrix addition
(#+#) :: (Vector p, Vector q, Num a) => q (p a) -> q (p a) -> q (p a)
(#+#) = liftA2 (^+^)

infixl 6 #-#

-- | Matrix-matrix subtraction
(#-#) :: (Vector p, Vector q, Num a) => q (p a) -> q (p a) -> q (p a)
(#-#) = liftA2 (^-^)

infix 7 ^*#

-- | Vector-matrix multiplication
(^*#) :: (Vector p, Vector q, Num a) => p a -> q (p a) -> q a
v ^*# m = fmap (`dot` v) m

infixl 7 ~*#

-- | Scalar-matrix multiplication
(~*#) :: (Vector p, Vector q, Num a) => a -> q (p a) -> q (p a)
s ~*# m = fmap (s *) <$> m

-- | Multiplicative inverse of a matrix /m/ (assumes @not $ det m ~= 0@)
inverse :: (Square v, Fractional a) => v (v a) -> v (v a)
inverse m = recip (det m) ~*# m

-- | Diagonal sum of a matrix
trace :: (Square v, Num a) => v (v a) -> a
trace = sum . view diagonal

-- | Transpose of a matrix
transpose :: (Vector p, Vector q) => q (p a) -> p (q a)
transpose qp = fmap (\o -> fmap (view o) qp) axes