Sudoku-15-Infer-Python : Resolution Probabiliste avec NumPyro

Navigation : << Choco | Index | Neural Network >>

Objectifs d’apprentissage

A la fin de ce notebook, vous saurez : 1. Comprendre les capacites et limites de NumPyro pour les problemes discrets 2. Implementer un modèle probabiliste avec distributions Dirichlet et contraintes douces 3. Utiliser SVI (Stochastic Variational Inference) pour l’inference 4. Comparer avec l’approche Infer.NET RobustProbabilisticSolver du notebook C# 5. Évaluer la v4 : facteurs AllDiff d’arité 9 (messages exacts par permanente) face au critère « Medium ≥ 80 % sans propagation déterministe »

Duree estimee : ~45 min | Prerequis : Sudoku-00 Environment, probabilites bayesiennes


Introduction : Programmation Probabiliste et Sudoku

Ce notebook implemente l’equivalent Python du RobustProbabilisticSolver du notebook C# Infer.NET. Les deux approches utilisent des distributions Dirichlet pour modeliser les probabilites des cellules.

Comparaison des approches

Aspect Infer.NET (Robust) NumPyro (Python)
Algorithme Expectation Propagation SVI (Variational Inference)
Contraintes Dures (ConstrainFalse) Douces (numpyro.factor)
Variables VariableArray<int> Dirichlet (continu)
Itérations 50 EP itérations 300+ SVI itérations
Performance (Easy) ~33ms (modèle compile) ~5-6s (CPU)

1. Installation et Imports

import numpy as np
import time
from typing import List, Tuple, Optional, Dict
import warnings
warnings.filterwarnings('ignore', category=UserWarning)

JAX_AVAILABLE = False
try:
    import jax
    import jax.numpy as jnp
    import numpyro
    import numpyro.distributions as dist
    from numpyro.infer import SVI, Trace_ELBO
    from numpyro.optim import Adam
    from jax import random
    JAX_AVAILABLE = True
    print(f"JAX version: {jax.__version__}")
    print(f"NumPyro version: {numpyro.__version__}")
    print(f"Backend: {jax.devices()}")
except ModuleNotFoundError as e:
    print(f"Module manquant : {e}")
    print("JAX requis pour ce notebook : pip install jax jaxlib numpyro")
    print("Les cellules de calcul seront ignorees avec un message d'information.")
JAX version: 0.9.0.1
NumPyro version: 0.20.0
Backend: [CpuDevice(id=0)]

Interpretation : Environnement NumPyro/JAX

Résultats obtenus : L’environnement est correctement configure : la cellule ci-dessus imprime les versions de JAX et NumPyro (execution sur CPU).

Composant Rôle
JAX Calcul differenciable et compilation JIT
NumPyro Programmation probabiliste sur JAX
Device (CPU) Exécution sur CPU (pas de GPU requis)

Points cles : 1. JAX comme backend : JAX fournit la differenciation automatique et la compilation JIT necessaires a NumPyro 2. Exécution CPU : Ce notebook fonctionne sur CPU, pas besoin de GPU (contrairement au Deep Learning) 3. Version recente : NumPyro (version stable recente) offre un SVI bien optimise 4. Compatibilite : L’environnement JAX/NumPyro est compatible avec les algorithmes EP/VI presentes dans ce notebook

Note technique : JAX est une bibliotheque de calcul numérique qui combine NumPy avec la differenciation automatique et la compilation JIT (Just-In-Time). NumPyro s’appuie sur JAX pour implementer des algorithmes d’inference probabiliste comme MCMC (NUTS) et VI (SVI). La distinction CPU/GPU est importante : JAX peut utiliser les GPU pour accelerer les calculs, mais les algorithmes SVI pour les CSP de petite taille comme Sudoku ne beneficient pas beaucoup du GPU (surcout de transfert de données).

2. Utilitaires

def load_puzzles(filepath: str, max_puzzles: int = None) -> List[str]:
    puzzles = []
    with open(filepath, 'r') as f:
        for line in f:
            line = line.strip()
            if len(line) >= 81:
                puzzles.append(line[:81])
                if max_puzzles and len(puzzles) >= max_puzzles:
                    break
    return puzzles

def puzzle_to_grid(puzzle_str: str) -> List[List[int]]:
    return [[int(puzzle_str[i * 9 + j]) if puzzle_str[i * 9 + j] in '123456789' else 0 
             for j in range(9)] for i in range(9)]

def grid_to_flat(grid: List[List[int]]) -> List[int]:
    """Version numpy (pas JAX) de grid_to_jax."""
    return np.array([cell for row in grid for cell in row])

if JAX_AVAILABLE:
    def grid_to_jax(grid: List[List[int]]):
        return jnp.array([cell for row in grid for cell in row])
else:
    def grid_to_jax(grid: List[List[int]]):
        return grid_to_flat(grid)

def verify_solution(grid: List[List[int]]) -> bool:
    for row in grid:
        if sorted(row) != list(range(1, 10)):
            return False
    for c in range(9):
        col = [grid[r][c] for r in range(9)]
        if sorted(col) != list(range(1, 10)):
            return False
    for br in range(3):
        for bc in range(3):
            box = [grid[br*3+r][bc*3+c] for r in range(3) for c in range(3)]
            if sorted(box) != list(range(1, 10)):
                return False
    return True

def count_errors(grid: List[List[int]]) -> int:
    errors = 0
    for row in grid:
        seen = set()
        for cell in row:
            if cell in seen:
                errors += 1
            elif cell > 0:
                seen.add(cell)
    for c in range(9):
        seen = set()
        for r in range(9):
            cell = grid[r][c]
            if cell in seen:
                errors += 1
            elif cell > 0:
                seen.add(cell)
    for br in range(3):
        for bc in range(3):
            seen = set()
            for r in range(3):
                for c in range(3):
                    cell = grid[br*3+r][bc*3+c]
                    if cell in seen:
                        errors += 1
                    elif cell > 0:
                        seen.add(cell)
    return errors

def print_grid(grid: List[List[int]], title: str = ""):
    if title:
        print(f"\n{title}")
    print("-" * 25)
    for i, row in enumerate(grid):
        if i % 3 == 0 and i > 0:
            print("|" + "-" * 23 + "|")
        line = "| "
        for j, cell in enumerate(row):
            if j % 3 == 0 and j > 0:
                line += "| "
            line += f"{cell if cell > 0 else ' '} "
        line += "|"
        print(line)
    print("-" * 25)

# Chargement des puzzles
possible_paths = ["Puzzles/Sudoku_Easy51.txt", "MyIA.AI.Notebooks/Sudoku/Puzzles/Sudoku_Easy51.txt"]
puzzles = []
for path in possible_paths:
    try:
        puzzles = load_puzzles(path, max_puzzles=5)
        if puzzles:
            print(f"{len(puzzles)} puzzles charges depuis {path}")
            break
    except FileNotFoundError:
        continue

if not puzzles:
    test_puzzle_str = "900200543100063025508407060026309001057010290090670530240530600705200304080041950"
    puzzles = [test_puzzle_str]
    print("Utilisation d'un puzzle de test par defaut")
5 puzzles charges depuis Puzzles/Sudoku_Easy51.txt

Interpretation : Chargement des Puzzles

Résultats obtenus : 5 puzzles Easy charges depuis le fichier Sudoku_Easy51.txt avec succes.

Aspect Valeur Signification
Puzzles charges 5 Limite arbitraire pour les tests (max_puzzles=5)
Source Sudoku_Easy51.txt Fichier de 51 puzzles Easy (difficulte faible)
Format 81 caractères Chaque puzzle est une ligne de 81 chiffres (0-9)
Cellules initiales ~25-35 Puzzles Easy ont ~30% de cases pre-remplies
Encodage Chiffres 0-9 0 = vide, 1-9 = valeurs connues

Points cles : 1. Stratégie de chargement robuste : Le code essaie deux chemins possibles (relatif et absolu) 2. Fallback : Si aucun fichier n’est trouve, un puzzle de test par defaut est utilise 3. Format compact : 81 caractères suffisent pour representer un grille 9x9 4. Conversion automatique : La fonction puzzle_to_grid convertit la chaîne en grille 2D (liste de listes)

Note technique : Le format “one-line” est standard pour les puzzles Sudoku : chaque ligne de 81 caractères represente la grille aplatie ligne par ligne. Les zeros representent les cases vides. Ce format est compact et facilite le chargement depuis des fichiers texte ou des bases de données.

3. Modèle NumPyro avec Dirichlet (Equivalent RobustProbabilisticSolver)

Ce modèle est l’equivalent Python du RobustSudokuModel du notebook C# Infer.NET. Il utilise :

  1. Distributions Dirichlet pour les probabilites des cellules (comme Infer.NET)
  2. Contraintes douces via numpyro.factor (au lieu de ConstrainFalse)
  3. SVI pour l’inference (au lieu d’Expectation Propagation)
if JAX_AVAILABLE:
    def compute_constraint_penalty(cell_probs):
        """
        Calcule une penalite pour les violations de contraintes Sudoku.
        Equivalent approximatif des ContraintFalse d'Infer.NET.
        """
        penalty = 0.0
        for r in range(9):
            row_probs = cell_probs[r * 9:(r + 1) * 9]
            for v in range(9):
                value_probs = row_probs[:, v]
                sum_probs = jnp.sum(value_probs)
                sum_sq = jnp.sum(value_probs ** 2)
                conflict = (sum_probs ** 2 - sum_sq) / 2
                penalty = penalty + conflict
        for c in range(9):
            col_probs = cell_probs[c::9]
            for v in range(9):
                value_probs = col_probs[:, v]
                sum_probs = jnp.sum(value_probs)
                sum_sq = jnp.sum(value_probs ** 2)
                conflict = (sum_probs ** 2 - sum_sq) / 2
                penalty = penalty + conflict
        for br in range(3):
            for bc in range(3):
                indices = [(br * 3 + i) * 9 + (bc * 3 + j) for i in range(3) for j in range(3)]
                box_probs = cell_probs[jnp.array(indices)]
                for v in range(9):
                    value_probs = box_probs[:, v]
                    sum_probs = jnp.sum(value_probs)
                    sum_sq = jnp.sum(value_probs ** 2)
                    conflict = (sum_probs ** 2 - sum_sq) / 2
                    penalty = penalty + conflict
        return penalty

    def dirichlet_sudoku_model(initial_grid, constraint_weight: float = 10.0):
        """Modele NumPyro equivalent au RobustSudokuModel d'Infer.NET."""
        n_cells, n_values = 81, 9
        alpha_base = numpyro.param("alpha_base", jnp.ones((n_cells, n_values)),
                                  constraint=dist.constraints.positive)
        epsilon, fixed_value = 1e-3, 1000.0
        is_known = (initial_grid > 0)[:, None]
        value_indices = jnp.arange(n_values)
        grid_expanded = initial_grid[:, None]
        values_expanded = value_indices[None, :] + 1
        known_values_onehot = (grid_expanded == values_expanded).astype(jnp.float32)
        known_values_onehot = known_values_onehot * is_known.astype(jnp.float32)
        alpha_known = known_values_onehot * fixed_value + (1 - known_values_onehot) * epsilon
        alpha = jnp.where(is_known.astype(jnp.bool_), alpha_known, alpha_base)
        with numpyro.plate("cells", n_cells):
            cell_probs = numpyro.sample("cell_probs", dist.Dirichlet(alpha))
        penalty = compute_constraint_penalty(cell_probs)
        numpyro.factor("constraint_penalty", -constraint_weight * penalty)
        return cell_probs

    def dirichlet_sudoku_guide(initial_grid, constraint_weight: float = 10.0):
        """Guide variationnel pour SVI."""
        n_cells, n_values = 81, 9
        alpha_var = numpyro.param("alpha_var", jnp.ones((n_cells, n_values)),
                                constraint=dist.constraints.positive)
        epsilon, fixed_value = 1e-3, 1000.0
        is_known = (initial_grid > 0)[:, None]
        value_indices = jnp.arange(n_values)
        grid_expanded = initial_grid[:, None]
        values_expanded = value_indices[None, :] + 1
        known_values_onehot = (grid_expanded == values_expanded).astype(jnp.float32)
        known_values_onehot = known_values_onehot * is_known.astype(jnp.float32)
        alpha_known = known_values_onehot * fixed_value + (1 - known_values_onehot) * epsilon
        alpha = jnp.where(is_known.astype(jnp.bool_), alpha_known, alpha_var)
        with numpyro.plate("cells", n_cells):
            numpyro.sample("cell_probs", dist.Dirichlet(alpha))

    print("Modeles NumPyro definis : compute_constraint_penalty, dirichlet_sudoku_model, dirichlet_sudoku_guide")
else:
    print("JAX requis : pip install jax jaxlib numpyro")
    print("Les modeles NumPyro (compute_constraint_penalty, dirichlet_sudoku_model, dirichlet_sudoku_guide) ne sont pas disponibles.")
Modeles NumPyro definis : compute_constraint_penalty, dirichlet_sudoku_model, dirichlet_sudoku_guide

4. Solveur Probabiliste Pur (Equivalent RobustProbabilisticSolver)

Ce solveur est l’equivalent Python direct du RobustProbabilisticSolver C# : - Une seule passe d’inference - Recupere le mode de chaque distribution Dirichlet - Pas de propagation déterministe

if JAX_AVAILABLE:
    class RobustProbabilisticSolverPy:
        """
        Equivalent Python du RobustProbabilisticSolver C# Infer.NET.
        Utilise SVI avec distributions Dirichlet pour inferer les probabilites
        des cellules, puis selectionne le mode de chaque distribution.
        """
        def __init__(self, n_iterations: int = 300, constraint_weight: float = 15.0,
                     learning_rate: float = 0.1):
            self.n_iterations = n_iterations
            self.constraint_weight = constraint_weight
            self.learning_rate = learning_rate

        def infer_probabilities(self, grid: List[List[int]]) -> np.ndarray:
            """Execute SVI pour inferer les probabilites des cellules."""
            initial_flat = grid_to_jax(grid)
            rng_key = random.PRNGKey(42)
            optimizer = Adam(step_size=self.learning_rate)
            svi = SVI(dirichlet_sudoku_model, dirichlet_sudoku_guide,
                      optimizer, loss=Trace_ELBO())
            svi_result = svi.run(rng_key, self.n_iterations, initial_flat,
                                self.constraint_weight, progress_bar=False)
            alpha_var = svi_result.params["alpha_var"]
            probs = alpha_var / jnp.sum(alpha_var, axis=-1, keepdims=True)
            return np.array(probs)

        def solve(self, grid: List[List[int]]) -> Tuple[List[List[int]], Dict]:
            """Resout un Sudoku en utilisant uniquement l'inference probabiliste."""
            grid = [row[:] for row in grid]
            start = time.time()
            probs = self.infer_probabilities(grid)
            inference_time = time.time() - start
            for i in range(9):
                for j in range(9):
                    if grid[i][j] == 0:
                        idx = i * 9 + j
                        grid[i][j] = int(np.argmax(probs[idx])) + 1
            metadata = {
                'inference_time': inference_time,
                'n_iterations': self.n_iterations,
                'converged': verify_solution(grid),
                'errors': count_errors(grid)
            }
            return grid, metadata

    print("Classe RobustProbabilisticSolverPy definie.")
else:
    print("JAX requis : pip install jax jaxlib numpyro")
    print("La classe RobustProbabilisticSolverPy n'est pas disponible sans JAX.")
Classe RobustProbabilisticSolverPy definie.

Test du solveur probabiliste pur

if JAX_AVAILABLE:
    # Test sur un puzzle facile
    test_grid = puzzle_to_grid(puzzles[0])
    print_grid(test_grid, "Puzzle initial:")

    # 300 iterations comme Infer.NET (50 EP iterations * 6 = 300 SVI)
    solver = RobustProbabilisticSolverPy(n_iterations=300, constraint_weight=15.0)

    start = time.time()
    solution, meta = solver.solve(test_grid)
    total_time = time.time() - start

    print_grid(solution, f"\nSolution (inference: {meta['inference_time']:.1f}s):")
    print(f"\nResultats:")
    print(f"  - Iterations SVI: {meta['n_iterations']}")
    print(f"  - Temps inference: {meta['inference_time']:.1f}s")
    print(f"  - Temps total: {total_time:.1f}s")
    print(f"  - Converge: {meta['converged']}")
    print(f"  - Erreurs: {meta['errors']}")
else:
    print("JAX requis : pip install jax jaxlib numpyro")
    print("Test du solveur probabiliste ignore. Installez JAX pour executer cette cellule.")

Puzzle initial:
-------------------------
| 9   2 |     5 | 4   3 |
| 1     |   6 3 |   2 5 |
| 5   8 | 4   7 |   6   |
|-----------------------|
|   2 6 | 3   9 |     1 |
|   5 7 |   1   | 2 9   |
|   9   | 6 7   | 5 3   |
|-----------------------|
| 2 4   | 5 3   | 6     |
| 7   5 | 2     | 3   4 |
|   8   |   4 1 | 9 5   |
-------------------------


Solution (inference: 14.3s):
-------------------------
| 9 6 2 | 1 8 5 | 4 7 3 |
| 1 7 4 | 9 6 3 | 8 2 5 |
| 5 3 8 | 4 2 7 | 1 6 9 |
|-----------------------|
| 8 2 6 | 3 5 9 | 7 4 1 |
| 3 5 7 | 8 1 4 | 2 9 6 |
| 4 9 1 | 6 7 2 | 5 3 8 |
|-----------------------|
| 2 4 9 | 5 3 8 | 6 1 7 |
| 7 1 5 | 2 9 6 | 3 8 4 |
| 6 8 3 | 7 4 1 | 9 5 2 |
-------------------------

Resultats:
  - Iterations SVI: 300
  - Temps inference: 14.3s
  - Temps total: 14.3s
  - Converge: True
  - Erreurs: 0

Interpretation : Solveur Probabiliste Pur (RobustProbabilisticSolver)

Résultats obtenus : Une seule passe d’inference SVI (300 itérations) suffit pour resoudre ce puzzle Easy avec succes.

Metrique Valeur Analyse
Itérations SVI 300 Equivalent a ~50 EP itérations Infer.NET
Temps inference ~6 s ~180x plus lent que C# (~34ms)
Convergence True Solution valide obtenue
Erreurs 0 Toutes les contraintes respectees
Cellules initiales 45 36 cellules a determiner

Points cles : 1. Mode des distributions : Le solveur selectionne la valeur la plus probable de chaque distribution Dirichlet (argmax) 2. Contraintes douces : Les penalites numpyro.factor encouragent la cohérence sans garantir les contraintes Sudoku 3. Performance acceptable : ~6 secondes pour un puzzle Easy est raisonnable pour un prototype pedagogique 4. Limitation evidente : Ce solveur echoue sur les puzzles Medium/Hard (voir benchmark plus loin)

Note technique : La distribution Dirichlet de dimension 9 modelise la probabilite de chaque valeur (1-9) pour chaque cellule. Le paramètre de concentration alpha est appris par SVI. Les cellules connues sont fixees avec alpha = 1000 (quasi-certain), les inconnues commencent avec alpha = 1 (uniforme). L’inference ajuste progressivement ces concentrations pour maximiser la vraisemblance tout en respectant les contraintes (via compute_constraint_penalty).

Exercice : Analyser la confiance du modèle probabiliste par cellule

Le solveur probabiliste produit pour chaque cellule une distribution de probabilites sur les 9 valeurs possibles (1-9). Certaines cellules ont une distribution très pointue (haute confiance) tandis que d’autres sont plus incertaines.

Objectif : Implementez deux fonctions pour analyser la confiance du modèle : 1. entropy(probs) : calculer l’entropie d’une distribution (mesure d’incertitude) 2. analyze_confidence(grid, probs) : identifier les cellules les plus certaines et les plus incertaines

Concepts cles : - Entropie : quantifie l’incertitude d’une distribution. Valeur faible = modèle certain, valeur elevee = modèle hesitant - Confiance : probabilite maximale de la distribution pour une cellule - Distribution uniforme : 9 valeurs equiprobables, entropie = ln(9) ~ 2.20 nats

Indices : - Étape 1 : Implementez entropy(probs) avec la formule -sum(p * log(p)) en ignorant les valeurs nulles - Étape 2 : Pour chaque cellule vide de la grille, extrayez confiance et entropie depuis la distribution - Étape 3 : Triez les résultats et affichez les extremes - Indice : np.log et un filtre probs > 0 pour eviter log(0)

# EXERCICE : Analyser la confiance du modele probabiliste par cellule
#
# Indications :
#   - Pour chaque cellule vide, la distribution Dirichlet donne 9 probabilites
#   - La confiance = max(probs) indique a quel point le modele est certain
#   - L'entropie mesure l'incertitude globale de la distribution
#
# Etape 1 : Implementez entropy(probs) avec la formule -sum(p * log(p))
#   Attention : ignorez les probabilites nulles (log(0) indefini)
# Etape 2 : Implementez analyze_confidence(grid, probs) qui retourne
#   une liste de (row, col, confidence, entropy, top_value) pour chaque cellule vide
#   triee par confiance decroissante
# Etape 3 : Affichez les 10 cellules les plus certaines et les 10 les plus incertaines
# Indice : pour l'entropie, utilisez np.log et filtrez les valeurs > 0
#   Entropie maximale pour 9 valeurs = log(9) ~ 2.20 (distribution uniforme)
#   Entropie minimale = 0 (une seule valeur a 100%)

import numpy as np


def entropy(probs):
    """Calcule l'entropie d'une distribution de probabilites.

    Args:
        probs: tableau de probabilites (somme = 1)

    Returns:
        Entropie en nats (>= 0)
    """
    # TODO etudiant : implementez le calcul d'entropie
    return 0.0  # TODO


def analyze_confidence(grid, probs):
    """Analyse la confiance du modele pour chaque cellule vide.

    Args:
        grid: grille 9x9 (0 = case vide)
        probs: tableau 81x9 de probabilites (sortie du solveur)

    Returns:
        Liste de (row, col, confidence, entropy, top_value)
        triee par confiance decroissante
    """
    # TODO etudiant : parcourez les cellules vides et collectez les metriques
    return []  # TODO


# Grille de test pour cet exercice
_ex_grid = [
    [0,0,0,2,6,0,7,0,1],
    [6,8,0,0,7,0,0,9,0],
    [1,9,0,0,0,4,5,0,0],
    [8,2,0,1,0,0,0,4,0],
    [0,0,4,6,0,2,9,0,0],
    [0,5,0,0,0,3,0,2,8],
    [0,0,9,3,0,0,0,7,4],
    [0,4,0,0,5,0,0,3,6],
    [7,0,3,0,1,8,0,0,0]
]

# Test avec des probabilites simulees (distribution Dirichlet aleatoire)
# Si JAX est disponible, decommentez pour utiliser les vraies probabilites :
# solver = RobustProbabilisticSolverPy(n_iterations=300)
# probs = solver.infer_probabilities(_ex_grid)
# analysis = analyze_confidence(_ex_grid, probs)

fake_probs = np.random.dirichlet(np.ones(9), size=81)
analysis = analyze_confidence(_ex_grid, fake_probs)

if analysis:
    print("Top 5 cellules les plus certaines :")
    for row, col, conf, ent, val in analysis[:5]:
        print(f"  ({row},{col}) : confiance={conf:.3f}, entropie={ent:.3f}, valeur={val}")
    print("\nTop 5 cellules les plus incertaines :")
    for row, col, conf, ent, val in analysis[-5:]:
        print(f"  ({row},{col}) : confiance={conf:.3f}, entropie={ent:.3f}, valeur={val}")
else:
    print("Analyse a completer : implementez analyze_confidence()")
Analyse a completer : implementez analyze_confidence()

Exercice : Compter les candidats possibles par case

Enonce

Avant d’utiliser un solveur probabiliste, il est utile de comprendre la complexite d’un puzzle en comptant les candidats possibles pour chaque case vide. Ce comptage est la base de nombreuses heuristiques de resolution (MRV, propagation de contraintes).

Implementez count_candidates(grid) qui retourne une matrice 9x9 du nombre de candidats (1-9) possibles pour chaque case vide (0 = case remplie).

Indices :

  • Étape 1 : Pour chaque case vide (valeur 0), determinez quels chiffres 1-9 sont absents de sa ligne
  • Étape 2 : Eliminez les chiffres presents dans sa colonne
  • Étape 3 : Eliminez les chiffres presents dans son bloc 3x3
  • Étape 4 : Comptez les candidats restants pour cette case
# EXERCICE : Compter les candidats possibles par case
#
# Indications :
#   - Pour chaque case vide, trouvez les chiffres absents de ligne+colonne+bloc
#   - Utilisez des ensembles (set) pour les operations ensemblistes
#   - Retournez une matrice 9x9 (0 pour les cases deja remplies)

def count_candidates(grid):
    """Compte les candidats possibles pour chaque case vide.

    Args:
        grid: liste 9x9 (0 = case vide)

    Returns:
        Matrice 9x9 du nombre de candidats (0 si case remplie)
    """
    # TODO etudiant : implementez le comptage des candidats
    return [[0]*9 for _ in range(9)]  # TODO


# Test sur un puzzle facile
test_grid_15 = [
    [0,0,0,2,6,0,7,0,1],
    [6,8,0,0,7,0,0,9,0],
    [1,9,0,0,0,4,5,0,0],
    [8,2,0,1,0,0,0,4,0],
    [0,0,4,6,0,2,9,0,0],
    [0,5,0,0,0,3,0,2,8],
    [0,0,9,3,0,0,0,7,4],
    [0,4,0,0,5,0,0,3,6],
    [7,0,3,0,1,8,0,0,0]
]

cand = count_candidates(test_grid_15)
print("Candidats par case vide :")
for row in cand:
    print(row)
Candidats par case vide :
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]
[0, 0, 0, 0, 0, 0, 0, 0, 0]

5. Solveur Iteratif (Equivalent IterativeSudokuModel)

Ce solveur est l’equivalent Python du IterativeSudokuModel C# qui fixe iterativement les cellules les plus certaines.

if JAX_AVAILABLE:
    class IterativeProbabilisticSolverPy:
        """
        Equivalent Python du IterativeSudokuModel C# Infer.NET.
        Fixe iterativement les cellules avec les probabilites les plus elevees,
        en reinjectant les valeurs fixees dans le modele.
        """
        def __init__(self, n_iterations_per_step: int = 100,
                     constraint_weight: float = 15.0,
                     confidence_threshold: float = 0.3,
                     cells_per_iteration: int = 1):
            self.n_iterations_per_step = n_iterations_per_step
            self.constraint_weight = constraint_weight
            self.confidence_threshold = confidence_threshold
            self.cells_per_iteration = cells_per_iteration

        def solve(self, grid: List[List[int]], max_steps: int = 100) -> Tuple[List[List[int]], Dict]:
            """Resout un Sudoku iterativement en fixant les cellules les plus certaines."""
            grid = [row[:] for row in grid]
            metadata = {
                'steps': 0, 'cells_fixed': 0, 'probabilistic_fixes': 0,
                'total_inference_time': 0, 'converged': False, 'errors': 0
            }
            n_empty = sum(1 for row in grid for cell in row if cell == 0)
            while n_empty > 0 and metadata['steps'] < max_steps:
                metadata['steps'] += 1
                initial_flat = grid_to_jax(grid)
                rng_key = random.PRNGKey(42 + metadata['steps'])
                optimizer = Adam(step_size=0.1)
                svi = SVI(dirichlet_sudoku_model, dirichlet_sudoku_guide,
                          optimizer, loss=Trace_ELBO())
                start = time.time()
                svi_result = svi.run(rng_key, self.n_iterations_per_step, initial_flat,
                                    self.constraint_weight, progress_bar=False)
                metadata['total_inference_time'] += time.time() - start
                alpha_var = svi_result.params["alpha_var"]
                probs = np.array(alpha_var / jnp.sum(alpha_var, axis=-1, keepdims=True))
                candidates = []
                for i in range(9):
                    for j in range(9):
                        if grid[i][j] == 0:
                            idx = i * 9 + j
                            confidence = float(np.max(probs[idx]))
                            if confidence >= self.confidence_threshold:
                                value = int(np.argmax(probs[idx])) + 1
                                candidates.append((confidence, i, j, value))
                if not candidates:
                    break
                candidates.sort(reverse=True)
                fixed_this_step = 0
                for conf, i, j, val in candidates[:self.cells_per_iteration]:
                    if grid[i][j] == 0:
                        grid[i][j] = val
                        n_empty -= 1
                        fixed_this_step += 1
                        metadata['probabilistic_fixes'] += 1
                metadata['cells_fixed'] += fixed_this_step
                if fixed_this_step == 0:
                    break
            metadata['converged'] = verify_solution(grid)
            metadata['errors'] = count_errors(grid)
            return grid, metadata

    print("Classe IterativeProbabilisticSolverPy definie.")
else:
    print("JAX requis : pip install jax jaxlib numpyro")
    print("La classe IterativeProbabilisticSolverPy n'est pas disponible sans JAX.")
Classe IterativeProbabilisticSolverPy definie.

Test du solveur iteratif

if JAX_AVAILABLE:
    print("=== Test IterativeProbabilisticSolverPy ===")

    iterative_solver = IterativeProbabilisticSolverPy(
        n_iterations_per_step=100,
        confidence_threshold=0.25,
        cells_per_iteration=1
    )

    for i, puzzle_str in enumerate(puzzles[:2]):
        test_grid = puzzle_to_grid(puzzle_str)
        start = time.time()
        solution, meta = iterative_solver.solve(test_grid)
        elapsed = time.time() - start
        status = "OK" if meta['converged'] else f"{meta['errors']} erreurs"
        print(f"\nPuzzle {i+1}: {status}")
        print(f"  - Etapes: {meta['steps']}")
        print(f"  - Cellules fixees (probabiliste): {meta['probabilistic_fixes']}")
        print(f"  - Temps inference: {meta['total_inference_time']:.1f}s")
        print(f"  - Temps total: {elapsed:.1f}s")
else:
    print("JAX requis : pip install jax jaxlib numpyro")
    print("Test du solveur iteratif ignore. Installez JAX pour executer cette cellule.")
=== Test IterativeProbabilisticSolverPy ===

Puzzle 1: OK
  - Etapes: 36
  - Cellules fixees (probabiliste): 36
  - Temps inference: 288.2s
  - Temps total: 309.2s

Puzzle 2: OK
  - Etapes: 49
  - Cellules fixees (probabiliste): 49
  - Temps inference: 310.6s
  - Temps total: 335.5s

Interpretation : Solveur Iteratif Probabiliste

Résultats obtenus : Le solveur iteratif fixe progressivement les cellules avec une confiance superieure au seuil de 0.25, en re-injectant les valeurs fixees dans le modèle.

Metrique Puzzle 1 Puzzle 2 Moyenne
Étapes requises 36 49 42.5
Cellules fixees 36 49 42.5
Temps inference cumule ~160s ~220s ~190s
Temps total ~170s ~230s ~200s
Temps par étape ~5s ~5s ~5s
Statut final OK OK 100% succes

Points cles : 1. Processus iteratif : Chaque étape fixe 1 cellule (paramètre cells_per_iteration=1) et relance l’inference 2. Convergence garantie : Les deux puzzles convergent vers une solution valide, contrairement au solveur robuste 3. Cout temporel : ~200 secondes par puzzle, soit ~30x plus lent que le solveur robuste (~6 s) 4. Seuil de confiance : 0.25 permet de fixer des cellules avec une confiance moderee mais suffisante

Note technique : L’approche iterative est plus robuste car chaque fixation reduit l’espace de recherche pour les étapes suivantes. Cependant, elle necessite de multiples passages d’inference (36-49 vs 1), ce qui explique le cout eleve. Le paramètre confidence_threshold est critique : trop eleve (>0.5), aucune cellule n’est fixee ; trop bas (<0.1), des erreurs s’accumulent.

Exercice : Verifier la reproductibilite d’un solveur

Enonce

Les solveurs probabilistes comme le solveur iteratif peuvent produire des résultats non déterministes (différentes exécutions donnent des résultats différents). Il est important de mesurer cette variabilite.

Implementez measure_reproducibility(solve_func, puzzle, n_runs=5) qui execute un solveur plusieurs fois et retourne le nombre de solutions uniques trouvees et le taux de succes.

Indices :

  • Étape 1 : Appelez solve_func(puzzle) dans une boucle de n_runs itérations
  • Étape 2 : Stockez chaque solution (ou None si echec) dans une liste
  • Étape 3 : Utilisez un set de tuples pour compter les solutions uniques
  • Étape 4 : Retournez {"success_rate": float, "unique_solutions": int, "results": list}
# EXERCICE : Verifier la reproductibilite d'un solveur
#
# Indications :
#   - Executez le solveur N fois sur le meme puzzle
#   - Comptez les solutions uniques avec un set de tuples
#   - Calculez le taux de succes

def measure_reproducibility(solve_func, puzzle, n_runs=5):
    """Mesure la reproductibilite d'un solveur.

    Args:
        solve_func: fonction prenant un puzzle et retournant une grille
        puzzle: grille 9x9 a resoudre
        n_runs: nombre d'executions

    Returns:
        dict avec cles "success_rate", "unique_solutions", "results"
    """
    # TODO etudiant : implementez la mesure de reproductibilite
    return {"success_rate": 0.0, "unique_solutions": 0, "results": []}  # TODO


# Test avec un solveur deterministe (backtracking simple)
def simple_backtrack_solve(grid):
    """Solveur backtracking simple pour test."""
    grid = [row[:] for row in grid]
    
    def find_empty(g):
        for i in range(9):
            for j in range(9):
                if g[i][j] == 0:
                    return (i, j)
        return None
    
    def is_valid(g, r, c, num):
        if num in g[r]: return False
        if any(g[i][c] == num for i in range(9)): return False
        br, bc = 3*(r//3), 3*(c//3)
        if any(g[br+i][bc+j] == num for i in range(3) for j in range(3)): return False
        return True
    
    def solve(g):
        empty = find_empty(g)
        if not empty:
            return True
        r, c = empty
        for num in range(1, 10):
            if is_valid(g, r, c, num):
                g[r][c] = num
                if solve(g):
                    return True
                g[r][c] = 0
        return False
    
    success = solve(grid)
    return grid if success else None


result = measure_reproducibility(simple_backtrack_solve, test_grid_15, n_runs=3)
success_rate = result["success_rate"]
unique_solutions = result["unique_solutions"]
print(f"Taux de succes : {success_rate:.1%}")
print(f"Solutions uniques : {unique_solutions}")
Taux de succes : 0.0%
Solutions uniques : 0

6. Solveur v3 : Belief Propagation (déclinaison non alignée avec Infer.NET)

La déclinaison Python de la v3 n’imite pas Infer.NET : elle remplace l’algorithme. Le solveur NumPyro (SVI sur pénalités molles, ~5 s par grille) optimise une approximation variationnelle ; un graphe de facteurs sur des contraintes dures appelle plutôt du message passing. Ici, du sum-product loopy BP vectorisé en NumPy pur (pas de JAX), sur le même graphe que le modèle C# : 81 cellules, une distribution par cellule, un facteur binaire cell_i != cell_j par paire de cellules partageant une unité (810 paires après déduplication).

Le message d’un facteur != vers une cellule se calcule en forme close : m_{f->i}(v) = 1 - p_j(v) (la masse de tout sauf v). Un sweep complet = quelques opérations matricielles sur des tableaux (810, 9) : ~3 ms, soit ~1000x plus rapide que le SVI, et les messages sont exacts sur chaque facteur (pas de gradient, pas de poids de contrainte à régler).

Deux étages, comme en C# : 1. Décimation : après convergence partielle, on fixe la cellule à la marge la plus nette (top1 - top2 de sa croyance) et on poursuit avec les messages conservés (warm start). Contradiction détectée par exclusions : une valeur est exclue quand le message d’un voisin pour elle est ~nul ; une cellule vide dont les 9 valeurs sont exclues est une contradiction. 2. Repli : les décisions sont révocables (essai des valeurs par croyance décroissante, remontée à la dernière décision sur contradiction), avec un budget de nœuds borné.

La comparaison honnête avec la version C# (décimation EP + propagation de contraintes) et les résultats mesurés ci-dessous montrent où cette déclinaison non hybride cale : exactement là où la relaxation par paires est aveugle (ensembles de Hall) — c’est le point de départ de la v4 (section suivante : facteurs AllDiff d’arité 9).

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.")
Graph BP : 810 facteurs binaires != sur 81 cellules.

Benchmark v3 : BP vs BP + repli

Le corpus facile (Easy51, déjà chargé plus haut) plus les deux corpus durs du dépôt (top95, hardest — les points y remplacent les zéros). Le repli est borné à 300 nœuds par grille : au-delà, échec honnête plutôt qu’attente indéterminée.

import pandas as pd
import time

bp_bench_sets = [
    ('Easy51', 'Puzzles/Sudoku_Easy51.txt', None),
    ('top95 (15 premieres)', 'Puzzles/Sudoku_top95.txt', 15),
    ('hardest (11)', 'Puzzles/Sudoku_hardest.txt', None),
]

rows = []
bp_decim_solver = BeliefPropagationSolver()
bp_dfs_solver = BPBacktrackingSolver()
for label, relpath, sub in bp_bench_sets:
    puzzle_strs = load_puzzles(relpath, max_puzzles=sub)
    flat_grids = [grid_to_flat(puzzle_to_grid(s)).tolist() for s in puzzle_strs]
    for method_name, solver in (('BP decimation', bp_decim_solver),
                                ('BP + repli borne', bp_dfs_solver)):
        t0 = time.time()
        ok = 0
        for k, p in enumerate(flat_grids):
            sol, st = (solver.solve(p, seed=k)
                       if isinstance(solver, BeliefPropagationSolver)
                       else solver.solve(p))
            if sol is not None and verify_solution_flat(sol):
                ok += 1
        dt = time.time() - t0
        rows.append({'corpus': label, 'solveur': method_name,
                     'resolus': f'{ok}/{len(flat_grids)}',
                     'taux': round(ok / len(flat_grids), 2),
                     'ms/grille': round(dt / len(flat_grids) * 1000)})
pd.DataFrame(rows)
corpus solveur resolus taux ms/grille
0 Easy51 BP decimation 36/51 0.71 645
1 Easy51 BP + repli borne 35/51 0.69 1034
2 top95 (15 premieres) BP decimation 0/15 0.00 981
3 top95 (15 premieres) BP + repli borne 0/15 0.00 2115
4 hardest (11) BP decimation 3/11 0.27 812
5 hardest (11) BP + repli borne 4/11 0.36 1651

Interpretation : le message passing domine le SVI, la relaxation par paires cale

Trois constats sur ces mesures (run de ce notebook, meme machine que les cellules SVI ci-dessus) :

  1. Vitesse : le BP decimation resout Easy51 a ~0,5 s par grille, contre ~5-8 s pour le SVI pur et ~3 minutes pour le SVI iteratif (qui re-execute 300 iterations a chaque cellule fixee). Messages exacts par facteur, warm start entre decisions, pas de taux d’apprentissage ni de poids de contrainte a regler.
  2. Couverture Easy : 36/51 pour la decimation pure, contre 1/3 pour le SVI pur sur sa propre selection. Le repli borne n’ajoute rien ici (35/51) : sur les grilles faciles, quand la decimation echoue, c’est la croyance qui est fausse, pas l’irrevocabilite de la decision.
  3. Le mur des corpus durs : 0/15 sur top95, 3-4/11 sur hardest. Ce n’est pas un defaut d’implementation mais une limite structurelle : avec 810 facteurs binaires !=, deux cellules d’une meme unite cantonnees a {3,7} s’envoient des messages symetriques qui s’annulent (ensemble de Hall de taille 2 invisible pour la relaxation par paires). La v3 C# franchit ce mur en hybrident avec de la propagation de contraintes (la “voie facile”) ; la suite principielle est la v4 (section 7) : des facteurs AllDiff d’arite 9 dont les messages exacts (permanente ou affectation hongroise sur une matrice 9x9) capturent les ensembles de Hall sans aucun apport deterministe externe.

7. Solveur v4 : facteurs AllDiff d’arité 9 — rendre l’ensemble de Hall visible au graphe

Le diagnostic de la v3 est graphique, pas algorithmique. La relaxation par paires modélise AllDiff(unité) par 810 facteurs binaires cell_i != cell_j (27 unités × C(9,2), dédoublonnés). Or un facteur binaire ne sait dire qu’une chose : « pas ma valeur ». Quand un ensemble de Hall de taille 2 apparaît — deux cellules d’une même unité toutes deux cantonnées à {3,7} — chaque membre de la paire envoie à l’autre un message symétrique qui s’annule, et les marginaux restent scindés à 50/50 indéfiniment.

La v4 remonte d’un cran sur le graphe de facteurs : un seul facteur AllDiff d’arité 9 par unité (27 facteurs au lieu de 810), dont les messages vers chaque cellule se calculent exactement. Pour un facteur AllDiff sur les cellules \(c_1,\dots,c_9\) avec messages entrants \(m_{c_j \to f}\), posons \(M_{jk} = m_{c_j \to f}(k)\) la matrice 9×9 des messages. Alors :

Sum-product (croyances) — le message est la somme sur toutes les affectations injectives restantes :

\[m^{\sum}_{f \to c_i}(v) = \operatorname{perm}\left(M_{-i,-v}\right)\]

où \(M_{-i,-v}\) est le mineur 8×8 de \(M\). La permanente compte (pondère) les affectations valides : si une valeur est devenue impossible pour l’unité, la permanente du mineur correspondant est exactement nulle — l’ensemble de Hall est capturé au lieu d’être moyenné. Calcul par la formule de Ryser : \(2^8 \times 8\) opérations, les 81 mineurs d’une unité batchés en un seul produit tensoriel NumPy.

Max-product (décisions) — le message est la meilleure affectation injective restante : \(m^{\max}_{f \to c_i}(v)\) = produit des messages du couplage optimal du mineur, calculé par l’algorithme hongrois (scipy.optimize.linear_sum_assignment, \(O(n^3)\)). C’est le même geste vu depuis la décision : la permanente somme, le hongrois maximise.

Parallèle structurant (fil rouge de la série) : la v2 est montée d’un cran sur le modèle (paramètre recompilé par instance → prior Dirichlet observé) ; la v4 monte d’un cran sur le graphe (relaxation par paires → facteur structurel d’arité 9). Les deux gestes restent dans le paradigme probabiliste : aucun apport déterministe externe, pas de propagation de contraintes, pas de repli/backtracking. Les redémarrages randomisés (bruit d’initialisation des messages, tirage selon la croyance sur les plateaux) sont de l’inférence probabiliste randomisée — l’équivalent exact des restarts de MCMC.

Trois gestes candidats existaient pour franchir ce mur : arité 9 natif, encodage one-hot + exactly-one, régions Kikuchi/GBP — la v4 implémente le premier.

import itertools
from scipy.optimize import linear_sum_assignment

UNIT_CELLS = np.array(BP_UNITS)               # (27, 9) cellule plate de chaque slot

# Formule de Ryser : perm(A) = somme sur les sous-ensembles S de lignes de
# (-1)^(n+|S|) * produit des sommes partielles de colonnes sur S. Iterer les
# 2^n sous-ensembles en Python serait trop lent ; on precalcule donc les
# tables une fois pour toutes : RYR_SUB[s, j] = la ligne j est-elle dans le
# sous-ensemble s (bits de s), RYR_SIGN[s] = le signe (-1)^(n+popcount(s)).
RYR_SUB = np.array([[(s >> j) & 1 for j in range(8)] for s in range(256)], dtype=np.float64)
RYR_SIGN = np.array([(-1.0) ** (8 + bin(s).count("1")) for s in range(256)])
RYR_SUB9 = np.array([[(s >> j) & 1 for j in range(9)] for s in range(512)], dtype=np.float64)
RYR_SIGN9 = np.array([(-1.0) ** (9 + bin(s).count("1")) for s in range(512)])
KEEP8 = np.array([[k for k in range(9) if k != i] for i in range(9)])


def batch_perm8(minors):
    """Permanentes d'un batch de matrices 8x8 (formule de Ryser, vectorisee).

    L'einsum "buv,sv->bsu" calcule, pour chaque batch b et chaque sous-ensemble
    s, la somme u des colonnes sur les lignes de s (dimension intermediaire :
    (batch, sous-ensemble, colonne)). Le produit des 8 colonnes puis la
    combinaison lineaire avec les signes ferment la formule de Ryser.
    """
    rs = np.einsum("buv,sv->bsu", minors, RYR_SUB)
    return RYR_SIGN @ rs.prod(axis=2).T


def batch_perm9(M):
    """Permanentes d'un batch de matrices 9x9."""
    rs = np.einsum("buv,sv->bsu", M, RYR_SUB9)
    return RYR_SIGN9 @ rs.prod(axis=2).T


def unit_minors(M):
    """M (B, 9, 9) -> mineurs (B, 9, 9, 8, 8) : [b, i, v] = M[b] sans ligne i ni colonne v.

    Le mineur (i, v) est la matrice 8x8 obtenue en retirant la ligne i (la
    cellule qu'on interroge) et la colonne v (la valeur candidate) : sa
    permanente = poids total des affectations injectives ou cellule i prend v.
    KEEP8[i] liste les 8 indices conserves quand on retire i. Le transpose+
    reshape reorganise les axes en (batch, cellule, valeur, 8, 8) sans copie
    par element.
    """
    Mr = M[:, KEEP8, :]          # retire la ligne i -> (B, 9, 9, 9, 8)
    Mc = Mr[:, :, :, KEEP8]      # retire la colonne v -> (B, 9, 9, 8, 8)
    return Mc.transpose(0, 1, 3, 2, 4).reshape(M.shape[0], 9, 9, 8, 8)


def hungarian_maxprod(minors):
    """minors (B, 8, 8) en espace probabiliste -> log du meilleur produit
    d'affectation injective du mineur (algorithme hongrois, max-product)."""
    B = minors.shape[0]
    out = np.empty(B)
    # Max-produit en espace log : le produit des entrees selectionnees devient
    # une somme de logs. linear_sum_assignment minimise un cout -> on lui donne
    # -log(message) ; le cout minimal correspond au produit maximal, et -cout
    # restitue le log du meilleur produit. BP_EPS plafonne les messages nuls
    # (une affectation interdite) sans creer de log(0).
    L = -np.log(np.maximum(minors, BP_EPS))
    for b in range(B):
        ri, ci = linear_sum_assignment(L[b])
        out[b] = -L[b, ri, ci].sum()
    return out


# Auto-validation des deux operateurs avant tout usage :
#  - permanente = brute force sur une 8x8 test
A_test = np.eye(8) * 0.5 + 0.01
brute = sum(np.prod([A_test[i, p[i]] for i in range(8)])
            for p in itertools.permutations(range(8)))
assert abs(batch_perm8(A_test[None])[0] - brute) < 1e-9
#  - perm(matrice de permutation) = 1, perm(matrice de 1) = n!
assert abs(batch_perm8(np.ones((8, 8))[None])[0] - 40320.0) < 1e-6
assert abs(batch_perm9(np.eye(9)[None])[0] - 1.0) < 1e-9
#  - mineur de I9 : perm > 0 ssi la ligne enlevee egale la colonne enlevee
pers_I9 = batch_perm8(unit_minors((np.eye(9) * 0.9)[None]).reshape(-1, 8, 8)).reshape(9, 9)
assert np.allclose(pers_I9 / pers_I9.max() - np.eye(9), 0.0, atol=1e-6)
#  - hongrois = brute force max-produit sur une 8x8 test
best_brute = max(np.prod([A_test[i, p[i]] for i in range(8)])
                 for p in itertools.permutations(range(8)))
assert abs(np.exp(hungarian_maxprod(A_test[None])[0]) - best_brute) < 1e-9
print("Operateurs exacts valides : Ryser (permanente) + hongrois (max-produit).")
Operateurs exacts valides : Ryser (permanente) + hongrois (max-produit).

Démonstration : un ensemble de Hall de taille 2 met la relaxation par paires en échec

Une unité (ligne) avec 6 cellules données (valeurs 1, 2, 4, 5, 6, 8), deux cellules cantonnées à {3,7} par leurs autres unités (colonnes/blocs — injectées ici comme evidence), et une neuvième cellule « spectatrice » libre sur {3,7,9}. L’ensemble de Hall {c1, c2} ⊂ {3,7} consomme 3 et 7 : la spectatrice vaut donc 9 avec certitude. C’est une inférence de niveau « candidat unique caché » — triviale pour un humain, invisible pour la relaxation par paires. Le message arité 9 utilisé ici est la permanente (sum-product) : c’est l’objet exact, le hongrois en est le dual de décision.

# Demonstration sur une unite isolee : meme evidence, deux couches facteur.
#
# Construction d'un ensemble de Hall de taille 2 : c1 et c2 voient toutes
# deux {3, 7} (deux cellules, deux valeurs : c'est exactement faisable,
# donc 3 et 7 sont CONSOMMES par c1/c2) ; la spectatrice c9 voit {3, 7, 9}.
# La valeur 9 est donc la seule possible pour c9 -- mais seule la couche
# arity-9 peut le VOIR, car il faut raisonner sur l'affectation injective
# complete de l'unite, pas sur des paires isolees. Les 6 autres cellules
# reoivent une valeur donnee pour fermer l'unite.
hall_ev = np.full((9, 9), BP_EPS)
for slot, val in zip(range(2, 8), [1, 2, 4, 5, 6, 8]):   # 6 cellules donnees
    hall_ev[slot, val - 1] = 1.0
hall_ev[0, 2] = hall_ev[0, 6] = 0.5                       # c1 : {3,7}
hall_ev[1, 2] = hall_ev[1, 6] = 0.5                       # c2 : {3,7}
hall_ev[8, 2] = hall_ev[8, 6] = hall_ev[8, 8] = 1.0 / 3.0 # spectatrice : {3,7,9}

# --- couche 1 : relaxation par paires (meme moteur que la v3, unite seule)
# On isole le moteur v3 : messages F (facteur->cellule) et C
# (cellule->facteur) entre cellules liees par un facteur != binaire.
_diag = np.arange(9)
F = np.full((9, 9, 9), 1.0 / 9)      # F[i, j] = message facteur(i,j) -> i
C = np.full((9, 9, 9), 1.0 / 9)      # C[i, j] = message cellule i -> facteur(i,j)
for _ in range(200):
    newF = 1.0 - C.transpose(1, 0, 2)
    newF[_diag, _diag] = 1.0 / 9
    F = _bp_normalize(0.5 * F + 0.5 * _bp_normalize(newF))
    log_tot = np.log(hall_ev + BP_EPS) + np.log(F + BP_EPS).sum(axis=1)
    raw = log_tot[:, None, :] - np.log(F + BP_EPS)
    raw = raw - raw.max(axis=2, keepdims=True)
    upd = _bp_normalize(np.exp(raw))
    upd[_diag, _diag] = 1.0 / 9
    C = _bp_normalize(0.5 * C + 0.5 * upd)
log_pair = np.log(hall_ev + BP_EPS) + np.log(F + BP_EPS).sum(axis=1)
beliefs_pair = np.exp(log_pair - log_pair.max(axis=1, keepdims=True))
beliefs_pair /= beliefs_pair.sum(axis=1, keepdims=True)

# --- couche 2 : facteur AllDiff d'arite 9 (permanente des mineurs, exact)
msg_ar9 = batch_perm8(unit_minors(hall_ev[None]).reshape(-1, 8, 8)).reshape(9, 9)
msg_ar9 = _bp_normalize(np.maximum(msg_ar9, 0.0))
log_ar9 = np.log(hall_ev + BP_EPS) + np.log(msg_ar9 + BP_EPS)
beliefs_ar9 = np.exp(log_ar9 - log_ar9.max(axis=1, keepdims=True))
beliefs_ar9 /= beliefs_ar9.sum(axis=1, keepdims=True)

df_hall = pd.DataFrame({
    "relaxation par paires": beliefs_pair[8][[2, 6, 8]],
    "facteur d'arite 9": beliefs_ar9[8][[2, 6, 8]],
}, index=["v=3", "v=7", "v=9"])
print("Spectatrice (cellule libre sur {3,7,9}), marginaux apres convergence :")
print(df_hall.round(3).to_string())
print()
print("Les deux membres de la paire de Hall {3,7} restent a 50/50 (arite 9) :",
      np.round(beliefs_ar9[0][[2, 6]], 3).tolist(),
      np.round(beliefs_ar9[1][[2, 6]], 3).tolist())
sp_pair = beliefs_pair[8]
sp_ar9 = beliefs_ar9[8]
print()
print(f"Certaine que la spectatrice vaut 9 ? paires : {sp_pair[8]:.3f} — arite 9 : {sp_ar9[8]:.3f}")
assert sp_ar9[8] > 0.999, "l'arite 9 doit decider (ensemble de Hall)"
assert sp_pair[8] < 0.9, "la relaxation par paires ne doit PAS decider"
print("Verifie : l'arite 9 tranche, la relaxation par paires reste myope.")
Spectatrice (cellule libre sur {3,7,9}), marginaux apres convergence :
     relaxation par paires  facteur d'arite 9
v=3                  0.167                0.0
v=7                  0.167                0.0
v=9                  0.667                1.0

Les deux membres de la paire de Hall {3,7} restent a 50/50 (arite 9) : [0.5, 0.5] [0.5, 0.5]

Certaine que la spectatrice vaut 9 ? paires : 0.667 — arite 9 : 1.000
Verifie : l'arite 9 tranche, la relaxation par paires reste myope.

Interprétation : la permanente voit l’affectation, la paire ne voit que le voisin

La relaxation par paires laisse la spectatrice à 0.667 sur v=9 : chaque membre de la paire {3,7} lui envoie « pas ma valeur » — mais comme aucun des deux ne s’est décidé, le message binaire 1 − m(v) reste à 0,5 sur 3 et 7, et 1,0 sur 9 : dilué, jamais certain (et il reste même 0.33 de masse sur des valeurs impossibles). Le facteur d’arité 9 calcule la permanente du mineur : affecter v=3 à la spectatrice laisserait les deux cellules de Hall se disputer la seule valeur 7 — permanente exactement nulle, message nul, et la croyance converge vers 1.000 en un seul passage de messages. Noter l’honnêteté du calcul : les deux membres de la paire {3,7} restent à 50/50 sous les DEUX couches — leur ambiguïté est réelle (deux solutions globales existent), c’est la certitude qu’ils consomment 3 et 7 qui se propage vers la spectatrice. C’est ce mécanisme — le « candidat unique caché » émergeant de la permanente — que la v3 pairwise ne pouvait pas produire sur Medium.

Solveur v4 : max-product arity-9 + décimation randomisée

Pour résoudre la grille complète, la v4 utilise le message max-product (hongrois) : la décimation a besoin de décisions, pas de moyennes. Deux randomisations probabilistes standard complètent le squelette v3 (evidence, damping, décimation par marge) : un bruit log-uniforme décroissant sur l’initialisation des messages à chaque redémarrage (casse les symétries du point fixe loopy), et un tirage selon la croyance quand la meilleure marge est sous le seuil de plateau. Toujours zéro couche déterministe : contradiction = permanente d’unité ≈ 0 ou grille finale invalide → restart, jamais de repli.

class Arity9MaxProductSolver:
    """Max-product loopy BP a facteurs AllDiff d'arite 9 + decimation (v4).

    Message facteur->cellule : meilleure affectation injective du mineur 8x8
    (algorithme hongrois). Redemarrages randomises : bruit d'init decroissant
    des messages + tirage stochastique de valeur sur plateaux. Aucune couche
    de propagation deterministe, aucun repli.
    """

    def __init__(self, damping=0.5, init_sweeps=30, decim_sweeps=4,
                 max_restarts=30, plateau=0.25, noise=0.6):
        self.damping = damping
        self.init_sweeps = init_sweeps
        self.decim_sweeps = decim_sweeps
        self.max_restarts = max_restarts
        self.plateau = plateau
        self.noise = noise

    def solve(self, grid81, seed=0):
        """Boucle externe : jusqu'a max_restarts tentatives, semee une fois pour toutes.

        Chaque tentative (_attempt) part d'un bruit d'init different ; si la
        decimation aboutit a une contradiction (grille finale invalide), la
        tentative suivante repart de l'evidence d'origine avec un autre bruit.
        """
        rng = np.random.default_rng(seed)
        stats = {"decisions": 0, "sweeps": 0, "contradictions": 0, "restarts": 0,
                 "ms": 0.0}
        t0 = time.perf_counter()
        for restart in range(self.max_restarts):
            res = self._attempt(np.array(grid81, dtype=int), rng, stats, restart)
            if res is not None:
                stats["restarts"] = restart
                stats["ms"] = (time.perf_counter() - t0) * 1000
                return res, stats
        stats["restarts"] = self.max_restarts
        stats["ms"] = (time.perf_counter() - t0) * 1000
        return None, stats

    def _attempt(self, grid, rng, stats, restart):
        """Une tentative complete : init des messages -> BP -> decimation.

        Renvoie la grille resolue (verifiee) ou None (contradiction ou
        decimation inherente). Toute la randomisation vient du rng commun.
        """
        # Evidence : les cellules donnees deviennent des deltas (0 partout
        # sauf leur valeur), les autres restent uniformes. C'est le SEUL
        # apport d'information exterieur -- ensuite tout est message-passing.
        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()
        # Restart 0 : messages uniformes (le point fixe canonique). Redemarrages
        # suivants : chaque message cellule->facteur reoit un bruit
        # multiplicatif uniforme en espace log. C'est le moteur de
        # l'exploration : un point fixe BP symetrique (ex. grille a deux
        # solutions) sera leve differemment selon le bruit, la ou le
        # pairwise n'avait qu'un seul bassin d'attraction.
        scale = 0.0 if restart == 0 else self.noise
        base = evidence[UNIT_CELLS]
        if scale > 0:
            m_c2f = _bp_normalize(
                np.exp(np.log(base + BP_EPS)
                       + rng.uniform(-scale, scale, base.shape)))
        else:
            m_c2f = base.copy()
        m_f2c = np.full((27, 9, 9), 1.0 / 9)
        # Optimisation cles : les mineurs ne sont calcules QUE pour les
        # cellules libres. Pour une cellule donnee, la ligne du mineur est
        # un delta (0 sauf sa valeur) : la permanente est exactement 0 pour
        # toute autre valeur -- le message est connu gratuitement, inutile
        # de payer un hongrois 8x8 pour lui. upd_mask fige aussi les
        # messages de ces cellules (leur evidence ne change plus).
        free_slots = [np.flatnonzero(grid[np.array(u)] == 0).astype(int)
                      for u in BP_UNITS]
        upd_mask = np.zeros((27, 9), dtype=bool)
        for u in range(27):
            upd_mask[u, free_slots[u]] = True

        def sweep(n):
            nonlocal m_f2c, m_c2f
            for _ in range(n):
                # --- passe facteur->cellule : LE coeur arity-9 ---
                # Mineurs de toutes les unites (cellules libres seulement),
                # regroupes en un seul batch (b, 8, 8) pour passer la
                # totalite des hongrois d'un balayage d'un coup.
                mins = unit_minors(m_c2f)
                sel = np.concatenate([mins[u][free_slots[u]] for u in range(27)])
                logs = hungarian_maxprod(sel.reshape(-1, 8, 8))
                # Chaque message f2c[u, cellule] = log du meilleur produit
                # d'affectation de l'unite u OU la cellule prend la valeur v
                # (= le hongrois du mineur). Normalisation par ligne (le
                # max a 0) pour rester en espace probabilite sans underflow.
                new_f2c = np.zeros((27, 9, 9))
                ofs = 0
                for u in range(27):
                    j = len(free_slots[u])
                    if j:
                        raw = logs[ofs:ofs + j * 9].reshape(j, 9)
                        raw = raw - raw.max(axis=1, keepdims=True)
                        new_f2c[u, free_slots[u]] = np.exp(raw)
                        ofs += j * 9
                # Damping : melange de l'ancien et du nouveau message (poids
                # 1-damping). Indispensable en loopy BP pour eviter
                # l'oscillation du a l'echange cyclique de messages.
                new_f2c = _bp_normalize(new_f2c)
                m_f2c = _bp_normalize(self.damping * m_f2c
                                      + (1 - self.damping) * new_f2c)
                # --- passe cellule->facteur ---
                # Le message c2f retire sa propre contribution f2c du
                # total (evidence + somme des messages recus) : c'est le
                # message extrinsique, il ne doit pas se re-envoyer
                # l'information qu'il vient de recevoir (contre-réaction).
                log_tot = np.log(evidence + BP_EPS).copy()
                np.add.at(log_tot, UNIT_CELLS.reshape(-1),
                          np.log(m_f2c + BP_EPS).reshape(-1, 9))
                raw = log_tot[UNIT_CELLS] - np.log(m_f2c + BP_EPS)
                raw = raw - raw.max(axis=2, keepdims=True)
                upd = _bp_normalize(np.exp(raw))
                m_c2f = np.where(upd_mask[:, :, None],
                                 _bp_normalize(self.damping * m_c2f
                                               + (1 - self.damping) * upd),
                                 m_c2f)
                stats["sweeps"] += 1
            log_b = np.log(evidence + BP_EPS).copy()
            np.add.at(log_b, UNIT_CELLS.reshape(-1),
                      np.log(m_f2c + BP_EPS).reshape(-1, 9))
            b = np.exp(log_b - log_b.max(axis=1, keepdims=True))
            return b / b.sum(axis=1, keepdims=True)

        # Croyances initiales apres init_sweeps balayages a effectif plein
        # (toutes cellules libres) : le point fixe approximatif du BP.
        beliefs = sweep(self.init_sweeps)
        # --- decimation : decider, figer, re-propager ---
        # A chaque etape on choisit la cellule la plus confiante (marge =
        # ecart entre les deux meilleures valeurs), on la FIGE dans
        # l'evidence (delta), puis on re-propage decim_sweeps balayages.
        # Figer reduit les unites : les mineurs rétrécissent d'une ligne a
        # chaque decision de leurs membres, la structure restante devient
        # plus contrainte -- c'est la que les ensembles de Hall sautent.
        while not assigned.all():
            part = np.partition(beliefs, -2, axis=1)
            margins = part[:, -1] - part[:, -2]
            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])
            # Plateau : si la marge est trop faible (croyance presque
            # uniforme, typique au coeur d'un ensemble de Hall symetrique),
            # on TIRE selon la croyance au lieu de prendre l'argmax -- un
            # tirage selon la croyance est de l'inférence probabiliste
            # randomisée, pas de la propagation déterministe : on casse la
            # symetrie en respectant les poids, sans regle exterieure.
            if margins[i] < self.plateau:
                v = int(rng.choice(9, p=beliefs[i] / beliefs[i].sum()))
            else:
                v = int(np.argmax(beliefs[i]))
            evidence[i, :] = BP_EPS
            evidence[i, v] = 1.0
            evidence[i] = evidence[i] / evidence[i].sum()
            assigned[i] = True
            # Propagation de la decision : le delta remplace le message
            # c2f de la cellule, et elle sort des free_slots des 3 unites
            # (lignes/colonnes/blocs) auxquelles elle appartient -- ses
            # mineurs ne seront plus calcules.
            m_c2f[UNIT_CELLS == i] = evidence[i]
            for u in np.flatnonzero([i in unit for unit in BP_UNITS]):
                free_slots[u] = np.array([s for s in free_slots[u]
                                          if BP_UNITS[u][s] != i], dtype=int)
                upd_mask[u] = False
                upd_mask[u, free_slots[u]] = True
            stats["decisions"] += 1
            beliefs = sweep(self.decim_sweeps)
        out = grid.copy()
        for i in range(81):
            if grid[i] == 0:
                out[i] = int(np.argmax(evidence[i])) + 1
        if not verify_solution_flat(out):
            stats["contradictions"] += 1
            return None
        return out


print("Arity9MaxProductSolver pret : 27 facteurs d'unite, hongrois + restarts randomises.")


def _sumprod_attempt(grid81, rng, restart=0, init_sweeps=60, decim_sweeps=6,
                     damping=0.5, noise=0.5):
    """Temoin sum-product : UN attempt de decimation a messages permanents
    (Ryser). Sert de comparateur au max-produit dans le bench."""
    grid = np.array(grid81, dtype=int)
    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()
    scale = 0.0 if restart == 0 else noise / (restart + 1)
    base = evidence[UNIT_CELLS]
    if scale > 0:
        m_c2f = _bp_normalize(
            np.exp(np.log(base + BP_EPS) + rng.uniform(-scale, scale, base.shape)))
    else:
        m_c2f = base.copy()
    m_f2c = np.full((27, 9, 9), 1.0 / 9)
    free_slots = [np.flatnonzero(grid[np.array(u)] == 0).astype(int) for u in BP_UNITS]
    upd_mask = np.zeros((27, 9), dtype=bool)
    for u in range(27):
        upd_mask[u, free_slots[u]] = True

    def sweep(n):
        nonlocal m_f2c, m_c2f
        for _ in range(n):
            allmin = unit_minors(m_c2f)
            sel = np.concatenate([allmin[u][free_slots[u]] for u in range(27)])
            pers = batch_perm8(sel.reshape(-1, 8, 8))
            new_f2c = np.zeros((27, 9, 9))
            ofs = 0
            for u in range(27):
                j = len(free_slots[u])
                if j:
                    new_f2c[u, free_slots[u]] = np.maximum(
                        pers[ofs:ofs + j * 9].reshape(j, 9), 0.0)
                    ofs += j * 9
            new_f2c = _bp_normalize(new_f2c)
            m_f2c = _bp_normalize(damping * m_f2c + (1 - damping) * new_f2c)
            log_tot = np.log(evidence + BP_EPS).copy()
            np.add.at(log_tot, UNIT_CELLS.reshape(-1),
                      np.log(m_f2c + BP_EPS).reshape(-1, 9))
            raw = log_tot[UNIT_CELLS] - np.log(m_f2c + BP_EPS)
            raw = raw - raw.max(axis=2, keepdims=True)
            upd = _bp_normalize(np.exp(raw))
            m_c2f = np.where(upd_mask[:, :, None],
                             _bp_normalize(damping * m_c2f + (1 - damping) * upd),
                             m_c2f)
        log_b = np.log(evidence + BP_EPS).copy()
        np.add.at(log_b, UNIT_CELLS.reshape(-1),
                  np.log(m_f2c + BP_EPS).reshape(-1, 9))
        b = np.exp(log_b - log_b.max(axis=1, keepdims=True))
        return b / b.sum(axis=1, keepdims=True)

    beliefs = sweep(init_sweeps)
    while not assigned.all():
        part = np.partition(beliefs, -2, axis=1)
        margins = part[:, -1] - part[:, -2]
        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
        evidence[i] = evidence[i] / evidence[i].sum()
        assigned[i] = True
        m_c2f[UNIT_CELLS == i] = evidence[i]
        for u in np.flatnonzero([i in unit for unit in BP_UNITS]):
            free_slots[u] = np.array([s for s in free_slots[u]
                                      if BP_UNITS[u][s] != i], dtype=int)
            upd_mask[u] = False
            upd_mask[u, free_slots[u]] = True
        beliefs = sweep(decim_sweeps)
    # Grille finale : les cellules donnees restent, les decimees
    # prennent leur valeur figee. Verif COMPLETE (lignes, colonnes,
    # blocs) : une decimation qui se contredit renvoie None -> le
    # solveur relancera une tentative avec un autre bruit.
    out = grid.copy()
    for i in range(81):
        if grid[i] == 0:
            out[i] = int(np.argmax(evidence[i])) + 1
    return out if verify_solution_flat(out) else None
Arity9MaxProductSolver pret : 27 facteurs d'unite, hongrois + restarts randomises.

Benchmark v4 : le critère Medium sans aucune couche déterministe

Même protocole que le bench v3 : Easy51, Medium (corpus Sudoku_hardest.txt, 11 grilles — la convention de difficulté du notebook C# 0-Environment) et les 15 premières de top95. Le solveur tourne en décimation pure — ni propagation de contraintes, ni repli/backtracking. Le critère d’acceptation : Medium ≥ 80 % résolu. Un témoin sum-product (permanente, 2 redémarrages seulement) est mesuré sur Medium pour ancrer la comparaison des deux opérateurs de message.

v4_solver = Arity9MaxProductSolver()   # parametres RR30 : les valeurs livrees
v4_rows = []                            # (plateau 0.25, bruit 0.6, 30 restarts)
for label, relpath, sub in bp_bench_sets:
    puzzle_strs = load_puzzles(relpath, max_puzzles=sub)
    flat_grids = [grid_to_flat(puzzle_to_grid(s)).tolist() for s in puzzle_strs]
    t0 = time.time()
    ok = 0
    # seed=k : la k-ieme grille de chaque corpus est TOUJOURS resolue avec
    # la meme semence -> bench reproductible d'une execution a l'autre.
    for k, p in enumerate(flat_grids):
        sol, st = v4_solver.solve(p, seed=k)
        if sol is not None and verify_solution_flat(sol):
            ok += 1
    dt = time.time() - t0
    v4_rows.append({"corpus": label, "solveur": "BP arite 9 max-produit + decimation",
                    "resolus": f"{ok}/{len(flat_grids)}",
                    "taux": round(ok / len(flat_grids), 2),
                    "ms/grille": round(dt / len(flat_grids) * 1000)})
    print(f"{label} {ok}/{len(flat_grids)} ({round(ok / len(flat_grids) * 100)}%) "
          f"{round(dt / len(flat_grids) * 1000)} ms/grille")
v4_df = pd.DataFrame(v4_rows)
medium_row = v4_df[v4_df["corpus"].str.contains("hardest")]
medium_ok = int(medium_row["resolus"].iloc[0].split("/")[0])
medium_tot = int(medium_row["resolus"].iloc[0].split("/")[1])
verdict = "ATTEINT" if medium_ok / medium_tot >= 0.8 else "MANQUE"
print(f"Critere Medium >= 80 % : {medium_ok}/{medium_tot} = "
      f"{round(medium_ok / medium_tot * 100)}% -> {verdict}")

# --- temoin sum-product (permanente) sur Medium, 2 redemarrages seulement
sumprod_rows = []
for label, relpath, sub in bp_bench_sets:
    if "hardest" not in label:
        continue
    puzzle_strs = load_puzzles(relpath, max_puzzles=sub)
    flat_grids = [grid_to_flat(puzzle_to_grid(s)).tolist() for s in puzzle_strs]
    t0 = time.time()
    ok = 0
    for k, p in enumerate(flat_grids):
        rng = np.random.default_rng(k)
        res = None
        for restart in range(2):
            res = _sumprod_attempt(p, rng, restart)
            if res is not None:
                break
        if res is not None:
            ok += 1
    dt = time.time() - t0
    sumprod_rows.append({"corpus": label, "solveur": "BP arite 9 sum-produit (temoin)",
                         "resolus": f"{ok}/{len(flat_grids)}",
                         "taux": round(ok / len(flat_grids), 2),
                         "ms/grille": round(dt / len(flat_grids) * 1000)})
    print(f"Temoin sum-product {label} : {ok}/{len(flat_grids)} "
          f"({round(ok / len(flat_grids) * 100)}%) "
          f"{round(dt / len(flat_grids) * 1000)} ms/grille")
pd.concat([v4_df, pd.DataFrame(sumprod_rows)], ignore_index=True)
Easy51 51/51 (100%) 1972 ms/grille
top95 (15 premieres) 9/15 (60%) 31752 ms/grille
hardest (11) 9/11 (82%) 21335 ms/grille
Critere Medium >= 80 % : 9/11 = 82% -> ATTEINT
Temoin sum-product hardest (11) : 1/11 (9%) 22049 ms/grille
corpus solveur resolus taux ms/grille
0 Easy51 BP arite 9 max-produit + decimation 51/51 1.00 1972
1 top95 (15 premieres) BP arite 9 max-produit + decimation 9/15 0.60 31752
2 hardest (11) BP arite 9 max-produit + decimation 9/11 0.82 21335
3 hardest (11) BP arite 9 sum-produit (temoin) 1/11 0.09 22049

Interprétation : le critère Medium ≥ 80 % est atteint sans aucune couche déterministe

Sur Medium (hardest, 11 grilles), la v4 en décimation pure résout 9/11 (82 %) — critère ≥ 80 % atteint — là où la v3 pairwise décimation plafonnait à 36/51 sur Easy et échouait sur la totalité de hardest. La hiérarchie des trois corpus garde sa pente : Easy 51/51 (100 %), Medium 9/11 (82 %), top95 9/15 (60 %) — l’arité 9 franchit le mur des ensembles de Hall qui bloquait Medium ; les corpus très clairsemés de top95 restent les plus durs pour une décimation sans repli. Le parallèle des « crans » se referme : v2 = un cran sur le modèle (Dirichlet), v4 = un cran sur le graphe (arité 9) — les deux restent dans le paradigme probabiliste pur ; la v3 C# reste seule à avoir payé ce mur avec de la propagation déterministe.

Coût par balayage : 810 facteurs O(1) contre 27 messages exacts d’arité 9

L’arité 9 n’est pas gratuite : un balayage v3 met à jour 810 messages de facteur en O(9) chacun ; un balayage v4 calcule, par cellule libre et par valeur candidate, une permanente 8×8 (Ryser, 2⁸ sous-ensembles — sum-product) ou un couplage optimal (hongrois — max-product). Mesure honnête des trois couches sur la même grille (première grille du corpus hardest, 20 balayages) :

# cout par balayage : meme grille, trois couches facteur, 20 balayages.
# On isole le cout d'UN balayage de chaque couche (hors decimation) :
# la v3 traite 810 facteurs O(9), la v4 sum traite ~1593 permanentes 8x8
# (Ryser 2^8 x 8), la v4 max ~1593 hongrois 8x8 -- meme probleme, trois
# prix differents pour le meme message exact.
cost_grid = grid_to_flat(puzzle_to_grid(load_puzzles("Puzzles/Sudoku_hardest.txt")[0]))
cost_evidence = np.full((81, 9), 1.0 / 9)
for i in np.flatnonzero(cost_grid > 0):
    cost_evidence[i, :] = BP_EPS
    cost_evidence[i, cost_grid[i] - 1] = 1.0
    cost_evidence[i] = cost_evidence[i] / cost_evidence[i].sum()

# --- couche v3 : 810 facteurs binaires
m_f2c_p = np.full((BP_F, 2, 9), 1.0 / 9)
m_c2f_p = np.full((BP_F, 2, 9), 1.0 / 9)
t0 = time.perf_counter()
for _ in range(20):
    new_f2c = np.empty((BP_F, 2, 9))
    new_f2c[:, 0] = _bp_normalize(1.0 - m_c2f_p[:, 1])
    new_f2c[:, 1] = _bp_normalize(1.0 - m_c2f_p[:, 0])
    m_f2c_p = _bp_normalize(0.5 * m_f2c_p + 0.5 * new_f2c)
    totals = np.log(cost_evidence + BP_EPS).copy()
    np.add.at(totals, BP_FI, np.log(m_f2c_p[:, 0] + BP_EPS))
    np.add.at(totals, BP_FJ, np.log(m_f2c_p[:, 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_p[:, s] + BP_EPS)
        new_c2f[:, s] = np.exp(raw - raw.max(axis=1, keepdims=True))
    m_c2f_p = _bp_normalize(0.5 * m_c2f_p + 0.5 * _bp_normalize(new_c2f))
t_pair = (time.perf_counter() - t0) / 20 * 1000

# --- couche v4 sum-product : permanentes des cellules libres (Ryser)
m_c2f_a = cost_evidence[UNIT_CELLS].copy()
free = [np.flatnonzero(cost_grid[np.array(u)] == 0).astype(int) for u in BP_UNITS]
t0 = time.perf_counter()
for _ in range(20):
    allmin = unit_minors(m_c2f_a)
    sel = np.concatenate([allmin[u][free[u]] for u in range(27)])
    batch_perm8(sel.reshape(-1, 8, 8))
t_perm = (time.perf_counter() - t0) / 20 * 1000

# --- couche v4 max-product : hongrois des cellules libres
t0 = time.perf_counter()
for _ in range(20):
    allmin = unit_minors(m_c2f_a)
    sel = np.concatenate([allmin[u][free[u]] for u in range(27)])
    hungarian_maxprod(sel.reshape(-1, 8, 8))
t_hung = (time.perf_counter() - t0) / 20 * 1000
n_minors = int(sum(len(f) * 9 for f in free))
print(f"Balayage v3 (paires)          : {t_pair:6.1f} ms  (810 facteurs x O(9))")
print(f"Balayage v4 sum-produit       : {t_perm:6.1f} ms  ({n_minors} permanentes 8x8, Ryser 2^8 x 8)")
print(f"Balayage v4 max-produit       : {t_hung:6.1f} ms  ({n_minors} couplages hongrois 8x8)")
print(f"Ratio v4-max/v3               : {t_hung/t_pair:.1f}x")
Balayage v3 (paires)          :    1.2 ms  (810 facteurs x O(9))
Balayage v4 sum-produit       :   51.4 ms  (1593 permanentes 8x8, Ryser 2^8 x 8)
Balayage v4 max-produit       :    9.2 ms  (1593 couplages hongrois 8x8)
Ratio v4-max/v3               : 7.9x

Interprétation : le prix exact de l’information structurelle

Un balayage v4 coûte 9.2 ms en max-product (hongrois) et 51.4 ms en sum-product (Ryser) contre 1.2 ms pour la v3 (7.9× le max-product) : 1593 couplages 8×8 ou 1593 permanentes 8×8 contre 810 facteurs O(9). Le surcoût est réel mais borné — quelques dizaines de millisecondes par balayage sur cette machine — et il n’achète pas « de la vitesse » mais de l’information : chaque message factor→cellule est désormais une quantité exacte sur les affectations injectives de l’unité (somme ou optimum), pas une approximation par paires. Sur un problème où la relaxation par paires est aveugle à une classe entière de structure (Hall), c’est le trade-off pertinent : payer un facteur borné par balayage pour convertir des grilles « impossibles » en résolues.

8. Comparaison avec Infer.NET RobustProbabilisticSolver (C#)

Le notebook C# Sudoku-15-Infer-CSharp.ipynb montre les résultats suivants pour le RobustProbabilisticSolver :

Puzzle Temps C# Statut C#
Easy #1 ~33ms (modèle compile) Resolu
Easy #2 ~34ms Resolu
Medium ~56ms 37 erreurs

Benchmark comparatif Python vs C

if JAX_AVAILABLE:
    def benchmark_comparison(puzzles: List[str], limit: int = 3):
        """Compare les solveurs Python avec les resultats C# Infer.NET."""
        print("=== Benchmark Python vs Infer.NET (RobustProbabilisticSolver) ===")
        print("\nResultats Python:")
        py_solver = RobustProbabilisticSolverPy(n_iterations=300, constraint_weight=15.0)
        results = []
        for i, puzzle_str in enumerate(puzzles[:limit]):
            grid = puzzle_to_grid(puzzle_str)
            start = time.time()
            solution, meta = py_solver.solve(grid)
            elapsed = time.time() - start
            result = {
                'puzzle': i + 1, 'time_python': elapsed,
                'converged': meta['converged'], 'errors': meta['errors']
            }
            results.append(result)
            status = "OK" if meta['converged'] else f"{meta['errors']} erreurs"
            print(f"  Puzzle {i+1}: {status} | Temps: {elapsed:.1f}s")
        print("\n" + "=" * 60)
        print("COMPARAISON PYTHON (NumPyro) vs C# (Infer.NET)")
        print("=" * 60)
        print(f"{'Puzzle':<10} {'Python (s)':<12} {'C# (ms)':<12} {'Python Status':<15} {'C# Status':<15}")
        print("-" * 60)
        csharp_results = [
            {'time_ms': 33.9, 'converged': True},
            {'time_ms': 33.9, 'converged': True},
            {'time_ms': 56.7, 'converged': False, 'errors': 37}
        ]
        for i, (py, cs) in enumerate(zip(results, csharp_results)):
            py_status = "Resolu" if py['converged'] else f"{py['errors']} err"
            cs_status = "Resolu" if cs['converged'] else f"{cs.get('errors', '?')} err"
            print(f"{i+1:<10} {py['time_python']:<12.1f} {cs['time_ms']:<12.1f} {py_status:<15} {cs_status:<15}")
        print("\nObservations:")
        print("  - Infer.NET est ~100x plus rapide grace a:")
        print("    * Expectation Propagation (plus efficace que SVI pour les CSP)")
        print("    * Contraintes dures (ConstrainFalse) vs douces")
        print("    * Modele precompile une seule fois")
        print("  - Les deux approches echouent sur les puzzles Medium")
        return results

    results = benchmark_comparison(puzzles, limit=3)
else:
    print("JAX requis : pip install jax jaxlib numpyro")
    print("Benchmark ignore. Installez JAX pour executer la comparaison.")
    print("\nResultats de reference (C# Infer.NET) :")
    print("  - Easy #1 : ~33ms, Resolu")
    print("  - Easy #2 : ~34ms, Resolu")
    print("  - Medium  : ~56ms, 37 erreurs")
=== Benchmark Python vs Infer.NET (RobustProbabilisticSolver) ===

Resultats Python:
  Puzzle 1: OK | Temps: 7.2s
  Puzzle 2: 8 erreurs | Temps: 5.9s
  Puzzle 3: 8 erreurs | Temps: 6.0s

============================================================
COMPARAISON PYTHON (NumPyro) vs C# (Infer.NET)
============================================================
Puzzle     Python (s)   C# (ms)      Python Status   C# Status      
------------------------------------------------------------
1          7.2          33.9         Resolu          Resolu         
2          5.9          33.9         8 err           Resolu         
3          6.0          56.7         8 err           37 err         

Observations:
  - Infer.NET est ~100x plus rapide grace a:
    * Expectation Propagation (plus efficace que SVI pour les CSP)
    * Contraintes dures (ConstrainFalse) vs douces
    * Modele precompile une seule fois
  - Les deux approches echouent sur les puzzles Medium

Interprétation : Comparaison Python vs C

Résultats obtenus : Le benchmark montre un écart de performance significatif entre Infer.NET (C#) et NumPyro (Python), de l’ordre de deux décades. Les durées ci-dessous sont des ordres de grandeur, reproductibles d’une machine à l’autre ; les mesures exactes de ce run restent visibles dans les cellules de mesure en amont.

Aspect Infer.NET (C#) NumPyro (Python) Ordre de grandeur
Temps Easy #1 ~30 ms ~5 s ×100
Temps Easy #2 ~30 ms ~5 s ×100
Temps Medium ~60 ms ~5 s ×100
Réussite Easy #1 Résolu Résolu Identique
Réussite Easy #2 Résolu 8 erreurs Infer.NET supérieur
Réussite Medium 37 erreurs 8 erreurs Python légèrement meilleur

Les ratios quantitatifs exacts sont machine-dépendants (compilation, JIT, GC, charge) : c’est pourquoi cette prose ne cite que des ordres de grandeur stables (mandat #9377/#9434 : seules les valeurs reproductibles d’une exécution à l’autre sont conservées en prose). Les résultats de résolution (résolu, N erreurs) sont déterministes et cités exactement.

Points clés : 1. Performance : Infer.NET est ~100× plus rapide grâce à la précompilation et à Expectation Propagation (ordre de grandeur stable d’une machine à l’autre) 2. Qualité de solution : Infer.NET réussit mieux les puzzles Easy (100% vs 50%), mais les deux moteurs montrent des difficultés sur Medium 3. Contraintes : Les contraintes dures d’Infer.NET (ConstrainFalse) sont plus efficaces que les contraintes douces de NumPyro 4. Variabilité : Les résultats Python varient plus (8 erreurs vs 0 pour Easy #2, 8 vs 37 pour Medium)

Note technique : L’écart de performance s’explique par trois facteurs : - EP vs SVI : Expectation Propagation est spécialement conçu pour les CSP, SVI est un algorithme généraliste - Contraintes : ConstrainFalse élimine directement les valeurs impossibles, alors que numpyro.factor ajoute seulement une pénalité - Compilation : Infer.NET génère une DLL compilée une seule fois, alors que JAX recompile à chaque exécution

9. Tableau comparatif final

import pandas as pd

comparison = {
    "Aspect": [
        "Bibliotheque",
        "Algorithme",
        "Contraintes",
        "Variables discretes",
        "Solveur robuste",
        "Solveur iteratif",
        "Solveur v3",
        "Solveur v4",
        "Performance (Easy)",
        "Performance (Medium)",
        "Precompilation"
    ],
    "Infer.NET Robust (C#)": [
        "Microsoft.ML.Probabilistic",
        "Expectation Propagation",
        "Dures (ConstrainFalse)",
        "Natives (Variable<int>)",
        "RobustProbabilisticSolver",
        "IterativeSudokuModel",
        "BacktrackingDecimationSolver (EP + propagation + repli)",
        "(stretch documente, non livre)",
        "~33ms",
        "Echec (37 erreurs)",
        "Oui (DLL generee)"
    ],
    "NumPyro (Python)": [
        "NumPyro + JAX",
        "SVI (Variational Inference)",
        "Douces (numpyro.factor)",
        "Via Dirichlet (continu)",
        "RobustProbabilisticSolverPy",
        "IterativeProbabilisticSolverPy",
        "BeliefPropagationSolver / BPBacktrackingSolver (non aligne)",
        "Arity9MaxProductSolver (facteurs AllDiff d'arite 9, hongrois + restarts randomises)",
        "~0.5s (BP, NumPy pur)",
        "Variable",
        "JIT compilation"
    ]
}

df = pd.DataFrame(comparison)
print(df.to_string(index=False))
              Aspect                                   Infer.NET Robust (C#)                                                                    NumPyro (Python)
        Bibliotheque                              Microsoft.ML.Probabilistic                                                                       NumPyro + JAX
          Algorithme                                 Expectation Propagation                                                         SVI (Variational Inference)
         Contraintes                                  Dures (ConstrainFalse)                                                             Douces (numpyro.factor)
 Variables discretes                                 Natives (Variable<int>)                                                             Via Dirichlet (continu)
     Solveur robuste                               RobustProbabilisticSolver                                                         RobustProbabilisticSolverPy
    Solveur iteratif                                    IterativeSudokuModel                                                      IterativeProbabilisticSolverPy
          Solveur v3 BacktrackingDecimationSolver (EP + propagation + repli)                         BeliefPropagationSolver / BPBacktrackingSolver (non aligne)
          Solveur v4                          (stretch documente, non livre) Arity9MaxProductSolver (facteurs AllDiff d'arite 9, hongrois + restarts randomises)
  Performance (Easy)                                                   ~33ms                                                               ~0.5s (BP, NumPy pur)
Performance (Medium)                                      Echec (37 erreurs)                                                                            Variable
      Precompilation                                       Oui (DLL generee)                                                                     JIT compilation

Conclusion

Ce que nous avons implemente

Classe Python Equivalent C# Description
RobustProbabilisticSolverPy RobustProbabilisticSolver Une seule inference, mode des Dirichlet
IterativeProbabilisticSolverPy IterativeSudokuModel Fixation iterative des cellules certaines
BeliefPropagationSolver (non aligne) Sum-product loopy BP + decimation par marge (NumPy pur)
BPBacktrackingSolver BacktrackingDecimationSolver (cousin) BP + repli borne sur contradiction
Arity9MaxProductSolver (non aligne) Facteurs AllDiff d’arité 9, messages exacts (permanente Ryser / hongrois scipy) + décimation randomisée

Lecons retenues

  1. Infer.NET reste superieur pour les CSP : Expectation Propagation + contraintes dures
  2. NumPyro fonctionne avec des limitations : Contraintes douces moins efficaces
  3. L’arité 9 franchit Medium sans propagation déterministe : 9/11 (82 %) sur hardest en décimation max-product randomisée
  4. Le BP vectorise domine le SVI sur ce probleme : messages exacts par facteur, ~10x plus rapide que le SVI pur (0,5 s vs 5-8 s) et ~350x que le SVI iteratif (0,5 s vs ~3 min), et la decimation couvre 71% de Easy51 la ou le SVI pur resout 1/3 de sa selection – mais la relaxation par paires cale sur les corpus durs (top95, hardest) : les ensembles de Hall sont invisibles pour des facteurs binaires
  5. Un cran sur le graphe, pas sur l’outil : la v4 (facteurs AllDiff d’arité 9, messages exacts permanente/hongrois + restarts randomisés) franchit le même mur que la propagation de contraintes en restant purement probabiliste — là où la v3 C# payait ce mur avec du déterminisme externe

Recommandation : Pour resoudre des Sudokus en production, utiliser OR-Tools, Z3 ou Choco. La programmation probabiliste est un outil pedagogique pour comprendre l’inference bayesienne.

Exercice : Solveur Hybride Probabiliste + Déterministe

Enonce

Le solveur probabiliste pur echoue souvent sur les puzzles Medium et Hard car l’inference SVI ne converge pas vers une solution exacte. Implementez un solveur hybride qui combine :

  1. Phase probabiliste : utiliser RobustProbabilisticSolverPy pour fixer les cellules avec une confiance superieure a threshold (ex : 0.90)
  2. Phase déterministe : appliquer un backtracking simple sur les cellules restantes

Structure a implementer :

class HybridProbabilisticSolver:
    def __init__(self, n_iterations=300, confidence_threshold=0.90):
        ...
    
    def solve(self, grid: List[List[int]]) -> Tuple[List[List[int]], dict]:
        # Phase 1 : inference probabiliste -> fixer les cellules certaines
        # Phase 2 : backtracking sur les cellules incertaines
        ...

Questions

  1. Quel seuil de confiance donne le meilleur compromis temps/precision ?
  2. Combien de cellules la phase probabiliste parvient-elle a fixer correctement ?
  3. Comparer le temps total avec le solveur iteratif (IterativeProbabilisticSolverPy)

Indice

Pour la phase probabiliste, utilisez np.max(probs[idx]) comme mesure de confiance. Pour le backtracking, vous pouvez reutiliser la logique de _backtrack du notebook Sudoku-01-Backtracking-Python.

class HybridProbabilisticSolver:
    """
    Solveur hybride : inference probabiliste (NumPyro) + backtracking deterministe.
    
    Phase 1 : inference probabiliste pour fixer les cellules certaines (confiance > threshold)
    Phase 2 : backtracking simple sur les cellules restantes
    """
    
    def __init__(self, n_iterations: int = 300, confidence_threshold: float = 0.90,
                 constraint_weight: float = 15.0):
        self.n_iterations = n_iterations
        self.confidence_threshold = confidence_threshold
        self.constraint_weight = constraint_weight
    
    def _backtrack_solve(self, grid: list) -> bool:
        """Backtracking simple sur les cellules restantes."""
        # TODO : Implementer le backtracking
        # Trouver la premiere case vide, essayer chaque valeur valide recursivitement
        pass
    
    def _is_valid(self, grid: list, row: int, col: int, val: int) -> bool:
        """Verifie si placer val en (row, col) est valide."""
        # TODO : Verifier ligne, colonne et bloc
        pass
    
    def solve(self, grid: list) -> tuple:
        """
        Resout un Sudoku avec le solveur hybride.
        
        Returns:
            (solution, metadata) avec metadata = {
                'probabilistic_fixes': int,  # cellules fixees par inference
                'inference_time': float,     # temps phase probabiliste
                'backtrack_time': float,     # temps phase backtracking
                'total_time': float,
                'converged': bool,
                'errors': int
            }
        """
        # TODO : Implementer le solveur hybride
        # 1. Appeler RobustProbabilisticSolverPy.infer_probabilities() sur le puzzle
        # 2. Pour chaque cellule vide :
        #    - Si confiance >= threshold : fixer la cellule avec la valeur la plus probable
        # 3. Appliquer _backtrack_solve sur les cellules encore vides
        pass


# Test (decommenter une fois les methodes implementees)
# test_grid = puzzle_to_grid(puzzles[0])
# hybrid_solver = HybridProbabilisticSolver(n_iterations=300, confidence_threshold=0.90)
# solution, meta = hybrid_solver.solve(test_grid)
# print(f"Hybride : {meta}")

print("Exercice solveur hybride a implementer !")
Exercice solveur hybride a implementer !

Navigation : << Choco | Index | Neural Network >>

Voir aussi : - Sudoku-15-Infer-CSharp - Version Infer.NET (C#) - Probas - Serie complete sur la programmation probabiliste - Sudoku-10-ORTools-Python - Approche CP-SAT (recommandee)

Retour au sommet