Solving regex crosswords with Z3

Nelson Elhage

用 Z3 求解正则表达式填字游戏

原文由 Nelson Elhage 发布,订阅该博客

一段时间以来,我一直对 Z3 以及更广义的 SMT 求解很感兴趣。最近休陪产假期间,我想起了正则表达式填字游戏的存在,于是忍不住写了一个基于 Z3 的求解器,让自己深陷其中。

我本以为花一个下午就能草草写出一个简易求解器;结果却深陷其中,开始钻研和调试 Z3 的性能,学到的关于 Z3 和 SMT 的东西远超预期。在本文中,我会介绍我的思路和最初的求解器,然后深入探讨我尝试过的各种改进和变体。

我的所有代码都已在 GitHub 上开源,如果你想跟着复现或查看最终成果,可以去看看。

正则表达式填字游戏

简要回顾一下:正则表达式填字游戏由一个待填的未知字符网格组成。这些字符受一组给定的正则表达式约束,填入后每一行或每一列都必须匹配对应的正则表达式。我选择求解的是 regexle 这一变体,它使用六边形网格,不过我的大部分代码也能很容易地适配到其他变体上。

把正则表达式看作 DFA

要想用 Z3 这类约束求解器来解题,我们需要找到一种方法来编码“这串字符匹配这个正则表达式”这一约束,然后把所有线索都断言进去,让求解器找出满足全部约束的网格。

我第一时间想到的是利用这样一个事实:正则表达式描述的是正则语言1,因此可以用确定性有限自动机来识别。于是,我可以把正则表达式编码成 DFA,然后在 Z3 中用它的转移函数来表达正则匹配2

也就是说,如果我们有一个三字符的字符串 S = [c0, c1, c2],并且手头有该正则表达式的转移表 T : State × Char → State,我们就可以把“该正则表达式匹配 S”表达为类似这样的形式:

state_0 = init_state
state_1 = T(state_0, c0)
state_2 = T(state_1, c1)
state_3 = T(state_2, c2)
assert(state_3 in accept_states)

将正则表达式转换为 DFA

稍作搜索后,我发现了 qntmgreenery 库,它恰好完全符合我的需求——一个相对简洁的 Python 库,用于操作正则表达式,包括在正则语法与 DFA(它称之为 FSM)之间相互转换。它会为我们构建 FSM,整体而言非常直观,优化目标是简洁、易用和便于试验,而非极致性能,这对我们来说正合适——重活都交给 Z3 去做。

一个小麻烦是,greenery 处理的是任意 Unicode 字符上的正则表达式;为了避免巨大的查找表,它用字符区间来表示转移,我觉得用起来稍微有点麻烦。对于这个项目,我愿意假设字母表就是英文字母(A-Z),于是我写了一个辅助函数,把 greenery 的结构转换成扁平的 numpy 转移表:

ALPHABET = string.ascii_uppercase

def flatten_fsm(fsm):
    nstate = len(fsm.states)
    nvocab = len(ALPHABET)

    assert fsm.states == set(range(nstate))
    state_map = np.full((nstate, nvocab), -1, dtype=np.int32)
    for st, map in fsm.map.items():
        for i, c in enumerate(ALPHABET):
            for cc, dst in map.items():
                if cc.accepts(c):
                    state_map[st, i] = dst
                    break
    return state_map

我还定义了一个小型的 dataclass,把单个正则表达式的各种表示和分析打包在一起:

@dataclass
class Regex:
    pattern: str
    parsed: greenery.Pattern

    # 2d: (state, vocab) -> new state
    transition: np.ndarray
    # 1d: state -> bool
    accept: np.ndarray

    @classmethod
    def from_pattern(cls, pattern: str):
        parsed = greenery.parse(pattern)
        fsm = parsed.to_fsm().reduce()

        transition = flatten_fsm(fsm)
        nstate = transition.shape[0]
        accept = np.zeros((nstate,), bool)
        for st in fsm.finals:
            accept[st] = True

        return cls(
            pattern=pattern,
            parsed=parsed,
            transition=transition,
            accept=accept,
        )

映射到 Z3

要把各部分拼起来,我们需要让 Z3 了解每条线索对应的转移函数,为网格中的每个字符定义 Z3 变量,然后利用相关的转移函数对这些字符施加相应的约束。

我选择把状态和字符都表示为整数(并加上断言确保它们在取值范围内),因为这看起来最简单。接下来,我们需要编码状态与字符之间的转移函数关系。

Z3 支持函数类型,可以在表达式和断言中声明和操作。因此,我声明了一个相应类型的函数,并通过对转移表中每一项添加断言来告知 Z3 它的行为:

def build_func(solv: z3.Solver, clue: Clue):
    ctx = solv.ctx
    state_func = z3.Function(
        clue.name + "_trans",
        z3.IntSort(ctx),
        z3.IntSort(ctx),
        z3.IntSort(ctx),
    )
    pat = clue.pattern

    for (state, char), next_state in pat.all_transitions():
        solv.add(state_func(state, char) == next_state)

    return state_func

有了这个定义,再给定一个字符串(以字符列表的形式表示),我们就能很容易地表达“正则匹配”约束:

def assert_matches(solv: z3.Solver, clue: Clue, chars: list[z3.AstRef]):
    state_func = self.build_func(solv, clue)
    pat = clue.pattern

    # Create constants for each state
    states = z3.IntVector(clue.name + "_state", 1 + nchar, ctx=solv.ctx)

    # Express the transition requirement
    for i, ch in enumerate(chars):
        solv.add(state_func(states[i], ch) == states[i + 1])

    # Each state must be valid
    for st in states:
        solv.add(0 <= st)
        solv.add(st < pat.nstate)
    # State zero is the initial state
    solv.add(states[0] == 0)

    # The final state must be accepting
    solv.add(z3.Or([
        states[-1] == i
        for i, ok in enumerate(pat.accept) if ok
    ])

整合起来

有了这些原语,接下来只需定义一个字符网格并添加相应的断言即可。我仍以 regexle 作为测试用例来源;它使用六边形网格,而在六边形网格与三条不同线索轴之间做坐标映射,说实话是整个任务中最难的部分之一!不过先略过这一块,代码本身相当简洁:

def make_char(solv: z3.Solver, name: str) -> z3.ExprRef:
    ch = z3.Int(name, solv.ctx)
    solv.add(0 <= ch)
    solv.add(ch < len(ALPHABET))
    return ch

grid = [
    [make_char(solv, f"grid_{x}_{y}") for y in range(maxdim)]
    for x in range(maxdim)
]

for clue in all_clues:
    coords = word_coords(clue.axis, clue.index)
    chars = [grid[x][y] for x, y in coords]

    assert_matches(
        solv,
        clue,
        chars,
    )

求解器性能

到这一步——把所有部分拼好之后——我已经能解出填字游戏了,但速度……慢得令人痛苦。在我的 M1 MacBook Air 上,有些 3×3 的谜题竟然要花上十分钟的 Z3 求解时间!我琢磨出了让 Z3 使用多核的方法,带来了一点常数级别的提速,但从根本上看还是太慢了。

凭着一种模糊的直觉,我原本预期 Z3 应该能在毫秒级内解出这些谜题。因此,我觉得有必要花些时间去调试和理解求解器为何如此缓慢,以及如何加速。

当时我完全是个 Z3 新手,并不清楚该尝试什么、该看什么;在随后的探索中,我学到了非常多的东西。最终我达到了预期的速度,但一路上遇到了不少误导性的线索和意外。

即使在这之前,我也隐约知道 SMT 求解器的性能出了名的晦涩且不稳定。尽管如此,我还是惊讶于这一点竟如此贴切,哪怕在我看来这是个相当简单的问题。我至今仍未完全理解所遇到的一些行为。

接下来我会先分享最终影响最大的两处改动——据我所知,它们造就了一个快速且可靠的求解器——然后再聊聊那些走不通的尝试和其他怪现象。

剪枝状态

如前所述,我不是(当时不是?)Z3 专家,但我对正则表达式相当了解。因此,我的第一个想法是通过领域相关的分析生成额外的约束来帮 Z3 一把。

举个启发性的例子,考虑我写本文时当天的 regexle 中的正则表达式 (NR|Q|I)+,来自当天的 regexle。一眼就能看出答案只能包含 N、R、Q 和 I 这几个字母。然而,我的代码对所有格子都硬编码了 A-Z 的字母表,需要靠 Z3 去排除其他字母。如果我们能检测出实际合法的字母表,并把它编码进 Z3 问题中,会怎么样呢?

人类通过查看正则表达式源文本就能轻易发现这一特性。我们也可以写代码做类似的分析,但一般来说问题并没那么直观。不过我意识到,我们完全可以在已有的转移表上做类似的分析。而且,因为我现在既是工程师又是科学家,我决定用几行 numpy 技巧就搞定它。

我们把 FSM 中不存在任何路径通往接受状态的状态定义为“死状态”。得益于一次 greenery 分析处理,我们处理的是最小化 FSM,因此知道至多只会有一个死状态,且该状态的所有出边都是自环。通过寻找这一条件,我们就能轻易检测出死状态:

# Method on the above class `Regex`
@cached_property
def dead_states(self) -> set[int]:
    looped = self.transition == np.arange(self.nstate)[:, None]
    return set(np.flatnonzero(looped.all(-1) & ~self.accept))

(注意,我们永远只会有零个或一个死状态,但使用 set[int] 可以让我们通过遍历集合来统一处理这两种情况,我觉得这比用 int | None 更简洁。)

有了死状态,我们就可以把死字符定义为:从任意起始状态出发都一定会转移到死状态的字符:

@cached_property
def dead_vocab(self) -> set[int]:
    dead = set()
    for d in self.dead_states:
        dead |= set(np.flatnonzero((self.transition == d).all(0)))
    return dead

我们还可以询问从任意给定状态出发,哪些字符是死的:

def dead_from(self, state: int) -> set[int]:
    dead = set()
    for d in self.dead_states:
        dead |= set(np.flatnonzero((self.transition == d)[state]))
    return dead

有了这些信息,我们就可以给 Z3 再加几条约束:

for d in pat.dead_vocab:
    for ch in chars:
        solv.add(ch != d)

for d in pat.dead_from(0):
    solv.add(chars[0] != d)

for d in pat.dead_states:
    for st in states:
        solv.add(st != d)

严格来说,这些约束是冗余的,但我们希望凭借对问题结构的了解,我们能比 Z3 更轻松地发现这些事实,从而让 Z3 更快地聚焦到搜索空间中“有趣”的部分。

事实上,我觉得这个剪枝练习是一个很有意思的例子,展示了如何使用 Z3 这样的通用证明器/求解器,同时又用领域相关的启发式来提升性能。我的感觉是,这种混合方式在实践中相当常见;求解器并非魔法,如果你能通过领域分析推导出额外的结构,往往能给求解器带来重要的提升。

如果愿意,我完全可以额外让 Z3 在某些具体谜题上验证这些约束确实是冗余的,这会是一个有用的正确性检查。不过对于这个玩具项目,我就没去做了。

显式定义转移函数

在上文中,我描述了将每个正则表达式的转移表编码为 Z3 函数,然后通过对每个 (state, character) 对的逐点断言来确定其行为。后来我发现,如果把转移表示为显式表达式而非未解释函数,性能会更好;我猜测这样做促使 Z3 通过等式推理来处理该关系,而不是回退到搜索,但说实话我也不太确定。

我们要写一个 Python 函数,把转移表表示为显式的 if-then 阶梯,对(状态,字符)输入进行分支判断。我们先从一个辅助函数开始,它实质上是对单个变量做“match”或“switch”语句:

def build_match(
    var: z3.AstRef,
    test: list[z3.AstRef],
    result: list[z3.AstRef],
) -> z3.AstRef:
    """Return a Z3 `if` ladder comparing var against each `test` value.

    If `var == test[i]`, the ladder evaluates to `result[i]`. If no
    `test` matches, evaluate to `result[-1]`; it is anticipated that
    normally the list of tests will be exhaustive.
    """
    expr = result[-1]
    for test_, then_ in zip(test[:-1], result[:-1], strict=True):
        expr = z3.If(var == test_, then_, expr)
    return expr

现在,给定对应于状态和字符的 Z3 表达式,我们就可以构建一个计算新状态的显式表达式:

def build_next_state(
    self, clue: Clue, st: z3.AstRef, ch: z3.AstRef
) -> z3.AstRef:
    pat = clue.pattern
    by_state = [
        build_match(
            ch,
            self.alphabet,
            [self.states[out] for out in pat.transition[i]],
        )
        for i in range(pat.nstate)
    ]

    return build_match(
        st,
        self.states[: pat.nstate],
        by_state,
    )

现在,之前我们写的是 state_func(states[i], ch) == states[i + 1],而我们可以改为为每次转移直接嵌入整个表达式:

for i, ch in enumerate(chars):
    solv.add(self.build_next_state(clue, states[i], ch) == states[i + 1])

这种方法带来了显著的加速,尤其让性能变得更加稳定;在较大的谜题尺寸下,相比旧的逐点断言方式,我遇到的“慢谜题”少了很多。

不幸的是,这种直接的方法有一个明显的缺点:构建庞大的 Z3 if-then 阶梯本身开销很大,以至于构建 Z3 表达式所花的时间远超过求解本身!

共享表达式

我们能兼得两者之长吗?我花了不少时间挖掘和试验,终于找到一种方法,既能获得与显式表达式相同的性能,又只需构建一次表达式。事实上,我找到了两种不同的策略,它们也展示了 Z3 的一些特性。

使用 Z3 lambda 表达式

我们不必使用 Z3 函数对象,而是可以把庞大的 if 阶梯包装在 z3.Lambda 表达式中,这样就能用不同的参数多次显式实例化它。基于上面的 build_next_state 函数,实现这一改动非常小:

st = z3.Int("st")
ch = z3.Int("ch")

lambda_ = z3.Lambda([st, ch], self.build_funcexpr(clue, st, ch))

for i, ch in enumerate(chars):
    solv.add(lambda_[states[i], ch] == states[i + 1])

使用 Z3 函数与 macro-finder

不过,在探索其他方案时,我在阅读SMT-LIB 规范时注意到,SMT-LIB 允许用显式函数体来定义函数:

(define-fun double ((x Int)) Int (* x 2))

;; evaluates to `8`
(simplify (double 4))
SMT-LIB 规范截图,其中将“define-fun”定义为等同于使用“declare-fun”声明具名函数,并通过对输入变量的“forall”来断言其行为。
SMT-LIB 规范 v2.7,第 66 页,定义了 define-fun 的语义

Z3 为我们提供了这些工具,于是我尝试这样使用:

state_expr = self.build_funcexpr(clue, st, ch)
solv.add(z3.ForAll([st, ch], state_func(st, ch) == state_expr))

这种方法在能产生正确解的意义上是可行的,但结果却是我尝试过的方法中最慢的之一!

不过,经过更多的挖掘,我找到了修复办法!Z3 有一个名为“macro-finder” 的 tactic。Z3 的 tactic 是一种变换,允许在核心 SMT 求解器之外进行用户主导的化简或变换处理。macro-finder 实现了多种变换,但最基本的功能是找出那些“定义”函数含义的 forall 断言,并有效地将其原地复制粘贴到该函数的调用处。本质上,它把我们的 forall 变体转换成了“每次都用显式表达式”的变体,但由于它是用 C++ 实现且更贴近 Z3 核心,所以执行得非常高效。

我发现这三种方法(显式表达式、Z3 lambda 以及 forall+macro-finder)在 Z3 中的求解时间相近,而 z3.Lambdamacro-finder 在定义问题的运行时开销上也都同样很快。

失败的尝试

现在,我们来深挖那些行不通、或者至少并非必要的东西。在此过程中,我们会了解到更多 Z3 特性,以及一些令人惊讶的性能表现。如果你愿意,也可以直接跳到我的总结

Z3 EnumSort

Z3 对整数非常了解,并且针对不同类别的整数运算有多种不同的求解器。我在想:能否通过使用其他数据类型,完全避开那些整数逻辑,从而加快求解速度呢?

于是,我尝试把状态和字符换成SMT-LIB 枚举类型(在 Python 中暴露为z3.EnumSort)。

我们可以为状态和字符分别定义一个枚举排序,创建一个除了恰好拥有 N 个不同取值之外没有任何其他行为的新类型:

nstates = max(c.pattern.nstate for c in all_clues)
state_sort, states = z3.EnumSort(
    "State",
    [f"S{i}" for i in range(nstates)],
)

char_sort, alphabet = z3.EnumSort("Char", list(ALPHABET))

其余代码也只需做极小的改动;我们需要在转移函数的声明中把 IntSort 换掉,并且在编码具体状态或字符时,把 i 换成 states[i]alphabet[i]

EnumSort 性能

当我第一次切换到 EnumSort 时,看到了巨大的提速!

Z3 求解时间:整数编码 vs EnumSort。EnumSort 能在不到 0.1 秒内解出所有边长为 3 的谜题,而整数编码至少需要 0.2 秒,最慢可达 3 秒
我早期绘制的边长为 3 的谜题的 Z3 求解时间小提琴图,使用了两种不同的表示方式。这些运行使用了上文所述的剪枝,但仍采用最初的逐点函数定义。

然而,事实证明我对于提速原因的理解充其量只是部分正确。下面是同一张图,但我新增了两个面板:一个仍使用逐点函数定义,但做了一个稍后会提到的小改动,另一个则使用了上文介绍的 z3.Lambda 编码:

同一张图,新增了“逐点(新)”和“z3.Lambda”两个面板。
同一张图,新增了两种转移函数表示方式对应的面板:“逐点(新)”和“z3.Lambda”。

在现有的代码下,使用逐点函数定义时 IntSort 反而更快,而显式编码转移函数带来的提速幅度远超其他任何改动。

一种诡异的性能不稳定性

“(旧)”和“(新)”之间我到底改了什么?

当我们用整数表示状态时,会通过额外的断言来约束整数的范围。上文我展示了最初为这些边界写的代码:

# Each state must be valid
for st in states:
    solv.add(0 <= st)
    solv.add(st < pat.nstate)

当我重构代码以试验其他表示方式时,无意中把行为改成了用所有线索中最大的状态数来约束状态,而不是仅用当前线索的状态数:

max_nstate = max(clue.pattern.nstate for clue in all_clues)

for st in states:
    solv.add(0 <= st)
    solv.add(st < max_nstate)

至少在我使用的 Z3 版本上,第二种方法——这可是更宽松的边界!——却要快得多。我认为,这一改动是“(旧)”和“(新)”的 IntSort 图表之间唯一有意义的差异!

此外,我其实也不明白为什么在逐点函数定义下 EnumSort 会更慢。我查看了 Z3 的跟踪和统计信息,看起来在枚举类型的情况下,Z3 似乎难以通过等式推理来处理 xfer(st, ch) == st_next 这类断言,因而需要更多的搜索和回溯,但我不理解其中的原因。

使用 Z3 正则表达式

在重构以支持多种表示方式的矩阵、并对其进行性能分析和绘图的某个阶段,我还尝试了一种完全不同的方法来解决问题!

事实证明,Z3 本身也有正则表达式理论!我们可以完全绕过状态机那一套,直接把正则线索编码为 Z3 正则表达式。

Z3 并没有自带正则表达式解析器,而是提供了用于组合构建模式的组合子,例如 z3.Re(char)z3.Range(start, end)z3.Star(re) 等等。因此,我写了一个简单的转换层,它递归遍历 greeneryPattern AST,并将其翻译成 Z3 正则表达式。接下来,我们只需把每个未知字符声明为 Z3 字符串对象,然后像这样为每条线索添加断言:

re = greenery_to_z3(clue.pattern)

string = z3.Concat(chars)
solv.add(z3.InRe(string, re))

可以说要简单得多!那它快吗?……还算过得去。

折线图,X 轴为“谜题边长”,Y 轴为“求解时间(秒)”。图中有三条线。我当前最快的求解器在边长为 3 时约为 10 毫秒,到边长为 8 时约为 30 毫秒。使用 Z3 正则表达式的求解器从约 150 毫秒增长到约 11 秒。我最初的慢速代码方差很大,到边长为 5 时就超过 10 秒,这是图中显示的最后一个点。
对比三种变体随谜题尺寸变化的性能。“当前最优”是我目前性能最好的方案,使用 IntSort、剪枝和 z3.Lambda。“逐点,无剪枝”则类似于我最初的实现。

一方面——尤其是对于大尺寸谜题——通过用上本文讨论的所有技巧,我们可以更快地解出谜题。另一方面,使用 Z3 正则表达式的求解器可能是我考虑或实现过的所有方案中最简单的,而且它比我最初那个朴素的实现快了将近 10 倍。

我猜这种模式是普遍适用的:如果 Z3 对你的问题领域提供了一等支持,那就值得从那里入手!不过,Z3 最闪光的地方首先在于它是一个非常通用的工具,能为众多不同类型的问题提供统一的接口;如果你愿意投入精力、做试验并运用领域专业知识,很有可能以增加自身复杂性为代价,更快地求解问题实例。

结论

这是一个很有趣的项目!我本来以为这会是一个轻松愉快的下午小挑战,事实上最初的脚本也只花了一两个小时就草草写好了。然而,我无法摆脱想要优化、想要更深入理解 Z3 的冲动,此后我花了相当多——甚至有点不合理——的时间来重构、探索各种变体、跑基准测试、绘制图表等等。不过,通过这些努力,我现在对 Z3 和 SMT-LIB 的理解都深入了许多,而这毕竟才是我最初的目标!我很期待未来能找到更多在实战中运用它们的机会。

最后,我想分享几点通过这次经历得到的关于使用 Z3 的体会与教训。

Z3 支持的功能远比我意识到的要多

过去我(偶尔)接触 Z3 时,主要把它当作处理整数、位向量、有时还有数组问题的求解器。这次我了解到它支持的许多新数据类型和理论,包括:

  • 字符串、序列和正则表达式
  • 代数数据类型
  • 未解释函数,包括递归函数

我也从未接触过 Z3 的 tactics 系统,它既可以用来解决超出基础 SMT 求解器能力范围的问题,也可以用来为特定类型的问题实现自定义的重写/化简策略。我了解到,我最喜欢的基于 Z3 的工具 Alive2大量使用了 Z3 tactics来针对 Alive2 生成的特定表达式进行优化。

Z3 的性能确实有时非常不稳定、难以预测

如前所述,我之前只是隐约听说过这一特性,但这个项目让我对此有了切身体会。通过调整我为这个问题想出的各种编码和方法的组合,即使不算领域相关的剪枝优化,我也能让求解器的速度相差约 100 倍。而且有些性能表现令人费解,比如仅仅放宽一下整数范围边界就能带来10 倍的提速

使用 Z3 枚举以获得更可预测的性能

这里给出一点关于性能的战术性/具体建议。

如果你要把某个问题编码到 Z3 中,需要表示“ N 选一”的情况,而这些选项没有自然的数值含义(例如,我们不会把它们当作整数来求和),那么我建议声明一个新的枚举排序,而不是仅仅用整数来标记它们。

以我的经验来看,这一改动往往不会带来什么区别,但有时却能通过绕开 Z3 对算术及其他数值特性专门的求解知识,避免神秘的性能不稳定和 10 倍的减速。

不过,关于可用性有一点需要注意:新的排序在每个 Z3 Context 中是全局的,因此如果你(例如)声明了一个名为 “State” 的 EnumSort,就无法在不重启程序或创建全新上下文对象的情况下,用相同的名称和不同的取值集合重新声明另一个。

Z3 的文档参差不齐,但确实存在

起初,我发现很难找到能深入解答 Z3 内部原理、策略或超出表面用法的良好文档,但最终我还是搜集到了一份不错的清单。我认为“请教专家”仍然是使用 Z3 最有效的方式(感谢Hillel Wayne 在我做这个项目期间回答了我的一些问题!),不过这里还有一些我找到的其他优秀资料:


  1. 我在此附上必要的免责声明:许多现代“正则表达式”库支持反向引用等非正则特性。我的策略对使用这些特性的正则填字游戏无效,不过我对此可以接受。↩︎

  2. 有些读者可能会想:“等等!Z3 不是原生就支持正则表达式吗??”当我开始这个项目时,我并不知道这一功能,不过后来我也尝试了一下↩︎

本文章由 muse-spark-1.2-contributor 进行翻译

评论