【问题标题】:Caching in Haskell and explicit parallelismHaskell 中的缓存和显式并行性
【发布时间】:2012-08-15 00:51:30
【问题描述】:

我目前正在尝试优化我在 Projet Euler 的problem 14 解决方案。 我真的很喜欢 Haskell,我认为它非常适合解决这类问题,这是我尝试过的三种不同的解决方案:

import Data.List (unfoldr, maximumBy)
import Data.Maybe (fromJust, isNothing)
import Data.Ord (comparing)
import Control.Parallel

next :: Integer -> Maybe (Integer)
next 1 = Nothing
next n
  | even n = Just (div n 2)
  | odd n  = Just (3 * n + 1)

get_sequence :: Integer -> [Integer]
get_sequence n = n : unfoldr (pack . next) n
  where pack n = if isNothing n then Nothing else Just (fromJust n, fromJust n)

get_sequence_length :: Integer -> Integer
get_sequence_length n
    | isNothing (next n) = 1
    | otherwise = 1 + (get_sequence_length $ fromJust (next n))

-- 8 seconds
main1 = print $ maximumBy (comparing length) $ map get_sequence [1..1000000]

-- 5 seconds
main2 = print $ maximum $ map (\n -> (get_sequence_length n, n)) [1..1000000]

-- Never finishes
main3 = print solution
  where
    s1 = maximumBy (comparing length) $ map get_sequence [1..500000]
    s2 = maximumBy (comparing length) $ map get_sequence [500001..10000000]
    solution = (s1 `par` s2) `pseq` max s1 s2

现在,如果您查看实际问题,则存在很大的缓存潜力,因为大多数新序列将包含之前已经计算过的子序列。

为了比较,我也用 C 写了一个版本:
缓存运行时间:0.03 秒
无缓存运行时间:0.3 秒

这简直是疯了!当然,缓存将时间减少了 10 倍,但即使没有缓存,它仍然比我的 Haskell 代码快至少 17 倍。

我的代码有什么问题? 为什么 Haskell 不为我缓存函数调用?由于函数是纯缓存,缓存不应该是微不足道的,只是可用内存的问题?

我的第三个并行版本有什么问题?为什么没有完成?

将 Haskell 视为一种语言,编译器是否会自动并行化某些代码(折叠、映射等),还是必须始终使用 Control.Parallel 显式完成?

编辑:我偶然发现了this 类似的问题。他们提到他的函数不是尾递归的。我的 get_sequence_length 是尾递归的吗?如果不是,我该怎么做?

Edit2:
致丹尼尔:
非常感谢回复,真的很棒。 我一直在玩弄你的改进,但发现了一些非常糟糕的问题。

我在 Windws 7(64 位)、3.3 GHZ 四核和 8GB RAM 上运行测试。
正如你所说,我做的第一件事是用 Int 替换所有 Integer,但是每当我运行任何主电源时,我的内存都用完了, 即使 +RTS kSize -RTS 设置得高得离谱。

最终我找到了this(stackoverflow 很棒...),这意味着由于 Windows 上的所有 Haskell 程序都以 32 位运行,因此 Int 溢出导致无限递归,哇...

我改为在 Linux 虚拟机(使用 64 位 ghc)中运行测试,得到了类似的结果。

【问题讨论】:

  • main3 中有一个额外的零...

标签: haskell caching optimization functional-programming


【解决方案1】:

好吧,让我们从头开始。首先重要的是给出你用来编译和运行的确切命令行;对于我的回答,我将使用这一行来计算所有程序的时间:

ghc -O2 -threaded -rtsopts test && time ./test +RTS -N

接下来:由于机器之间的时间差异很大,我们将为我的机器和您的程序提供一些基准时间。这是我电脑上uname -a 的输出:

Linux sorghum 3.4.4-2-ARCH #1 SMP PREEMPT Sun Jun 24 18:59:47 CEST 2012 x86_64 Intel(R) Core(TM)2 Quad CPU Q6600 @ 2.40GHz GenuineIntel GNU/Linux

亮点是:四核、2.4GHz、64 位。

使用main1: 30.42s user 2.61s system 149% cpu 22.025 total
使用main2:21.42s user 1.18s system 129% cpu 17.416 total
使用main322.71s user 2.02s system 220% cpu 11.237 total

实际上,我通过两种方式修改了main3:首先,从s2 的范围末尾删除一个零,其次,将max s1 s2 更改为maximumBy (comparing length) [s1, s2],因为只有前者意外地计算出正确的答案。 =)

我现在将重点介绍串行速度。 (要回答您的一个直接问题:不,GHC 不会自动并行化或记忆您的程序。这两件事的开销都很难估计,因此很难决定什么时候做它们会是有益的。我有不知道为什么这个答案中的串行解决方案的 CPU 利用率也超过 100%;也许在另一个线程中发生了一些垃圾收集或类似的事情。)我们将从 main2 开始,因为它是两者中更快的串行实现。获得一点提升的最便宜的方法是将所有类型签名从 Integer 更改为 Int

使用Int:11.17s user 0.50s system 129% cpu 8.986 total(大约快一倍)

下一个提升来自于减少内部循环中的分配(消除中间的Maybe 值)。

import Data.List
import Data.Ord

get_sequence_length :: Int -> Int
get_sequence_length 1 = 1
get_sequence_length n
    | even n = 1 + get_sequence_length (n `div` 2)
    | odd  n = 1 + get_sequence_length (3 * n + 1)

lengths :: [(Int,Int)]
lengths = map (\n -> (get_sequence_length n, n)) [1..1000000]

main = print (maximumBy (comparing fst) lengths)

使用这个:4.84s user 0.03s system 101% cpu 4.777 total

下一个提升来自使用比evendiv 更快的操作:

import Data.Bits
import Data.List
import Data.Ord

even' n = n .&. 1 == 0

get_sequence_length :: Int -> Int
get_sequence_length 1 = 1
get_sequence_length n = 1 + get_sequence_length next where
    next = if even' n then n `quot` 2 else 3 * n + 1

lengths :: [(Int,Int)]
lengths = map (\n -> (get_sequence_length n, n)) [1..1000000]

main = print (maximumBy (comparing fst) lengths)

使用这个:1.27s user 0.03s system 105% cpu 1.232 total

对于那些在家跟随的人来说,这比我们开始使用的 main2 快了大约 17 倍 - 改用 C 后的竞争性改进。

对于记忆,有几个选择。最简单的方法是使用像data-memocombinators 这样的预先存在的包来创建一个不可变数组并从中读取。时间对于为这个数组选择一个合适的大小是相当敏感的。对于这个问题,我发现50000 是一个很好的上限。

import Data.Bits
import Data.MemoCombinators
import Data.List
import Data.Ord

even' n = n .&. 1 == 0

pre_length :: (Int -> Int) -> (Int -> Int)
pre_length f 1 = 1
pre_length f n = 1 + f next where
    next = if even' n then n `quot` 2 else 3 * n + 1

get_sequence_length :: Int -> Int
get_sequence_length = arrayRange (1,50000) (pre_length get_sequence_length)

lengths :: [(Int,Int)]
lengths = map (\n -> (get_sequence_length n, n)) [1..1000000]

main = print (maximumBy (comparing fst) lengths)

有了这个:0.53s user 0.10s system 149% cpu 0.421 total

最快的方法是使用一个可变的、未装箱的数组作为记忆位。它的惯用语要少得多,但它是裸机速度。速度对这个数组的大小不那么敏感,只要数组大约与您想要答案的最大事物一样大。

import Control.Monad
import Control.Monad.ST
import Data.Array.Base
import Data.Array.ST
import Data.Bits
import Data.List
import Data.Ord

even' n = n .&. 1 == 0
next  n = if even' n then n `quot` 2 else 3 * n + 1

get_sequence_length :: STUArray s Int Int -> Int -> ST s Int
get_sequence_length arr n = do
    bounds@(lo,hi) <- getBounds arr
    if not (inRange bounds n) then (+1) `fmap` get_sequence_length arr (next n) else do
        let ix = n-lo
        v <- unsafeRead arr ix
        if v > 0 then return v else do
            v' <- get_sequence_length arr (next n)
            unsafeWrite arr ix (v'+1)
            return (v'+1)

maxLength :: (Int,Int)
maxLength = runST $ do
    arr <- newArray (1,1000000) 0
    writeArray arr 1 1
    loop arr 1 1 1000000
    where
    loop arr n len 1  = return (n,len)
    loop arr n len n' = do
        len' <- get_sequence_length arr n'
        if len' > len then loop arr n' len' (n'-1) else loop arr n len (n'-1)

main = print maxLength

有了这个:0.16s user 0.02s system 138% cpu 0.130 total(与记忆的 C 版本竞争)

【讨论】:

  • 不错的进展和最终结果。在这一点上,整个优化顺序感觉都已编纂。编辑:一个问题,你为什么使用数组而不是Vector?个人喜好,就是受不了Array这个界面。
  • 非常感谢,非常直接的回答。然而,我没有得到的是您的第一个代码示例如何消除子列表。长度函数不只是按顺序运行 get_sequence_length 吗?我看不出它与原来的 main2 有什么不同,除了它的一部分已被分解为 lenghts 函数。 (另外,请参阅我的编辑以获得更长的回复)
  • @user1599468 哎呀,32位的东西有点烦人。至于消除列表——你说得对,我说得不准确。我将很快在线更新我的答案,但简短的答案是它消除了在每次循环迭代期间分配两个 JustNothing 值。
  • 如果您追求原始速度,您应该将n `quot` 2 替换为n `shiftR` 1。在我的盒子上,这要快得多。此外,数组版本(我尝试过的唯一一个)在这里更快非线程。最后,您可以通过避开getBoundsinRangelet ix = n - lo 来获得一些额外的速度;从索引 0 开始数组并将上限作为参数传递给 get_sequence_length,这样您就可以直接在 n 处比较 nhiunsafeRead/Write
【解决方案2】:

GHC 不会自动为您并行化任何内容。正如您猜想的那样,get_sequence_length 不是尾递归的。见here。并考虑编译器(除非它为您做了一些很好的优化)如何在您完成之前无法评估所有这些递归加法;您正在“建立 thunk”,这通常不是一件好事。

尝试调用递归辅助函数并传递一个累加器,或者尝试根据foldr 定义它。

【讨论】:

    猜你喜欢
    • 2019-12-18
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2013-02-06
    • 1970-01-01
    • 2014-11-22
    • 2011-01-13
    • 2013-09-12
    相关资源
    最近更新 更多