如何在Python中实现SKI组合子的快速简单类型推断?
嘿,我来帮你搞定SKI组合子的类型推断问题!咱们先理清核心思路,一步步实现你要的两个函数,再解答你的理论疑问,最后结合你提供的基础代码给出完整实现。
一、核心思路:类型推断的本质是合一
SKI组合子的类型推断属于简单类型lambda演算范畴,核心逻辑是给每个组合子分配初始类型模板,然后通过**合一(Unification)**求解类型变量之间的约束。Hindley-Milner算法的核心合一过程完全适用这里,甚至不需要用到它的多态泛化部分(SKI本身是单态的,主类型就是最一般的类型实例)。
二、完整代码实现
我们把类型表示、合一算法、类型推断和你提供的基础代码整合在一起:
from __future__ import annotations import typing from dataclasses import dataclass, field from typing import Dict, Set, Union # ---------------------- 你提供的基础组合子与工具函数 ---------------------- @dataclass(eq=True, frozen=True) class S: def __str__(self): return "S" def __len__(self): return 1 @dataclass(eq=True, frozen=True) class K: def __str__(self): return "K" def __len__(self): return 1 @dataclass(eq=True, frozen=True) class App: left: Term right: Term def __str__(self): return f"({self.left}{self.right})" def __len__(self): return len(str(self)) Term = typing.Union[S, K, App] def parse_ski_string(s): s = ''.join(s.split()) stack = [] for c in s: if c == '(': pass elif c == 'S': stack.append(S()) elif c == 'K': stack.append(K()) elif c == ')': x = stack.pop() if len(stack) > 0: f = stack.pop() stack.append(App(f, x)) else: stack.append(x) else: raise Exception(f'无效字符: {c}') if len(stack) != 1: raise Exception(f'解析失败,栈状态异常: {str(stack)}') return stack[0] def simplify(expr: Term): if isinstance(expr, S) or isinstance(expr, K): return expr elif isinstance(expr, App) and isinstance(expr.left, App) and isinstance(expr.left.left, K): return simplify(expr.left.right) elif isinstance(expr, App) and isinstance(expr.left, App) and isinstance(expr.left.left, App) and isinstance(expr.left.left.left, S): return simplify(App(App(expr.left.left.right, expr.right), App(expr.left.right, expr.right))) elif isinstance(expr, App): l2 = simplify(expr.left) r2 = simplify(expr.right) if expr.left == l2 and expr.right == r2: return App(expr.left, expr.right) else: return simplify(App(l2, r2)) else: raise Exception(f'未知组合子类型: {expr}') # ---------------------- 类型推断核心实现 ---------------------- # 类型表达式表示 @dataclass(eq=True, frozen=True) class TVar: name: str def __str__(self): return self.name @dataclass(eq=True, frozen=True) class Arrow: left: TypeExpr right: TypeExpr def __str__(self): left_str = f"({str(self.left)})" if isinstance(self.left, Arrow) else str(self.left) return f"{left_str} -> {str(self.right)}" TypeExpr = Union[TVar, Arrow] # 生成唯一类型变量的计数器 _var_counter = 0 def fresh_var() -> TVar: global _var_counter var = TVar(f"t{_var_counter}") _var_counter += 1 return var def occurs_check(var: TVar, expr: TypeExpr) -> bool: """检查类型变量是否出现在类型表达式中(防止循环合一)""" if isinstance(expr, TVar): return var == expr elif isinstance(expr, Arrow): return occurs_check(var, expr.left) or occurs_check(var, expr.right) return False def apply_subst(expr: TypeExpr, subst: Dict[TVar, TypeExpr]) -> TypeExpr: """将替换规则应用到类型表达式上""" if isinstance(expr, TVar): return subst.get(expr, expr) elif isinstance(expr, Arrow): return Arrow(apply_subst(expr.left, subst), apply_subst(expr.right, subst)) return expr def unify(a: TypeExpr, b: TypeExpr, subst: Dict[TVar, TypeExpr]) -> Dict[TVar, TypeExpr]: """合一两个类型表达式,返回替换规则;无法合一则抛出异常""" a = apply_subst(a, subst) b = apply_subst(b, subst) if isinstance(a, TVar): if a == b: return subst if occurs_check(a, b): raise ValueError("循环类型约束,无法类型化") subst[a] = b return subst elif isinstance(b, TVar): return unify(b, a, subst) elif isinstance(a, Arrow) and isinstance(b, Arrow): subst = unify(a.left, b.left, subst) return unify(a.right, b.right, subst) else: raise ValueError(f"无法合一类型: {a} 和 {b}") def infer_type(term: Term) -> TypeExpr: """推断SKI项的主类型,无法类型化则抛出异常""" subst = {} def _infer(t: Term) -> TypeExpr: nonlocal subst if isinstance(t, S): # S的标准类型模板: (α→β→γ) → (α→β) → α→γ α, β, γ = fresh_var(), fresh_var(), fresh_var() return Arrow(Arrow(α, Arrow(β, γ)), Arrow(Arrow(α, β), Arrow(α, γ))) elif isinstance(t, K): # K的标准类型模板: α→β→α α, β = fresh_var(), fresh_var() return Arrow(α, Arrow(β, α)) elif isinstance(t, App): # 应用约束:函数类型为A→B,参数类型为A,结果类型为B func_type = _infer(t.left) arg_type = _infer(t.right) result_var = fresh_var() subst = unify(func_type, Arrow(arg_type, result_var), subst) return apply_subst(result_var, subst) else: raise ValueError(f"未知组合子: {t}") raw_type = _infer(term) return apply_subst(raw_type, subst) # ---------------------- 你需要的两个函数 ---------------------- def typable(term: Term) -> bool: try: infer_type(term) return True except ValueError: return False def principle_type(term: Term) -> Union[TypeExpr, bool]: try: return infer_type(term) except ValueError: return False # ---------------------- 测试验证 ---------------------- if __name__ == "__main__": SKK = App(App(S(), K()), K()) SII = App(App(S(), SKK), SKK) # I = SKK print(f"typable(SKK) = {typable(SKK)}") # 输出: True print(f"typable(SII) = {typable(SII)}") # 输出: False print(f"principle_type(S) = {principle_type(S())}") print(f"principle_type(K) = {principle_type(K())}") print(f"principle_type(SK) = {principle_type(App(S(), K()))}") print(f"principle_type(SKK) = {principle_type(SKK)}")
三、解答你的理论疑问
1. Hindley-Milner与System F的关系
你的理解完全正确:Algorithm J和W是Hindley-Milner算法的实现,后者确实支持System F的参数多态。而SKI组合子属于简单类型lambda演算,是System F的子集。
确实存在System F中可类型化但简单类型系统中不行的组合子,比如Y递归组合子:在简单类型系统中,它的类型会要求t → t和t合一,导致循环约束无法满足;但在System F中,我们可以用多态类型∀t. (t→t)→t来表示它。
2. 用SMT求解器简化算法
完全可以!类型推断的约束本质是一阶逻辑等式,Z3这类SMT求解器能高效处理这类合一问题。对于SKI这种简单场景,自己实现合一算法更轻量、速度更快;但如果要扩展到更复杂的类型系统(比如带交集类型、递归类型),用SMT求解器可以省去自己处理复杂约束的麻烦,代码会更简洁。
比如用Z3实现的话,你可以把每个类型变量映射成Z3的变量,箭头类型表示为函数,然后添加S、K的类型模板和应用的类型约束,最后检查是否有解,并提取解作为主类型。
内容的提问来源于stack exchange,提问作者Oleg Dats

