Refactor terms module

dev
Bram van den Heuvel 2026-07-09 22:59:14 +02:00
parent fa9d4f0a83
commit eeac39a4fa
3 changed files with 360 additions and 164 deletions

View File

@ -4,30 +4,32 @@
# from ..proof import Context, Proof # from ..proof import Context, Proof
from ..props import Eq, Or, Prop from ..props import Eq, Or, Prop
from ..terms import App, Const, EqT, Lambda, Term, Var, register_known_const from ..terms import App, Const, A1, A2, L1, L2, L3, Lambda, Term, Var, register_known_const
# ----------------------------------------------------------------------------- # -----------------------------------------------------------------------------
# Identity function # Identity function
# ----------------------------------------------------------------------------- # -----------------------------------------------------------------------------
# Always returns its input. # Always returns its input.
register_known_const("identity", Lambda(Var(0))) id_func = Lambda(lambda x : x)
register_known_const("id", id_func)
register_known_const("identity", id_func)
# ----------------------------------------------------------------------------- # # -----------------------------------------------------------------------------
# Equality # # Equality
# ----------------------------------------------------------------------------- # # -----------------------------------------------------------------------------
equality = Lambda(Lambda(EqT(lhs=Var(1), rhs=Var(0)))) # equality = Lambda(Lambda(EqT(lhs=Var(1), rhs=Var(0))))
register_known_const("eq", equality) # register_known_const("eq", equality)
register_known_const("==", equality) # register_known_const("==", equality)
# ----------------------------------------------------------------------------- # -----------------------------------------------------------------------------
# Bool # Bool
# ----------------------------------------------------------------------------- # -----------------------------------------------------------------------------
bool_t = Lambda(Lambda(Var(1))) bool_t = L2(lambda x, y : x).normalize()
bool_f = Lambda(Lambda(Var(0))) bool_f = L2(lambda x, y : y).normalize()
if_statement = Lambda(Var(0)) if_statement = L3(lambda x, y, z : A2(x, y, z)).normalize()
not_statement = Lambda(App(f=App(f=Var(0), x=bool_f), x=bool_t)) not_statement = L1(lambda b : A2(b, bool_f, bool_t)).normalize()
register_known_const("Bool.False", bool_f) register_known_const("Bool.False", bool_f)
register_known_const("Bool.True", bool_t) register_known_const("Bool.True", bool_t)
@ -50,25 +52,15 @@ def kernel_from_int(n : int) -> Term:
raise ValueError( raise ValueError(
"Int is not Nat" "Int is not Nat"
) )
elif n == 0:
return zero
else:
return App(f=succ, x=kernel_from_int(n-1))
def is_nat(x : Term) -> Prop: t = zero
return Or( for _ in range(n):
Eq(x=x, y=zero), t = A1(succ, t)
Eq(bool_f, bool_t) # TODO: Construct x = Succ n where isNat(n),
)
register_known_const("Nat.add", Lambda(Lambda( return t
App(f=App( # If-statement: is the first argument zero?
f=EqT(lhs=Var(1), rhs=zero), # def is_nat(x : Term) -> Prop:
# If so, return the second argument # return Or(
x=Var(0), # Eq(x=x, y=zero),
), # Eq(bool_f, bool_t) # TODO: Construct x = Succ n where isNat(n),
# If not, return a recursive solution of the add function # )
# TODO: Not correct yet
x=App(f=succ, x=App(f=App(f=Const("Nat.add"), x=Var(1)), x=Var(0)))
)
)))

View File

@ -238,9 +238,10 @@ def reflexivity(eq : Eq) -> Proof:
statement=eq, statement=eq,
) )
else: else:
diff_l, diff_r = l.reduce_on_equality(r) # diff_l, diff_r = l.reduce_on_equality(r)
raise ProofCheckerFailedException( raise ProofCheckerFailedException(
f"Could not normalize {eq.x} = {eq.y}, got stuck at {diff_l} = {diff_r}" f"Could not normalize {eq.x} = {eq.y}, got stuck at {l} = {r}"
# f"Could not normalize {eq.x} = {eq.y}, got stuck at {diff_l} = {diff_r}"
) )
def suppose(ctx : Context, proof : Proof) -> Proof: def suppose(ctx : Context, proof : Proof) -> Proof:

View File

@ -4,43 +4,170 @@
""" """
from __future__ import annotations from __future__ import annotations
from typing import Callable
import copy import copy
from dataclasses import dataclass from dataclasses import dataclass
# Known const values that can be substituted into other values.
_known_consts : dict[str, Term] = {} _known_consts : dict[str, Term] = {}
# Known functions to normalize certain specific terms with given shapes.
_known_norms : list[Callable[[Term], Term]] = []
USABLE_VAR_NAMES = [
'x', 'y', 'z', 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l',
'm', 'n', 'o', 'p', 'q', 'r', 's', 't', 'u', 'v', 'w',
'α', 'β', 'γ', 'δ', 'ε', 'ζ', 'η', 'θ', 'ι', 'κ',
# 'λ', # Avoid confusion
'μ', 'ν', 'ξ', 'ο', 'π', 'ρ', 'σ', 'τ', 'υ', 'φ', 'χ', 'ψ', 'ω'
]
def normalize(term : Term) -> Term:
"""
Create a normalization function. This function simplifies functions as
much as possible.
"""
def norm_iter(term : Term) -> Term:
"""
Normalize a single term. Trust that all children values have
already been normalized, or that they will be normalized later on.
:param term: The term to normalize.
:type term: Term
:return: A normalized version of the term.
:rtype: Term
"""
match term:
case App():
match term.f:
case Lambda():
return term.f.substitute(value=norm_iter(term.x))
case _:
return term
case Const():
return _known_consts.get(term.name, term)
case Lambda():
return term
case Var():
return term
case Term():
return term
return traverse(f=norm_iter, term=term)
def register_known_const(name : str, value : Term) -> None: def register_known_const(name : str, value : Term) -> None:
_known_consts[name] = value _known_consts[name] = value
def traverse(f : Callable[[Term], Term], term : Term) -> Term:
"""
Traverse through the tree back-and-forth, aiming to run the given
function as thoroughly as possible to manipulate the term tree.
:param f: The traversal function that replaces terms.
:type f: Callable[[Term], Term]
:param term: The term to traverse.
:type term: Term
:return: The result of the traversed term.
:rtype: Term
"""
return traverse_bottom_up(f=f, term=traverse_top_down(f=f, term=term))
def traverse_bottom_up(f : Callable[[Term], Term], term : Term) -> Term:
"""
Traverse a term from the bottom up, replacing values with whatever the
inserted function returns.
:param f: The traversal function that replaces terms.
:type f: Callable[[Term], Term]
:param term: The term to traverse.
:type term: Term
:return: The result of the traversed term.
:rtype: Term
"""
match term:
case App():
return f(App(
f=traverse_bottom_up(f=f, term=term.f),
x=traverse_bottom_up(f=f, term=term.x)),
)
case Const():
return f(term)
case Lambda():
return term.replace_content(
traverse_bottom_up(f=f, term=term.out)
)
case Var():
return f(term)
case Term():
return f(term)
def traverse_top_down(f : Callable[[Term], Term], term : Term) -> Term:
"""
Traverse a term from the top down, replacing values with whatever the
inserted function returns.
:param f: The traversal function that replaces terms.
:type f: Callable[[Term], Term]
:param term: The term to traverse.
:type term: Term
:return: The result of the traversed term.
:rtype: Term
"""
t = f(term)
match t:
case App():
return App(
f=traverse_top_down(f=f, term=t.f),
x=traverse_top_down(f=f, term=t.x),
)
case Const():
return t
case Lambda():
return t.replace_content(traverse_top_down(f=f, term=t.out))
case Var():
return t
case Term():
return t
class Term: class Term:
""" """
Base class for all terms. Base class for all terms.
""" """
def __eq__(self, other) -> bool:
"""
Determine whether two terms are equal.
"""
return str(self) == str(other)
def __repr__(self) -> str:
"""
Create a representation of a term.
"""
return self.to_str(vars={})
def normalize(self) -> Term: def normalize(self) -> Term:
""" return normalize(term=self)
Function to return a simplified version of the term.
"""
return self
def reduce_on_equality(self, other : Term) -> tuple[Term, Term]:
"""
Reduce parts that the two items are equal on.
"""
return self, other
def substitute(self, index : int, value : Term) -> Term:
"""
Substitute all variables of a given index number.
:param index: The variable's index. def to_str(self, vars : dict[str, Var]) -> str:
:type index: int return "<undefined>"
:param value: The value to replace the variable with.
:type value:
"""
return self
@dataclass(frozen=True) @dataclass(frozen=True)
class App(Term): class App(Term):
@ -52,61 +179,26 @@ class App(Term):
f : Term f : Term
x : Term x : Term
def normalize(self) -> Term: def __repr__(self) -> str:
f = self.f.normalize()
x = self.x.normalize()
default = App(f=f, x=x)
match f:
case App():
return default
case Const():
return default
case Lambda():
return f.resolve(x)
case Term():
return default
case Var():
return default
def reduce_on_equality(self, other: Term) -> tuple[Term, Term]:
if not isinstance(other, App):
return self, other
f1, f2 = self.f.normalize(), other.f.normalize()
x1, x2 = self.x.normalize(), other.x.normalize()
match ( f1 == f2, x1 == x2 ):
case ( True, True ):
return Term(), Term()
case ( True, False ):
return x1.reduce_on_equality(x2)
case ( False, True ):
return f1.reduce_on_equality(f2)
case ( False, False ):
return self, other
def substitute(self, index : int, value : Term) -> Term:
""" """
Substitute all variables of a given index number. Create a representation of a term.
:param index: The variable's index.
:type index: int
:param value: The value to replace the variable with.
:type value:
""" """
return App( return self.to_str(vars={})
f=self.f.substitute(index=index, value=value),
x=self.x.substitute(index=index, value=value), def to_str(self, vars: dict[str, Var]) -> str:
) f_str = self.f.to_str(vars=vars)
x_str = self.x.to_str(vars=vars)
return f_str + " " + (f"({x_str})" if " " in x_str else x_str)
A1 = lambda fn, a : App(f=fn, x=a)
A2 = lambda fn, a, b : App(f=A1(fn, a), x=b)
A3 = lambda fn, a, b, c : App(f=A2(fn, a, b), x=c)
A4 = lambda fn, a, b, c, d : App(f=A3(fn, a, b, c), x=d)
A5 = lambda fn, a, b, c, d, e : App(f=A4(fn, a, b, c, d), x=e)
A6 = lambda fn, a, b, c, d, e, f : App(f=A5(fn, a, b, c, d, e), x=f)
A7 = lambda fn, a, b, c, d, e, f, g : App(f=A6(fn, a, b, c, d, e, f), x=g)
A8 = lambda fn, a, b, c, d, e, f, g, h : App(f=A7(fn, a, b, c, d, e, f, g), x=h)
@dataclass(frozen=True) @dataclass(frozen=True)
class Const(Term): class Const(Term):
@ -117,90 +209,201 @@ class Const(Term):
name : str name : str
def normalize(self) -> Term: def __repr__(self) -> str:
if self.name in _known_consts: """
return copy.deepcopy(_known_consts.get(self.name, self)) Create a representation of a term.
else: """
return self return self.to_str(vars={})
@dataclass(frozen=True) def to_str(self, vars: dict[str, Var]) -> str:
class EqT(Term): return self.name
"""
Equality operator. Possibly the only function that is implemented into
the core of the term definitions.
"""
lhs : Term
rhs : Term
def normalize(self) -> Term:
l = self.lhs.normalize()
r = self.rhs.normalize()
if l == r:
return Const("Bool.True")
else:
return Const("Bool.False")
@dataclass(frozen=True)
class Lambda(Term): class Lambda(Term):
""" """
Lambdas a nameless 1-ary functions. Lambdas a nameless 1-ary functions.
""" """
out : Term out : Term
var : Var
def normalize(self) -> Term: def __init__(self, f : Callable[[Var], Term]) -> None:
return Lambda(out=self.out.normalize()) self.var = Var()
self.out = f(self.var)
def __repr__(self) -> str:
"""
Create a representation of a term.
"""
return self.to_str(vars={})
def reduce_on_equality(self, other: Term) -> tuple[Term, Term]: def replace_content(self, term : Term) -> Lambda:
match other:
case Lambda():
return self.out, other.out
case _:
return self, other
def resolve(self, value : Term) -> Term:
""" """
Resolve this lambda function by inserting a value. Replace the content of the lambda function with the given term.
The new term may still contain the lambda's var as an argument.
:param value: The value to substitute. :param term: The term to insert in a lambda function.
:type term: Term
:return: The new Lambda function.
:rtype: Lambda
""" """
return self.substitute(index=-1, value=value) return Lambda(
lambda x : traverse_bottom_up(
f=lambda t : x if t is self.var else t,
term=term,
)
)
def substitute(self, index : int, value : Term) -> Term: def substitute(self, value : Term) -> Term:
""" """
Substitute all variables of a given index number. Substitute all variables of a given index number.
:param index: The variable's index.
:type index: int
:param value: The value to replace the variable with. :param value: The value to replace the variable with.
:type value: :type value: Term
:return: A resolved lambda where the term has been substituted.
:rtype: Term
""" """
if index == -1: return traverse_bottom_up(
# We're resolving this lambda function! f=lambda t : value if t is self.var else t,
return self.out.substitute(index=index + 1, value=value) term=self.out,
)
def to_str(self, vars: dict[str, Var]) -> str:
char = ''
for c in USABLE_VAR_NAMES:
if c not in vars:
char = c
vars[c] = self.var
break
else: else:
return Lambda(out=self.out.substitute(index=index + 1, value=value)) raise ValueError(
f"Cannot represent a lambda function that's over {len(USABLE_VAR_NAMES)} layers deep!"
)
s = self.out.to_str(vars=vars)
del vars[char]
return "λ" + char + "." + (f"({s})" if " " in s else s)
L1 = lambda f : Lambda(lambda x : f(x))
L2 = lambda f : Lambda(lambda x : L1(lambda y : f(x, y)))
L3 = lambda f : Lambda(lambda x : L2(lambda y, z : f(x, y, z)))
L4 = lambda f : Lambda(lambda x : L3(lambda y, z, a : f(x, y, z, a)))
L5 = lambda f : Lambda(lambda x : L4(lambda y, z, a, b : f(x, y, z, a, b)))
L6 = lambda f : Lambda(lambda x : L5(lambda y, z, a, b, c : f(x, y, z, a, b, c)))
L7 = lambda f : Lambda(lambda x : L6(lambda y, z, a, b, c, d : f(x, y, z, a, b, c, d)))
L8 = lambda f : Lambda(lambda x : L7(lambda y, z, a, b, c, d, e : f(x, y, z, a, b, c, d, e)))
@dataclass(frozen=True) @dataclass(frozen=True)
class Var(Term): class Var(Term):
""" """
A variable is a value that can be substituted by a lambda function. A variable is a value that can be substituted by a lambda function.
Each lambda function stores a unique var, which can then point to a new
structure later on.
""" """
index : int def __repr__(self) -> str:
def substitute(self, index : int, value : Term) -> Term:
""" """
Substitute all variables of a given index number. Create a representation of a term.
:param index: The variable's index.
:type index: int
:param value: The value to replace the variable with.
:type value:
""" """
if self.index == index: return self.to_str(vars={})
return copy.deepcopy(value)
def to_str(self, vars: dict[str, Var]) -> str:
for i, v in vars.items():
if v is self:
return i
else: else:
return self return "[ERROR:UNKNOWN_VAR]"
if __name__ == "__main__":
print("REPRESENTATIONS")
print(30 * "=")
# x
print(Const("x"))
# λx.(f x)
print(Lambda(lambda a : App(f=Const("f"), x=a)))
# λx.(λy.(x y))
print(Lambda(lambda b : Lambda(lambda c : App(f=b, x=c))))
# Nat.Succ (Nat.Succ Nat.Zero)
print(App(f=Const("Nat.Succ"), x=App(f=Const("Nat.Succ"), x=Const("Nat.Zero"))))
print("\nNORMALIZATIONS")
print(30 * "=")
# foo
a = App(f=Lambda(lambda x : x), x=Const("foo"))
print(a)
print(a.normalize())
# f x
b = App(
f=App(
f=Lambda(
lambda f : Lambda(
lambda x : App(
f=f, x=x
)
)
),
x=Const("f"),
),
x= Const("x"),
)
print(b)
print(b.normalize())
# a b c
c = App(
f=App(
f=App(
f=Lambda(
lambda x : Lambda(
lambda y : Lambda(
lambda z : App(
f=App(f=x, x=y), x=z
)
)
)
),
x=Const("a"),
),
x=Const("b"),
),
x=Const("c"),
)
print(c)
print(c.normalize())
#
d = L3(lambda x, y, z : A2(x, y, z))
print(d)
print(d.normalize())
print("\nSHORTHAND NOTATION")
print(30 * "=")
# λx.(λy.(λz.(x y z))
e = L3(lambda a, b, c : App(f=App(f=a, x=b), x=c))
print(e)
# λx.(λy.(λz.(x y z))
f = L3(lambda a, b, c : A3(L3(lambda a, b, c : A2(a, b, c)), a, b, c))
print(f)
print(f.normalize())
assert e == f.normalize()
# not
bool_t = L2(lambda x, y : x).normalize()
bool_f = L2(lambda x, y : y).normalize()
g = L1(lambda b : A2(b, bool_f, bool_t))
print(g)
print(g.normalize())