Some content in this article was created with AI assistance. Please verify as needed.
问题与目标 本文用一个很小的 Python 模型说明形式化证明中的几类对象:表达式表示数学项,命题由表达式组成,证明是带有规则和子证明的树,tactic 只是构造这棵树的程序。在这个 demo 中,checker 独立检查显式 Proof 树。
在 Lean 中,kernel 检查 tactic 生成的 proof term 是基于类型系统的,而不是这个示例程序中的 Proof 树。
示例脚本处理如下代数问题:设 $0<p,q,r<1$,并且
$$
(1-p^2)(1-q^2)(1-r^2)=8p^2q^2r^2, \tag{1}
$$
分别证明
$$
1<p+q+r,
\qquad
p+q+r<2.
$$
注意这里的做法是形式化证明,不是将 $p,q,r$ 赋浮点数然后抽样验证结论。
表达式、命题和证明都被组织为树状的数据,再由一个小型 checker 逐步检查:
1 2 3 4 5 6 7 数学写法 ↓ Expr / Prop 抽象语法树 ↓ 规则或 tactic 构造 Proof 树 ↓ Kernel 递归检查
这个 demo 与 Lean 的最本质的差异,可能就在于它完全没有采用类型论中的 “命题是类型、证明是项” 的 Curry–Howard 结构,而是选择了“显式证明树 + 外部规则检查”的架构。在这个 demo 中,Prop 只是保存命题语法的普通数据,Proof 也只是等待 check() 检查的证明树;但是在 Lean 中,一个证明是相应命题类型的项。
下面的代码需要一些头文件
1 2 3 4 5 6 7 8 9 """A one-file Lean-style proof demo for two algebraic inequalities.""" from __future__ import annotationsfrom dataclasses import dataclassfrom fractions import Fractionfrom itertools import productfrom pathlib import Pathfrom typing import Any
表达式 Expr 首先需要实现一个最基础的表达式类型 Expr,它的实例就是一棵简单的抽象语法树(AST)。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 @dataclass(frozen=True , eq=False ) class Expr : op: str args: tuple [Any , ...] @staticmethod def var (name: str ) -> Expr: return Expr("var" , (name,)) @staticmethod def const (value: int | Fraction ) -> Expr: return Expr("const" , (Fraction(value),)) def __add__ (self, other ): return Expr("add" , (self , as_expr(other))) def __radd__ (self, other ): return as_expr(other) + self def __neg__ (self ): return Expr("neg" , (self ,)) def __sub__ (self, other ): return self + -as_expr(other) def __rsub__ (self, other ): return as_expr(other) - self def __mul__ (self, other ): return Expr("mul" , (self , as_expr(other))) def __rmul__ (self, other ): return as_expr(other) * self def __pow__ (self, exponent: int ): if not isinstance (exponent, int ) or exponent < 0 : raise ValueError("only non-negative integer powers are supported" ) return Expr("pow" , (self , exponent)) def __lt__ (self, other ): return Prop("<" , self , as_expr(other)) def __le__ (self, other ): return Prop("≤" , self , as_expr(other)) def __gt__ (self, other ): return Prop("<" , as_expr(other), self ) def __ge__ (self, other ): return Prop("≤" , as_expr(other), self ) def __eq__ (self, other ): return Prop("=" , self , as_expr(other)) def __hash__ (self ) -> int : raise TypeError("formal expressions are intentionally unhashable" ) def __str__ (self ): return show_expr(self ) def as_expr (value ) -> Expr: """Convert an integer to a constant expression.""" if isinstance (value, Expr): return value if isinstance (value, (int , Fraction)): return Expr.const(value) raise TypeError(f"cannot use {type (value).__name__} as an expression" ) def show_expr (expr: Expr, outer=0 ) -> str : """Format an expression with minimal parentheses.""" if expr.op == "var" : return expr.args[0 ] if expr.op == "const" : value = expr.args[0 ] return str (value.numerator) if value.denominator == 1 else f"({value} )" precedence = {"add" : 1 , "mul" : 2 , "neg" : 3 , "pow" : 3 }[expr.op] if expr.op == "add" : left, right = expr.args text = ( f"{show_expr(left, 1 )} - {show_expr(right.args[0 ], 2 )} " if right.op == "neg" else f"{show_expr(left, 1 )} + {show_expr(right, 1 )} " ) elif expr.op == "mul" : text = f"{show_expr(expr.args[0 ], 2 )} * {show_expr(expr.args[1 ], 2 )} " elif expr.op == "neg" : text = f"-{show_expr(expr.args[0 ], 3 )} " else : text = f"{show_expr(expr.args[0 ], 3 )} ^{expr.args[1 ]} " return f"({text} )" if precedence < outer else text
下面是一些说明:
用形式节点支持未知量,例如 Expr.var("p") 不表示一个等待赋值的普通 Python 变量,而表示形式节点 Var("p"),未知量可以和已知的整数分数等一起运算。
as_expr() 负责把 Python 整数和 Fraction 提升为形式常数,后续的多项式运算不会引入浮点误差。
show_expr() 只影响输出展示,不参与证明正确性。它根据运算优先级决定是否加括号,提供更接近普通数学写法的展示。
加法、乘法和幂的运算符重载不做数值计算。例如 1 - p**2 内部保存成 Add(Const(1), Neg(Pow(Var("p"), 2)))。比较运算也不返回 Python bool,而是构造后面定义的 Prop。
一个小例子:如果 p = Expr.var("p"),那么 1 - p**2 对应的结构是
1 2 3 1 - p**2 = add(const(1), neg(pow(var("p"), 2)))
命题 Prop 我们还需要一个命题类型 Prop,让每一个命题都是它的实例。
代码实现只需要三个成员,方法的实现不是重点,例如 __bool__ 是保护性的,__str__ 只是展示用途。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 @dataclass(frozen=True , eq=False ) class Prop : kind: str left: Any = None right: Any = None def __bool__ (self ): raise TypeError("a formal proposition is not a Python bool" ) def __str__ (self ): if self .kind == "not" : return f"¬({self.left} )" if self .kind == "false" : return "False" return f"{self.left} {self.kind} {self.right} "
这里采用了简化实现,只支持下面两类命题:
由两个表达式生成命题:例如 Prop("<", left, right) 表示 left < right,同理还有其他二元比较运算符(为了省事,后面的相关代码只实现了对 =,< 和 <= 的支持)
由命题生成命题:Prop("not", proposition) 表示 not proposition
例如对于表达式 p(类型是 Expr),p < 1 会构造出一个命题:
1 Prop(kind="<", left=p, right=1)
这个命题只是在说“目标形状是 p 小于 1”,并没有说明这个命题为什么成立,甚至这个命题本身可能就是错误的,还可能是无法判断的。换言之,命题自身只携带自己的内容,不携带自己的真假信息,也不支持对此做出判断。
注意 0 < p 的类型是 Prop 而不是 bool,而且直接调用 bool(0 < p) 会抛出异常。
多项式运算 我们的命题需要大量的多项式运算,因此主要支持这类表达式运算就够了。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 def poly (expr: Expr ): """Normalize to {monomial: coefficient}; a monomial is ((name, power), ...).""" if expr.op == "const" : return {} if expr.args[0 ] == 0 else {(): expr.args[0 ]} if expr.op == "var" : return {((expr.args[0 ], 1 ),): Fraction(1 )} if expr.op == "neg" : return scale(poly(expr.args[0 ]), -1 ) if expr.op == "add" : return add_poly(poly(expr.args[0 ]), poly(expr.args[1 ])) if expr.op == "mul" : return mul_poly(poly(expr.args[0 ]), poly(expr.args[1 ])) result = {(): Fraction(1 )} for _ in range (expr.args[1 ]): result = mul_poly(result, poly(expr.args[0 ])) return result def scale (value, coefficient ): coefficient = Fraction(coefficient) return {} if coefficient == 0 else {m: coefficient * c for m, c in value.items()} def add_poly (left, right ): result = dict (left) for monomial, value in right.items(): result[monomial] = result.get(monomial, 0 ) + value if result[monomial] == 0 : del result[monomial] return result def mul_poly (left, right ): result = {} for lm, lc in left.items(): for rm, rc in right.items(): powers = {} for name, exponent in lm + rm: powers[name] = powers.get(name, 0 ) + exponent monomial = tuple (sorted (powers.items())) result[monomial] = result.get(monomial, 0 ) + lc * rc return {m: c for m, c in result.items() if c}
其中 poly() 函数负责对多项式进行标准化处理,具体来说,就是把表达式展开为字典 {单项式: 系数}。例如 $8p^2q$ 的表示为:
1 {(("p", 2), ("q", 1)): Fraction(8)}
递归分支分别处理常数、变量、负号、加法、乘法和非负整数幂。这样 (1-p)*(1+p) 和 1-p**2 虽然 AST 不同,却可以归到相同正规形。
1 2 poly((1 - p) * (1 + p)) = {(): 1, (("p", 2),): -1} poly(1 - p**2) = {(): 1, (("p", 2),): -1}
后面几个函数则是负责基于这个字典进行多项式运算(多项式加法、数乘、乘法),返回的仍然是这个形式的字典。
证明 Proof 我们还需要一个证明类型 Proof。
这里只有几个成员有意义,剩下的方法都只是附带的:tree() 展开缩进树,size() 递归统计节点数。二者都只做展示用途。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 @dataclass(frozen=True ) class Proof : proposition: Prop rule: str premises: tuple [Proof, ...] = () certificate: tuple [Fraction, ...] | None = None name: str | None = None def tree (self, indent="" ) -> str : """Render the proof tree as indented text.""" label = f" [{self.name} ]" if self .name else "" lines = [f"{indent} {self.rule} {label} : {self.proposition} " ] lines += [premise.tree(indent + " " ) for premise in self .premises] return "\n" .join(lines) def size (self ) -> int : """Return the number of nodes in the proof tree.""" return 1 + sum (premise.size() for premise in self .premises)
Proof 实例的成员包括:
名称代号
当前证明的结论(一个命题)
规则 rule,代表从前提推出当前节点的依据,但是也包括假设(rule = "assumption")
若干的子证明组成的前提元组 premises
若干的证书组成的证书元组 certificate
关于规则和证书的内容,见下面的 kernel 部分。
一个证明对象和它的前提元组中的若干证明对象(也就是当前证明需要的前提)一起组成树结构。我们希望证明的命题通常就是树的根节点那个证明对象的结论,但是显然我们要从叶子一步步推到根,才能逻辑严密。
树最底层的叶子节点通常是规则为假设的证明对象,例如一个假设 $0 < p$,对应如下的证明对象
1 2 3 4 5 6 Proof( proposition = 0 < p, rule = "assumption", premises = (), name = "hp0", )
下面是可视化的辅助函数 proof_tree_text(),它会把嵌套的 tuple 存储的证明对象在 proof_trees.txt 中渲染树形文本。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 def proof_tree_text (proofs: tuple [tuple [str , Proof], ...] ) -> str : """Render named proof trees as readable indented text.""" lines = [] def visit (node: Proof, prefix="" , branch="" , is_root=False ): label = f"{node.rule} " if node.name: label += f" [{node.name} ]" lines.append(f"{prefix} {branch} {label} : {node.proposition} " ) child_prefix = prefix if is_root else prefix + (" " if branch == "`-- " else "| " ) for index, premise in enumerate (node.premises): is_last = index == len (node.premises) - 1 visit(premise, child_prefix, "`-- " if is_last else "|-- " ) for index, (name, proof) in enumerate (proofs): if index: lines.append("" ) lines.append(name) visit(proof, is_root=True ) return "\n" .join(lines) + "\n"
下面会关注 Proof 树的组装,也就是将一些基础的证明步骤标准化,还包括常见的证明策略。注意这里的组装只是按照我们认为是对的推导步骤,虽然在组装时我们也进行了基本的校验,但是内核还是会进行独立检查,重新验证。
基础规则构造器 正性规则 1 2 def positive_const (n ): return Proof(Expr.const(0 ) < n, "positive_const" )
例如给定一个常数 1,positive_const(1) 构造得到的 Proof 对应的树状结构如下
1 2 3 4 def mul_pos (a, b ): return Proof( Expr.const(0 ) < a.proposition.right * b.proposition.right, "mul_pos" , (a, b) )
例如已有两个证明 hp0 : 0 < p 和 hq0 : 0 < q 时,mul_pos(hp0, hq0) 构造得到的 Proof 对应的树状结构如下
1 2 3 mul_pos: 0 < p * q <rule> [hp0]: 0 < p <rule> [hq0]: 0 < q
1 2 3 4 def sub_pos (h ): return Proof( Expr.const(0 ) < h.proposition.right - h.proposition.left, "sub_pos" , (h,) )
例如已有证明 hq1 : q < 1,sub_pos(hq1) 构造得到的 Proof 对应的树状结构如下
1 2 sub_pos: 0 < 1 - q <rule> [hq1]: q < 1
不等式传递与同乘 1 2 def lt_trans (a, b ): return Proof(a.proposition.left < b.proposition.right, "lt_trans" , (a, b))
例如已有 hpq : p < q 和 hqr : q < r,lt_trans(hpq, hqr) 构造得到的 Proof 对应的树状结构如下
1 2 3 lt_trans: p < r <rule> [hpq]: p < q <rule> [hqr]: q < r
1 2 3 4 5 6 7 def mul_lt_right (h, hp ): return Proof( h.proposition.left * hp.proposition.right < h.proposition.right * hp.proposition.right, "mul_lt_right" , (h, hp), )
例如已有 hpq : p < q 和 hr0 : 0 < r,mul_lt_right(hpq, hr0) 构造得到的 Proof 对应的树状结构如下
1 2 3 mul_lt_right: p * r < q * r <rule> [hpq]: p < q <rule> [hr0]: 0 < r
1 2 3 4 5 6 7 def mul_lt_left (h, hp ): return Proof( hp.proposition.right * h.proposition.left < hp.proposition.right * h.proposition.right, "mul_lt_left" , (h, hp), )
例如已有 hpq : p < q 和 hr0 : 0 < r,mul_lt_left(hpq, hr0) 构造得到的 Proof 对应的树状结构如下
1 2 3 mul_lt_left: r * p < r * q <rule> [hpq]: p < q <rule> [hr0]: 0 < r
改写、反设和矛盾构造 1 2 def eq_then_lt (eq, lt ): return Proof(eq.proposition.left < lt.proposition.right, "eq_then_lt" , (eq, lt))
例如已有 hpq : p = q 和 hqr : q < r,eq_then_lt(hpq, hqr) 构造得到的 Proof 对应的树状结构如下
1 2 3 eq_then_lt: p < r <rule> [hpq]: p = q <rule> [hqr]: q < r
1 2 def lt_then_eq (lt, eq ): return Proof(lt.proposition.left < eq.proposition.right, "lt_then_eq" , (lt, eq))
例如已有 hpq : p < q 和 hqr : q = r,lt_then_eq(hpq, hqr) 构造得到的 Proof 对应的树状结构如下
1 2 3 lt_then_eq: p < r <rule> [hpq]: p < q <rule> [hqr]: q = r
1 2 3 def le_of_not_gt (h ): negated = h.proposition.left return Proof(negated.right <= negated.left, "le_of_not_gt" , (h,))
例如已有证明 h : ¬(p < q),le_of_not_gt(h) 构造得到的 Proof 对应的树状结构如下
1 2 le_of_not_gt: q ≤ p <rule> [h]: ¬(p < q)
1 2 def contradiction (eq, lt ): return Proof(Prop("false" ), "eq_lt_false" , (eq, lt))
例如已有 hpq : p = q 和 hpq_lt : p < q,contradiction(hpq, hpq_lt) 构造得到的 Proof 对应的树状结构如下
1 2 3 eq_lt_false: False <rule> [hpq]: p = q <rule> [hpq_lt]: p < q
1 2 def by_contra (goal, false_proof ): return Proof(goal, "by_contra" , (false_proof,))
例如目标是 p < q,并且反设 h : ¬(p < q) 后已经构造出下面的矛盾,by_contra(p < q, hfalse) 会把它包装成目标的 Proof:
1 2 3 4 by_contra: p < q eq_lt_false: False <rule> [heq]: p = q <rule> [hpq]: p < q
三个简单 tactic 除了基础规则,这里还准备了三个需要用到的 tactic:ring、positivity、linarith。
它们读取当前目标和已有证明,计算出新的 Proof 节点,但不会让命题直接变成真。
ring ring 做一次多项式正规形比较。如果两边正规形不同,它直接失败抛异常;若成功,返回一个等式 Proof。
1 2 3 4 5 def ring (left, right ): """Build a polynomial equality proof and let the kernel recheck it.""" if poly(left) != poly(right): raise ValueError("ring failed" ) return Proof(left == right, "ring" )
例如 ring(1 - p**2, (1 - p) * (1 + p)),传入的是两个表达式,成功生成的是一个结论为等式的 Proof 节点,规则标记为 ring。构造得到的 Proof 对应的树状结构如下
1 ring: 1 - p^2 = (1 - p) * (1 + p)
这里为了简化,将 ring 的规则也直接放入 checker,因此 ring 看起来就不像 tactic 而是更像内置规则了。但是在真实 Lean 中,ring 是生成证明项的 tactic。
positivity 1 2 3 4 5 6 7 8 9 10 def positivity (expr, known ): """Prove positivity by reusing assumptions and splitting products.""" for proof in known: if prop_key(proof.proposition) == prop_key(Expr.const(0 ) < expr): return proof if expr.op == "const" and expr.args[0 ] > 0 : return positive_const(expr.args[0 ]) if expr.op == "mul" : return mul_pos(positivity(expr.args[0 ], known), positivity(expr.args[1 ], known)) raise ValueError(f"positivity cannot prove 0 < {expr} " )
positivity() 沿乘法 AST 递归,在 hs 中寻找已有的正性证明,并为正常数构造 positive_const。结果是一棵由已有证明、positive_const 和 mul_pos 组成的树。
例如 hs 中已有证明 hp0 : 0 < p 和 hr0 : 0 < r,positivity(2 * p * r, hs) 构造得到的 Proof 对应的树状结构如下
1 2 3 4 5 mul_pos: 0 < 2 * p * r mul_pos: 0 < 2 * p positive_const: 0 < 2 <rule> [hp0]: 0 < p <rule> [hr0]: 0 < r
linarith 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 def linarith (goal, premises ): """Search coefficients 0, 1, and 2 for a valid linear certificate. This is a deliberately small Lean-style model, not full ``linarith``. """ for weights in product((0 , 1 , 2 ), repeat=len (premises)): total, strict = {}, False for weight, premise in zip (weights, premises): total = add_poly( total, scale( add_poly( poly(premise.proposition.right), scale(poly(premise.proposition.left), -1 ), ), weight, ), ) strict |= premise.proposition.kind == "<" and weight > 0 if total == add_poly(poly(goal.right), scale(poly(goal.left), -1 )) and ( goal.kind == "≤" or strict ): return Proof(goal, "linear" , tuple (premises), tuple (map (Fraction, weights))) raise ValueError(f"tiny linarith found no certificate for {goal} " )
这里的极简版 linarith 通过枚举系数 $0,1,2$ 来寻找线性组合证书。它先把每个前提统一写成“右边减左边”的形式:
$$
d_i=\text{premise.right}_i-\text{premise.left}_i.
$$
例如 $a<b$ 会变成 $0<b-a$。外层循环枚举每个前提对应的 weight,内层循环用 add_poly() 和 scale() 计算
$$
\text{total}=\sum_i \text{weight}_i d_i.
$$
如果 total 恰好等于目标的“右边减左边”,这些 weight 就构成一个证书。strict 记录线性组合中是否实际使用了严格不等式:证明 < 时至少要有一个系数为正的严格不等式,证明 ≤ 则没有这个要求。找到证书后,tactic 把前提和系数保存进 Proof;如果所有组合都不匹配,才抛出异常。
例如已有 hqr : 0 < q * (1 - r) 和 hrq : 0 < r * (1 - q),linarith(2 * q * r < q + r, (hqr, hrq)) 构造得到的 Proof 对应的树状结构如下
1 2 3 linear: 2 * q * r < q + r <rule> [hqr]: 0 < q * (1 - r) <rule> [hrq]: 0 < r * (1 - q)
tactic 会把 (1, 1) 存入 certificate,Kernel 再做一次同样的线性组合检查。
这里的证书可以理解为 tactic 提交给 Kernel 的验证提示或补充信息。(1, 1) 明确表示“第一个前提乘 1,再加上第二个前提乘 1”。Kernel 不需要再重新搜索系数,只需按照证书复算并核对结果;如果证书与前提或目标不匹配,整个 Proof 就会被拒绝。
一些准备 由于 Expr.__eq__(也就是 ==)已被用来构造数学等式,不能再承担表达式的比较,我们需要单独实现函数来比较表达式之间或者命题之间是否相同,做法也很简单,直接把内容拆成 tuple,后面对比 tuple 即可比较。
1 2 3 4 5 6 7 def expr_key (expr: Expr ): if expr.op in {"var" , "const" }: return (expr.op, expr.args[0 ]) if expr.op == "pow" : return ("pow" , expr_key(expr.args[0 ]), expr.args[1 ]) return (expr.op, *(expr_key(arg) for arg in expr.args))
这里用 expr_key() 把 AST 转为只含 tuple、字符串和有理数的 key,便于后面用这些 key 核对命题。
注意这里比较的是结构完全相同,而不是代数等价。例如:
1 2 expr_key(p + q) 和 expr_key(p + q) 相同 expr_key(p + q) 和 expr_key(q + p) 不同
与之不同的是,前面定义的 poly() 函数已经根据多项式的特点拆分 tuple 了,所以可以识别简单的代数等价性。
1 poly(p + q) 和 poly(q + p) 相同
下面是把同样的做法迁移到命题上,把命题拆解为 tuple 便于进行结构性的比较。
1 2 3 4 5 6 def prop_key (prop: Prop ): if prop.kind in {"=" , "<" , "≤" }: return prop.kind, expr_key(prop.left), expr_key(prop.right) if prop.kind == "not" : return "not" , prop_key(prop.left) return (prop.kind,)
kernel 部分也需要一些辅助,主要是统一错误类型和报错处理,其中:
require() 函数要求 condition 为真,否则报错。
expect() 函数要求两个命题是完全一致的,否则报错。
1 2 3 4 5 6 7 8 9 10 11 class KernelError (ValueError ): pass def require (condition, rule ): if not condition: raise KernelError(f"invalid use of {rule} " ) def expect (actual, expected ): require(prop_key(actual) == prop_key(expected), "wrong conclusion" )
Kernel check 终于来到了最核心的证明验证的部分,前面的部分只是把命题和证明构造出来,现在才是检查证明是否通过的核心环节。
check() 函数是当前简单实现的证明检查内核,它读取证明对象,返回证明是否通过。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 def check (candidate: Proof, assumptions=( ) ) -> bool : """Recursively check a proof tree against the declared assumptions.""" allowed = tuple (prop_key(prop) for prop in assumptions) def visit (node: Proof, local=allowed ): prop, rule, ps = node.proposition, node.rule, node.premises if rule == "assumption" : if prop_key(prop) not in local: raise KernelError(f"unavailable assumption: {prop} " ) return if rule == "by_contra" : require( len (ps) == 1 and prop.kind == "<" and ps[0 ].proposition.kind == "false" , rule, ) visit(ps[0 ], local + (prop_key(Prop("not" , prop)),)) return for premise in ps: visit(premise, local) if rule == "positive_const" : require( not ps and prop.kind == "<" and expr_key(prop.left) == expr_key(Expr.const(0 )) and prop.right.op == "const" and prop.right.args[0 ] > 0 , rule, ) elif rule == "mul_pos" : ps0_positive = ps[0 ].proposition.kind == "<" and expr_key( ps[0 ].proposition.left ) == expr_key(Expr.const(0 )) ps1_positive = ps[1 ].proposition.kind == "<" and expr_key( ps[1 ].proposition.left ) == expr_key(Expr.const(0 )) require(len (ps) == 2 and ps0_positive and ps1_positive, rule) expect( prop, Expr.const(0 ) < ps[0 ].proposition.right * ps[1 ].proposition.right ) elif rule == "sub_pos" : require(len (ps) == 1 and ps[0 ].proposition.kind == "<" , rule) expect( prop, Expr.const(0 ) < ps[0 ].proposition.right - ps[0 ].proposition.left ) elif rule == "lt_trans" : require(len (ps) == 2 and all (p.proposition.kind == "<" for p in ps), rule) require( expr_key(ps[0 ].proposition.right) == expr_key(ps[1 ].proposition.left), rule, ) expect(prop, ps[0 ].proposition.left < ps[1 ].proposition.right) elif rule in {"mul_lt_right" , "mul_lt_left" }: ps1_positive = ps[1 ].proposition.kind == "<" and expr_key( ps[1 ].proposition.left ) == expr_key(Expr.const(0 )) require( len (ps) == 2 and ps[0 ].proposition.kind == "<" and ps1_positive, rule ) relation, term = ps[0 ].proposition, ps[1 ].proposition.right expected = ( (relation.left * term < relation.right * term) if rule.endswith("right" ) else (term * relation.left < term * relation.right) ) expect(prop, expected) elif rule == "ring" : require( not ps and prop.kind == "=" and poly(prop.left) == poly(prop.right), rule, ) elif rule == "linear" : certificate = node.certificate require(prop.kind in {"<" , "≤" } and certificate is not None , rule) if certificate is None : raise KernelError("linear proof has no certificate" ) require(len (certificate) == len (ps), rule) total, strict = {}, False for weight, premise in zip (certificate, ps): relation = premise.proposition require(relation.kind in {"=" , "<" , "≤" }, rule) require(relation.kind == "=" or weight >= 0 , rule) total = add_poly( total, scale( add_poly(poly(relation.right), scale(poly(relation.left), -1 )), weight, ), ) strict |= relation.kind == "<" and weight > 0 require( total == add_poly(poly(prop.right), scale(poly(prop.left), -1 )) and (prop.kind == "≤" or strict), rule, ) elif rule in {"eq_then_lt" , "lt_then_eq" }: require(len (ps) == 2 , rule) eq, lt = ( (ps[0 ].proposition, ps[1 ].proposition) if rule == "eq_then_lt" else (ps[1 ].proposition, ps[0 ].proposition) ) require(eq.kind == "=" and lt.kind == "<" , rule) if rule == "eq_then_lt" : require(expr_key(eq.right) == expr_key(lt.left), rule) expect(prop, eq.left < lt.right) else : require(expr_key(lt.right) == expr_key(eq.left), rule) expect(prop, lt.left < eq.right) elif rule == "le_of_not_gt" : require( len (ps) == 1 and ps[0 ].proposition.kind == "not" and ps[0 ].proposition.left.kind == "<" , rule, ) negated = ps[0 ].proposition.left expect(prop, negated.right <= negated.left) elif rule == "eq_lt_false" : require(len (ps) == 2 and prop.kind == "false" , rule) eq, lt = ps[0 ].proposition, ps[1 ].proposition require(eq.kind == "=" and lt.kind == "<" , rule) endpoints = {expr_key(eq.left), expr_key(eq.right)} require(endpoints == {expr_key(lt.left), expr_key(lt.right)}, rule) else : raise KernelError(f"unknown rule: {rule} " ) visit(candidate) return True
check() 函数的核心行为如下:
递归检查 Proof 树:先确认每个子证明能通过,再根据当前节点的 rule 重新计算它应该得到的结论。
假设节点必须出现在 allowed 中;把任意命题包装成 Proof(prop, "assumption") 不能凭空获得结论。
by_contra 是唯一会临时扩展局部假设的地方:检查反证子树时加入 goal 的否定作为假设,离开子树后这个假设消失。
主体的 if/elif 分支就是各条规则的检查逻辑。mul_pos 要求两个前提都是正性证明,mul_lt_left 和 mul_lt_right 额外要求乘数为正,ring 会重新比较两边多项式正规形。linear 则复查 tactic 给出的系数证书:
$$
\sum_i c_i(\text{right}_i-\text{left}_i)
=\text{goal.right}-\text{goal.left}.
$$
这里对应 Lean 的核心分工:自动化可以负责搜索,但检查器仍然坚持独立审查。只要求最终 Proof 树的每一步都能被 check() 独立复算通过,给出证明的过程本身是不必被信任的。
这里为了简化模型,把 kernel checker 的代码写得比各种 tratic 更复杂。
但是实际对于 Lean 来说,kernel 很小,tactic 的代码反而可以很复杂,而且容易扩展,但是 kernel 完全不信任 tactic。
上面的代码虽然是demo性质的,但是并不是面向具体问题的,下面终于进入了这两个具体问题所对应的代码部分。
公共假设与小引理 problem_assumptions() 创建六个范围假设和等式 (1),并为每个假设附上名字。后面的 hp0、hq1、hprod 都是这些假设对应的 Proof 叶子。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 def problem_assumptions (p, q, r ): assumptions = ( (0 < p, "hp0" ), (p < 1 , "hp1" ), (0 < q, "hq0" ), (q < 1 , "hq1" ), (0 < r, "hr0" ), (r < 1 , "hr1" ), ( (1 - p**2 ) * (1 - q**2 ) * (1 - r**2 ) == 8 * p**2 * q**2 * r**2 , "hprod" , ), ) return tuple (Proof(prop, "assumption" , name=name) for prop, name in assumptions)
square_positive() 将 $1-x^2$ 写成 $(1-x)(1+x)$,再由 $0<x<1$ 证明两个因子为正,得到 $0<1-x^2$。
1 2 3 4 def square_positive (x, hx0, hx1 ): factors = mul_pos(sub_pos(hx1), linarith(0 < 1 + x, (positive_const(1 ), hx0))) return lt_then_eq(factors, ring((1 - x) * (1 + x), 1 - x**2 ))
sum_pair() 从 $x(1-y)>0$、$y(1-x)>0$ 得 $2xy<x+y$。这个小引理用于第一个目标。
1 2 3 4 5 def sum_pair (x, y, hx0, hx1, hy0, hy1 ): a = mul_pos(hx0, sub_pos(hy1)) b = mul_pos(hy0, sub_pos(hx1)) return linarith(2 * x * y < x + y, (a, b))
product_pair() 从 $(1-x)(1-y)>0$ 得 $x+y-1<xy$。这个小引理用于第二个目标。
1 2 3 def product_pair (x, y, hx1, hy1 ): return linarith(x + y - 1 < x * y, (mul_pos(sub_pos(hx1), sub_pos(hy1)),))
第一个目标 第一个目标函数 prove_goal1 如下
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 def prove_goal1 (p, q, r, hs ): hp0, hp1, hq0, hq1, hr0, hr1, hprod = hs hp_sq, hq_sq = square_positive(p, hp0, hp1), square_positive(q, hq0, hq1) h2pr = positivity(2 * p * r, hs) h2pq = positivity(2 * p * q, hs) goal = 1 < p + q + r hnot = Proof(Prop("not" , goal), "assumption" , name="h" ) hsum = le_of_not_gt(hnot) hqr = sum_pair(q, r, hq0, hq1, hr0, hr1) hpr = sum_pair(p, r, hp0, hp1, hr0, hr1) hpq = sum_pair(p, q, hp0, hp1, hq0, hq1) hp1a = linarith(2 * q * r < 1 - p, (hqr, hsum)) hq1a = linarith(2 * p * r < 1 - q, (hpr, hsum)) hr1a = linarith(2 * p * q < 1 - r, (hpq, hsum)) hp = linarith(2 * q * r < 1 - p**2 , (hp1a, mul_pos(hp0, sub_pos(hp1)))) hq = linarith(2 * p * r < 1 - q**2 , (hq1a, mul_pos(hq0, sub_pos(hq1)))) hr = linarith(2 * p * q < 1 - r**2 , (hr1a, mul_pos(hr0, sub_pos(hr1)))) first = lt_trans(mul_lt_right(hp, h2pr), mul_lt_left(hq, hp_sq)) triple = lt_trans(mul_lt_right(first, h2pq), mul_lt_left(hr, mul_pos(hp_sq, hq_sq))) expanded = (2 * q * r) * (2 * p * r) * (2 * p * q) strict = eq_then_lt(ring(8 * p**2 * q**2 * r**2 , expanded), triple) return by_contra(goal, contradiction(hprod, strict))
第一个目标使用反证法。hnot 表示 $
eg(1<p+q+r)$,le_of_not_gt() 得到 $p+q+r\le1$。三次调用 sum_pair() 后,通过线性组合得到
$$
2qr<1-p,
\quad 2pr<1-q,
\quad 2pq<1-r.
$$
因为 $p(1-p)>0$,有 $(1-p)+p(1-p)=1-p^2>1-p$,另外两组同理。于是代码中的 hp、hq、hr 分别证明
$$
2qr<1-p^2,
\quad 2pr<1-q^2,
\quad 2pq<1-r^2.
$$
三个不等式分两次相乘。每次调用 mul_lt_left 或 mul_lt_right 都显式提供正性证明,不能无条件把两个严格不等式相乘。ring 将左侧整理为 $8p^2q^2r^2$,所得严格不等式与 hprod 冲突,by_contra 返回目标 Proof。
第二个目标 首先是辅助函数 square_lt
1 2 3 4 5 6 7 def square_lt (x, pair, canonical, hprime, hx0, hx1, pair_pos ): """Multiply the bounds to derive 1-x^2 < 2*pair.""" first = mul_lt_right(hprime, linarith(0 < 1 + x, (positive_const(1 ), hx0))) second = mul_lt_left(linarith(1 + x < 2 , (hx1,)), pair_pos) chain = lt_trans(first, second) chain = eq_then_lt(ring(1 - x**2 , (1 - x) * (1 + x)), chain) return lt_then_eq(chain, ring(pair * 2 , canonical))
square_lt() 把 $1-x<pair$ 和 $1+x<2$ 同乘正数,得到 $1-x^2<2\cdot pair$。参数 canonical 用来指定最终想要的乘法顺序,再由 ring 做改写。
然后是第二个目标函数 prove_goal2
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 def prove_goal2 (p, q, r, hs ): hp0, hp1, hq0, hq1, hr0, hr1, hprod = hs hq_sq = square_positive(q, hq0, hq1) hr_sq = square_positive(r, hr0, hr1) h2qr, h2pr = positivity(2 * q * r, hs), positivity(2 * p * r, hs) goal = p + q + r < 2 hnot = Proof(Prop("not" , goal), "assumption" , name="h" ) hsum = le_of_not_gt(hnot) hp1a = linarith(1 - p < q * r, (hsum, product_pair(q, r, hq1, hr1))) hq1a = linarith(1 - q < p * r, (hsum, product_pair(p, r, hp1, hr1))) hr1a = linarith(1 - r < p * q, (hsum, product_pair(p, q, hp1, hq1))) hp = square_lt(p, q * r, 2 * q * r, hp1a, hp0, hp1, positivity(q * r, hs)) hq = square_lt(q, p * r, 2 * p * r, hq1a, hq0, hq1, positivity(p * r, hs)) hr = square_lt(r, p * q, 2 * p * q, hr1a, hr0, hr1, positivity(p * q, hs)) first = lt_trans(mul_lt_right(hp, hq_sq), mul_lt_left(hq, h2qr)) triple = lt_trans(mul_lt_right(first, hr_sq), mul_lt_left(hr, mul_pos(h2qr, h2pr))) expanded = (2 * q * r) * (2 * p * r) * (2 * p * q) strict = lt_then_eq(triple, ring(expanded, 8 * p**2 * q**2 * r**2 )) return by_contra(goal, contradiction(hprod, strict))
第二个目标的反设是 $2\le p+q+r$。由 $(1-q)(1-r)>0$ 得 $q+r-1<qr$,结合反设可得 $1-p<qr$。循环处理三组变量后,分两次相乘得到
$$
(1-p^2)(1-q^2)(1-r^2)<8p^2q^2r^2,
$$
这同样与 hprod 矛盾。
运行最终检查 main 函数如下
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 def main (): p = Expr.var("p" ) q = Expr.var("q" ) r = Expr.var("r" ) hs = problem_assumptions(p, q, r) goal1_proof = prove_goal1(p, q, r, hs) goal2_proof = prove_goal2(p, q, r, hs) print ("Kernel-checked proofs:" ) print (f" First goal: {goal1_proof.proposition} " ) print (f" check: {check(goal1_proof, tuple (h.proposition for h in hs))} " ) print (f" root rule = {goal1_proof.rule} " ) print (f" proof nodes = {goal1_proof.size()} " ) print (f" Second goal: {goal2_proof.proposition} " ) print (f" check: {check(goal2_proof, tuple (h.proposition for h in hs))} " ) print (f" root rule = {goal2_proof.rule} " ) print (f" proof nodes = {goal2_proof.size()} " ) tree_path = Path(__file__).with_name("proof_trees.txt" ) tree_path.write_text( proof_tree_text((("First goal" , goal1_proof), ("Second goal" , goal2_proof))), encoding="utf-8" , newline="\n" , ) print (f"\nProof trees written to: {tree_path} " ) return goal1_proof, goal2_proof if __name__ == "__main__" : goal1_proof, goal2_proof = main()
运行输出为:
1 2 3 4 5 6 7 8 9 10 11 Kernel-checked proofs: First goal: 1 < p + q + r status: True root rule = by_contra proof nodes = 97 Second goal: p + q + r < 2 status: True root rule = by_contra proof nodes = 115 Proof trees written to: proof_trees.txt
完整证明树写入的 proof_trees.txt 内容很大,局部例如:
1 2 3 4 mul_pos: 0 < q * (1 - r) |-- assumption [hq0]: 0 < q `-- sub_pos: 0 < 1 - r `-- assumption [hr1]: r < 1
补充:demo 的局限及其与 Lean 的差异 这个 demo 只是一段面向代数不等式的微型模型。它的对象语言只有多项式表达式、等式、不等式、否定和 False,没有函数、量词、定义展开、归纳类型、类型类、隐式参数和 universe。因此它能表达的问题非常少,也没有 Lean 中“命题是类型、证明是项”的完整类型论结构。
这里的 Kernel 也比 Lean 粗很多。check() 直接把 ring、linear、mul_pos 等规则当作内建规则检查;真实 Lean 的 kernel 不认识这些 tactic 名字,而是检查 tactic 最终生成的低层 proof term。也就是说,Lean 的自动化可以很复杂,但最后仍要落到更小的一组核心类型检查规则上。
tactic 部分同样是教学版。linarith 只枚举系数 0, 1, 2,ring 只正规化当前 AST 支持的多项式,positivity 只会拆乘法和读取已有正性假设。真实 Lean 的相应 tactic 会处理更广的语法、上下文和归约规则,也会生成更复杂的证明项。
所以这份脚本不试图复刻 Lean。它保留的是一条最关键的主线:证明自动化负责寻找或拼装证明,可信检查器负责独立复查证明。理解这条分界线,再看 Lean 里的 ring、linarith、positivity,就不会把 tactic 的“算出来了”和定理的“被证明了”混为一谈。