packages feed

eigen-hhlo-0.1.0.0: examples/02-svd.hs

{-# LANGUAGE DataKinds #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeApplications #-}

-- | Example 2: Singular Value Decomposition of a 3x2 matrix.
--
-- Demonstrates a multi-output custom call (U, S, Vt).
--
-- Build and run with:
--   cabal run example-svd --flag=examples

module Main where

import qualified Data.Text as T

import HHLO.Core.Types (DType(..))
import HHLO.IR.AST (FuncArg(..), TensorType(..))
import HHLO.EDSL.Ops (returnTuple3)
import HHLO.IR.Builder (arg, moduleFromBuilder3)
import HHLO.IR.Pretty (render)
import HHLO.Session (HostTensor, compile, hostFromList, hostToList, run)

import EigenHHLO.Core.Types (eigenSession)
import EigenHHLO.EDSL.Decomposition (svd)
import EigenHHLO.Runtime.Session (withEigenGPU)

main :: IO ()
main = withEigenGPU $ \esess -> do
    putStrLn "=== Example 2: Singular Value Decomposition (GPU) ==="

    let modu = moduleFromBuilder3 @'[3,2] @'F64 @'[2] @'F64 @'[2,2] @'F64 "main"
            [ FuncArg "a" (TensorType [3,2] F64) ]
            $ do a <- arg @'[3,2] @'F64
                 (u, s, vt) <- svd esess a
                 returnTuple3 u s vt

    putStrLn "\nGenerated MLIR:"
    putStrLn (T.unpack $ render modu)

    compiled <- compile (eigenSession esess) modu

    -- A = [[3, 0],
    --      [0, 2],
    --      [0, 0]]  (3x2, column-major)
    let input = hostFromList @'[3,2] @'F64
            [3, 0, 0, 0, 2, 0]

    putStrLn "\nInput matrix (column-major):"
    print (hostToList input)

    (uResult, sResult, vtResult) <- run (eigenSession esess) compiled input
        :: IO (HostTensor '[3,3] 'F64, HostTensor '[2] 'F64, HostTensor '[2,2] 'F64)

    putStrLn "\nU (3x3, column-major):"
    print (hostToList uResult)
    putStrLn "S (2,):"
    print (hostToList sResult)
    putStrLn "Vt (2x2, column-major):"
    print (hostToList vtResult)