|
94 | 94 | "source": [ |
95 | 95 | "# Configuration for our referential game\n", |
96 | 96 | "config = CommunicationConfig(\n", |
97 | | - " vocabulary_size=10, # Size of the message vocabulary\n", |
98 | | - " message_length=1, # Length of messages (1 token for simplicity)\n", |
99 | | - " hidden_size=64, # Hidden layer size\n", |
| 97 | + " vocabulary_size=10, # Size of the message vocabulary\n", |
| 98 | + " message_length=1, # Length of messages (1 token for simplicity)\n", |
| 99 | + " hidden_size=64, # Hidden layer size\n", |
100 | 100 | ")\n", |
101 | 101 | "\n", |
102 | 102 | "# Learning parameters\n", |
|
151 | 151 | "# Simple training loop for referential game\n", |
152 | 152 | "from src.langlab.data.world import encode_object\n", |
153 | 153 | "\n", |
| 154 | + "\n", |
154 | 155 | "def train_step(speaker, listener, scene_objects, target_idx, speaker_opt, listener_opt):\n", |
155 | 156 | " \"\"\"Single training step for the referential game.\"\"\"\n", |
156 | 157 | " speaker.train()\n", |
157 | 158 | " listener.train()\n", |
158 | | - " \n", |
| 159 | + "\n", |
159 | 160 | " # Encode the scene objects\n", |
160 | 161 | " scene_tensor = torch.stack([encode_object(obj) for obj in scene_objects]).to(device)\n", |
161 | | - " target_tensor = torch.tensor([target_idx], dtype=torch.long).to(device) # Add batch dimension\n", |
162 | | - " \n", |
| 162 | + " target_tensor = torch.tensor([target_idx], dtype=torch.long).to(\n", |
| 163 | + " device\n", |
| 164 | + " ) # Add batch dimension\n", |
| 165 | + "\n", |
163 | 166 | " # Speaker generates message about target object\n", |
164 | | - " target_object = scene_tensor[target_idx:target_idx+1]\n", |
165 | | - " message_logits, message_tokens, gesture_logits, gesture_tokens = speaker(target_object)\n", |
166 | | - " \n", |
| 167 | + " target_object = scene_tensor[target_idx : target_idx + 1]\n", |
| 168 | + " message_logits, message_tokens, gesture_logits, gesture_tokens = speaker(\n", |
| 169 | + " target_object\n", |
| 170 | + " )\n", |
| 171 | + "\n", |
167 | 172 | " # Add batch dimension to scene_tensor for Listener (batch_size=1, num_candidates=3, object_dim)\n", |
168 | 173 | " scene_tensor_batched = scene_tensor.unsqueeze(0) # Shape: (1, 3, object_dim)\n", |
169 | | - " \n", |
| 174 | + "\n", |
170 | 175 | " # Listener tries to identify target from message and scene\n", |
171 | 176 | " listener_logits = listener(message_tokens, scene_tensor_batched)\n", |
172 | | - " \n", |
| 177 | + "\n", |
173 | 178 | " # Compute listener loss (cross-entropy)\n", |
174 | 179 | " listener_loss = torch.nn.functional.cross_entropy(listener_logits, target_tensor)\n", |
175 | | - " \n", |
| 180 | + "\n", |
176 | 181 | " # Compute speaker loss (REINFORCE-style)\n", |
177 | 182 | " # Get the predicted target from listener\n", |
178 | 183 | " predicted_target = torch.argmax(listener_logits, dim=1)\n", |
179 | 184 | " # Reward is 1 if correct, 0 if incorrect\n", |
180 | 185 | " reward = (predicted_target == target_tensor).float()\n", |
181 | | - " \n", |
| 186 | + "\n", |
182 | 187 | " # Speaker loss: encourage generating messages that lead to correct predictions\n", |
183 | 188 | " # Use the message logits to compute log probabilities\n", |
184 | 189 | " message_probs = torch.softmax(message_logits, dim=-1)\n", |
185 | | - " message_log_probs = torch.log(message_probs + 1e-8) # Add small epsilon for numerical stability\n", |
186 | | - " \n", |
| 190 | + " message_log_probs = torch.log(\n", |
| 191 | + " message_probs + 1e-8\n", |
| 192 | + " ) # Add small epsilon for numerical stability\n", |
| 193 | + "\n", |
187 | 194 | " # Get the log probability of the generated message\n", |
188 | | - " speaker_log_prob = message_log_probs.gather(2, message_tokens.unsqueeze(-1)).squeeze(-1)\n", |
| 195 | + " speaker_log_prob = message_log_probs.gather(\n", |
| 196 | + " 2, message_tokens.unsqueeze(-1)\n", |
| 197 | + " ).squeeze(-1)\n", |
189 | 198 | " speaker_log_prob = speaker_log_prob.sum(dim=1) # Sum over message length\n", |
190 | | - " \n", |
| 199 | + "\n", |
191 | 200 | " # REINFORCE loss: -log_prob * reward\n", |
192 | 201 | " speaker_loss = -(speaker_log_prob * reward).mean()\n", |
193 | | - " \n", |
| 202 | + "\n", |
194 | 203 | " # Total loss\n", |
195 | 204 | " total_loss = listener_loss + 0.1 * speaker_loss # Weight speaker loss less\n", |
196 | | - " \n", |
| 205 | + "\n", |
197 | 206 | " # Backpropagation\n", |
198 | 207 | " speaker_opt.zero_grad()\n", |
199 | 208 | " listener_opt.zero_grad()\n", |
200 | 209 | " total_loss.backward()\n", |
201 | 210 | " speaker_opt.step()\n", |
202 | 211 | " listener_opt.step()\n", |
203 | | - " \n", |
| 212 | + "\n", |
204 | 213 | " return total_loss.item()\n", |
205 | 214 | "\n", |
| 215 | + "\n", |
206 | 216 | "# Train for a few steps\n", |
207 | 217 | "print(\"Training agents...\")\n", |
208 | 218 | "losses = []\n", |
|
211 | 221 | "for step in tqdm(range(100), desc=\"Training\"):\n", |
212 | 222 | " # Sample a new scene each step\n", |
213 | 223 | " scene_objects, target_idx = sample_scene(k=3, seed=step)\n", |
214 | | - " \n", |
| 224 | + "\n", |
215 | 225 | " # Training step\n", |
216 | | - " loss = train_step(speaker, listener, scene_objects, target_idx, \n", |
217 | | - " speaker_optimizer, listener_optimizer)\n", |
| 226 | + " loss = train_step(\n", |
| 227 | + " speaker,\n", |
| 228 | + " listener,\n", |
| 229 | + " scene_objects,\n", |
| 230 | + " target_idx,\n", |
| 231 | + " speaker_optimizer,\n", |
| 232 | + " listener_optimizer,\n", |
| 233 | + " )\n", |
218 | 234 | " losses.append(loss)\n", |
219 | | - " \n", |
| 235 | + "\n", |
220 | 236 | " # Evaluate accuracy every 10 steps\n", |
221 | 237 | " if step % 10 == 0:\n", |
222 | 238 | " speaker.eval()\n", |
223 | 239 | " listener.eval()\n", |
224 | 240 | " with torch.no_grad():\n", |
225 | | - " scene_tensor = torch.stack([encode_object(obj) for obj in scene_objects]).to(device)\n", |
226 | | - " target_object = scene_tensor[target_idx:target_idx+1]\n", |
227 | | - " message_logits, message_tokens, gesture_logits, gesture_tokens = speaker(target_object)\n", |
| 241 | + " scene_tensor = torch.stack(\n", |
| 242 | + " [encode_object(obj) for obj in scene_objects]\n", |
| 243 | + " ).to(device)\n", |
| 244 | + " target_object = scene_tensor[target_idx : target_idx + 1]\n", |
| 245 | + " message_logits, message_tokens, gesture_logits, gesture_tokens = speaker(\n", |
| 246 | + " target_object\n", |
| 247 | + " )\n", |
228 | 248 | " # Add batch dimension to scene_tensor for Listener\n", |
229 | | - " scene_tensor_batched = scene_tensor.unsqueeze(0) # Shape: (1, 3, object_dim)\n", |
| 249 | + " scene_tensor_batched = scene_tensor.unsqueeze(\n", |
| 250 | + " 0\n", |
| 251 | + " ) # Shape: (1, 3, object_dim)\n", |
230 | 252 | " listener_logits = listener(message_tokens, scene_tensor_batched)\n", |
231 | 253 | " predicted = torch.argmax(listener_logits, dim=1)\n", |
232 | 254 | " accuracy = (predicted == target_idx).float().mean().item()\n", |
|
247 | 269 | "fig, ax = plt.subplots(1, 1, figsize=(10, 6))\n", |
248 | 270 | "\n", |
249 | 271 | "# Plot raw loss\n", |
250 | | - "ax.plot(losses, alpha=0.3, color='lightblue', label='Raw Loss')\n", |
| 272 | + "ax.plot(losses, alpha=0.3, color=\"lightblue\", label=\"Raw Loss\")\n", |
251 | 273 | "\n", |
252 | 274 | "# Calculate and plot smoothed loss\n", |
253 | 275 | "window_size = 10\n", |
254 | 276 | "if len(losses) >= window_size:\n", |
255 | 277 | " smoothed_losses = []\n", |
256 | 278 | " for i in range(len(losses)):\n", |
257 | 279 | " start_idx = max(0, i - window_size + 1)\n", |
258 | | - " smoothed_loss = sum(losses[start_idx:i+1]) / (i - start_idx + 1)\n", |
| 280 | + " smoothed_loss = sum(losses[start_idx : i + 1]) / (i - start_idx + 1)\n", |
259 | 281 | " smoothed_losses.append(smoothed_loss)\n", |
260 | | - " \n", |
261 | | - " ax.plot(smoothed_losses, color='blue', linewidth=2, label=f'Smoothed Loss (window={window_size})')\n", |
262 | 282 | "\n", |
263 | | - "ax.set_title('Training Loss')\n", |
264 | | - "ax.set_xlabel('Step')\n", |
265 | | - "ax.set_ylabel('Loss')\n", |
| 283 | + " ax.plot(\n", |
| 284 | + " smoothed_losses,\n", |
| 285 | + " color=\"blue\",\n", |
| 286 | + " linewidth=2,\n", |
| 287 | + " label=f\"Smoothed Loss (window={window_size})\",\n", |
| 288 | + " )\n", |
| 289 | + "\n", |
| 290 | + "ax.set_title(\"Training Loss\")\n", |
| 291 | + "ax.set_xlabel(\"Step\")\n", |
| 292 | + "ax.set_ylabel(\"Loss\")\n", |
266 | 293 | "ax.grid(True)\n", |
267 | 294 | "ax.legend()\n", |
268 | 295 | "\n", |
269 | 296 | "plt.tight_layout()\n", |
270 | 297 | "plt.show()\n", |
271 | 298 | "\n", |
272 | | - "print(f\"Random baseline accuracy: {1/3:.2%} (3 objects)\")\n", |
| 299 | + "print(f\"Random baseline accuracy: {1 / 3:.2%} (3 objects)\")\n", |
273 | 300 | "print(f\"Final accuracy: {accuracies[-1]:.2%}\")\n", |
274 | | - "print(f\"Improvement over random: {(accuracies[-1] - 1/3)*100:.1f} percentage points\")" |
| 301 | + "print(\n", |
| 302 | + " f\"Improvement over random: {(accuracies[-1] - 1 / 3) * 100:.1f} percentage points\"\n", |
| 303 | + ")" |
275 | 304 | ] |
276 | 305 | }, |
277 | 306 | { |
|
0 commit comments