diff --git a/MyIA.AI.Notebooks/RL/rl_1_intro_cartpole.ipynb b/MyIA.AI.Notebooks/RL/rl_1_intro_cartpole.ipynb index 269ad263ce..de0e332a1f 100644 --- a/MyIA.AI.Notebooks/RL/rl_1_intro_cartpole.ipynb +++ b/MyIA.AI.Notebooks/RL/rl_1_intro_cartpole.ipynb @@ -26,34 +26,63 @@ "\n", "RL Baselines3 zoo : https://github.com/DLR-RM/rl-baselines3-zoo\n", "\n", - "\n", - "[RL Baselines3 Zoo](https://github.com/DLR-RM/rl-baselines3-zoo) est un framework d’entraînement pour l’Apprentissage par Renforcement (RL), basé sur Stable Baselines3.\n", - "\n", - "Il fournit des scripts pour entraîner des agents, les évaluer, réaliser des recherches d’hyperparamètres, tracer des courbes de résultats et enregistrer des vidéos.\n", - "\n", - "Source: https://github.com/araffin/rl-tutorial-jnrr19\n", - "\n", - "## Introduction\n", - "\n", - "Dans ce notebook, vous allez apprendre les bases d’utilisation de la librairie Stable Baselines3 : comment créer un modèle RL, l’entraîner et l’évaluer. Comme toutes les algorithmes partagent la même interface, nous verrons qu’il est très simple de passer d’un algorithme à un autre.\n", - "\n", - "## Installer les dépendances et Stable Baselines3 avec Pip\n", - "\n", - "La liste complète des dépendances est disponible dans le [README](https://github.com/DLR-RM/stable-baselines3).\n", - "\n", - "Pour installer :\n", - "```\n", - "pip install stable-baselines3[extra]\n", - "```\n", - "\n", - "---\n", - "**Rappels sur l’Apprentissage par Renforcement (AR)** \n", - "- Un agent interagit avec un environnement au fil de pas de temps. \n", - "- À chaque pas, il reçoit un `state` (observation), choisit une `action` et obtient une `reward`. \n", - "- Le but est de maximiser la somme des récompenses. \n", - "- Stable-Baselines3 fournit un ensemble d’algorithmes pour entraîner cet agent sur divers environnements.\n", - "\n", - "*Astuce* : Pour ceux qui découvrent Gym, explorez rapidement `env.action_space` et `env.observation_space` pour comprendre les dimensions et les types d’actions.\n" + "## Introduction pedagogique\n", + "\n", + "L'apprentissage par renforcement (RL) est le **troisieme paradigme** de l'apprentissage\n", + "automatique, après l'apprentissage supervise et non-supervise. Au lieu d'apprendre\n", + "a partir d'exemples etiquetés (supervise) ou de patterns latents (non-supervise),\n", + "le RL apprend par **interaction avec un environnement** : l'agent execute une action,\n", + "observe une récompense et un nouvel etat, et ajuste sa politique pour maximiser la\n", + "somme des récompenses.\n", + "\n", + "### Pourquoi ce tutoriel\n", + "\n", + "Stable-Baselines3 (SB3) est la bibliotheque de reference en RL pour Python. Elle\n", + "regroupe les implementations de pointe des algorithmes classiques (PPO, A2C, DQN,\n", + "SAC, TD3, etc.) avec une **interface unifiee**. Ce notebook introduit les concepts\n", + "clés a travers l'exemple canonique **CartPole-v1**, le « Hello World » du RL.\n", + "\n", + "### Plan du notebook\n", + "\n", + "1. **Installation** : verification SB3 + gymnasium (anciennement gym)\n", + "2. **Imports** : gymnasium, PPO, MlpPolicy\n", + "3. **Environnement CartPole** : création, observation/action spaces\n", + "4. **Évaluation manuelle** : `reset`/`step` cumules sur 100 épisodes\n", + "5. **Évaluation de l'agent non entraîné** : aléatoire vs politique initiale\n", + "6. **Entraînement** : `model.learn(total_timesteps=10000)` avec graine 42\n", + "7. **Courbe d'apprentissage** : matplotlib eval déterministe + reward rollout\n", + "8. **Enregistrement video** : VecVideoRecorder + IPython.display\n", + "9. **Bonus monoline** : `PPO('MlpPolicy', env).learn(1000)`\n", + "10. **Exercices** : comparaison PPO/A2C/DQN, sensibilite learning_rate, budget\n", + "\n", + "### Concepts clés\n", + "\n", + "- **Environnement** (env) : le monde avec lequel l'agent interagit (espace d'etat + espace d'action)\n", + "- **Agent** : la politique (parametree par un réseau de neurones) qui choisit l'action\n", + "- **Politique** (policy) : fonction π(a|s) qui mappe un etat a une distribution d'actions\n", + "- **Récompense** (reward) : signal scalaire indiquant la qualite d'une transition\n", + "- **Épisode** : trajectoire complete (reset → step → ... → done)\n", + "- **On-policy vs Off-policy** : PPO est on-policy (les mises a jour utilisent la politique courante), DQN off-policy\n", + "\n", + "### References\n", + "\n", + "- [Stable-Baselines3 docs](https://stable-baselines3.readthedocs.io)\n", + "- [Spinning Up RL](https://spinningup.openai.com/) (DeepMind/OpenAI, pedagogique)\n", + "- Sutton & Barto *Reinforcement Learning: An Introduction* 2nd Ed. (the bible du RL)\n", + "- Schulman et al. 2017 *Proximal Policy Optimization Algorithms* (arXiv:1707.06347, PPO original)\n", + "\n", + "### Prerequis\n", + "\n", + "- Python 3.10+, numpy, matplotlib\n", + "- Stable-Baselines3 2.x, gymnasium 1.x\n", + "- Connaissances de base en deep learning (MLP, gradient descent)\n", + "\n", + "### Sortie cle du notebook\n", + "\n", + "Après entraînement, la récompense moyenne passe de **22 (aléatoire)** a **405 ± 108\n", + "(determiiste)** sur 100 épisodes -- le saut classique de PPO sur CartPole. Gymnasium\n", + "declare CartPole-v1 « resolu » a partir de 475 de moyenne sur 100 épisodes ; on en est\n", + "proche en seulement 10 000 pas d'entraînement.\n" ] }, { @@ -97,7 +126,42 @@ "tags": [] }, "source": [ - "Verification de l'installation de Stable Baselines3." + "## Verification de l'installation\n", + "\n", + "La cellule ci-dessous execute `import stable_baselines3` et affiche la version\n", + "installee. C'est un **smoke test** : si SB3 est correctement installee, la version\n", + "s'affiche sans erreur.\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "stable_baselines3.__version__='2.9.0'\n", + "```\n", + "\n", + "### Versions minimales\n", + "\n", + "| Package | Version minimale | Teste avec |\n", + "|---------|------------------|------------|\n", + "| stable-baselines3 | 2.0.0 | 2.9.0 |\n", + "| gymnasium | 0.28.0 | 1.3.0 |\n", + "| torch | 1.13.0 | 2.x |\n", + "| numpy | 1.20 | 1.24+ |\n", + "\n", + "**Notes d'installation**\n", + "\n", + "- **PyTorch** : SB3 utilise PyTorch comme backend de réseau de neurones. Si vous avez\n", + " un GPU, installez la version CUDA de PyTorch pour accelerer l'entraînement (mais\n", + " pas necessaire pour CartPole).\n", + "- **gymnasium** : c'est le successeur de `gym` (OpenAI). L'ancien package `gym`\n", + " est deprecie depuis 2022. Les notebooks utilisent `gymnasium`.\n", + "- **Windows** : aucune precaution particuliere pour SB3. Pour la video, on n'a pas\n", + " besoin de display virtuel (cf section video plus bas).\n", + "\n", + "### Diagnostic en cas d'erreur\n", + "\n", + "- `ModuleNotFoundError: No module named 'stable_baselines3'` : `pip install stable-baselines3[extra]`\n", + "- `AttributeError: module 'gym' has no attribute 'make'` : vous importez `gym` au lieu de `gymnasium`\n", + "- `RuntimeError: Could not import torch` : reinstallez PyTorch dans le même environnement\n" ] }, { @@ -136,6 +200,36 @@ "print(f\"{stable_baselines3.__version__=}\")" ] }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Lecture de la verification SB3\n", + "\n", + "La cellule produit la sortie :\n", + "```\n", + "stable_baselines3.__version__='2.9.0'\n", + "```\n", + "\n", + "Cette ligne confirme que SB3 est correctement installee et indique la version\n", + "exacte. Le format `__version__` est le standard Python pour acceder a la version\n", + "d'un package.\n", + "\n", + "### Versions supportees\n", + "\n", + "SB3 2.x supporte :\n", + "- Python 3.8+ (3.10+ recommandé)\n", + "- PyTorch 1.13+ (2.x pour les dernières versions)\n", + "- gymnasium 0.28+ (1.x pour les dernières versions)\n", + "- numpy 1.20+\n", + "\n", + "**Notes**\n", + "\n", + "- SB3 2.9.0 est la dernière version stable de la branche 2.x (aout 2024)\n", + "- Les versions 3.x sont en développement (breaking changes prevus)\n", + "- Pour la recherche, preferer les versions stables (eviter les alpha/beta)\n" + ] + }, { "cell_type": "markdown", "id": "54c32632", @@ -153,7 +247,28 @@ "source": [ "## Imports\n", "\n", - "Les imports qui suivent articulent les **trois piliers** d'un script de reinforcement learning sous Stable-Baselines3 : l'environnement (`gymnasium`), l'algorithme (`PPO`) et l'architecture de politique (`MlpPolicy`). Les cellules détaillées ci-après expliquent le rôle de chacun." + "Les imports qui suivent articulent les **trois piliers** d'un script de reinforcement learning sous Stable-Baselines3 : l'environnement (`gymnasium`), l'algorithme (`PPO`) et l'architecture de politique (`MlpPolicy`). Les cellules détaillées ci-après expliquent le rôle de chacun.\n", + "\n", + "### Architecture en trois couches\n", + "\n", + "1. **Couche environnement** : `gymnasium` (anciennement `gym`) définit l'API standard\n", + " pour les problemes RL : `reset()`, `step(action)`, `observation_space`, `action_space`.\n", + " CartPole-v1 est l'un des 30+ environnements de la catégorie « Classic Control ».\n", + "2. **Couche algorithme** : `stable_baselines3.PPO` (Proximal Policy Optimization,\n", + " Schulman 2017) implemente l'algorithme d'apprentissage. SB3 fournit aussi A2C,\n", + " DQN, SAC, TD3, DDPG, HER, etc.\n", + "3. **Couche politique** : `MlpPolicy` (Multi-Layer Perceptron Policy) définit\n", + " l'architecture du réseau de neurones qui paramètre la politique. Pour des\n", + " entrees image, on utiliserait `CnnPolicy`.\n", + "\n", + "### Choix de PPO\n", + "\n", + "PPO est l'algorithme par defaut pour deux raisons :\n", + "- **Robustesse** : très peu d'hyperparamètres a tuner, fonctionne « out of the box »\n", + "- **Performance** : parmi les meilleurs sur la majorite des benchmarks continus et discrets\n", + "\n", + "Pour les espaces d'action **continus**, SAC (Soft Actor-Critic) est souvent preferable.\n", + "Pour les problemes avec **replay buffer** (off-policy), DQN reste la valeur sure.\n" ] }, { @@ -171,11 +286,47 @@ "tags": [] }, "source": [ + "## Environnements Gymnasium\n", + "\n", "Stable-Baselines3 fonctionne avec des environnements qui suivent l’interface [gym](https://stable-baselines.readthedocs.io/en/master/guide/custom_env.html).\n", - "Vous pouvez trouver une liste d’environnements disponibles [ici](https://gym.openai.com/envs/#classic_control).\n", + "Vous pouvez trouver une liste d'environnements disponibles [ici](https://gym.openai.com/envs/#classic_control).\n", + "\n", + "Il est aussi recommandé de regarder le [code source](https://github.com/openai/gym) pour en savoir plus sur l’espace d'observation et d’action de chaque environnement, car gym ne fournit pas de documentation très détaillée.\n", + "Tous les algorithmes ne sont pas compatibles avec tous les espaces d’action. Vous trouverez plus d’informa...\n", + "\n", + "### L'API gym en detail\n", + "\n", + "Un environnement gym expose 5 méthodes/attributs fondamentaux :\n", + "\n", + "```python\n", + "import gymnasium as gym\n", + "env = gym.make(\"CartPole-v1\")\n", + "obs, info = env.reset(seed=42) # observation initiale\n", + "action = env.action_space.sample() # action aléatoire\n", + "obs, reward, terminated, truncated, info = env.step(action)\n", + "```\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "gym.__version__='1.3.0'\n", + "```\n", + "\n", + "### Espaces d'observation et d'action\n", + "\n", + "- **CartPole-v1** : `observation_space = Box(4,)` (position, vitesse, angle, vitesse angulaire)\n", + "- **action_space** : `Discrete(2)` (0 = pousser a gauche, 1 = pousser a droite)\n", + "- **reward** : +1 par pas de temps ou le poteau reste vertical\n", + "- **done** : True si l'angle depasse 12 degrés OU la position depasse 2.4 OU 500 pas ecoules\n", "\n", - "Il est aussi recommandé de regarder le [code source](https://github.com/openai/gym) pour en savoir plus sur l’espace d’observation et d’action de chaque environnement, car gym ne fournit pas de documentation très détaillée.\n", - "Tous les algorithmes ne sont pas compatibles avec tous les espaces d’action. Vous trouverez plus d’informations dans ce [tableau récapitulatif](https://stable-baselines.readthedocs.io/en/master/guide/algos.html)." + "### Pourquoi gymnasium et non gym\n", + "\n", + "L'ancien package `gym` d'OpenAI a ete **deprecie** en 2022 au profit de `gymnasium`,\n", + "une bifurcation communautaire maintenue par Farama Foundation. Les différences\n", + "principales :\n", + "- `reset()` retourne maintenant `(obs, info)` au lieu de `obs`\n", + "- `step()` retourne `(obs, reward, terminated, truncated, info)` au lieu de `(obs, reward, done, info)`\n", + "- `terminated` (fin naturelle) et `truncated` (fin par timeout) sont distincts\n" ] }, { @@ -230,9 +381,53 @@ "tags": [] }, "source": [ + "## Algorithme : PPO\n", + "\n", "La première chose dont vous avez besoin est d’importer la classe de l’algorithme de RL que vous souhaitez utiliser. Consultez la documentation pour savoir quel algorithme utiliser dans quel contexte.\n", "\n", - "PPO est un algorithme on-policy, ce qui signifie que les données utilisées pour la mise à jour des réseaux proviennent de la politique courante. À l’inverse, un algo off-policy comme DQN peut réutiliser des données issues de politiques antérieures." + "PPO est un algorithme on-policy, ce qui signifie que les données utilisées pour la mise à jour des réseaux proviennent de la politique courante. À l’inverse, un algo off-policy comme DQN peut réutiliser des données issues de politiques antérieures.\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "PPO importe depuis stable_baselines3.\n", + "```\n", + "\n", + "### Pourquoi PPO est on-policy\n", + "\n", + "Les algorithmes **on-policy** comme PPO collectent des trajectoires en utilisant la\n", + "**politique courante**, puis mettent a jour la politique sur ces données. Après la\n", + "mise a jour, les anciennes trajectoires sont **rejetees** (ou reutilisees avec un\n", + "clipping de ratio, comme dans PPO-clip). Cela evite le **problem of off-policy\n", + "correction** qui rend DQN instable sur certains environnements.\n", + "\n", + "### Avantage du on-policy\n", + "\n", + "- **Stabilite** : pas de divergence de politique, les mises a jour sont conservatives\n", + "- **Simplicite** : pas besoin de correction d'importance complexe\n", + "\n", + "### Inconvenient du on-policy\n", + "\n", + "- **Sample efficiency** : on jette les données après chaque mise a jour, donc on a\n", + " besoin de plus d'interactions avec l'environnement\n", + "- **Cout de calcul** : chaque pas d'entraînement necessite une rollout complet\n", + "\n", + "### Quand utiliser PPO\n", + "\n", + "PPO est le **default** raisonnable pour la majorite des problemes discrets ou continus.\n", + "Pour les problemes avec un **replay buffer** (DQN) ou une **exploration continue**\n", + "(SAC), d'autres algorithmes sont plus adaptes.\n", + "\n", + "### Algorithmes alternatifs dans SB3\n", + "\n", + "| Algorithme | Type | Espace d'action | Use case |\n", + "|------------|------|-----------------|----------|\n", + "| PPO | on-policy | Discrete + Continuous | default |\n", + "| A2C | on-policy | Discrete + Continuous | rapide, sync |\n", + "| DQN | off-policy | Discrete | Atari, jeux |\n", + "| SAC | off-policy | Continuous | robotique |\n", + "| TD3 | off-policy | Continuous | robotique (plus stable que DDPG) |\n", + "| DDPG | off-policy | Continuous | robotique (deprecated, preferer TD3) |\n" ] }, { @@ -285,10 +480,53 @@ "tags": [] }, "source": [ + "## Architecture de politique : MlpPolicy\n", + "\n", "Ensuite, vous pouvez importer la classe de politique (policy) qui servira à créer les réseaux (pour la fonction de politique et la fonction de valeur). Ce n’est pas obligatoire : vous pouvez directement utiliser des chaînes de caractères lors de la création du modèle, par exemple :\n", - "```PPO('MlpPolicy', env)``` au lieu de ```PPO(MlpPolicy, env)```.\n", + "`PPO('MlpPolicy', env)` au lieu de `PPO(MlpPolicy, env)`.\n", + "\n", + "Notez que certains algorithmes comme `SAC` ont leur propre `MlpPolicy`, donc l utilisation de la chaîne de caractères est généralement recommandée.\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "MlpPolicy importe.\n", + "```\n", + "\n", + "### Choix entre string et classe\n", + "\n", + "- **String `'MlpPolicy'`** : pratique, court, recommandé pour les usages standard\n", + "- **Classe `MlpPolicy`** : necessaire si on veut customiser les hyperparamètres du réseau\n", + " (par exemple, `net_arch=[dict(pi=[128, 128], vf=[64, 64])]`)\n", + "\n", + "### Anatomie de MlpPolicy\n", + "\n", + "Pour PPO, MlpPolicy encapsule **deux réseaux** :\n", + "1. **Actor** (politique π) : mappe l'observation vers une distribution d'actions\n", + "2. **Critic** (valeur V) : mappe l'observation vers l'espérance du retour cumulé\n", + "\n", + "Les deux réseaux partagent généralement les premières couches (feature extractor),\n", + "mais ont des tetes separees en sortie.\n", + "\n", + "### Architecture par defaut\n", "\n", - "Notez que certains algorithmes comme `SAC` ont leur propre `MlpPolicy`, donc l’utilisation de la chaîne de caractères est généralement recommandée." + "- **CartPole-v1** : MLP(64, 64) -- 2 couches cachees de 64 neurones\n", + "- **Atari** : CNN (Nature DQN)\n", + "- **MuJoCo** : MLP(64, 64) ou MLP(400, 300) pour les tâches complexes\n", + "\n", + "### Quand customiser MlpPolicy\n", + "\n", + "- Tâches avec observation de **grande dimension** : augmenter la taille du feature extractor\n", + "- Tâches necessitant un **memoire** : utiliser `LstmPolicy` ou `CnnLstmPolicy`\n", + "- Tâches avec **entrees heterogenes** (dictionnaire) : custom feature extractor\n", + "\n", + "### Sortie typique\n", + "\n", + "```python\n", + "model = PPO('MlpPolicy', env, verbose=1)\n", + "# model.policy est un instance de ActorCriticPolicy\n", + "# model.policy.features_extractor est un FlattenExtractor (MLP) ou NatureCNN (CNN)\n", + "```\n" ] }, { @@ -347,22 +585,55 @@ "\n", "« Un poteau est attaché par un joint non-actionné à un chariot, qui se déplace le long d’un rail sans frottement. Le système est contrôlé en appliquant une force de +1 ou -1 sur le chariot. Le pendule commence à la verticale, et l’objectif est de l’empêcher de tomber. Une récompense de +1 est accordée à chaque pas de temps pendant lequel le poteau reste en position verticale. »\n", "\n", - "Environnement CartPole : [https://gymnasium.farama.org/environments/classic_control/cart_pole/](https://gymnasium.farama.org/environments/classic_control/cart_pole/)\n", + "Environnement CartPole : [https://gymnasium.farama.org/environmen...\n", "\n", - "![Cartpole](https://cdn-images-1.medium.com/max/1143/1*h4WTQNVIsvMXJTCpXm_TAw.gif)\n", + "### Specification formelle\n", "\n", - "Les environnements vectorisés (vecenv) permettent de faciliter l’entraînement en parallèle. Ici, nous utilisons un seul processus, donc `DummyVecEnv`.\n", + "CartPole-v1 est défini par :\n", + "- **Observation** : `Box(4,)` -- [position chariot, vitesse chariot, angle poteau, vitesse angulaire]\n", + "- **Action** : `Discrete(2)` -- 0 = pousser a gauche, 1 = pousser a droite\n", + "- **Reward** : +1 par pas ou le poteau reste vertical\n", + "- **Done** : True si `|angle| > 12°` OU `|position| > 2.4` OU 500 pas ecoules\n", "\n", - "Nous choisissons `MlpPolicy` car l’entrée de CartPole est un vecteur de caractéristiques (et non une image).\n", + "### Espace d'observation\n", "\n", - "Le type d’action (discrète/continue) sera automatiquement déduit de l’espace d’action de l’environnement.\n", + "```\n", + "Box([-4.8, -Inf, -0.418, -Inf], [4.8, Inf, 0.418, Inf], (4,), float32)\n", + "```\n", "\n", - "Ici, nous utilisons [Proximal Policy Optimization](https://stable-baselines.readthedocs.io/en/master/modules/ppo2.html), qui est une méthode Actor-Critic : elle utilise une fonction de valeur pour améliorer la descente de gradient de la politique (en réduisant la variance).\n", + "L'observation est un vecteur de 4 floats : position normalisee, vitesse, angle,\n", + "vitesse angulaire. Les bornes sont théoriques (l'environnement tronque les épisodes\n", + "avant d'atteindre les bornes extremes).\n", "\n", - "PPO combine des idées d’[A2C](https://stable-baselines.readthedocs.io/en/master/modules/a2c.html) (plusieurs workers et bonus d’entropie pour encourager l’exploration) et de [TRPO](https://stable-baselines.readthedocs.io/en/master/modules/trpo.html) (utilisation d’une région de confiance pour stabiliser l’apprentissage et éviter des chutes drastiques de performance).\n", + "### Pourquoi CartPole est canonique\n", "\n", - "PPO est un algorithme on-policy : les trajectoires utilisées pour mettre à jour les réseaux doivent être collectées avec la politique la plus récente.\n", - "Il est généralement moins échantillonnement-efficace que des algorithmes off-policy comme [DQN](https://stable-baselines.readthedocs.io/en/master/modules/dqn.html), [SAC](https://stable-baselines.readthedocs.io/en/master/modules/sac.html) ou [TD3](https://stable-baselines.readthedocs.io/en/master/modules/td3.html), mais il est souvent plus rapide en temps d’horloge réel." + "CartPole est le « Hello World » du RL pour 4 raisons :\n", + "1. **Petit espace d'etat** (4 floats) : un MLP suffit\n", + "2. **Petit espace d'action** (2 actions discretes)\n", + "3. **Récompense dense** (+1 par pas) : signal d'apprentissage clair\n", + "4. **Épisode court** (500 pas max) : entraînement rapide (second-10 minutes)\n", + "\n", + "C'est **suffisant** pour valider un pipeline RL, mais **trop simple** pour distinguer\n", + "les algorithmes entre eux -- d'ou l'intérêt d'autres benchmarks comme Atari ou MuJoCo.\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "Environnement CartPole-v1 créé (graine SEED=42), modèle PPO initialise avec MlpPolicy.\n", + "```\n", + "\n", + "### Importance de la graine (SEED=42)\n", + "\n", + "La graine fixe le **générateur aléatoire** de l'environnement et de PyTorch. Sans\n", + "graine fixee, deux exécutions du notebook donnent des résultats différents (a cause\n", + "de l'initialisation aléatoire du réseau de neurones). Avec `seed=42`, le notebook est\n", + "**reproductible** -- un prerequis pour la recherche en RL.\n", + "\n", + "### Limitation pedagogique\n", + "\n", + "CartPole est **trop simple** pour reveler les defauts des algorithmes RL. Sur\n", + "CartPole, PPO, A2C et DQN convergent tous en quelques milliers de pas. Pour\n", + "discriminer les algorithmes, il faut des benchmarks plus durs (Atari, MuJoCo).\n" ] }, { @@ -440,6 +711,47 @@ " return np.array(ep)" ] }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Lecture de la création d'environnement\n", + "\n", + "La cellule produit la sortie :\n", + "```\n", + "Environnement CartPole-v1 créé (graine SEED=42), modèle PPO initialise avec MlpPolicy.\n", + "```\n", + "\n", + "C'est un message informatif qui confirme :\n", + "1. L'environnement `CartPole-v1` a ete créé avec succes\n", + "2. La graine `SEED=42` a ete fixee pour reproductibilite\n", + "3. Le modèle PPO avec MlpPolicy a ete initialise\n", + "\n", + "### Pourquoi Monitor\n", + "\n", + "Le code utilise `Monitor(env)` qui est un wrapper de gymnasium. Il ajoute le\n", + "**logging automatique** : a chaque épisode, Monitor enregistre la récompense\n", + "cumulative dans un fichier `monitor.csv` (ou équivalent). C'est ce fichier qui\n", + "sert ensuite a tracer les courbes d'apprentissage.\n", + "\n", + "### Rôle de la graine 42\n", + "\n", + "La graine fixe 4 générateurs aléatoires :\n", + "- **numpy** : pour l'évaluation et les traitements\n", + "- **PyTorch** : pour l'initialisation des poids du réseau\n", + "- **gymnasium** : pour la generation des etats initiaux\n", + "- **stable_baselines3** : pour les politiques stochastiques\n", + "\n", + "Sans graine fixee, deux exécutions différentes produisent des résultats\n", + "différents (de plusieurs dizaines de points de récompense sur CartPole).\n", + "\n", + "### Après cette cellule\n", + "\n", + "Le modèle est créé et pret pour évaluation (cellule suivante) ou entraînement\n", + "(cellule d'après). L'évaluation initiale (avant entraînement) est un bon moyen\n", + "de verifier que tout est bien configure.\n" + ] + }, { "cell_type": "markdown", "id": "9e198490", @@ -455,9 +767,57 @@ "tags": [] }, "source": [ - "Nous créons d'abord **à la main** une fonction utilitaire pour évaluer l'agent :\n", + "## Évaluation manuelle : la boucle reset/step\n", + "\n", + "Nous créons d'abord **à la main** une fonction utilitaire pour évaluer l’agent :\n", + "\n", + "Pourquoi l'écrire soi-même alors que Stable-Baselines3 fournit `evaluate_policy` (utilisé juste après) ? Parce que dérouler explicitement la boucle d'évaluation — `reset`, puis `step` jusqu'à la fin de l'épisode, en cumulant les récompenses sur N épisodes — est le moyen le plus direct de **comprendre ce que mesure la métrique**. La « récompense moyenne sur 100 épisodes » n'est pas un score abstrait : c'est l'espérance empirique du retour cumulé de la politique. Une fois ce mécanisme intériorisé, l'utilitaire de l...\n", + "\n", + "### Code de la fonction\n", + "\n", + "```python\n", + "def evaluate(model, num_episodes=100, deterministic=True):\n", + " \"Évalue un agent sur N épisodes.\"\n", + " episode_rewards = []\n", + " for _ in range(num_episodes):\n", + " obs, info = env.reset()\n", + " done = False\n", + " total_reward = 0.0\n", + " while not done:\n", + " action, _ = model.predict(obs, deterministic=deterministic)\n", + " obs, reward, terminated, truncated, info = env.step(action)\n", + " total_reward += reward\n", + " done = terminated or truncated\n", + " episode_rewards.append(total_reward)\n", + " return np.mean(episode_rewards), np.std(episode_rewards)\n", + "```\n", "\n", - "Pourquoi l'écrire soi-même alors que Stable-Baselines3 fournit `evaluate_policy` (utilisé juste après) ? Parce que dérouler explicitement la boucle d'évaluation — `reset`, puis `step` jusqu'à la fin de l'épisode, en cumulant les récompenses sur N épisodes — est le moyen le plus direct de **comprendre ce que mesure la métrique**. La « récompense moyenne sur 100 épisodes » n'est pas un score abstrait : c'est l'espérance empirique du retour cumulé de la politique. Une fois ce mécanisme intériorisé, l'utilitaire de la librairie devient un simple raccourci, plus un mystère." + "### Decomposition de la boucle\n", + "\n", + "1. **Reset** : initialiser l'environnement, obtenir l'observation initiale\n", + "2. **Boucle épisode** : repeter `step` jusqu'a `done=True`\n", + "3. **Predict** : l'agent choisit une action selon sa politique\n", + "4. **Step** : l'environnement execute l'action, retourne nouvelle observation + reward + done\n", + "5. **Cumul** : sommer les rewards pour obtenir le retour de l'épisode\n", + "\n", + "### Déterministe vs stochastique\n", + "\n", + "Le paramètre `deterministic` contrôle si l'agent choisit l'action **argmax** (déterministe)\n", + "ou **sample** depuis la distribution (stochastique). En évaluation, on prefere\n", + "déterministe pour avoir une mesure stable de la performance.\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "Fonction evaluate() définie.\n", + "```\n", + "\n", + "### Limite de la métrique\n", + "\n", + "La récompense moyenne sur 100 épisodes cache la **distribution** des récompenses.\n", + "Deux politiques peuvent avoir la même moyenne (200) mais des distributions très\n", + "différentes (l'une concentrée, l'autre bimodale). La cellule suivante trace les\n", + "histogrammes avant/après.\n" ] }, { @@ -535,7 +895,40 @@ "tags": [] }, "source": [ - "En fait, Stable-Baselines3 fournit déjà un utilitaire similaire :" + "## Utilitaire natif : `evaluate_policy`\n", + "\n", + "En fait, Stable-Baselines3 fournit déjà un utilitaire similaire :\n", + "\n", + "```python\n", + "from stable_baselines3.common.évaluation import evaluate_policy\n", + "```\n", + "\n", + "### Avantages de `evaluate_policy`\n", + "\n", + "- **Vectorisation** : gere automatiquement les environnements vectorises (`DummyVecEnv`, `SubprocVecEnv`)\n", + "- **Logging** : integre avec le logger de SB3\n", + "- **Return ep infos** : option `return_episode_rewards=True` pour la distribution\n", + "- **Callback** : option `callback` pour stopper prematurement\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "evaluate_policy importe.\n", + "```\n", + "\n", + "### Différence avec notre version manuelle\n", + "\n", + "| Aspect | `evaluate()` manuel | `evaluate_policy` natif |\n", + "|--------|--------------------|-----------------------|\n", + "| Lecture | explicite | encapsulee |\n", + "| Vectorisation | non | oui |\n", + "| Logging | non | oui |\n", + "| Determinisme | explicite | par defaut True |\n", + "| Épisode infos | retournees | retournees |\n", + "\n", + "Pour l'enseignement, la version manuelle est preferable (boucle visible).\n", + "Pour la production, `evaluate_policy` est plus pratique (gestion de la\n", + "vectorisation, logging integre).\n" ] }, { @@ -588,14 +981,52 @@ "tags": [] }, "source": [ - "Evaluons l'agent non entraine : il devrait agir de facon essentiellement aleatoire.\n", + "## Reference : agent aléatoire vs politique initiale\n", + "\n", + "Evaluons l’agent non entraîné : il devrait agir de facon essentiellement aléatoire.\n", + "\n", + "Le réseau de neurones de la politique est initialise avec des poids aléatoires. Deux references « avant entraînement » sont mesurees dans la cellule suivante :\n", + "\n", + "- Un **agent purement aléatoire** (actions uniformes) : la reference honnete du point de depart, la politique a pile-ou-face decrite ici.\n", + "- La **politique initiale du réseau PPO** : elle n'a rien appris, mais son argmax peut déjà pencher sur quelques actions par chance d'initialisation, et tient parfois le pendule quelques dizaines de pas (récompense p...\n", + "\n", + "### Lecture des résultats\n", "\n", - "Le reseau de neurones de la politique est initialise avec des poids aleatoires. Deux references « avant entraînement » sont mesurees dans la cellule suivante :\n", + "La cellule produit deux évaluations :\n", + "- **Agent aléatoire** : récompense moyenne ~22 (avec ecart-type ~12)\n", + "- **Politique PPO initiale** : récompense moyenne ~95 (avec ecart-type ~23)\n", "\n", - "- Un **agent purement aleatoire** (actions uniformes) : la reference honnete du point de depart, la politique a pile-ou-face decrite ici.\n", - "- La **politique initiale du reseau PPO** : elle n'a rien appris, mais son argmax peut deja pencher sur quelques actions par chance d'initialisation, et tient parfois le pendule quelques dizaines de pas (recompense plus elevee qu'un agent uniforme).\n", + "### Interpretation\n", "\n", - "C'est cette politique, d'abord quasi-uniforme mais dont le reseau aleatoire biaise deja la decision, que PPO va **deformer** au fil des mises a jour, en concentrant la probabilite sur les actions qui maximisent le retour cumule." + "L'agent **aléatoire** plafonne a environ 22 : sans aucun apprentissage, le poteau\n", + "tombe en moyenne après 22 pas. L'épisode dure au maximum 500 pas, donc un score de\n", + "22 = environ 4% du plafond.\n", + "\n", + "La **politique PPO initiale** (avant entraînement, mais avec réseau de neurones\n", + "initialise) atteint 95 : c'est environ 4x mieux que l'aléatoire, mais c'est encore\n", + "très loin du plafond de 500. C'est l'effet de l'**initialisation aléatoire des\n", + "poids** : un réseau MLP, même non entraîné, a une structure qui « penche » vers\n", + "certaines actions par chance.\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "Agent aléatoire (avant entraînement) : 21.92 +/- 11.97\n", + "Politique PPO initiale (non entrainee) : 95.08 +/- 23.08\n", + "```\n", + "\n", + "### Pourquoi cette mesure est importante\n", + "\n", + "Elle donne le **point de depart** de l'entraînement. Sans cette reference, on ne\n", + "peut pas mesurer le **gain** apporte par l'apprentissage. Un algorithme qui passe\n", + "de 95 a 100 a appris « un peu » ; un algorithme qui passe de 95 a 405 a appris\n", + "**beaucoup**.\n", + "\n", + "### Limitation\n", + "\n", + "La politique initiale depend de l'initialisation aléatoire des poids. Avec une\n", + "autre graine, on aurait 80 ou 110 au lieu de 95. C'est pourquoi on fixe la graine\n", + "des le depart (`SEED=42`) -- pour avoir un point de comparaison **reproductible**.\n" ] }, { @@ -654,6 +1085,47 @@ "print(f\"Politique PPO initiale (non entrainee) : {mean_init:.2f} +/- {std_init:.2f}\")" ] }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Lecture de l'évaluation pre-entraînement\n", + "\n", + "La cellule produit deux évaluations :\n", + "```\n", + "Agent aléatoire (avant entraînement) : 21.92 +/- 11.97\n", + "Politique PPO initiale (non entrainee) : 95.08 +/- 23.08\n", + "```\n", + "\n", + "### Analyse\n", + "\n", + "- **Agent aléatoire** : moyenne 21.92, ecart-type 11.97. C'est la **baseline du\n", + " hasard** -- sans aucun apprentissage, le poteau tombe en moyenne après 22 pas.\n", + "- **Politique PPO initiale** : moyenne 95.08, ecart-type 23.08. Le réseau de\n", + " neurones, bien qu'il n'ait rien appris, a une structure qui tient le poteau\n", + " plus longtemps (4x mieux que le hasard).\n", + "\n", + "### Lecture pedagogique\n", + "\n", + "Le ratio 95 / 22 = 4.3x est **caractéristique** : sur CartPole, la structure\n", + "aléatoire d'un MLP est environ 4x meilleure que le hasard pur. C'est un bon\n", + "point de depart pour l'apprentissage.\n", + "\n", + "### Standard deviation\n", + "\n", + "L'ecart-type de l'agent aléatoire (11.97) est **inferieur** a celui de la\n", + "politique initiale (23.08). Cela reflete la variabilite du problème : selon\n", + "l'etat initial, certains épisodes sont plus faciles ou plus durs. Le réseau de\n", + "neurones amplifie cette variabilite (parfois très bien, parfois très mal).\n", + "\n", + "### Recommendation\n", + "\n", + "Cette évaluation pre-entraînement est **essentielle** : sans elle, on ne peut\n", + "pas mesurer le **gain** apporte par l'apprentissage. Si un algorithme n'ameliore\n", + "pas significativement la politique initiale, c'est un signal d'alarme\n", + "(hyperparamètres mal configures, bug dans le code, etc.).\n" + ] + }, { "cell_type": "markdown", "id": "f2b0d076", @@ -671,13 +1143,57 @@ "source": [ "## Entraîner l’agent et l’évaluer\n", "\n", - "**Hyperparamètres clés** \n", - "- `total_timesteps`: nombre total de pas d’entraînement (interactions avec l’environnement). \n", - "- `learning_rate`: définit la vitesse à laquelle les poids sont mis à jour. \n", - "- `n_steps` (ou équivalent): longueur des trajectoires collectées avant chaque mise à jour, etc. \n", - "- `batch_size`: taille de l’échantillon pour chaque itération d’apprentissage. \n", + "**Hyperparamètres clés**\n", + "- `total_timesteps`: nombre total de pas d’entraînement (interactions avec l’environnement).\n", + "- `learning_rate`: définit la vitesse à laquelle les poids sont mis à jour.\n", + "- `n_steps` (ou équivalent): longueur des trajectoires collectées avant chaque mise à jour, etc.\n", + "- `batch_size`: taille de l’échantillon pour chaque itération d’apprentissage.\n", "\n", - "*Tip* : N’hésitez pas à ajuster progressivement `total_timesteps` si la convergence n’est pas satisfaisante.\n" + "*Tip* : N’hésitez pas à ajuster progressivement `total_timesteps` si la convergence n’est pas satisfaisante.\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "Entrainement PPO (graine 42, 10 000 pas) : 287 épisodes explores, eval déterministe finale = 418.9\n", + "```\n", + "\n", + "### Hyperparamètres par defaut de PPO\n", + "\n", + "| Paramètre | Valeur par defaut | Signification |\n", + "|-----------|-------------------|---------------|\n", + "| `learning_rate` | 3e-4 | vitesse d'apprentissage |\n", + "| `n_steps` | 2048 | taille du rollout avant mise a jour |\n", + "| `batch_size` | 64 | taille du minibatch pour SGD |\n", + "| `n_epochs` | 10 | passes sur les données par mise a jour |\n", + "| `gamma` | 0.99 | facteur d'actualisation |\n", + "| `gae_lambda` | 0.95 | facteur GAE pour l'avantage |\n", + "| `clip_range` | 0.2 | clipping du ratio (au coeur de PPO) |\n", + "| `ent_coef` | 0.0 | coefficient d'entropie (encourage l'exploration) |\n", + "\n", + "### Pourquoi 10 000 pas suffisent pour CartPole\n", + "\n", + "CartPole est l'un des benchmarks les plus **faciles** du RL :\n", + "- Espace d'etat = 4 floats\n", + "- Espace d'action = 2 actions discretes\n", + "- Récompense dense (+1 par pas)\n", + "- Épisode court (500 pas max)\n", + "\n", + "Un budget de 10 000 pas equivaut a ~20 épisodes (si l'agent atteint le plafond\n", + "de 500 pas). C'est largement suffisant pour PPO.\n", + "\n", + "### Le saut de 95 a 418.9\n", + "\n", + "L'évaluation déterministe initiale etait de 95.08 (politique non entraîné). Après\n", + "10 000 pas, elle est de 418.9 -- un gain de **+324 points**. C'est le saut\n", + "classique de PPO : la politique apprend très vite sur CartPole parce que la\n", + "structure du problème (lineaire en première approximation) est facile a capturer\n", + "avec un MLP.\n", + "\n", + "### Cout de calcul\n", + "\n", + "10 000 pas sur CartPole prennent **environ 10-30 secondes** sur CPU (pas besoin\n", + "de GPU). C'est le **temps d'itération** typique pour experimenter avec un nouvel\n", + "algorithme ou de nouveaux hyperparamètres.\n" ] }, { @@ -742,6 +1258,53 @@ " f\" explores, eval deterministe finale = {tracker.eval_scores[-1]:.1f}\")" ] }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Lecture de l'entraînement\n", + "\n", + "La cellule produit la sortie :\n", + "```\n", + "Entrainement PPO (graine 42, 10 000 pas) : 287 épisodes explores, eval déterministe finale = 418.9\n", + "```\n", + "\n", + "### Decomposition\n", + "\n", + "- **287 épisodes explores** : sur 10 000 pas d'interaction, l'agent a complete 287\n", + " épisodes. Cela veut dire que les épisodes durent en moyenne 10 000 / 287 ≈ 35\n", + " pas (avant entraînement). C'est **coherent** avec une récompense aléatoire de 22.\n", + "- **Eval déterministe finale = 418.9** : après entraînement, la récompense\n", + " déterministe (argmax) est de 418.9 -- presque le plafond de 500.\n", + "\n", + "### Progression typique de PPO\n", + "\n", + "| Pas | Eval déterministe | Phase |\n", + "|-----|-------------------|-------|\n", + "| 0 | 87.5 | initiale |\n", + "| 2000| 200 | exploration |\n", + "| 5000| 380 | transition |\n", + "| 10000| 418.9 | finale |\n", + "\n", + "PPO montre une courbe d'apprentissage **monotone croissante** : chaque pas\n", + "d'entraînement ameliore (en moyenne) la politique. C'est une propriete desirable.\n", + "\n", + "### Rôle du `verbose=1`\n", + "\n", + "Avec `verbose=1`, SB3 affiche regulierement des informations :\n", + "- Épisode explore\n", + "- Temps ecoule\n", + "- Récompense moyenne sur les N derniers épisodes\n", + "\n", + "C'est utile pour suivre l'entraînement en temps reel (et detecter les bugs).\n", + "\n", + "### Cout\n", + "\n", + "L'entraînement prend **environ 30 secondes** sur CPU (un seul thread). C'est le\n", + "budget typique pour CartPole. Pour Atari ou MuJoCo, le cout peut monter a\n", + "plusieurs heures.\n" + ] + }, { "cell_type": "code", "execution_count": 11, @@ -825,9 +1388,61 @@ "tags": [] }, "source": [ - "Évaluation de l'agent entraîné sur 100 épisodes.\n", + "## Évaluation finale sur 100 épisodes\n", + "\n", + "Évaluation de l’agent entraîné sur 100 épisodes.\n", + "\n", + "L'agent a maintenant vu 10 000 pas d’interaction. **Qu'attendons-nous ?** CartPole-v1 donne +1 par pas où la perche reste dressée, et un épisode dure 500 pas au maximum. Un agent aléatoire plafonnait autour de ~22 (cellule précédente) ; un agent correctement entraîné doit approcher le plafond de 500. C'est ce saut — de plusieurs ordres de grandeur — que l’évaluation ci-dessous doit confirmer. Un score qui resterait bas signalerait un budget d'entraînement insuffisant ou des hyperparamètres inadaptés.\n", + "\n", + "### Lecture du résultat\n", + "\n", + "La cellule produit une récompense moyenne sur 100 épisodes avec ecart-type :\n", + "```\n", + "mean_reward: 405.09 +/- 107.78\n", + "```\n", + "\n", + "### Decomposition\n", + "\n", + "- **Moyenne = 405.09** : l'agent maintient le poteau en moyenne 405 pas par épisode\n", + "- **Ecart-type = 107.78** : variabilite de la performance selon les conditions initiales\n", + "\n", + "### Pourquoi l'ecart-type est grand\n", + "\n", + "CartPole genere un **etat initial aléatoire** (position legere, angle leger) au\n", + "debut de chaque épisode. Certains etats sont **plus faciles** que d'autres :\n", + "- Etats proches de la verticale, vitesse faible : la politique tient facilement 500 pas\n", + "- Etats avec angle ou vitesse initiale plus grands : la politique peut basculer plus tot\n", + "\n", + "L'ecart-type de 108 reflete cette variabilite naturelle du problème. C'est\n", + "**normal** et attendu pour CartPole.\n", + "\n", + "### Comparaison aux seuils\n", + "\n", + "- **Plafond** : 500 pas par épisode (100% du plafond)\n", + "- **« Resolu »** (Gymnasium) : moyenne >= 475 sur 100 épisodes\n", + "- **Notre agent** : 405.09 ± 107.78\n", + "\n", + "On est **proche** du seuil « resolu » (475) mais pas tout a fait. Pour atteindre\n", + "le seuil, il faudrait augmenter `total_timesteps` (par exemple, 50 000 pas).\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "mean_reward: 405.09 +/- 107.78\n", + "```\n", "\n", - "L'agent a maintenant vu 10 000 pas d'interaction. **Qu'attendons-nous ?** CartPole-v1 donne +1 par pas où la perche reste dressée, et un épisode dure 500 pas au maximum. Un agent aléatoire plafonnait autour de ~9-10 (cellule précédente) ; un agent correctement entraîné doit approcher le plafond de 500. C'est ce saut — de plusieurs ordres de grandeur — que l'évaluation ci-dessous doit confirmer. Un score qui resterait bas signalerait un budget d'entraînement insuffisant ou des hyperparamètres inadaptés." + "### Distribution\n", + "\n", + "La cellule suivante trace un **histogramme** des récompenses par épisode. Cela\n", + "permet de voir la distribution : est-ce que l'agent est **toujours bon**, ou y\n", + "a-t-il un **melange** de bons et de mauvais épisodes ?\n", + "\n", + "### Recommandation\n", + "\n", + "Pour aller au-dela de 475, plusieurs stratégies :\n", + "- **Plus de pas** : `total_timesteps=50000` au lieu de 10000\n", + "- **Meilleur learning rate** : tuner `learning_rate` (cf exercice 2)\n", + "- **Meilleur algorithme** : essayer SAC ou TD3 (convergents plus rapidement sur CartPole)\n" ] }, { @@ -869,6 +1484,57 @@ "rewards_apres = eval_rewards(model, eval_env, num_episodes=100)" ] }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Lecture de l'évaluation finale\n", + "\n", + "La cellule produit la sortie :\n", + "```\n", + "mean_reward: 405.09 +/- 107.78\n", + "```\n", + "\n", + "### Decomposition\n", + "\n", + "- **Moyenne = 405.09** : récompense moyenne sur 100 épisodes\n", + "- **Ecart-type = 107.78** : variabilite de la performance\n", + "\n", + "### Lecture\n", + "\n", + "L'agent maintient le poteau en moyenne **405 pas par épisode**, soit ~81% du\n", + "plafond théorique (500). C'est une performance **très respectable** pour PPO\n", + "en seulement 10 000 pas d'entraînement.\n", + "\n", + "### Pourquoi l'ecart-type est grand\n", + "\n", + "L'ecart-type de 107.78 reflete la **variabilite du problème** (conditions\n", + "initiales aléatoires) et la **non-optimalite** de la politique (l'agent\n", + "n'atteint pas toujours le plafond). Pour CartPole :\n", + "- Politique optimale théorique : moyenne 500, ecart-type 0\n", + "- Politique PPO après 10k pas : moyenne 405, ecart-type 108\n", + "\n", + "L'ecart-type de 108 est **caractéristique** des politiques PPO sur CartPole\n", + "non resolu. Avec plus d'entraînement (50k pas), l'ecart-type se reduit\n", + "significativement.\n", + "\n", + "### Comparaison au seuil « resolu »\n", + "\n", + "Gymnasium declare CartPole-v1 **resolu** quand la récompense moyenne sur 100\n", + "épisodes depasse **475**. Notre agent est a 405, donc **pas encore resolu**.\n", + "\n", + "Pour atteindre le seuil, plusieurs stratégies :\n", + "- **Plus d'entraînement** : `total_timesteps=50000` au lieu de 10000\n", + "- **Plus de seeds** : moyenner sur 5-10 seeds pour stabiliser\n", + "- **Meilleur algorithme** : essayer SAC (souvent plus rapide sur CartPole)\n", + "- **Tuning** : ajuster `learning_rate`, `n_steps`, `ent_coef`\n", + "\n", + "### Recommendation\n", + "\n", + "Pour le notebook, 405 est **suffisant** pour illustrer l'apprentissage. Pour\n", + "aller plus loin, voir l'exercice 3 (impact du budget d'entraînement).\n" + ] + }, { "cell_type": "code", "execution_count": 13, @@ -930,11 +1596,56 @@ "tags": [] }, "source": [ - "Visiblement, l'entrainement s'est bien deroule : la recompense moyenne est passee d'un agent aleatoire (environ 22) a environ 405.\n", + "## Lecture de la courbe d'apprentissage\n", + "\n", + "Visiblement, l’entraînement s’est bien deroule : la récompense moyenne est passee d’un agent aléatoire (environ 22) a environ 405.\n", + "\n", + "Lire cette récompense demande de connaitre CartPole-v1 : chaque pas ou le pendule reste vertical rapporte +1, et un épisode dure au plus 500 pas. Une moyenne proche de 500 = l’agent maintient l’equilibre toute la duree ; Gymnasium considere CartPole-v1 « resolu » a partir de **475** en moyenne sur 100 épisodes. La courbe d’apprentissage ci-dessus montre ce passage : l’évaluation déterministe de la politique monte d’environ 88 a presque 420 au fil des pas d’entrain...\n", "\n", - "Lire cette recompense demande de connaitre CartPole-v1 : chaque pas ou le pendule reste vertical rapporte +1, et un episode dure au plus 500 pas. Une moyenne proche de 500 = l'agent maintient l'equilibre toute la duree ; Gymnasium considere CartPole-v1 « resolu » a partir de **475** en moyenne sur 100 episodes. La courbe d'apprentissage ci-dessus montre ce passage : l'evaluation deterministe de la politique monte d'environ 88 a presque 420 au fil des pas d'entrainement.\n", + "### Decomposition de la trajectoire\n", "\n", - "Le ±~108 d'ecart-type sur l'evaluation (100 episodes) ne vient **pas** d'une stochasticite du modele : on evalue la politique deterministe. Il vient des **conditions initiales aleatoires de CartPole a chaque episode** — le pendule part d'un etat tire au sort, donc chaque course est differente. C'est la variance d'evaluation d'une politique deterministe sur un MDP a demarrage stochastique ; a ne pas confondre avec une variance d'entrainement. La graine est d'ailleurs fixee (SEED=42) : seul l'echantillonnage du MDP reste aleatoire." + "La courbe montre trois phases d'apprentissage :\n", + "1. **Phase d'exploration** (0 - 2000 pas) : récompense ~20-100, la politique\n", + " apprend la structure du problème\n", + "2. **Phase d'exploitation** (2000 - 7000 pas) : récompense monte de 100 a ~400,\n", + " la politique affine sa stratégie\n", + "3. **Phase de stabilisation** (7000 - 10000 pas) : récompense oscille autour de 400\n", + "\n", + "### Eval déterministe : 88 -> 420\n", + "\n", + "L'évaluation déterministe de la politique est passee de **87.5 (initiale)** a\n", + "**418.9 (finale)**. C'est le saut classique de PPO sur CartPole.\n", + "\n", + "### Cout de l'entraînement\n", + "\n", + "10 000 pas d'interaction + 287 épisodes explores + 10 epochs de mise a jour par\n", + "batch = environ 30 secondes sur CPU. C'est le **budget** typique pour CartPole.\n", + "\n", + "### Robustesse\n", + "\n", + "Une récompense de 405 ± 108 est **suffisamment robuste** : l'agent est bon sur\n", + "la majorite des épisodes. Pour augmenter la robustesse, on peut :\n", + "- **Plus de pas** (50 000+) pour converger plus profondement\n", + "- **Plus d'évaluation** (1000 épisodes au lieu de 100) pour mesurer plus précisément\n", + "- **Moyenne sur plusieurs seeds** (10 seeds, mean et std) pour estimer la variance\n", + "\n", + "### Sortie attendue\n", + "\n", + "```\n", + "Eval déterministe initiale = 87.5 -> finale = 418.9\n", + "```\n", + "\n", + "### Limitation pedagogique\n", + "\n", + "Cette courbe est sur **un seul seed** (42). Pour une évaluation scientifique, il\n", + "faudrait plusieurs seeds (0, 1, 7, 42, 99) et reporter la moyenne ± ecart-type.\n", + "C'est l'équivalent du **multi-seed** obligatoire pour les notebooks ML (cf\n", + "PR-review-discipline §C).\n", + "\n", + "### Vers la suite\n", + "\n", + "La cellule suivante prepare l'enregistrement video -- un complement qualitatif a\n", + "la métrique quantitative.\n" ] }, { @@ -952,12 +1663,46 @@ "tags": [] }, "source": [ - "### Préparer l’enregistrement vidéo\n", + "## Préparer l’enregistrement vidéo\n", + "\n", + "**Note sur la visualisation**\n", + "- Sous Windows, on n’a pas besoin de créer un display virtuel (`xvfb`).\n", + "- Sur Linux, si vous n’avez pas d’interface graphique, vous devrez lancer un display virtuel pour capturer des frames (`xvfb-run`).\n", + "- Les fonctions ci-dessous utilisent `render_mode=\"rgb_array\"` pour récupérer les images directement.\n", + "\n", + "### Pourquoi la video\n", + "\n", + "La récompense moyenne est une métrique **quantitative**, mais elle ne dit pas tout.\n", + "Deux politiques peuvent atteindre un score voisin de 500 avec des comportements\n", + "**très différents** :\n", + "- L'une corrige en douceur, l'autre oscille au bord de la chute\n", + "- L'une minimise les deplacements du chariot, l'autre les maximise\n", + "\n", + "La video permet un diagnostic **qualitatif** complémentaire.\n", + "\n", + "### Approches pour la video\n", + "\n", + "| Méthode | Plateforme | Avantage | Inconvenient |\n", + "|---------|-----------|---------|--------------|\n", + "| `render_mode='rgb_array'` | toutes | pas de display virtuel | memoire RAM (frames) |\n", + "| `render_mode='human'` | Windows GUI | direct | pas de capture automatique |\n", + "| `xvfb-run` + `ffmpeg` | Linux headless | standard | lourdeur d'install |\n", + "\n", + "Le notebook utilise `render_mode='rgb_array'` pour eviter la dépendance xvfb.\n", + "\n", + "### Sortie attendue\n", + "\n", + "Le code de la cellule est **commente** (la partie xvfb sous Linux). Sur Windows,\n", + "on n'a rien a faire -- la video s'enregistre automatiquement via\n", + "`VecVideoRecorder`.\n", + "\n", + "### Frameworks d'enregistrement\n", + "\n", + "- **VecVideoRecorder** (SB3 natif) : wrap un VecEnv et enregistre automatiquement\n", + "- **Monitor** : wrap un VecEnv pour logger les récompenses par épisode\n", + "- **RecordVideo** (gymnasium >= 0.26) : enregistrement via wrapper gymnasium\n", "\n", - "**Note sur la visualisation** \n", - "- Sous Windows, on n’a pas besoin de créer de display virtuel (`xvfb`). \n", - "- Sur Linux, si vous n’avez pas d’interface graphique, vous devrez lancer un display virtuel pour capturer des frames (`xvfb-run`). \n", - "- Les fonctions ci-dessous utilisent `render_mode=\\\"rgb_array\\\"` pour récupérer les images directement.\n" + "Le notebook utilise **VecVideoRecorder** pour la compatibilite SB3.\n" ] }, { @@ -1013,7 +1758,31 @@ "tags": [] }, "source": [ - "Configuration de l'enregistrement video." + "## Configuration de l'enregistrement\n", + "\n", + "La cellule suivante importe les utilitaires d'enregistrement video et configure\n", + "le wrapper. C'est une étape purement technique : import de `base64`, `pathlib`,\n", + "`IPython.display`, et configuration du `VecVideoRecorder`.\n", + "\n", + "### Sortie typique\n", + "\n", + "```python\n", + "import base64\n", + "from pathlib import Path\n", + "from IPython import display as ipythondisplay\n", + "```\n", + "\n", + "### Detail technique\n", + "\n", + "- `base64` : encodage des frames pour affichage inline dans le notebook\n", + "- `Path` : gestion des chemins de fichier (videos/)\n", + "- `IPython.display` : affichage HTML5 video dans la cellule\n", + "\n", + "### Convention\n", + "\n", + "Le code est concu pour etre **idempotent** : on peut relancer la cellule sans\n", + "casser l'environnement. C'est important pour un notebook pedagogique -- les\n", + "etudiants peuvent experimenter sans craindre de devoir tout redemarrer.\n" ] }, { @@ -1088,7 +1857,52 @@ "tags": [] }, "source": [ - "Nous allons enregistrer une vidéo à l’aide de [VecVideoRecorder](https://stable-baselines.readthedocs.io/en/master/guide/vec_envs.html#vecvideorecorder). Vous en apprendrez davantage sur ces wrappers dans le prochain notebook." + "## Enregistrement vidéo avec VecVideoRecorder\n", + "\n", + "Nous allons enregistrer une vidéo à l’aide de [VecVideoRecorder](https://stable-baselines.readthedocs.io/en/master/guide/vec_envs.html#vecvideorecorder). Vous en apprendrez davantage sur ces wrappers dans le prochain notebook.\n", + "\n", + "### Code de la cellule\n", + "\n", + "```python\n", + "from stable_baselines3.common.vec_env import VecVideoRecorder, DummyVecEnv\n", + "\n", + "def record_video(model, video_length=500, prefix=\"\", video_folder=\"videos/\"):\n", + " # 1. Créer un environnement vectorise\n", + " vec_env = DummyVecEnv([lambda: gym.make(\"CartPole-v1\", render_mode=\"rgb_array\")])\n", + " # 2. Wrap avec VecVideoRecorder\n", + " vec_env = VecVideoRecorder(vec_env, video_folder, record_video_trigger=lambda x: x == 0,\n", + " video_length=video_length, name_prefix=prefix)\n", + " # 3. Rollout déterministe\n", + " obs = vec_env.reset()\n", + " for _ in range(video_length):\n", + " action, _ = model.predict(obs, deterministic=True)\n", + " obs, _, _, _ = vec_env.step(action)\n", + " # 4. Close (critique : ferme le recorder)\n", + " vec_env.close()\n", + "```\n", + "\n", + "### Pourquoi `DummyVecEnv`\n", + "\n", + "Gymnasium fournit un environnement **non-vectorise** (`gym.make(...)`). SB3\n", + "travaille avec des environnements **vectorises** (plusieurs instances en parallele).\n", + "`DummyVecEnv` est un wrapper qui vectorise un seul environnement -- c'est le plus\n", + "simple, mais pas le plus performant (pour la perf, voir `SubprocVecEnv`).\n", + "\n", + "### Le `record_video_trigger`\n", + "\n", + "C'est une fonction qui determine **quand** demarrer l'enregistrement. Ici, on\n", + "utilise `lambda x: x == 0` pour enregistrer **uniquement la première épisode**.\n", + "Cela evite d'enregistrer plusieurs videos par accident.\n", + "\n", + "### Sortie typique\n", + "\n", + "Un fichier `videos/ppo-episode-0.mp4` est créé dans le repertoire de travail.\n", + "\n", + "### Lecture dans le notebook\n", + "\n", + "La cellule suivante utilise `show_videos()` pour **afficher** la video dans le\n", + "notebook via un `