Skip to content

Commit e22122c

Browse files
committed
Apply code formatting and style improvements
- Applied black formatting to test files for consistent style - Improved import formatting in test files (multi-line imports) - Updated .gitignore and README.md - Applied formatting to core modules and experiments - All changes maintain functionality while improving code readability
1 parent a516ffd commit e22122c

12 files changed

Lines changed: 85 additions & 56 deletions

File tree

.gitignore

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -76,4 +76,5 @@ Desktop.ini
7676

7777
# Project-specific
7878
outputs/
79-
mlruns/emergent/
79+
mlruns/
80+
emergent/

README.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -59,6 +59,8 @@ The framework achieves state-of-the-art performance with advanced training techn
5959

6060
## Quick Start
6161

62+
[![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/bangyen/zsharp/blob/main/zsharp_demo.ipynb)
63+
6264
### Installation
6365

6466
```bash

emergent_demo.ipynb

Lines changed: 66 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -94,9 +94,9 @@
9494
"source": [
9595
"# Configuration for our referential game\n",
9696
"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",
100100
")\n",
101101
"\n",
102102
"# Learning parameters\n",
@@ -151,58 +151,68 @@
151151
"# Simple training loop for referential game\n",
152152
"from src.langlab.data.world import encode_object\n",
153153
"\n",
154+
"\n",
154155
"def train_step(speaker, listener, scene_objects, target_idx, speaker_opt, listener_opt):\n",
155156
" \"\"\"Single training step for the referential game.\"\"\"\n",
156157
" speaker.train()\n",
157158
" listener.train()\n",
158-
" \n",
159+
"\n",
159160
" # Encode the scene objects\n",
160161
" 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",
163166
" # 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",
167172
" # Add batch dimension to scene_tensor for Listener (batch_size=1, num_candidates=3, object_dim)\n",
168173
" scene_tensor_batched = scene_tensor.unsqueeze(0) # Shape: (1, 3, object_dim)\n",
169-
" \n",
174+
"\n",
170175
" # Listener tries to identify target from message and scene\n",
171176
" listener_logits = listener(message_tokens, scene_tensor_batched)\n",
172-
" \n",
177+
"\n",
173178
" # Compute listener loss (cross-entropy)\n",
174179
" listener_loss = torch.nn.functional.cross_entropy(listener_logits, target_tensor)\n",
175-
" \n",
180+
"\n",
176181
" # Compute speaker loss (REINFORCE-style)\n",
177182
" # Get the predicted target from listener\n",
178183
" predicted_target = torch.argmax(listener_logits, dim=1)\n",
179184
" # Reward is 1 if correct, 0 if incorrect\n",
180185
" reward = (predicted_target == target_tensor).float()\n",
181-
" \n",
186+
"\n",
182187
" # Speaker loss: encourage generating messages that lead to correct predictions\n",
183188
" # Use the message logits to compute log probabilities\n",
184189
" 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",
187194
" # 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",
189198
" speaker_log_prob = speaker_log_prob.sum(dim=1) # Sum over message length\n",
190-
" \n",
199+
"\n",
191200
" # REINFORCE loss: -log_prob * reward\n",
192201
" speaker_loss = -(speaker_log_prob * reward).mean()\n",
193-
" \n",
202+
"\n",
194203
" # Total loss\n",
195204
" total_loss = listener_loss + 0.1 * speaker_loss # Weight speaker loss less\n",
196-
" \n",
205+
"\n",
197206
" # Backpropagation\n",
198207
" speaker_opt.zero_grad()\n",
199208
" listener_opt.zero_grad()\n",
200209
" total_loss.backward()\n",
201210
" speaker_opt.step()\n",
202211
" listener_opt.step()\n",
203-
" \n",
212+
"\n",
204213
" return total_loss.item()\n",
205214
"\n",
215+
"\n",
206216
"# Train for a few steps\n",
207217
"print(\"Training agents...\")\n",
208218
"losses = []\n",
@@ -211,22 +221,34 @@
211221
"for step in tqdm(range(100), desc=\"Training\"):\n",
212222
" # Sample a new scene each step\n",
213223
" scene_objects, target_idx = sample_scene(k=3, seed=step)\n",
214-
" \n",
224+
"\n",
215225
" # 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",
218234
" losses.append(loss)\n",
219-
" \n",
235+
"\n",
220236
" # Evaluate accuracy every 10 steps\n",
221237
" if step % 10 == 0:\n",
222238
" speaker.eval()\n",
223239
" listener.eval()\n",
224240
" 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",
228248
" # 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",
230252
" listener_logits = listener(message_tokens, scene_tensor_batched)\n",
231253
" predicted = torch.argmax(listener_logits, dim=1)\n",
232254
" accuracy = (predicted == target_idx).float().mean().item()\n",
@@ -247,31 +269,38 @@
247269
"fig, ax = plt.subplots(1, 1, figsize=(10, 6))\n",
248270
"\n",
249271
"# 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",
251273
"\n",
252274
"# Calculate and plot smoothed loss\n",
253275
"window_size = 10\n",
254276
"if len(losses) >= window_size:\n",
255277
" smoothed_losses = []\n",
256278
" for i in range(len(losses)):\n",
257279
" 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",
259281
" 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",
262282
"\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",
266293
"ax.grid(True)\n",
267294
"ax.legend()\n",
268295
"\n",
269296
"plt.tight_layout()\n",
270297
"plt.show()\n",
271298
"\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",
273300
"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+
")"
275304
]
276305
},
277306
{

src/langlab/core/agents.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -330,8 +330,9 @@ def forward(
330330
Returns:
331331
Tensor of shape (batch_size, num_candidates) with probabilities over candidates.
332332
"""
333-
batch_size, num_candidates = candidate_objects.size(0), candidate_objects.size(
334-
1
333+
batch_size, num_candidates = (
334+
candidate_objects.size(0),
335+
candidate_objects.size(1),
335336
)
336337

337338
# One-hot encode message tokens
@@ -710,8 +711,9 @@ def forward(
710711
Returns:
711712
Tensor of shape (batch_size, num_candidates) with probabilities over candidates.
712713
"""
713-
batch_size, num_candidates = candidate_objects.size(0), candidate_objects.size(
714-
1
714+
batch_size, num_candidates = (
715+
candidate_objects.size(0),
716+
candidate_objects.size(1),
715717
)
716718

717719
# Embed message tokens

src/langlab/core/channel.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -130,8 +130,7 @@ def send_multimodal(
130130
)
131131
if gesture_size != self.config.gesture_size:
132132
raise ValueError(
133-
f"Expected gesture_size={self.config.gesture_size}, "
134-
f"got {gesture_size}"
133+
f"Expected gesture_size={self.config.gesture_size}, got {gesture_size}"
135134
)
136135

137136
# Sample tokens

src/langlab/experiments/contact.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -283,7 +283,9 @@ def measure_intelligibility(self) -> None:
283283

284284
# Create evaluation dataset
285285
eval_dataset = ReferentialGameDataset(
286-
n_scenes=1000, k=self.config.k, seed=42 # Fixed size for evaluation
286+
n_scenes=1000,
287+
k=self.config.k,
288+
seed=42, # Fixed size for evaluation
287289
)
288290
eval_dataloader = DataLoader(
289291
eval_dataset, batch_size=self.config.batch_size, shuffle=False

src/langlab/experiments/ensemble_training.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -137,7 +137,7 @@ def train_ensemble(
137137

138138
for i in range(n_models):
139139
seed = base_seed + i * 1000
140-
logger.info(f"Training model {i+1}/{n_models} with seed {seed}")
140+
logger.info(f"Training model {i + 1}/{n_models} with seed {seed}")
141141

142142
# Train individual model
143143
train(
@@ -202,9 +202,9 @@ def train_ensemble(
202202
)
203203
speaker.load_state_dict(checkpoint["speaker_state_dict"])
204204
listener.load_state_dict(checkpoint["listener_state_dict"])
205-
logger.info(f"Loaded checkpoint for model {i+1}")
205+
logger.info(f"Loaded checkpoint for model {i + 1}")
206206
else:
207-
logger.warning(f"No checkpoint found for model {i+1}")
207+
logger.warning(f"No checkpoint found for model {i + 1}")
208208

209209
speakers.append(speaker)
210210
listeners.append(listener)

src/langlab/training/train_grounded.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -100,7 +100,6 @@ def update(self, success: bool) -> bool:
100100
and self.level_successes / self.episodes_at_level >= self.success_threshold
101101
and self.current_level < len(self.curriculum_grids) - 1
102102
):
103-
104103
logger.info(f"Advancing to curriculum level {self.current_level + 1}")
105104
self.current_level += 1
106105
self.episodes_at_level = 0

tests/integration/test_ablate.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -174,7 +174,6 @@ def test_run_ablation_suite_experiment_id_format() -> None:
174174
) as mock_zipf, patch(
175175
"torch.load"
176176
) as mock_load:
177-
178177
mock_eval.return_value = {
179178
"train": {"acc": 0.8},
180179
"iid": {"acc": 0.75},

tests/integration/test_cli.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -214,7 +214,8 @@ def test_heldout_parsing(self) -> None:
214214
"""Test heldout pair parsing in train command."""
215215
runner = CliRunner()
216216
result = runner.invoke(
217-
main, ["train", "--steps", "10", "--heldout", "red,circle"] # Valid format
217+
main,
218+
["train", "--steps", "10", "--heldout", "red,circle"], # Valid format
218219
)
219220
# Should not crash on parsing
220221
assert result.exit_code == 0 or result.exit_code == 1 # May fail on training

0 commit comments

Comments
 (0)