import numpy as np
if not JAX_AVAILABLE:
print("Note : la section BP ne depend pas de JAX, elle fonctionne sous NumPy pur.")
# 27 unites (9 lignes, 9 colonnes, 9 blocs)
BP_UNITS = []
for r in range(9):
BP_UNITS.append([r * 9 + c for c in range(9)])
for c in range(9):
BP_UNITS.append([r * 9 + c for r in range(9)])
for br in range(0, 9, 3):
for bc in range(0, 9, 3):
BP_UNITS.append([(br + i) * 9 + (bc + j) for i in range(3) for j in range(3)])
# une paire de cellules partageant une unite = un facteur binaire !=
_bp_pairs = set()
for unit in BP_UNITS:
for a in range(9):
for b in range(a + 1, 9):
_bp_pairs.add((min(unit[a], unit[b]), max(unit[a], unit[b])))
BP_PAIRS = np.array(sorted(_bp_pairs))
BP_FI, BP_FJ = BP_PAIRS[:, 0], BP_PAIRS[:, 1]
BP_F = len(BP_PAIRS)
BP_EPS = 1e-12
# valeur "exclue" par un voisin : message < EXCL_EPS
BP_EXCL_EPS = 1e-8
def _bp_normalize(x):
x = np.maximum(x, BP_EPS)
return x / x.sum(axis=-1, keepdims=True)
def verify_solution_flat(sol81):
"""Verifie les 27 unites d'une grille plate (0..80, valeurs 1..9)."""
g = np.asarray(sol81)
for unit in BP_UNITS:
if sorted(g[np.array(unit)].tolist()) != list(range(1, 10)):
return False
return True
class BeliefPropagationSolver:
"""Sum-product loopy BP + decimation par marge (v3 Python, etage 1).
Tous les messages vivent en espace de probabilites (vecteurs de taille 9
normalises). Facteur != : m_{f->i}(v) = 1 - m_{j->f}(v). Cellule -> facteur :
croyance totale divisee par le message inverse (en espace log). Amortissement
(damping) sur les deux directions pour stabiliser le graphe boucle.
"""
def __init__(self, damping=0.5, init_sweeps=50, decim_sweeps=3, max_restarts=4):
self.damping = damping
self.init_sweeps = init_sweeps
self.decim_sweeps = decim_sweeps
self.max_restarts = max_restarts
def solve(self, grid81, seed=0):
"""Renvoie (grille plate ou None, stats)."""
rng = np.random.default_rng(seed)
stats = {'decisions': 0, 'sweeps': 0, 'contradictions': 0}
for restart in range(self.max_restarts):
res = self._attempt(np.array(grid81, dtype=int), rng, stats)
if res is not None:
stats['restarts'] = restart
return res, stats
stats['restarts'] = self.max_restarts
return None, stats
def _attempt(self, grid, rng, stats):
m_f2c = np.full((BP_F, 2, 9), 1.0 / 9)
evidence = np.full((81, 9), 1.0 / 9)
assigned = grid > 0
for i in np.flatnonzero(assigned):
evidence[i, :] = BP_EPS
evidence[i, grid[i] - 1] = 1.0
evidence[i] = evidence[i] / evidence[i].sum()
def sweep(msgs, n):
m_f2c, m_c2f = msgs
for _ in range(n):
new_f2c = np.empty((BP_F, 2, 9))
new_f2c[:, 0] = _bp_normalize(1.0 - m_c2f[:, 1])
new_f2c[:, 1] = _bp_normalize(1.0 - m_c2f[:, 0])
m_f2c = _bp_normalize(self.damping * m_f2c + (1 - self.damping) * new_f2c)
totals = np.log(evidence + BP_EPS).copy()
np.add.at(totals, BP_FI, np.log(m_f2c[:, 0] + BP_EPS))
np.add.at(totals, BP_FJ, np.log(m_f2c[:, 1] + BP_EPS))
new_c2f = np.empty((BP_F, 2, 9))
for s, fs in ((0, BP_FI), (1, BP_FJ)):
raw = totals[fs] - np.log(m_f2c[:, s] + BP_EPS)
new_c2f[:, s] = np.exp(raw - raw.max(axis=1, keepdims=True))
new_c2f = _bp_normalize(new_c2f)
m_c2f = _bp_normalize(self.damping * m_c2f + (1 - self.damping) * new_c2f)
stats['sweeps'] += 1
beliefs = np.log(evidence + BP_EPS).copy()
np.add.at(beliefs, BP_FI, np.log(m_f2c[:, 0] + BP_EPS))
np.add.at(beliefs, BP_FJ, np.log(m_f2c[:, 1] + BP_EPS))
beliefs = np.exp(beliefs - beliefs.max(axis=1, keepdims=True))
excluded = np.zeros((81, 9))
np.add.at(excluded, BP_FI, (m_f2c[:, 0] < BP_EXCL_EPS))
np.add.at(excluded, BP_FJ, (m_f2c[:, 1] < BP_EXCL_EPS))
return beliefs, excluded, (m_f2c, m_c2f)
msgs = (m_f2c, np.full((BP_F, 2, 9), 1.0 / 9))
beliefs, excluded, msgs = sweep(msgs, self.init_sweeps)
while not assigned.all():
# contradiction : cellule vide dont les 9 valeurs sont exclues
if (excluded.min(axis=1)[~assigned] >= 1).any():
stats['contradictions'] += 1
return None
part = np.partition(beliefs, -2, axis=1)
margins = np.where(np.isfinite(part[:, -1] - part[:, -2]),
part[:, -1] - part[:, -2], -np.inf)
margins[assigned] = -np.inf
top = np.flatnonzero(margins >= margins.max() - 1e-12)
i = int(rng.choice(top)) if len(top) > 1 else int(top[0])
v = int(np.argmax(beliefs[i]))
evidence[i, :] = BP_EPS
evidence[i, v] = 1.0
assigned[i] = True
stats['decisions'] += 1
beliefs, excluded, msgs = sweep(msgs, self.decim_sweeps)
out = grid.copy()
for i in range(81):
if grid[i] == 0:
out[i] = int(np.argmax(evidence[i])) + 1
return out
class BPBacktrackingSolver:
"""BP + repli (v3 Python, etage 2) : decisions revocables.
Au chaque noeud : quelques sweeps BP (messages herites du parent = warm
start), detection de contradiction par exclusions, branchement sur la
cellule a la marge la plus nette, valeurs tentees par croyance decroissante
en excluant celles deja exclues par les messages. Budget de noeuds borne.
"""
def __init__(self, damping=0.5, root_sweeps=50, node_sweeps=4, max_nodes=300):
self.damping = damping
self.root_sweeps = root_sweeps
self.node_sweeps = node_sweeps
self.max_nodes = max_nodes
def solve(self, grid81):
stats = {'nodes': 0, 'decisions': 0, 'contradictions': 0}
sol = self._dfs(np.array(grid81, dtype=int), self.root_sweeps, stats)
return sol, stats
def _sweep(self, evidence, msgs, n):
m_f2c, m_c2f = msgs
for _ in range(n):
new_f2c = np.empty((BP_F, 2, 9))
new_f2c[:, 0] = _bp_normalize(1.0 - m_c2f[:, 1])
new_f2c[:, 1] = _bp_normalize(1.0 - m_c2f[:, 0])
m_f2c = _bp_normalize(self.damping * m_f2c + (1 - self.damping) * new_f2c)
totals = np.log(evidence + BP_EPS).copy()
np.add.at(totals, BP_FI, np.log(m_f2c[:, 0] + BP_EPS))
np.add.at(totals, BP_FJ, np.log(m_f2c[:, 1] + BP_EPS))
new_c2f = np.empty((BP_F, 2, 9))
for s, fs in ((0, BP_FI), (1, BP_FJ)):
raw = totals[fs] - np.log(m_f2c[:, s] + BP_EPS)
new_c2f[:, s] = np.exp(raw - raw.max(axis=1, keepdims=True))
m_c2f = _bp_normalize(self.damping * m_c2f + (1 - self.damping) * _bp_normalize(new_c2f))
beliefs = np.log(evidence + BP_EPS).copy()
np.add.at(beliefs, BP_FI, np.log(m_f2c[:, 0] + BP_EPS))
np.add.at(beliefs, BP_FJ, np.log(m_f2c[:, 1] + BP_EPS))
beliefs = np.exp(beliefs - beliefs.max(axis=1, keepdims=True))
excluded = np.zeros((81, 9))
np.add.at(excluded, BP_FI, (m_f2c[:, 0] < BP_EXCL_EPS))
np.add.at(excluded, BP_FJ, (m_f2c[:, 1] < BP_EXCL_EPS))
return beliefs, excluded, (m_f2c, m_c2f)
def _dfs(self, grid, sweeps, stats, evidence=None, msgs=None):
if stats['nodes'] > self.max_nodes:
return None
if evidence is None:
evidence = np.full((81, 9), 1.0 / 9)
for i in np.flatnonzero(grid > 0):
evidence[i, :] = BP_EPS
evidence[i, grid[i] - 1] = 1.0
evidence[i] = evidence[i] / evidence[i].sum()
msgs = (np.full((BP_F, 2, 9), 1.0 / 9), np.full((BP_F, 2, 9), 1.0 / 9))
stats['nodes'] += 1
beliefs, excluded, msgs = self._sweep(evidence, msgs, sweeps)
unassigned = grid == 0
if not unassigned.any():
out = grid.copy()
# La detection de contradiction BP n'est pas complete : une grille
# completee peut encore violer une unite -> verification finale.
return out if verify_solution_flat(out) else None
if (excluded.min(axis=1)[unassigned] >= 1).any():
stats['contradictions'] += 1
return None
part = np.partition(beliefs, -2, axis=1)
margins = part[:, -1] - part[:, -2]
margins[~unassigned] = -np.inf
i = int(np.argmax(margins))
order = np.argsort(-beliefs[i])
order = order[excluded[i, order] == 0]
for v in order:
ev = evidence.copy()
ev[i, :] = BP_EPS
ev[i, v] = 1.0
g2 = grid.copy()
g2[i] = v + 1
stats['decisions'] += 1
res = self._dfs(g2, self.node_sweeps, stats, ev,
(msgs[0].copy(), msgs[1].copy()))
if res is not None:
return res
return None
print(f"Graph BP : {BP_F} facteurs binaires != sur 81 cellules.")