Problème de base : Les grands modèles de langage (LLMs) génèrent des requêtes « hallucinées » pour des bases de données. Par exemple, pour une question concernant les modifications d'un contrat de projet, le modèle peut interroegr directement la table contract_change_log en ignorant que la relation passe par les tables project et contract. Cela mène à des résultats vides ou erronés.
Solution : Intégrer une couche de routage (Router) qui détermine d'abord le chemin logique à travers le schéma de la base de données, puis génère la requête.
Approche en deux étapes
- Router : Produit un triplet
(Source, Cible, Via)définissant le chemin à emprunter. - Générateur : Construit la requête MongoDB (MQL) en suivant le chemin indiqué.
Validation préliminaire
Une expérience A/B sur des scénarios de $lookup a montré une amélioration de la précision de 80 % à 95 % et une réduction des requêtes vides de 20 à 5.
Mise en route avec LunarLander (RL)
Pour comprendre les bases du renforcement par essai-erreur, nous utilisons l'environnement LunarLander. Ci-dessous un wrapper configurable pour l'expérimentation des récompenses :
from dataclasses import dataclass
import numpy as np
import gymnasium as gym
@dataclass
class RewardSetup:
distance_weight: float = 0.0
velocity_weight: float = 0.0
angle_weight: float = 0.0
legs_weight: float = 0.0
main_thrust_penalty: float = 0.0
side_thrust_penalty: float = 0.0
use_custom: bool = False
reward_scale: float = 1.0
class ShapedRewardEnv(gym.Wrapper):
def __init__(self, env: gym.Env, cfg: RewardSetup):
super().__init__(env)
self.config = cfg
def _get_components(self, observation, action):
# Décompose l'observation en composantes de récompense
x, y, vx, vy, theta, omega, left_contact, right_contact = observation[:8]
distance = -np.sqrt(x**2 + y**2)
velocity = -np.sqrt(vx**2 + vy**2)
angle = -abs(theta)
legs = float(left_contact > 0.5) + float(right_contact > 0.5)
thrust_main = float(action == 2)
thrust_side = float(action in [1, 3])
return {
'dist': distance,
'vel': velocity,
'ang': angle,
'legs': legs,
'main_pen': thrust_main,
'side_pen': thrust_side
}
def step(self, action):
obs, original_reward, terminated, truncated, info = self.env.step(action)
components = self._get_components(obs, action)
shaped_reward = (
self.config.distance_weight * components['dist'] +
self.config.velocity_weight * components['vel'] +
self.config.angle_weight * components['ang'] +
self.config.legs_weight * components['legs'] -
self.config.main_thrust_penalty * components['main_pen'] -
self.config.side_thrust_penalty * components['side_pen']
) * self.config.reward_scale
info['reward_components'] = {**components, 'original': original_reward, 'shaped': shaped_reward}
final_reward = shaped_reward if self.config.use_custom else original_reward
return obs, final_reward, terminated, truncated, info
# Configurations d'exemple
default_cfg = RewardSetup() # Enregistre seulement
custom_cfg = RewardSetup(
distance_weight=100.0,
velocity_weight=150.0,
angle_weight=50.0,
legs_weight=10.0,
main_thrust_penalty=0.3,
side_thrust_penalty=0.03,
use_custom=True,
)
Ce code permet de comparer l'impact de différentes fonctions de récompense sur l'apprentissage.
Conception de l'environnement et des données
L'environnement de simulation repose sur des règles métier :
- 12 tables principales dans le schéma de la base de données.
- 42 chemins valides prédéfinis (l'espace d'action du Router).
Pipeline de données : génération initiale (Gemini) → augmentation sémantique (Qwen-72B) → remplissage d'entités → 566 échantillons finaux avec une répartition 60 % réels / 40 % ambigus.
Stratégie d'entraînement en deux phases
Phase 1 : Démarrage à froid avec SFT (Supervised Fine-Tuning)
Un entraînement RL direct échoue souvent dans les premières étapes à cause d'un espace d'action complexe (42 chemins). Le SFT fournit une initialisation solide (ex. précision ~80 %).
Phase 2 : Optimisation avec PPO (Proximal Policy Optimization)
PPO améliore la robustesse, surtout sur les requêtes complexes multi-sauts.
Architecture du modèle
Un backbone partiellement gelé (10 premières couches) et trois têtes de classification (Source, Cible, Via) ainsi qu'une tête de valeur (Critic).
Trick crucial : Initialiser les poids du Critic à partir du modèle SFT pour éviter l'instabilité.
Configuration stable pour PPO
Après de nombreux essais, une configuration conservatrice s'est avérée nécessaire :
| Paramètre | Plage recommandée | Raison |
|---|---|---|
learning_rate |
1e-6 à 5e-6 | Éviter les mises à jour brutales |
vf_coef |
0.01 à 0.1 | Limiter l'influence du Critic |
clip_range |
0.1 à 0.15 | Contraint les mises à jour de politique |
target_kl |
0.02 à 0.05 | Seuil de divergence pour arrêt |
| Couches gelées | 10 premières | Protéger les capacités linguistiques |
Ingénierie des récompenses
La récompense est basée sur des règles vérifiables, pas sur un modèle appris.
Conception par niveaux
- Niveau 1 (dense) : Récompenses pour les composants corrects (+1.0) et pénalités pour les incorrects (-0.5).
- Niveau 2 (contrainte) : Bonus pour les chemins légaux (+0.2), pénalité forte pour les chemins illégaux (-2.0).
- Niveau 3 (sparse) : Bonus important pour une correspondance parfaite (+10.0).
Anti-triche (Reward Hacking)
Problème : Le modèle prédisait systématiquement null pour le champ Via pour obtenir une récompense moyenne élevée.
Solution :
- Pondération dynamique : Augmenter fortement le poids (x10) des récompenses pour les champs
Vianon nuls. - Distribution conditionnelle : La récompense pour
Vian'est attribuée que siSourceest correct.
# Exemple de distribution conditionnelle
if source_est_correcte:
recompense += recompense_via * poids_dynamique
else:
recompense += 0
Évaluation des résultats
Comparaison de trois approches :
| Méthode | Correspondance complète | Précision Via |
|---|---|---|
| LLM général (baseline) | 35.71 % | - |
| SFT seul | 89.29 % | 87.50 % |
| RL (PPO après SFT) | 89.29 % | 89.29 % |
Résultat princpial : Le taux d'exécution bout-en-bout est passé de 80 % à 90 %, et le taux de requêtes vides/erreurs a baissé de 20 % à 10 %. L'apport du PPO se manifeste davantage sur les cas complexes et la stabilité en production.
Guide de dépannage
| Symptôme | Cause probable | Solution |
|---|---|---|
Divergence KL (approx_kl > 0.1) |
Pas de mise à jour trop important | Réduire learning_rate, clip_range ; abaisser target_kl |
| Récompense stagnante | Poids SFT non chargés | Vérifier l'initialisation du modèle |
| Champ Via toujours nul | Reward Hacking | Activer la pondération dynamique |
| Value Loss instable | Critic déséquilibré | Réduire vf_coef, geler plus de couches |
| Dégradation en cours d'entraînement | Oubli catastrophique | Arrêt anticipé, réduire n_epochs |
Indicateur clé à surveiller : approx_kl. Si elle dépasse 0.1, arrêter immédiatement pour diagnostic.
Leçons fondamentales
- Toujours démarrer par SFT avant le RL pour une initialisation efficace.
- PPO doit être conservateur : faible taux d'apprentissage, contrainte sur le Critic, divergence KL stricte.
- Concevoir des récompenses anti-triche avec pondération dynamique et conditions.