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