# LTNtorch : memes faits, vraie librairie (pip install LTNtorch)
import warnings
warnings.filterwarnings("ignore", message=".*pynvml.*") # bruit cosmetique torch.cuda/Windows
import torch
import ltn
torch.manual_seed(42)
# Memes individus et memes faits que la section 4 (reutilises du kernel)
n = len(entities)
emb = torch.eye(n)
# Negatifs closed-world pour grandparent : tout sauf (Marie, Paul).
# On inclut (Marie, Pierre) -- un VRAI parent -- comme negatif difficile.
gp_negative = [('Marie', 'Pierre'), ('Pierre', 'Marie'), ('Sophie', 'Paul'),
('Paul', 'Luc'), ('Luc', 'Sophie'), ('Pierre', 'Paul')]
class PairPredicate(torch.nn.Module):
"""Predicat binaire : MLP sur la concatenation des embeddings (s, o).
Contrairement a notre NeuralPredicate (1 couche, gradient manuel),
n'importe quelle architecture differentiable convient : l'autograd
de PyTorch se charge du calcul du gradient a travers les operateurs flous.
"""
def __init__(self):
super().__init__()
self.net = torch.nn.Sequential(
torch.nn.Linear(2 * n, 16), torch.nn.ELU(),
torch.nn.Linear(16, 1), torch.nn.Sigmoid())
def forward(self, s, o):
return self.net(torch.cat([s, o], dim=-1)).squeeze(-1)
Parent = ltn.Predicate(PairPredicate())
Grandparent = ltn.Predicate(PairPredicate())
# Semantique floue : les memes operateurs que notre section 2, version librairie
Not = ltn.Connective(ltn.fuzzy_ops.NotStandard())
And = ltn.Connective(ltn.fuzzy_ops.AndProd()) # t-norme produit
Implies = ltn.Connective(ltn.fuzzy_ops.ImpliesReichenbach())
Forall = ltn.Quantifier(ltn.fuzzy_ops.AggregPMeanError(p=2), quantifier="f")
sat_agg = ltn.fuzzy_ops.SatAgg() # agregation des axiomes
def pairs_to_vars(pairs, prefix):
"""Transforme une liste de paires (sujet, objet) en variables LTN batchees."""
s = ltn.Variable(prefix + "_s", torch.stack([emb[entities[a]] for a, b in pairs]))
o = ltn.Variable(prefix + "_o", torch.stack([emb[entities[b]] for a, b in pairs]))
return s, o
pos_s, pos_o = pairs_to_vars(positive, "pos") # faits parent (section 4)
neg_s, neg_o = pairs_to_vars(negative, "neg") # negatifs closed-world (section 4)
gpn_s, gpn_o = pairs_to_vars(gp_negative, "gpn")
# Variables libres sur TOUS les individus : le Forall de l'axiome 3 est
# evalue sur le produit cartesien complet (5x5x5 = 125 triplets) par
# broadcasting -- la ou notre jouet ne traitait qu'une paire codee en dur.
x = ltn.Variable("x", emb)
y = ltn.Variable("y", emb)
z = ltn.Variable("z", emb)
params = list(Parent.parameters()) + list(Grandparent.parameters())
opt = torch.optim.Adam(params, lr=0.05)
print("=== Entrainement LTNtorch (4 axiomes, satisfaction a maximiser) ===")
for epoch in range(400):
opt.zero_grad()
axioms = [
# 1-2. Faits parent : positifs vrais, negatifs faux (ltn.diag = paires
# alignees element par element, pas de produit cartesien)
Forall(ltn.diag(pos_s, pos_o), Parent(pos_s, pos_o)),
Forall(ltn.diag(neg_s, neg_o), Not(Parent(neg_s, neg_o))),
# 3. LA regle : aucun fait grandparent n'est fourni, seule cette
# contrainte universelle relie les deux predicats
Forall([x, y, z],
Implies(And(Parent(x, y), Parent(y, z)), Grandparent(x, z))),
# 4. Negatifs closed-world pour grandparent (sinon "tout est
# grandparent" satisferait trivialement l'implication)
Forall(ltn.diag(gpn_s, gpn_o), Not(Grandparent(gpn_s, gpn_o))),
]
sat = sat_agg(*axioms)
loss = 1.0 - sat
loss.backward()
opt.step()
if (epoch + 1) % 100 == 0:
print(f" Epoch {epoch+1:3d} : satisfaction globale = {sat.item():.4f}")
def query(P, a, b):
"""Interroge un predicat appris sur une paire d'individus."""
return P(ltn.Constant(emb[entities[a]]), ltn.Constant(emb[entities[b]])).value.item()
print()
print("=== Predicat parent (faits fournis, comme section 4) ===")
for s, o in positive:
print(f" parent({s}, {o}) = {query(Parent, s, o):.3f} [positif]")
for s, o in negative[:3]:
print(f" parent({s}, {o}) = {query(Parent, s, o):.3f} [negatif]")
print()
print("=== Predicat grandparent (AUCUN fait fourni, appris via la regle) ===")
print(f" grandparent(Marie, Paul) = {query(Grandparent, 'Marie', 'Paul'):.3f} [seul vrai grandparent]")
for s, o in [('Marie', 'Pierre'), ('Sophie', 'Paul'), ('Luc', 'Marie')]:
print(f" grandparent({s}, {o}) = {query(Grandparent, s, o):.3f} [faux]")