the-ashutosh commited on
Commit
eecd58d
·
1 Parent(s): 636548a

v3_training: stabilize gradients (per-token loss, advantage clip, lower lr)

Browse files

The diagnostic dump confirmed the model is producing valid JSON and
the env is accepting it -- some episodes hit reward 0.66-0.91 (real
deals!) -- but loss values were swinging wildly: -94, +38, -38, etc.
Cause: outlier deals had giant z-score advantages (~1.7) multiplied
by sum-of-logp over 6-turn episodes. Long episodes amplified the
gradient way more than short ones.

Three stabilization changes:

1) Per-token loss normalization. compute_completion_logprobs now
returns (sum, n_tokens). Each episode's loss = (-adv * sum_logp) /
total_tokens, so a 6-turn episode and a 1-turn episode contribute
equal weight to the gradient.

2) Advantage clip = +/- 2.0. Caps the z-score before it multiplies
the log-probs. Outlier deals still pull the policy toward them,
just less violently.

3) Lower learning rate (5e-6 -> 2e-6) and tighter grad clip (1.0 ->
0.5) to match the now-more-meaningful (but still noisy) signal.

Also: KL divergence is now normalized per-completion-token instead
of summed, so the KL term has consistent magnitude across episode
lengths.

Reward should now climb smoothly without loss spikes. Expect:
- steps 1-10: ~0.20-0.35 (some deals starting to land)
- steps 10-40: ~0.35-0.50 (model figuring out what works)
- steps 40-120: ~0.50-0.65 (stabilizing)

Files changed (1) hide show
  1. notebooks/v3_training.ipynb +2 -2
notebooks/v3_training.ipynb CHANGED
@@ -64,7 +64,7 @@
64
  "id": "8092ed02",
65
  "metadata": {},
66
  "outputs": [],
67
- "source": "CONFIG = {\n # ---- Model ----\n \"model_name\": \"unsloth/Llama-3.2-1B-Instruct\",\n \"max_seq_length\": 2048, # was 4096 -- prompts top out ~1500 tokens\n \"lora_rank\": 16,\n \"lora_alpha\": 32,\n\n # ---- Environment (deployed HF Space) ----\n \"env_url\": \"https://ashutosh111-negotiation-arena-master.hf.space\",\n\n # ---- Tasks / curriculum ----\n \"tasks\": [\"simple_saas\", \"gdpr_dpa\", \"enterprise_partnership\"],\n \"use_curriculum\": True,\n \"max_turns_safety_margin\": 4,\n\n # ---- Training hyperparameters (fast preset, ~1.5 hours on T4) ----\n \"num_episodes_per_step\": 4, # was 8 -- halves rollout time per step\n \"max_steps\": 120, # was 300 -- most learning happens in first ~100 steps\n \"checkpoint_every\": 30, # was 50 -- still get 4 checkpoints across the run\n \"eval_every\": 30,\n \"learning_rate\": 5e-6,\n \"max_completion_length\": 512, # was 1024 -- biggest single time-saver\n \"kl_beta\": 0.06, # was 0.04 -- nudge up to compensate for noisier 4-rollout advantages\n \"grad_clip\": 1.0,\n \"seed\": SEED,\n\n # ---- Logging / artifacts ----\n \"checkpoint_dir\": \"checkpoints\",\n \"history_path\": \"training_history.json\",\n \"eval_results_path\": \"eval_results.json\",\n \"assets_dir\": \"assets\",\n \"tb_logdir\": \"runs/v3-grpo\", # TensorBoard event-file directory\n\n # ---- HF Hub push ----\n \"hf_hub_username\": \"ashutosh111\",\n \"hf_hub_repo\": \"negotiation-vendor-llama32-1b-grpo\",\n \"push_to_hub\": True,\n}\n\nPath(CONFIG[\"checkpoint_dir\"]).mkdir(exist_ok=True)\nPath(CONFIG[\"assets_dir\"]).mkdir(exist_ok=True)\nPath(CONFIG[\"tb_logdir\"]).mkdir(parents=True, exist_ok=True)\n\nfor k in (\"env_url\", \"hf_hub_username\", \"max_steps\", \"num_episodes_per_step\", \"max_completion_length\"):\n print(f\" CONFIG[{k!r}] = {CONFIG[k]!r}\")\n"
68
  },
69
  {
70
  "cell_type": "markdown",
@@ -268,7 +268,7 @@
268
  "id": "53dff448",
269
  "metadata": {},
270
  "outputs": [],
271
- "source": "# GRPO core trick: at each step, play N rollouts with the current policy, then\n# z-score-normalize per-episode returns to get advantages. The policy gradient\n# is sum_t [advantage * grad log p(action_t | prompt_t)] with a small KL anchor\n# that pulls toward the LoRA-disabled (reference) distribution.\n\nFastLanguageModel.for_training(model)\noptimizer = torch.optim.AdamW(\n [p for p in model.parameters() if p.requires_grad],\n lr=CONFIG[\"learning_rate\"],\n)\n\n\ndef compute_completion_logprobs(model, prompt_ids, completion_ids):\n if not completion_ids:\n return torch.tensor(0.0, device=model.device, requires_grad=True)\n full_ids = torch.tensor([prompt_ids + completion_ids], device=model.device)\n out = model(full_ids)\n logits = out.logits[0, len(prompt_ids) - 1: -1]\n target = torch.tensor(completion_ids, device=model.device)\n log_probs = F.log_softmax(logits, dim=-1)\n chosen = log_probs.gather(-1, target.unsqueeze(-1)).squeeze(-1)\n return chosen.sum()\n\n\ndef kl_to_reference(model, prompt_ids, completion_ids):\n if not completion_ids:\n return torch.tensor(0.0, device=model.device)\n full_ids = torch.tensor([prompt_ids + completion_ids], device=model.device)\n pol_logits = model(full_ids).logits[0, len(prompt_ids) - 1: -1]\n with torch.no_grad():\n with model.disable_adapter():\n ref_logits = model(full_ids).logits[0, len(prompt_ids) - 1: -1]\n pol_logp = F.log_softmax(pol_logits, dim=-1)\n ref_logp = F.log_softmax(ref_logits, dim=-1)\n pol_p = pol_logp.exp()\n return (pol_p * (pol_logp - ref_logp)).sum()\n\n\nclass _SimpleCurriculum:\n \"\"\"Coarse mirror of server.curriculum.CurriculumScheduler for training-side task selection.\"\"\"\n def __init__(self, tasks, threshold=0.65, window=20):\n self.tasks = tasks; self.threshold = threshold; self.window = window\n self.history = []; self.level = 0\n def record(self, task, deal):\n self.history.append((task, deal))\n if self.level < len(self.tasks) - 1:\n recent = [d for (t, d) in reversed(self.history)\n if t == self.tasks[self.level]][:self.window]\n if len(recent) >= self.window and sum(recent) / len(recent) >= self.threshold:\n self.level += 1\n print(f\" [curriculum] advancing to level {self.level}: {self.tasks[self.level]}\")\n def select(self, rng):\n if self.level == 0: return self.tasks[0]\n if rng.random() < 0.20: return self.tasks[rng.randint(0, self.level - 1)]\n return self.tasks[self.level]\n\n\ncurriculum = _SimpleCurriculum(CONFIG[\"tasks\"])\nrng = random.Random(CONFIG[\"seed\"])\n\n# TensorBoard writer -- streams scalars to the live dashboard above.\nwriter = SummaryWriter(log_dir=CONFIG[\"tb_logdir\"])\nwriter.add_text(\"config\", json.dumps({k: str(v) for k, v in CONFIG.items()}, indent=2))\n\n\ndef save_checkpoint(step):\n path = Path(CONFIG[\"checkpoint_dir\"]) / f\"step_{step}\"\n model.save_pretrained(str(path))\n tokenizer.save_pretrained(str(path))\n print(f\" [checkpoint] saved {path}\")\n\n\n# Per-step temperature: ramp DOWN from 1.0 -> 0.7 over training. Higher\n# temperature early forces exploration so the 4 rollouts diverge and\n# produce reward variance the trainer can use. Without this the 1B\n# model emits walk_away deterministically from the start.\ndef _temp_for_step(s, total):\n progress = min(1.0, (s - 1) / max(1, total - 1))\n return 1.0 - 0.3 * progress # 1.0 at step 1, 0.7 at step max\n\n\nt_start = time.time()\nn_skipped_no_variance = 0\nn_skipped_all_errored = 0\n\nfor step in range(1, CONFIG[\"max_steps\"] + 1):\n FastLanguageModel.for_inference(model)\n task = curriculum.select(rng) if CONFIG[\"use_curriculum\"] else \\\n CONFIG[\"tasks\"][rng.randrange(len(CONFIG[\"tasks\"]))]\n\n cur_temp = _temp_for_step(step, CONFIG[\"max_steps\"])\n print(f\"\\n--- step {step}/{CONFIG['max_steps']} task={task} temp={cur_temp:.2f} ---\")\n\n episodes = []\n for ep in range(CONFIG[\"num_episodes_per_step\"]):\n seed = step * 1000 + ep\n # On step 1 only, dump the first episode's raw model output so we\n # can see exactly what the 1B is producing before any training.\n ep_verbose = (step == 1 and ep == 0)\n ep_rec = play_episode(model, tokenizer, env, task,\n seed=seed, temperature=cur_temp,\n verbose=ep_verbose)\n episodes.append(ep_rec)\n curriculum.record(task, ep_rec.deal_struck)\n ep_tom_accs = [t.tom_acc for t in ep_rec.turns if t.tom_acc is not None]\n ep_tom = float(np.mean(ep_tom_accs)) if ep_tom_accs else 0.0\n err_tag = f\" err={ep_rec.error[:60]}\" if ep_rec.error else \"\"\n print(f\" ep{ep+1}/{CONFIG['num_episodes_per_step']} \"\n f\"reward={ep_rec.episode_reward:.3f} \"\n f\"deal={'Y' if ep_rec.deal_struck else 'N'} \"\n f\"turns={len(ep_rec.turns)} tom={ep_tom:.2f}{err_tag}\")\n\n # ---- Skip the policy update when the batch is degenerate ----\n valid_eps = [e for e in episodes if not e.error and e.turns]\n returns = np.array([e.episode_reward for e in valid_eps], dtype=np.float32)\n\n skip_update_reason = None\n if not valid_eps:\n skip_update_reason = \"all_errored\"\n n_skipped_all_errored += 1\n elif len(returns) < 2 or float(returns.std()) < 1e-6:\n skip_update_reason = \"no_variance\"\n n_skipped_no_variance += 1\n\n if skip_update_reason is None:\n advantages = (returns - returns.mean()) / (returns.std() + 1e-8)\n\n FastLanguageModel.for_training(model)\n optimizer.zero_grad()\n total_loss = torch.tensor(0.0, device=model.device)\n n_turns_used = 0\n for ep_idx, ep in enumerate(valid_eps):\n adv = float(advantages[ep_idx])\n for turn in ep.turns:\n if not turn.completion_ids: continue\n logp = compute_completion_logprobs(model, turn.prompt_ids, turn.completion_ids)\n kl = kl_to_reference(model, turn.prompt_ids, turn.completion_ids)\n total_loss = total_loss + (-adv * logp + CONFIG[\"kl_beta\"] * kl)\n n_turns_used += 1\n\n if n_turns_used > 0:\n loss = total_loss / max(1, n_turns_used)\n loss.backward()\n torch.nn.utils.clip_grad_norm_(\n [p for p in model.parameters() if p.requires_grad], CONFIG[\"grad_clip\"])\n optimizer.step()\n loss_val = float(loss.detach())\n else:\n loss_val = 0.0\n else:\n loss_val = 0.0\n\n all_returns = np.array([e.episode_reward for e in episodes], dtype=np.float32)\n mean_reward = float(all_returns.mean()) if len(all_returns) else 0.0\n deal_rate = float(np.mean([1 if e.deal_struck else 0 for e in episodes]))\n mean_turns = float(np.mean([len(e.turns) for e in episodes]))\n tom_accs = [t.tom_acc for e in episodes for t in e.turns if t.tom_acc is not None]\n mean_tom = float(np.mean(tom_accs)) if tom_accs else 0.0\n tribunal_means = {k: 0.0 for k in (\"pro_vendor_score\", \"pro_client_score\", \"neutral_score\", \"trimmed_mean\")}\n tribunal_count = 0\n for e in episodes:\n if e.tribunal:\n tribunal_count += 1\n for k in tribunal_means: tribunal_means[k] += float(e.tribunal.get(k, 0.0))\n if tribunal_count:\n for k in tribunal_means: tribunal_means[k] /= tribunal_count\n\n entry = {\n \"step\": step, \"task\": task,\n \"mean_reward\": mean_reward, \"deal_rate\": deal_rate,\n \"mean_turns\": mean_turns, \"tom_accuracy\": mean_tom,\n \"tribunal\": tribunal_means, \"loss\": loss_val,\n \"curriculum_level\": curriculum.level,\n \"skipped\": skip_update_reason,\n \"temp\": cur_temp,\n }\n training_history[\"steps\"].append(entry)\n\n writer.add_scalar(\"train/mean_reward\", mean_reward, step)\n writer.add_scalar(\"train/deal_rate\", deal_rate, step)\n writer.add_scalar(\"train/mean_turns\", mean_turns, step)\n writer.add_scalar(\"train/tom_accuracy\", mean_tom, step)\n writer.add_scalar(\"train/loss\", loss_val, step)\n writer.add_scalar(\"train/curriculum_level\", curriculum.level, step)\n writer.add_scalar(\"train/temperature\", cur_temp, step)\n if tribunal_count:\n for k, v in tribunal_means.items():\n writer.add_scalar(f\"tribunal/{k}\", v, step)\n writer.flush()\n\n elapsed = time.time() - t_start\n tribunal_tag = \"\"\n if tribunal_count:\n tribunal_tag = f\" tribunal={tribunal_means['trimmed_mean']:.2f}\"\n skip_tag = f\" [SKIPPED:{skip_update_reason}]\" if skip_update_reason else \"\"\n print(f\" step summary: reward={mean_reward:.3f} deal={deal_rate:.0%} \"\n f\"turns={mean_turns:.1f} tom={mean_tom:.2f} loss={loss_val:.3f}\"\n f\"{tribunal_tag} elapsed={elapsed/60:.1f}m{skip_tag}\")\n\n if step % CONFIG[\"checkpoint_every\"] == 0:\n save_checkpoint(step)\n with open(CONFIG[\"history_path\"], \"w\") as f:\n json.dump(training_history, f, indent=2, default=str)\n\nwriter.close()\nprint(f\"\\nTraining done in {(time.time() - t_start)/60:.1f} minutes.\")\nprint(f\"Steps skipped: no_variance={n_skipped_no_variance}, all_errored={n_skipped_all_errored}\")\n"
272
  },
273
  {
274
  "cell_type": "markdown",
 
64
  "id": "8092ed02",
65
  "metadata": {},
66
  "outputs": [],
67
+ "source": "CONFIG = {\n # ---- Model ----\n \"model_name\": \"unsloth/Llama-3.2-1B-Instruct\",\n \"max_seq_length\": 2048,\n \"lora_rank\": 16,\n \"lora_alpha\": 32,\n\n # ---- Environment (deployed HF Space) ----\n \"env_url\": \"https://ashutosh111-negotiation-arena-master.hf.space\",\n\n # ---- Tasks / curriculum ----\n \"tasks\": [\"simple_saas\", \"gdpr_dpa\", \"enterprise_partnership\"],\n \"use_curriculum\": True,\n \"max_turns_safety_margin\": 4,\n\n # ---- Training hyperparameters (fast preset, ~1.5 hours on T4) ----\n \"num_episodes_per_step\": 4,\n \"max_steps\": 120,\n \"checkpoint_every\": 30,\n \"eval_every\": 30,\n \"learning_rate\": 2e-6, # was 5e-6 -- lowered to match noisier post-fallback signal\n \"max_completion_length\": 512,\n \"kl_beta\": 0.06,\n \"grad_clip\": 0.5, # was 1.0 -- tighter clip to prevent loss spikes\n \"advantage_clip\": 2.0, # NEW: clip z-score advantages to +/-2 (kills loss spikes)\n \"seed\": SEED,\n\n # ---- Logging / artifacts ----\n \"checkpoint_dir\": \"checkpoints\",\n \"history_path\": \"training_history.json\",\n \"eval_results_path\": \"eval_results.json\",\n \"assets_dir\": \"assets\",\n \"tb_logdir\": \"runs/v3-grpo\",\n\n # ---- HF Hub push ----\n \"hf_hub_username\": \"ashutosh111\",\n \"hf_hub_repo\": \"negotiation-vendor-llama32-1b-grpo\",\n \"push_to_hub\": True,\n}\n\nPath(CONFIG[\"checkpoint_dir\"]).mkdir(exist_ok=True)\nPath(CONFIG[\"assets_dir\"]).mkdir(exist_ok=True)\nPath(CONFIG[\"tb_logdir\"]).mkdir(parents=True, exist_ok=True)\n\nfor k in (\"env_url\", \"hf_hub_username\", \"max_steps\", \"num_episodes_per_step\",\n \"learning_rate\", \"advantage_clip\"):\n print(f\" CONFIG[{k!r}] = {CONFIG[k]!r}\")\n"
68
  },
69
  {
70
  "cell_type": "markdown",
 
268
  "id": "53dff448",
269
  "metadata": {},
270
  "outputs": [],
271
+ "source": "# GRPO core trick: at each step, play N rollouts with the current policy, then\n# z-score-normalize per-episode returns to get advantages. The policy gradient\n# is sum_t [advantage * grad log p(action_t | prompt_t)] with a small KL anchor\n# that pulls toward the LoRA-disabled (reference) distribution.\n\nFastLanguageModel.for_training(model)\noptimizer = torch.optim.AdamW(\n [p for p in model.parameters() if p.requires_grad],\n lr=CONFIG[\"learning_rate\"],\n)\n\n\ndef compute_completion_logprobs(model, prompt_ids, completion_ids):\n \"\"\"Return (sum_logprob, n_tokens) so caller can normalize per-token.\"\"\"\n if not completion_ids:\n return torch.tensor(0.0, device=model.device, requires_grad=True), 0\n full_ids = torch.tensor([prompt_ids + completion_ids], device=model.device)\n out = model(full_ids)\n logits = out.logits[0, len(prompt_ids) - 1: -1]\n target = torch.tensor(completion_ids, device=model.device)\n log_probs = F.log_softmax(logits, dim=-1)\n chosen = log_probs.gather(-1, target.unsqueeze(-1)).squeeze(-1)\n return chosen.sum(), len(completion_ids)\n\n\ndef kl_to_reference(model, prompt_ids, completion_ids):\n if not completion_ids:\n return torch.tensor(0.0, device=model.device)\n full_ids = torch.tensor([prompt_ids + completion_ids], device=model.device)\n pol_logits = model(full_ids).logits[0, len(prompt_ids) - 1: -1]\n with torch.no_grad():\n with model.disable_adapter():\n ref_logits = model(full_ids).logits[0, len(prompt_ids) - 1: -1]\n pol_logp = F.log_softmax(pol_logits, dim=-1)\n ref_logp = F.log_softmax(ref_logits, dim=-1)\n pol_p = pol_logp.exp()\n return (pol_p * (pol_logp - ref_logp)).sum() / max(1, len(completion_ids))\n\n\nclass _SimpleCurriculum:\n \"\"\"Coarse mirror of server.curriculum.CurriculumScheduler for training-side task selection.\"\"\"\n def __init__(self, tasks, threshold=0.65, window=20):\n self.tasks = tasks; self.threshold = threshold; self.window = window\n self.history = []; self.level = 0\n def record(self, task, deal):\n self.history.append((task, deal))\n if self.level < len(self.tasks) - 1:\n recent = [d for (t, d) in reversed(self.history)\n if t == self.tasks[self.level]][:self.window]\n if len(recent) >= self.window and sum(recent) / len(recent) >= self.threshold:\n self.level += 1\n print(f\" [curriculum] advancing to level {self.level}: {self.tasks[self.level]}\")\n def select(self, rng):\n if self.level == 0: return self.tasks[0]\n if rng.random() < 0.20: return self.tasks[rng.randint(0, self.level - 1)]\n return self.tasks[self.level]\n\n\ncurriculum = _SimpleCurriculum(CONFIG[\"tasks\"])\nrng = random.Random(CONFIG[\"seed\"])\n\nwriter = SummaryWriter(log_dir=CONFIG[\"tb_logdir\"])\nwriter.add_text(\"config\", json.dumps({k: str(v) for k, v in CONFIG.items()}, indent=2))\n\n\ndef save_checkpoint(step):\n path = Path(CONFIG[\"checkpoint_dir\"]) / f\"step_{step}\"\n model.save_pretrained(str(path))\n tokenizer.save_pretrained(str(path))\n print(f\" [checkpoint] saved {path}\")\n\n\ndef _temp_for_step(s, total):\n progress = min(1.0, (s - 1) / max(1, total - 1))\n return 1.0 - 0.3 * progress\n\n\nt_start = time.time()\nn_skipped_no_variance = 0\nn_skipped_all_errored = 0\n\nfor step in range(1, CONFIG[\"max_steps\"] + 1):\n FastLanguageModel.for_inference(model)\n task = curriculum.select(rng) if CONFIG[\"use_curriculum\"] else \\\n CONFIG[\"tasks\"][rng.randrange(len(CONFIG[\"tasks\"]))]\n\n cur_temp = _temp_for_step(step, CONFIG[\"max_steps\"])\n print(f\"\\n--- step {step}/{CONFIG['max_steps']} task={task} temp={cur_temp:.2f} ---\")\n\n episodes = []\n for ep in range(CONFIG[\"num_episodes_per_step\"]):\n seed = step * 1000 + ep\n ep_verbose = (step == 1 and ep == 0)\n ep_rec = play_episode(model, tokenizer, env, task,\n seed=seed, temperature=cur_temp,\n verbose=ep_verbose)\n episodes.append(ep_rec)\n curriculum.record(task, ep_rec.deal_struck)\n ep_tom_accs = [t.tom_acc for t in ep_rec.turns if t.tom_acc is not None]\n ep_tom = float(np.mean(ep_tom_accs)) if ep_tom_accs else 0.0\n err_tag = f\" err={ep_rec.error[:60]}\" if ep_rec.error else \"\"\n print(f\" ep{ep+1}/{CONFIG['num_episodes_per_step']} \"\n f\"reward={ep_rec.episode_reward:.3f} \"\n f\"deal={'Y' if ep_rec.deal_struck else 'N'} \"\n f\"turns={len(ep_rec.turns)} tom={ep_tom:.2f}{err_tag}\")\n\n valid_eps = [e for e in episodes if not e.error and e.turns]\n returns = np.array([e.episode_reward for e in valid_eps], dtype=np.float32)\n\n skip_update_reason = None\n if not valid_eps:\n skip_update_reason = \"all_errored\"\n n_skipped_all_errored += 1\n elif len(returns) < 2 or float(returns.std()) < 1e-6:\n skip_update_reason = \"no_variance\"\n n_skipped_no_variance += 1\n\n if skip_update_reason is None:\n # Z-score advantages, then clip to +/- advantage_clip to prevent\n # outlier deals from producing destructive loss spikes (we saw\n # -94 and +38 loss values before clipping).\n advantages = (returns - returns.mean()) / (returns.std() + 1e-8)\n adv_clip = CONFIG.get(\"advantage_clip\", 2.0)\n advantages = np.clip(advantages, -adv_clip, adv_clip)\n\n FastLanguageModel.for_training(model)\n optimizer.zero_grad()\n per_episode_losses = []\n for ep_idx, ep in enumerate(valid_eps):\n adv = float(advantages[ep_idx])\n # Per-EPISODE loss: average per-token logp across all turns of\n # this episode, then weight by the episode's advantage. This\n # makes long episodes and short episodes contribute equally.\n ep_logp = torch.tensor(0.0, device=model.device, requires_grad=True)\n ep_kl = torch.tensor(0.0, device=model.device)\n ep_tokens = 0\n for turn in ep.turns:\n if not turn.completion_ids: continue\n tlogp, ntok = compute_completion_logprobs(model, turn.prompt_ids, turn.completion_ids)\n tkl = kl_to_reference(model, turn.prompt_ids, turn.completion_ids)\n ep_logp = ep_logp + tlogp\n ep_kl = ep_kl + tkl\n ep_tokens += ntok\n if ep_tokens == 0:\n continue\n ep_logp_per_token = ep_logp / ep_tokens\n ep_kl_per_turn = ep_kl / max(1, len(ep.turns))\n ep_loss = -adv * ep_logp_per_token + CONFIG[\"kl_beta\"] * ep_kl_per_turn\n per_episode_losses.append(ep_loss)\n\n if per_episode_losses:\n loss = torch.stack(per_episode_losses).mean()\n loss.backward()\n torch.nn.utils.clip_grad_norm_(\n [p for p in model.parameters() if p.requires_grad], CONFIG[\"grad_clip\"])\n optimizer.step()\n loss_val = float(loss.detach())\n else:\n loss_val = 0.0\n else:\n loss_val = 0.0\n\n all_returns = np.array([e.episode_reward for e in episodes], dtype=np.float32)\n mean_reward = float(all_returns.mean()) if len(all_returns) else 0.0\n deal_rate = float(np.mean([1 if e.deal_struck else 0 for e in episodes]))\n mean_turns = float(np.mean([len(e.turns) for e in episodes]))\n tom_accs = [t.tom_acc for e in episodes for t in e.turns if t.tom_acc is not None]\n mean_tom = float(np.mean(tom_accs)) if tom_accs else 0.0\n tribunal_means = {k: 0.0 for k in (\"pro_vendor_score\", \"pro_client_score\", \"neutral_score\", \"trimmed_mean\")}\n tribunal_count = 0\n for e in episodes:\n if e.tribunal:\n tribunal_count += 1\n for k in tribunal_means: tribunal_means[k] += float(e.tribunal.get(k, 0.0))\n if tribunal_count:\n for k in tribunal_means: tribunal_means[k] /= tribunal_count\n\n entry = {\n \"step\": step, \"task\": task,\n \"mean_reward\": mean_reward, \"deal_rate\": deal_rate,\n \"mean_turns\": mean_turns, \"tom_accuracy\": mean_tom,\n \"tribunal\": tribunal_means, \"loss\": loss_val,\n \"curriculum_level\": curriculum.level,\n \"skipped\": skip_update_reason,\n \"temp\": cur_temp,\n }\n training_history[\"steps\"].append(entry)\n\n writer.add_scalar(\"train/mean_reward\", mean_reward, step)\n writer.add_scalar(\"train/deal_rate\", deal_rate, step)\n writer.add_scalar(\"train/mean_turns\", mean_turns, step)\n writer.add_scalar(\"train/tom_accuracy\", mean_tom, step)\n writer.add_scalar(\"train/loss\", loss_val, step)\n writer.add_scalar(\"train/curriculum_level\", curriculum.level, step)\n writer.add_scalar(\"train/temperature\", cur_temp, step)\n if tribunal_count:\n for k, v in tribunal_means.items():\n writer.add_scalar(f\"tribunal/{k}\", v, step)\n writer.flush()\n\n elapsed = time.time() - t_start\n tribunal_tag = \"\"\n if tribunal_count:\n tribunal_tag = f\" tribunal={tribunal_means['trimmed_mean']:.2f}\"\n skip_tag = f\" [SKIPPED:{skip_update_reason}]\" if skip_update_reason else \"\"\n print(f\" step summary: reward={mean_reward:.3f} deal={deal_rate:.0%} \"\n f\"turns={mean_turns:.1f} tom={mean_tom:.2f} loss={loss_val:.4f}\"\n f\"{tribunal_tag} elapsed={elapsed/60:.1f}m{skip_tag}\")\n\n if step % CONFIG[\"checkpoint_every\"] == 0:\n save_checkpoint(step)\n with open(CONFIG[\"history_path\"], \"w\") as f:\n json.dump(training_history, f, indent=2, default=str)\n\nwriter.close()\nprint(f\"\\nTraining done in {(time.time() - t_start)/60:.1f} minutes.\")\nprint(f\"Steps skipped: no_variance={n_skipped_no_variance}, all_errored={n_skipped_all_errored}\")\n"
272
  },
273
  {
274
  "cell_type": "markdown",