packages feed

proarrow-0.1.0.0: src/Proarrow/Category/Instance/ZX.hs

{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE GeneralizedNewtypeDeriving #-}
{-# OPTIONS_GHC -Wno-orphans #-}

-- | The __ZX calculus__ for reasoning about quantum computations: objects are numbers of qubits
-- and a morphism @'ZX' i o@ is a complex matrix between the corresponding state spaces, stored
-- sparsely. Provides the generators ('zSpider', 'xSpider' and 'hadamard') as a dagger monoidal
-- category.
module Proarrow.Category.Instance.ZX where

import Data.Bits (Bits (..), shiftL, (.|.))
import Data.Char (chr)
import Data.Complex (Complex (..), conjugate, magnitude, mkPolar)
import Data.Functor ((<&>))
import Data.List (intercalate, sort)
import Data.Map.Strict qualified as Map
import Data.Proxy (Proxy (..))
import Data.Type.Nat qualified as M
import Data.Vec.Lazy (Vec (..), reifyList)
import GHC.TypeNats (KnownNat, Nat, natVal, type (+), type (-))
import Numeric (showFFloat)
import Unsafe.Coerce (unsafeCoerce)
import Prelude hiding (Monoid, id, (**), (.))

import Proarrow.Category.Enriched.Dagger (DaggerProfunctor (..))
import Proarrow.Category.Instance.Cost (withPlusIsNat)
import Proarrow.Category.Monoidal (Monoidal (..), MonoidalProfunctor (..), SymMonoidal (..))
import Proarrow.Category.Monoidal.Action (MonoidalAction)
import Proarrow.Category.Monoidal.Closed (Closed (..))
import Proarrow.Category.Monoidal.CompactClosed (CompactClosed (..), coactCC)
import Proarrow.Category.Monoidal.CopyDiscard (CopyDiscard)
import Proarrow.Category.Monoidal.Hypergraph (Frobenius, Hypergraph, cap, cup)
import Proarrow.Category.Monoidal.StarAutonomous (ExpSA, StarAutonomous (..), applySA, currySA, expSA)
import Proarrow.Category.Monoidal.Strength (Costrong (..))
import Proarrow.Core (CAT, CategoryOf (..), Profunctor (..), Promonad (..), dimapDefault, obj, type (+->))
import Proarrow.Monoid (CocommutativeComonoid, CommutativeMonoid, Comonoid (..), Monoid (..))

newtype Bitstring (n :: Nat) = BS Int
  deriving (Eq, Ord)
  deriving newtype (Num)

instance (KnownNat n) => Bounded (Bitstring n) where
  minBound = BS 0
  maxBound = BS ((1 `shiftL` nat @n) - 1)
instance Enum (Bitstring n) where
  fromEnum (BS x) = x
  toEnum x = BS x

-- | Split n + m bits into two parts: the lower n bits and the higher m bits.
split :: (KnownNat n) => Bitstring (n + m) -> (Bitstring n, Bitstring m)
split @n (BS x) = let (m, n) = x `divMod` (1 `shiftL` nat @n) in (BS n, BS m)

-- | Combine two bitstrings of lengths n and m into one bitstring with the n lower bits or m higher bits.
combine :: (KnownNat n) => Bitstring n -> Bitstring m -> Bitstring (n + m)
combine @n (BS x) (BS y) = BS ((y `shiftL` nat @n) .|. x)

-- The order is (output, input)!
type SparseMatrix o i = Map.Map (Bitstring o, Bitstring i) (Complex Double)

epsilon :: Double
epsilon = 1e-12

isZero :: Complex Double -> Bool
isZero z = magnitude z <= epsilon

filterSparse :: SparseMatrix o i -> SparseMatrix o i
filterSparse = Map.filter (Prelude.not . isZero)

transpose :: SparseMatrix o i -> SparseMatrix i o
transpose = Map.mapKeys \(o, i) -> (i, o)

mirror :: (KnownNat n) => Bitstring n -> Bitstring (n + n)
mirror @n (BS x) = BS (go (nat @n) x x)
  where
    go 0 _ acc = acc
    go k y acc = go (k - 1) (y `shiftR` 1) ((acc `shiftL` 1) .|. (y .&. 1))

enumAll :: (KnownNat n) => [Bitstring n]
enumAll = [minBound .. maxBound]

nat :: (KnownNat n) => Int
nat @n = fromIntegral $ natVal (Proxy @n)

type ZX :: CAT Nat
data ZX i o where
  ZX :: (KnownNat i, KnownNat o) => SparseMatrix o i -> ZX i o

instance (KnownNat n) => Show (Bitstring n) where
  show (BS x) = go (nat @n) x ""
    where
      go 0 _ acc = acc
      go n bs acc = case bs `divMod` 2 of (r, b) -> go (n - 1) r (chr (48 + b) : acc)

instance Show (ZX a b) where
  show (ZX m) = intercalate ", " (sort (fmt <$> Map.toList m))
    where
      fmt ((o, i), r :+ c) =
        show i
          ++ "->"
          ++ show o
          ++ "="
          ++ (if abs (1 - r) < epsilon then "1" else showFFloat (Just 3) r "")
          ++ (if abs c < epsilon then "" else " :+ " ++ showFFloat (Just 3) c "")

type family MatrixSize (n :: Nat) :: M.Nat where
  MatrixSize 0 = M.Nat1
  MatrixSize n = M.Mult M.Nat2 (MatrixSize (n - 1))

toMatrix :: forall o i. ZX i o -> Vec (MatrixSize o) (Vec (MatrixSize i) (Complex Double))
toMatrix (ZX m) =
  reifyList (enumAll @o) \vo ->
    reifyList (enumAll @i) \vi ->
      unsafeCoerce $ vo <&> \o -> vi <&> \i -> Map.findWithDefault 0 (o, i) m

instance Profunctor ZX where
  dimap = dimapDefault
  r \\ ZX _ = r
instance Promonad ZX where
  id @n = ZX $ Map.fromList [((i, i), 1) | i <- enumAll @n]
  ZX n . ZX m =
    ZX $
      filterSparse $
        Map.fromListWith
          (+)
          [ ((c, a), mv * nv)
          | ((c, b1), mv) <- Map.toList n
          , ((b2, a), nv) <- Map.toList m
          , b1 == b2
          ]

-- | The category of qubits, to implement ZX calculus from quantum computing.
instance CategoryOf Nat where
  type (~>) = ZX
  type Ob a = KnownNat a

instance DaggerProfunctor ZX where
  dagger (ZX m) = ZX $ Map.fromList [((i, o), conjugate v) | ((o, i), v) <- Map.toList m]

instance MonoidalProfunctor ZX where
  one = id
  ZX @ni @no n ** ZX @mi @mo m =
    withOb2 @_ @ni @mi $
      withOb2 @_ @no @mo $
        ZX $
          Map.fromList
            [ ((combine no mo, combine ni mi), nv * mv)
            | ((no, ni), nv) <- Map.toList n
            , ((mo, mi), mv) <- Map.toList m
            ]

-- | Addition of the number of qubits as monoidal tensor. This is the Kronecker product of the matrices.
instance Monoidal Nat where
  type Unit = 0
  type p ** q = p + q
  withOb2 @a @b r = withPlusIsNat @a @b r
  associator @a @b @c = unsafeCoerce (withOb2 @_ @a @b (withOb2 @_ @(a + b) @c (obj @(a + b + c))))
  associatorInv @a @b @c = unsafeCoerce (withOb2 @_ @a @b (withOb2 @_ @(a + b) @c (obj @(a + b + c))))

instance SymMonoidal Nat where
  swap @m @n =
    withOb2 @_ @m @n $
      withOb2 @_ @n @m $
        ZX $
          Map.fromList
            [ ((combine n m, combine m n), 1)
            | n <- enumAll @n
            , m <- enumAll @m
            ]

instance Closed Nat where
  type x ~~> y = ExpSA x y
  withObExp @a @b r = withOb2 @_ @a @b r
  curry @x @y = currySA @x @y
  apply @y @z = applySA @y @z
  (^^^) = expSA

instance StarAutonomous Nat where
  type Dual x = x
  withObDual r = r
  dual (ZX m) = ZX (transpose m)
  dualInv = dual
  linDist @_ @b @c (ZX m) =
    withOb2 @_ @b @c $
      ZX (Map.mapKeys (\(c, ab) -> case split ab of (a, b) -> (combine b c, a)) m)
  linDistInv @a @b (ZX m) =
    withOb2 @_ @a @b $
      ZX (Map.mapKeys (\(bc, a) -> case split bc of (b, c) -> (c, combine a b)) m)
  doubleNeg = id
  doubleNegInv = id

instance CompactClosed Nat where
  distribDual @a @b = withOb2 @_ @a @b id
  dualUnit = id
  dualityUnit @a = cup @a
  dualityCounit @a = cap @a

instance (MonoidalAction (t :: (Nat, Nat) +-> Nat)) => Costrong t ZX where
  coact @x = coactCC @t @x

-- No terminal or initial object: @hom(n, m)@ is the space of @2^m x 2^n@ complex matrices, which
-- is a singleton for no @n@ and @m@ at all. @0@, the monoidal unit, is not terminal either: the
-- zero matrix is an arrow @1 ~> 0@, but so is @zSpider 0 :: ZX 1 0@, and the two differ. The zero
-- matrix is a zero morphism, which wants a class of its own, not a fake zero object.

-- No binary(co)products, since that would need 2^n + 2^m = 2^(x :: nat)

instance (KnownNat n) => Monoid (n :: Nat) where
  mempty = ZX $ Map.fromList [((o, minBound), 1) | o <- enumAll @n]
  mappend = withPlusIsNat @n @n $ ZX $ Map.fromList [((o, combine o o), 1) | o <- enumAll @n]

instance (KnownNat n) => Comonoid (n :: Nat) where
  counit = ZX $ Map.fromList [((minBound, i), 1) | i <- enumAll @n]
  comult = withPlusIsNat @n @n $ ZX $ Map.fromList [((combine i i, i), 1) | i <- enumAll @n]
instance (KnownNat n) => CocommutativeComonoid (n :: Nat)

instance (KnownNat a) => Frobenius (a :: Nat)
instance (KnownNat a) => CommutativeMonoid (a :: Nat)
instance Hypergraph Nat
instance CopyDiscard Nat

zSpider :: (KnownNat i, KnownNat o) => Double -> ZX i o
zSpider alpha = ZX $ Map.fromListWith (+) [(minBound, 1), (maxBound, mkPolar 1 alpha)]

xSpider :: (KnownNat i, KnownNat o) => Double -> ZX i o
xSpider alpha = hadamard . zSpider alpha . hadamard

zCopy :: ZX 1 2
zCopy = zSpider 0

zDisc :: ZX 1 0
zDisc = zSpider 0

xCopy :: ZX 1 2
xCopy = xSpider 0

xDisc :: ZX 1 0
xDisc = xSpider 0

zeroState :: ZX 0 1
zeroState = xSpider 0

oneState :: ZX 0 1
oneState = xSpider pi

plusState :: ZX 0 1
plusState = zSpider 0

minusState :: ZX 0 1
minusState = zSpider pi

not :: ZX 1 1
not = xSpider pi

-- | Controlled NOT gate
cnot :: ZX 2 2
cnot = (id ** xSpider @2 0) . (zCopy ** id)

-- | Greenberger–Horne–Zeilinger state
ghzState :: ZX 0 3
ghzState = zSpider 0

hadamard :: (KnownNat n) => ZX n n
hadamard @n = ZX $ Map.fromList [((i, j), (sign i j * amp) :+ 0) | i <- enumAll @n, j <- enumAll @n]
  where
    amp = 1 / sqrt (fromIntegral (shiftL (1 :: Int) (nat @n)))
    sign (BS a) (BS b) = if even (popCount (a .&. b)) then 1 else -1