在Haskell中对基于熵的Wordle求解器进行优化
我的Wordle求解器 似乎运行得很慢,我想听听你们的意见,看看还能怎么进一步加速。
我做了一些我能想到的优化,但速度仍然看起来很差。
据我理解,瓶颈在 filterByResult 函数。
具体在 eqZip 和 zippedEq 的“子”函数中。

我暂时看不出有什么方法可以优化 filterByResult,也许至多再考虑一种更好的实现 passesCounts 的方式。
但 passesChars 是 && 的第一个参数,程序以这种方式运行得更快;因此我认为现在没必要再处理计数方面的事情。
类型:
import qualified Data.ByteString as BS
import qualified Data.Map.Strict as Map
-- Allowed guesses, allowed answers
type GuessCtx = ([WordleWord], [WordleWord])
data GuessResult = GRWrong | GROtherPlace | GRCorrect
type WordleWord = BS.ByteString
filterByResult 的代码
filterByResult :: WordleWord -> [GuessResult] -> (WordleWord -> Bool)
filterByResult guess res w = passesChars && passesCounts
where
zippedGuess :: [(GuessResult, Word8)]
zippedGuess = zip res (BS.unpack guess)
eqZip :: [Bool]
eqZip = BS.zipWith (==) guess w
countChars :: Map.Map Word8 Int -> (GuessResult, Word8) -> Map.Map Word8 Int
countChars m (GRWrong, gc) = Map.insertWith (+) gc 0 m
countChars m (_, gc) = Map.insertWith (+) gc 1 m
countedChars :: Map.Map Word8 Int
countedChars = foldl' countChars Map.empty zippedGuess
passesCounts :: Bool
passesCounts = Map.foldlWithKey' filterWithMap True countedChars
where
filterWithMap :: Bool -> Word8 -> Int -> Bool
filterWithMap False _ _ = False
filterWithMap True c n
| n == 0 = c `BS.notElem` w
| otherwise = BS.count c w >= n
passesChars :: Bool
passesChars = foldl' filterWithChars True zippedEq
where
zippedEq :: [(GuessResult, Bool)]
zippedEq = zip res eqZip
filterWithChars :: Bool -> (GuessResult, Bool) -> Bool
filterWithChars False _ = False
filterWithChars True (GRCorrect, isEq) = isEq
filterWithChars True (_, isEq) = not isEq
熵计算的另一个方向:
wordEntropy :: WordleWord -> [WordleWord] -> Double
wordEntropy w gs = sum $ do
res <- possibleResults
let newGuessList = filter (filterByResult w res) gs
let ngCount :: Double = fromIntegral $ length newGuessList
guard $ ngCount > 0
let gCount :: Double = fromIntegral $ length gs
let probability = ngCount / gCount
let entropy = -logBase 2 probability
pure $ probability * entropy
avg :: (Fractional r) => [r] -> r
avg vals = sum vals / fromIntegral (length vals)
你可以看到这里我使用了 sum,而不是 avg。
之所以这样做,是因为 avg 让这段代码变慢得多(十倍或更多)。
在我当前的测试中 sum 在求解Wordle方面(至少在字典前100个答案的情况下)工作还算正常。
其实我也不清楚为什么 avg 会带来如此大的性能惩罚。
我原以为是因为遍历列表两次(求和和长度),但通过用 foldl' 重写 avg,使其一次遍历,我在性能上没有提升(而且,GHC很可能已经优化了它)。
也许在调用 filterByResult 之前添加一些检查来过滤 possibleResults(以过滤掉不可行的结果)是有意义的。
但我担心实现起来也会与 filterByResult 一样耗时,所以现在我只是把 guard 放在 filterByResult 之后使用。
selectBestNextWord :: GuessCtx -> WordleWord
selectBestNextWord (gs, _) = snd $ foldl' entIter (0, "") gs
where
entIter :: (Double, WordleWord) -> WordleWord -> (Double, WordleWord)
entIter prev@(maxEnt, _) curW =
let
wEnt = wordEntropy curW gs
in
if wEnt > maxEnt
then (wEnt, curW)
else prev
Main只是从文件中读取字典,计算要使用的第一个单词,然后遍历字典中的所有答案,并尝试在最多6 次猜测内解决每一个。
我想提高速度,至少在20分钟内看到结果,并且在使用包含约2000个单词的字典时,至少有像“交互式”一样的解决单词速度(见 iterateWord 函数)。
也许通过使用已经预先计算的2-3个单词来实现更高的速度是可能的,因此之后我们将对较小的单词范围计算熵——但我想避免使用程序自己计算不到的“外部”单词。
在这里向有经验的朋友请教时,如果你对代码风格,或对这段代码的任何方面有建议/评论,我将不胜感激。
现在的时间(带剖析,使用 +RTS -sstderr -p 测量):
100个单词 - 1.9s
200个单词 - 24.5s
500个单词 - 102.1s
未做剖析的时间(使用 time 测量):
100个单词 - 0.8s
200个单词 - 8.2s
500个单词 - 33.7s
此外,我确实尝试在 wordEntropy 中使用 parallel 模块来并行化列表求值,但我所有的尝试都产生了相同或更差的性能(而使用 top 时,我看到一次核心达到100% 占用,另一颗核心只有1-5%)。
PS:我确实在运行时对程序使用了 -threaded GHC标志并传递了 -N1/-N2 标志,可能只是我的老旧CPU的表现而已。
相关代码是这样的:
wordEntropy w gs = sum $ withStrategy (parList rpar) $ do 和
wordEntropy w gs = sum $ withStrategy (parList rseq) $ do
我对 parallel 不太熟练,所以也许这是错误的用法?
你可以在我的GitHub仓库看到完整代码:https://github.com/qwaszx000/fictional-octo-memory/blob/master/app/Main.hs
我使用cabal 3.14.2.0和 GHC 9.6.7,并带有 -O2 标志
解决方案
为了让它成为一个可编译成单一、独立的可执行文件,我在文件中加入了以下内容,以确保能单独编译:
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
import Control.Monad
import Data.List
import Data.Word
possibleResults :: [[GuessResult]]
possibleResults = replicateM 6 [GRWrong, GROtherPlace, GRCorrect]
main :: IO ()
main = do
gs <- BS.split 10 <$> BS.getContents
print (selectBestNextWord (gs, []))
我的实验将使用ghc-9.8.1。为了得到一个单词表,我把 /usr/share/dict/words 过滤为六字母单词,将结果分成500个等大小的连续块,并从每个块中取出第一个单词。我将使用 ghc -O2 test 进行编译(并使用 ./test <test.txt 来执行)。
用问题中代码进行的初始运行时间是16.1秒。我认为这在某种程度上也是一次学习Haskell的练习,因此在优化之前,我想提出一些使用库函数而不是自己手动重新实现的清理方法,这样你就能看到地道的Haskell是什么样子。首先,我们可以使用 fromListWith 来定义 countedChars。
-- was
data GuessResult = GRWrong | GROtherPlace | GRCorrect
zippedGuess :: [(GuessResult, Word8)]
zippedGuess = zip res (BS.unpack guess)
countChars :: Map.Map Word8 Int -> (GuessResult, Word8) -> Map.Map Word8 Int
countChars m (GRWrong, gc) = Map.insertWith (+) gc 0 m
countChars m (_, gc) = Map.insertWith (+) gc 1 m
countedChars :: Map.Map Word8 Int
countedChars = foldl' countChars Map.empty zippedGuess
-- suggestion
data GuessResult = GRWrong | GROtherPlace | GRCorrect deriving (Eq, Ord, Read, Show)
countedChars :: Map.Map Word8 Int
countedChars = Map.fromListWith (+) $ zipWith
(\gc singleRes -> (gc, if singleRes == GRWrong then 0 else 1))
(BS.unpack guess) res
我们可以使用 and 和 zipWith 来定义 passesChars。
-- was
eqZip :: [Bool]
eqZip = BS.zipWith (==) guess w
passesChars :: Bool
passesChars = foldl' filterWithChars True zippedEq
where
zippedEq :: [(GuessResult, Bool)]
zippedEq = zip res eqZip
filterWithChars :: Bool -> (GuessResult, Bool) -> Bool
filterWithChars False _ = False
filterWithChars True (GRCorrect, isEq) = isEq
filterWithChars True (_, isEq) = not isEq
-- suggestion
passesChars :: Bool
passesChars = and $ zipWith
(\singleRes singleEq -> (singleRes == GRCorrect) == singleEq)
res (BS.zipWith (==) guess w)
我们也可以用 and 来定义 passesCounts;mapWithKey 有点像 zipWith!
-- was
passesCounts :: Bool
passesCounts = Map.foldlWithKey' filterWithMap True countedChars
where
filterWithMap :: Bool -> Word8 -> Int -> Bool
filterWithMap False _ _ = False
filterWithMap True c n
| n == 0 = c `BS.notElem` w
| otherwise = BS.count c w >= n
-- suggestion
passesCounts :: Bool
passesCounts = and $ Map.mapWithKey
(\c n -> if n == 0 then c `BS.notElem` w else BS.count c w >= n)
countedChars
你的 selectBestNextWord 是一个秘密的 maximumOn。……好吧,应该是这样的,只是标准库没有为列表提供 maximumOn。自从引入 sortOn 以来,我一直觉得这是个疏忽。但在这种情况下,通常更好的是定义一个通用的东西,其行为可以从名称和标准库的约定中推断出来,然后对其进行特化,而不是写一个需要一次性读完的完全特殊的东西。就像这样:
-- was
selectBestNextWord :: GuessCtx -> WordleWord
selectBestNextWord (gs, _) = snd $ foldl' entIter (0, "") gs
where
entIter :: (Double, WordleWord) -> WordleWord -> (Double, WordleWord)
entIter prev@(maxEnt, _) curW =
let
wEnt = wordEntropy curW gs
in
if wEnt > maxEnt
then (wEnt, curW)
else prev
-- suggestion
import Data.Ord
maximumOn :: Ord b => (a -> b) -> [a] -> a
maximumOn f = fst . maximumBy (comparing snd) . map (\a -> (a, f a))
selectBestNextWord :: GuessCtx -> WordleWord
selectBestNextWord (gs, _) = maximumOn (flip wordEntropy gs) gs
有了这些改动,程序速度略有提升,达到9.25s,且并未本质改变算法。我要强调的是,这些改动是为了可读性;它也提升了性能这是件好事,但并非目标。
现在进入优化阶段。我首先注意到的是,你对每个可能的结果都要遍历整份单词表一次,等于3^6 = 729次。可能更好的做法是只遍历一次,计算正确的结果,然后用它对列表进行分区。没错,这意味着我们会丢弃上面几乎所有的改动…… 在迭代改进的世界里,这就是现实!首先是计算结果的函数:
computeResult :: WordleWord -> WordleWord -> [GuessResult]
computeResult guess w = go mismatches0 (BS.unpack guess) (BS.unpack w) where
mismatches0 :: Map.Map Word8 Int
mismatches0 = Map.fromListWith (+) $ BS.zipWith
(\singleGuess singleW -> (singleW, if singleGuess == singleW then 0 else 1))
guess w
go :: Map.Map Word8 Int -> [Word8] -> [Word8] -> [GuessResult]
go mismatches (g:gs) (w:ws)
| g == w = GRCorrect : go mismatches gs ws
| otherwise = case Map.findWithDefault 0 g mismatches of
0 -> GRWrong : go mismatches gs ws
n -> GROtherPlace : go (Map.insert g (n-1) mismatches) gs ws
go _ _ _ = []
我不太喜欢重复的 unpack,但稍后再说。至于 go……好吧,即便是我也不总是使用内置的fold。=)(有一个这样的函数,但我不认为它比那样更易读。)有了这个函数后,我们可以重写 wordEntropy 的开头:
-- was
wordEntropy w gs = sum $ do
res <- possibleResults
let newGuessList = filter (filterByResult w res) gs
-- suggestion
wordEntropy w gs = sum $ do
(res, newGuessList) <- Map.toList . Map.fromListWith (++) $
[(computeResult w g, [g]) | g <- gs]
这一改动带来显著的差异,运行时降至0.17s。等等!它计算出一个不同的结果!发生了什么?
在我看来,你原来的 filterByResult 存在一个错误:它会在本不应该出现 GROtherPlace 的地方接受它们。快速示例:
> mapM_ print $ filter (\res -> filterByResult "aaaaab" res "ccccca") possibleResults
[GRWrong,GRWrong,GRWrong,GRWrong,GROtherPlace,GRWrong]
[GRWrong,GRWrong,GRWrong,GROtherPlace,GRWrong,GRWrong]
[GRWrong,GRWrong,GROtherPlace,GRWrong,GRWrong,GRWrong]
[GRWrong,GROtherPlace,GRWrong,GRWrong,GRWrong,GRWrong]
[GROtherPlace,GRWrong,GRWrong,GRWrong,GRWrong,GRWrong]
再次强调,这只是我的看法;游戏会选择一个特定的 GROtherPlace 放置以给出线索,因此你应该只从 filterByResult 得到一个匹配。但如果你的看法不同,也并非没有解决办法:computeResult 可以修改为返回所有可能的结果,而不仅仅是一个;这些都可以包括在 Map 产生于 wordEntropy 的结果中。无论如何,我将继续使用我的版本的程序;这对性能并没有实质性的影响。
现在我注意到我们计算了一个列表,newGuessList,然后为了它的长度就立即把它丢弃。其实完全可以直接计算它的长度。我们也不再需要 guard,因为我们只对第一步就有匹配的猜测结果进行遍历。在这个清理过程中,我还将 logBase 2 改为 log——它们之间只是一个常量因子之差,这不会改变哪个对象具有最大的熵;并且 logBase 使用两次对 log 的调用。对数概率的名字 entropy 也不太恰当,更准确的 logProbability 也不比它的定义更易读,所以我们把它内联吧:
-- was
wordEntropy w gs = sum $ do
(res, newGuessList) <- Map.toList . Map.fromListWith (++) $
[(computeResult w g, [g]) | g <- gs]
let ngCount :: Double = fromIntegral $ length newGuessList
guard $ ngCount > 0
let gCount :: Double = fromIntegral $ length gs
let probability = ngCount / gCount
let entropy = -logBase 2 probability
pure $ probability * entropy
-- suggestion
wordEntropy w gs = sum $ do
(_, ngCount) <- Map.toList . Map.fromListWith (+) $
[(computeResult w g, 1) | g <- gs]
let gCount :: Double = fromIntegral $ length gs
let probability = ngCount / gCount
pure $ -probability * log probability
这对性能没有影响,我们仍然是0.17s。既然你表示有兴趣让它并行化,一个便宜的做法是把 maximumOn 修改成如下:
-- was
maximumOn :: Ord b => (a -> b) -> [a] -> a
maximumOn f = fst . maximumBy (comparing snd) . map (\a -> (a, f a))
-- suggestion
import Control.Parallel.Strategies
maximumOn :: Ord b => (a -> b) -> [a] -> a
maximumOn f = fst . maximumBy (comparing snd) . parMap (\ab -> ab <$ rseq (snd ab)) (\a -> (a, f a))
要获得收益,你需要在这一个上进行略微不同的编译和运行:
% ghc -O2 -threaded test
% ./test +RTS -N6 <test.txt
我从中得到的速度提升有点让人失望,降到了0.07s,但至少付出的努力不大!此时完整程序看起来是这样的:
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE OverloadedStrings #-}
import Control.Parallel.Strategies
import Data.List
import Data.Ord
import Data.Word
import qualified Data.ByteString as BS
import qualified Data.Map.Strict as Map
-- Allowed guesses, allowed answers
type GuessCtx = ([WordleWord], [WordleWord])
data GuessResult = GRWrong | GROtherPlace | GRCorrect deriving (Eq, Ord, Read, Show)
type WordleWord = BS.ByteString
computeResult :: WordleWord -> WordleWord -> [GuessResult]
computeResult guess w = go mismatches0 (BS.unpack guess) (BS.unpack w) where
mismatches0 :: Map.Map Word8 Int
mismatches0 = Map.fromListWith (+) $ BS.zipWith
(\singleGuess singleW -> (singleW, if singleGuess == singleW then 0 else 1))
guess w
go :: Map.Map Word8 Int -> [Word8] -> [Word8] -> [GuessResult]
go mismatches (g:gs) (w:ws)
| g == w = GRCorrect : go mismatches gs ws
| otherwise = case Map.findWithDefault 0 g mismatches of
0 -> GRWrong : go mismatches gs ws
n -> GROtherPlace : go (Map.insert g (n-1) mismatches) gs ws
go _ _ _ = []
wordEntropy :: WordleWord -> [WordleWord] -> Double
wordEntropy w gs = sum $ do
(_, ngCount) <- Map.toList . Map.fromListWith (+) $
[(computeResult w g, 1) | g <- gs]
let gCount :: Double = fromIntegral $ length gs
let probability = ngCount / gCount
pure $ -probability * log probability
maximumOn :: Ord b => (a -> b) -> [a] -> a
maximumOn f = fst . maximumBy (comparing snd) . parMap (\pair -> pair <$ rseq (snd pair)) (\a -> (a, f a))
selectBestNextWord :: GuessCtx -> WordleWord
selectBestNextWord (gs, _) = maximumOn (flip wordEntropy gs) gs
main :: IO ()
main = do
gs <- BS.split 10 <$> BS.getContents
print (selectBestNextWord (gs, []))
我打算宣布230倍的加速为成功并止步于此,但如果你还需要更快的速度,仍有很多可能性,包括:
IntMap在可能使用的时候通常比Map快。把它用于所有的Map Word8可能会有收获。也许尝试使用长度为256的Vector或MVector也会有帮助,但映射的稀疏性很可能使这变得更差。但嘿,在这类问题上,实际测量胜过概率判断。- 作为更具侵入性的改动,我们在把
[GuessResult]作为映射键使用,这有点尴尬。因为单个GuessResult只需要编码成两位,长度为6 的[GuessResult]也可以很容易地编码成一个Int,我们就可以再次切换到IntMap。 - 甚至更具侵入性的改动是,字母只有6 个,而可选字母不到2^5,因此你可以把每个单词编码成30位,这几乎可以放在任何支持ghc的系统上的单个
Int中。我们在迭代和比较时需要的操作可以放在一个独立模块中,通过可怕的位运算来实现,然后高层算法可以用这些操作以更易读的方式使用。过去尝试过,我强烈不推荐;对于几乎任何目的,它都需要足够快,且它们写起来、读起来、调试和维护都很困难。 - 如上所述,我不太喜欢在
computeResult的解包。你可以直接切换到[Word8]作为表示方式,因为你目前没有从使用ByteString中获得太多好处;或者尝试改写computeResult以避免解包的需要。 - 我们将用这来构建一棵树。由于我们需要计算每对(猜测,秘密)的结果才能知道树的顶节点,因此我们不妨预计算所有配对,然后在构建树时使用该查找表——前提是我们有足够的内存!
- 我们已经做了并行化的简单实现,但也可能有更简单的办法获得更好的提升;例如,我会尝试将输入列表按大约与线程数相同的块数划分,让每个线程处理一个块中的所有猜测。这样做的希望是降低记账开销,使你能更接近你线程数的理想加速。