{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Reward Model Demo: Generation Evaluation\n",
    "\n",
    "This notebook demonstrates:\n",
    "1. Generating music with different CFG settings (low vs high)\n",
    "2. Scoring generations with trained reward model\n",
    "3. Visualizing token-level rewards over time\n",
    "4. Comparing rewards across different configurations\n",
    "\n",
    "**Goal**: Validate that reward model aligns with generation quality intuitions\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 1. Setup and Imports\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import sys\n",
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\"  # Use GPU 0\n",
    "\n",
    "import torch\n",
    "import matplotlib.pyplot as plt\n",
    "import numpy as np\n",
    "\n",
    "# Add paths\n",
    "sys.path.insert(0, \"/home/tony/Work/neon_2/sunoGPT\")\n",
    "\n",
    "from scripts.reward_eval_utils import load_reward_model, extract_scalar_rewards\n",
    "from suno_utils.gpt.generation import GenerationConfig\n",
    "from suno_utils.gpt.engine import Engine\n",
    "from suno_utils.gpt.generation_engine import make_request\n",
    "\n",
    "print(\"✓ Imports successful\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 2. Load Reward Model\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Load trained reward model\n",
    "reward_model_path = (\n",
    "    \"/app2/suno/checkpoints/2025-11-05_13-12-29/last_ckpt_infer.pt\"  # Update with your checkpoint path\n",
    ")\n",
    "\n",
    "reward_model, model_args = load_reward_model(reward_model_path, device=\"cuda\")\n",
    "print(f\"✓ Reward model loaded from {reward_model_path}\")\n",
    "print(f\"  use_reward_head: {reward_model.config.use_reward_head}\")\n",
    "print(f\"  Output modules: {list(reward_model.output_modules.keys())}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 3. Load Generation Engine\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Load GPT engine for music generation\n",
    "gpt_checkpoint = \"/app2/suno/checkpoints/2025-09-10_04-35-21/last_ckpt_infer.pt\"  # Bluejay SFT model\n",
    "tokenizer_path = \"/app/suno/models/chirp_v2/tokenizer_60k.json\"\n",
    "\n",
    "engine = Engine(\n",
    "    gpt_checkpoint,\n",
    "    tokenizer_path,\n",
    "    max_sequences=8,\n",
    "    compile=False,\n",
    ")\n",
    "\n",
    "cfg_model = engine.model.config\n",
    "print(f\"✓ Generation engine loaded from {gpt_checkpoint}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 4. Helper Function: Get Rewards for Generation\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_reward_for_generation(tokens, reward_model):\n",
    "    \"\"\"Get token-level and scalar rewards for generated semantic tokens.\n",
    "\n",
    "    Args:\n",
    "        tokens: Generated semantic tokens, shape (seq_len,)\n",
    "        reward_model: Trained reward model\n",
    "\n",
    "    Returns:\n",
    "        token_rewards: Reward for each token, shape (seq_len,)\n",
    "        scalar_reward: Overall reward score (float)\n",
    "    \"\"\"\n",
    "    batch_size = 1\n",
    "    n_streams = 1 + reward_model.config.semantic_n_codebooks + reward_model.config.coarse_n_codebooks\n",
    "    seq_len = len(tokens)\n",
    "\n",
    "    # Create input tensor (batch, n_streams, seq_len)\n",
    "    X = torch.zeros(batch_size, n_streams, seq_len, dtype=torch.long, device=\"cuda\")\n",
    "    X[0, 0, :] = (\n",
    "        reward_model.config.text_pad_token\n",
    "    )  # Text stream (padding - no text for pure generation)\n",
    "    X[0, 1, :] = tokens.cuda()  # Semantic tokens\n",
    "\n",
    "    # Coarse streams (if any) - use pad tokens\n",
    "    for i in range(reward_model.config.coarse_n_codebooks):\n",
    "        X[0, 2 + i, :] = reward_model.config.coarse_pad_token\n",
    "\n",
    "    # Get token-level rewards\n",
    "    with torch.no_grad():\n",
    "        with torch.amp.autocast(device_type=\"cuda\", dtype=torch.bfloat16):\n",
    "            output = reward_model(X, return_logits=True)\n",
    "\n",
    "    # Extract reward_logits from dict\n",
    "    reward_logits = output[\"reward_logits\"]  # (1, seq_len)\n",
    "\n",
    "    # Create Y for masking (all valid - no padding for generated tokens)\n",
    "    Y = torch.ones(\n",
    "        batch_size,\n",
    "        reward_model.config.semantic_n_codebooks,\n",
    "        seq_len - 1,\n",
    "        dtype=torch.long,\n",
    "        device=\"cuda\",\n",
    "    )\n",
    "\n",
    "    # Find where semantic generation ACTUALLY starts (after semantic_infer_token)\n",
    "    # The semantic stream (index 1) contains the infer token marking generation start\n",
    "    semantic_infer_token = reward_model.config.semantic_infer_token\n",
    "    semantic_stream = X[0, 1, :]\n",
    "\n",
    "    if (semantic_stream == semantic_infer_token).any():\n",
    "        loss_start_index = (semantic_stream == semantic_infer_token).nonzero(as_tuple=True)[0][0].item()\n",
    "    else:\n",
    "        # Fallback: use t_text from config (text portion length)\n",
    "        loss_start_index = reward_model.config.t_text\n",
    "\n",
    "    # Get scalar reward (no padding in generated tokens, so end_index=None)\n",
    "    scalar_reward = extract_scalar_rewards(reward_logits, [loss_start_index], None)\n",
    "\n",
    "    # Return float tensors (convert from bfloat16 for numpy compatibility)\n",
    "    return reward_logits[0].float().cpu(), scalar_reward.item()\n",
    "\n",
    "\n",
    "print(\"✓ Helper function defined\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 5. Define Test Prompt\n",
    "\n",
    "Using a simple upbeat pop song from the inference examples:\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "text = \"\"\"\n",
    "[Verse]\n",
    "Walking down the street, feeling so alive\n",
    "Got my head in the clouds, got a gleam in my eye\n",
    "Every step I take, it's like a brand new start\n",
    "No matter where I'm going, I'll always find my part\n",
    "\n",
    "[Chorus]\n",
    "Life is like a high-wire act, we're dancing in the sky\n",
    "No need to worry, no need to ask why\n",
    "With a little bit of courage, we can chase our dreams\n",
    "No matter what comes our way, we'll always be a team\n",
    "\"\"\"\n",
    "\n",
    "tags = \"upbeat pop\"\n",
    "print(f\"✓ Prompt defined: {len(text)} characters\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 6. Generate with Different CFG Settings\n",
    "\n",
    "Generate the same prompt with LOW CFG vs HIGH CFG to compare reward scores:\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Configuration for LOW CFG (weak guidance)\n",
    "cfg_low = GenerationConfig(\n",
    "    text=text,\n",
    "    text_tags=tags,\n",
    "    cfg_coef=1.0,  # Low CFG\n",
    "    cfg_coef_tags=0.1,\n",
    "    n_batch=1,\n",
    "    max_gen_duration_s=120,\n",
    "    cfg_coef_tags_max_steps=25 * 10,\n",
    "    random_seed=42,\n",
    ")\n",
    "\n",
    "# Configuration for HIGH CFG (strong guidance)\n",
    "cfg_high = GenerationConfig(\n",
    "    text=text,\n",
    "    text_tags=tags,\n",
    "    cfg_coef=2.0,  # High CFG\n",
    "    cfg_coef_tags=2.0,\n",
    "    n_batch=1,\n",
    "    max_gen_duration_s=120,\n",
    "    cfg_coef_tags_max_steps=25 * 120,\n",
    "    random_seed=42,  # Same seed for fair comparison\n",
    ")\n",
    "\n",
    "# Generate both\n",
    "print(\"Generating with LOW CFG...\")\n",
    "request_low = make_request(\"low\", cfg_low, engine.model.config, engine.tokenizer)\n",
    "job_low = engine.run_request([request_low])[0]\n",
    "tokens_low = torch.stack(list(engine.token_generator(job_low)))[:, 1]\n",
    "print(f\"  Generated: {tokens_low.shape[0]} tokens ({tokens_low.shape[0]/25:.1f}s)\")\n",
    "\n",
    "print(\"Generating with HIGH CFG...\")\n",
    "request_high = make_request(\"high\", cfg_high, engine.model.config, engine.tokenizer)\n",
    "job_high = engine.run_request([request_high])[0]\n",
    "tokens_high = torch.stack(list(engine.token_generator(job_high)))[:, 1]\n",
    "print(f\"  Generated: {tokens_high.shape[0]} tokens ({tokens_high.shape[0]/25:.1f}s)\")\n",
    "\n",
    "print(\"\\n✓ Both generations complete\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 7. Score Generations with Reward Model\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Compute rewards for both generations\n",
    "print(\"Computing rewards...\")\n",
    "\n",
    "token_rewards_low, scalar_reward_low = get_reward_for_generation(tokens_low, reward_model)\n",
    "print(f\"  LOW CFG:  Scalar reward = {scalar_reward_low:.4f}\")\n",
    "\n",
    "token_rewards_high, scalar_reward_high = get_reward_for_generation(tokens_high, reward_model)\n",
    "print(f\"  HIGH CFG: Scalar reward = {scalar_reward_high:.4f}\")\n",
    "\n",
    "reward_diff = scalar_reward_high - scalar_reward_low\n",
    "print(f\"\\n  Difference (HIGH - LOW): {reward_diff:+.4f}\")\n",
    "\n",
    "if reward_diff > 0:\n",
    "    print(\"  → High CFG gets higher reward ✓\")\n",
    "elif reward_diff < 0:\n",
    "    print(\"  → Low CFG gets higher reward (unexpected!)\")\n",
    "else:\n",
    "    print(\"  → Same reward\")\n",
    "\n",
    "# Store results\n",
    "results = {\n",
    "    \"low_cfg\": {\n",
    "        \"tokens\": tokens_low,\n",
    "        \"token_rewards\": token_rewards_low,\n",
    "        \"scalar_reward\": scalar_reward_low,\n",
    "        \"length_s\": len(tokens_low) / 25.0,\n",
    "    },\n",
    "    \"high_cfg\": {\n",
    "        \"tokens\": tokens_high,\n",
    "        \"token_rewards\": token_rewards_high,\n",
    "        \"scalar_reward\": scalar_reward_high,\n",
    "        \"length_s\": len(tokens_high) / 25.0,\n",
    "    },\n",
    "}"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 8. Visualize Token-Level Rewards\n",
    "\n",
    "Show how reward changes for each token across the sequence:\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, axes = plt.subplots(2, 1, figsize=(14, 8))\n",
    "\n",
    "for ax, (name, data) in zip(axes, results.items()):\n",
    "    # Already float32 from helper function\n",
    "    token_rewards = data[\"token_rewards\"].numpy()\n",
    "    time_axis = np.arange(len(token_rewards)) / 25.0  # Convert to seconds\n",
    "\n",
    "    ax.plot(time_axis, token_rewards, linewidth=2, alpha=0.8, label=f\"{name} (per-token)\")\n",
    "    ax.axhline(y=0, color=\"k\", linestyle=\"--\", alpha=0.3, label=\"Zero\")\n",
    "    ax.axhline(\n",
    "        y=data[\"scalar_reward\"],\n",
    "        color=\"r\",\n",
    "        linestyle=\"-\",\n",
    "        label=f'Mean: {data[\"scalar_reward\"]:.3f}',\n",
    "        linewidth=2,\n",
    "        alpha=0.6,\n",
    "    )\n",
    "    ax.set_ylabel(\"Reward\", fontsize=12)\n",
    "    ax.set_title(\n",
    "        f'{name.upper()}: Scalar Reward = {data[\"scalar_reward\"]:.4f}', fontsize=13, fontweight=\"bold\"\n",
    "    )\n",
    "    ax.legend(fontsize=10)\n",
    "    ax.grid(True, alpha=0.3)\n",
    "\n",
    "axes[-1].set_xlabel(\"Time (seconds)\", fontsize=12)\n",
    "fig.suptitle(\"Token-Level Rewards: LOW CFG vs HIGH CFG\", fontsize=15, fontweight=\"bold\")\n",
    "plt.tight_layout()\n",
    "# plt.savefig(\"reward_comparison_cfg.png\", dpi=150, bbox_inches=\"tight\")\n",
    "plt.show()\n",
    "\n",
    "print(f\"✓ Saved to reward_comparison_cfg.png\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 9. Progressive Reward: How Score Evolves\n",
    "\n",
    "Show cumulative average reward at steps (0-750, 0-1500, etc.):\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Compute progressive rewards (cumulative average)\n",
    "fig, ax = plt.subplots(1, 1, figsize=(14, 6))\n",
    "\n",
    "for name, data in results.items():\n",
    "    # Already float32 from helper function\n",
    "    token_rewards = data[\"token_rewards\"]\n",
    "\n",
    "    # Compute progressive rewards every 30 tokens (~1.2s)\n",
    "    progressive_rewards = []\n",
    "    progressive_times = []\n",
    "\n",
    "    for end_pos in range(30, len(token_rewards), 30):\n",
    "        progressive_reward = token_rewards[:end_pos].mean().item()\n",
    "        progressive_rewards.append(progressive_reward)\n",
    "        progressive_times.append(end_pos / 25.0)\n",
    "\n",
    "    # Add final point\n",
    "    if len(token_rewards) % 30 != 0:\n",
    "        progressive_rewards.append(token_rewards.mean().item())\n",
    "        progressive_times.append(len(token_rewards) / 25.0)\n",
    "\n",
    "    marker = \"o\" if name == \"low_cfg\" else \"s\"\n",
    "    ax.plot(\n",
    "        progressive_times,\n",
    "        progressive_rewards,\n",
    "        marker=marker,\n",
    "        linewidth=2.5,\n",
    "        label=f'{name.upper()} (final: {data[\"scalar_reward\"]:.3f})',\n",
    "        markersize=6,\n",
    "        alpha=0.8,\n",
    "    )\n",
    "\n",
    "ax.set_xlabel(\"Sequence Length (seconds)\", fontsize=12)\n",
    "ax.set_ylabel(\"Cumulative Average Reward\", fontsize=12)\n",
    "ax.set_title(\n",
    "    \"Progressive Reward: How Assessment Changes as Sequence Grows\", fontsize=14, fontweight=\"bold\"\n",
    ")\n",
    "ax.legend(fontsize=11, loc=\"best\")\n",
    "ax.grid(True, alpha=0.3)\n",
    "ax.axhline(y=0, color=\"k\", linestyle=\"--\", alpha=0.3)\n",
    "plt.tight_layout()\n",
    "# plt.savefig(\"progressive_reward_comparison.png\", dpi=150, bbox_inches=\"tight\")\n",
    "plt.show()\n",
    "\n",
    "print(\"✓ Saved to progressive_reward_comparison.png\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 10. Summary and Analysis\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"=\" * 80)\n",
    "print(\"REWARD MODEL DEMO SUMMARY\")\n",
    "print(\"=\" * 80)\n",
    "print()\n",
    "print(f\"Test prompt: '{text[:60]}...'\")\n",
    "print(f\"Tags: '{tags}'\")\n",
    "print()\n",
    "print(\"Results:\")\n",
    "print(\"-\" * 80)\n",
    "for name, data in results.items():\n",
    "    print(f\"{name.upper():15s}: Reward = {data['scalar_reward']:.4f}, Length = {data['length_s']:.1f}s\")\n",
    "\n",
    "print()\n",
    "print(f\"Reward Difference (HIGH - LOW): {reward_diff:+.4f}\")\n",
    "print()\n",
    "\n",
    "if abs(reward_diff) > 0.1:\n",
    "    if reward_diff > 0:\n",
    "        print(\"✓ Higher CFG → Higher Reward (as expected)\")\n",
    "        print(\"  Interpretation: Reward model prefers stronger tag guidance\")\n",
    "    else:\n",
    "        print(\"⚠ Lower CFG → Higher Reward (unexpected)\")\n",
    "        print(\"  Interpretation: Reward model may prefer more diverse/creative outputs\")\n",
    "else:\n",
    "    print(\"≈ Similar rewards regardless of CFG\")\n",
    "    print(\"  Interpretation: CFG doesn't strongly affect perceived quality\")\n",
    "\n",
    "print()\n",
    "print(\"Generated Files:\")\n",
    "print(\"  - reward_comparison_cfg.png: Token-level rewards over time\")\n",
    "print(\"  - progressive_reward_comparison.png: Cumulative reward evolution\")\n",
    "print()\n",
    "print(\"=\" * 80)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 11. Experiment: Does Text Conditioning Affect Rewards?\n",
    "\n",
    "Test whether the same audio tokens get different rewards with different text prompts.\n",
    "\n",
    "**Hypothesis**: If model learned text-audio alignment, rewards should change with different prompts.\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Take tokens from HIGH CFG generation (use as test audio)\n",
    "test_tokens = tokens_high.clone()\n",
    "print(f\"Using {len(test_tokens)} tokens ({len(test_tokens)/25:.1f}s) for text conditioning test\")\n",
    "\n",
    "# Define different text prompts to test\n",
    "text_variations = [\n",
    "    (\"no_text\", None, \"No text (baseline)\"),\n",
    "    (\"original\", text, \"Original lyrics (upbeat pop)\"),\n",
    "    (\n",
    "        \"happy\",\n",
    "        \"[Verse]\\nDancing in the sunshine, feeling so free\\nJumping with joy, happy as can be\",\n",
    "        \"Happy lyrics\",\n",
    "    ),\n",
    "    (\n",
    "        \"sad\",\n",
    "        \"[Verse]\\nTears falling down, heart feeling blue\\nLost and alone, missing you\",\n",
    "        \"Sad lyrics\",\n",
    "    ),\n",
    "    (\"instrumental\", \"\", \"Empty text (instrumental)\"),\n",
    "    (\n",
    "        \"mismatched\",\n",
    "        \"[Verse]\\nDeath and destruction, apocalypse now\\nScreaming in terror, end of days\",\n",
    "        \"Dark/mismatched lyrics\",\n",
    "    ),\n",
    "]\n",
    "\n",
    "print(f\"\\nTesting {len(text_variations)} different text conditions...\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_reward_with_text(tokens, reward_model, text=None, tokenizer=None):\n",
    "    \"\"\"Get reward for tokens with specific text conditioning.\n",
    "\n",
    "    Args:\n",
    "        tokens: Semantic tokens to evaluate\n",
    "        reward_model: Trained reward model\n",
    "        text: Text prompt (None = all padding, \"\" = empty, or actual text)\n",
    "        tokenizer: Tokenizer for encoding text\n",
    "\n",
    "    Returns:\n",
    "        reward_logits, scalar_reward\n",
    "    \"\"\"\n",
    "    batch_size = 1\n",
    "    n_streams = 1 + reward_model.config.semantic_n_codebooks + reward_model.config.coarse_n_codebooks\n",
    "    seq_len = len(tokens)\n",
    "\n",
    "    # Create input tensor\n",
    "    X = torch.zeros(batch_size, n_streams, seq_len, dtype=torch.long, device=\"cuda\")\n",
    "\n",
    "    # Text stream - condition on text if provided\n",
    "    if text is not None and text != \"\" and tokenizer is not None:\n",
    "        # Encode text\n",
    "        text_tokens = tokenizer.encode(text)\n",
    "        # Put text tokens at the beginning\n",
    "        text_len = min(len(text_tokens), seq_len)\n",
    "        X[0, 0, :text_len] = torch.tensor(text_tokens[:text_len], device=\"cuda\")\n",
    "        # Rest is padding\n",
    "        X[0, 0, text_len:] = reward_model.config.text_pad_token\n",
    "    else:\n",
    "        # No text - all padding\n",
    "        X[0, 0, :] = reward_model.config.text_pad_token\n",
    "\n",
    "    # Semantic tokens (same for all variations)\n",
    "    X[0, 1, :] = tokens.cuda()\n",
    "\n",
    "    # Coarse streams\n",
    "    for i in range(reward_model.config.coarse_n_codebooks):\n",
    "        X[0, 2 + i, :] = reward_model.config.coarse_pad_token\n",
    "\n",
    "    # Get rewards\n",
    "    with torch.no_grad():\n",
    "        with torch.amp.autocast(device_type=\"cuda\", dtype=torch.bfloat16):\n",
    "            output = reward_model(X, return_logits=True)\n",
    "\n",
    "    reward_logits = output[\"reward_logits\"]\n",
    "\n",
    "    # Create Y for masking\n",
    "    Y = torch.ones(\n",
    "        batch_size,\n",
    "        reward_model.config.semantic_n_codebooks,\n",
    "        seq_len - 1,\n",
    "        dtype=torch.long,\n",
    "        device=\"cuda\",\n",
    "    )\n",
    "\n",
    "    # Find where semantic generation starts\n",
    "    semantic_infer_token = reward_model.config.semantic_infer_token\n",
    "    semantic_stream = X[0, 1, :]\n",
    "\n",
    "    if (semantic_stream == semantic_infer_token).any():\n",
    "        loss_start_index = (semantic_stream == semantic_infer_token).nonzero(as_tuple=True)[0][0].item()\n",
    "    else:\n",
    "        loss_start_index = reward_model.config.t_text\n",
    "\n",
    "    # No padding in generated tokens, so end_index=None\n",
    "    scalar_reward = extract_scalar_rewards(reward_logits, [loss_start_index], None)\n",
    "\n",
    "    return reward_logits[0].float().cpu(), scalar_reward.item()\n",
    "\n",
    "\n",
    "print(\"✓ Enhanced helper function defined\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Compute rewards with different text conditioning\n",
    "text_test_results = {}\n",
    "\n",
    "for name, test_text, description in text_variations:\n",
    "    token_rewards, scalar_reward = get_reward_with_text(\n",
    "        test_tokens,\n",
    "        reward_model,\n",
    "        text=test_text,\n",
    "        tokenizer=engine.tokenizer if test_text is not None else None,\n",
    "    )\n",
    "\n",
    "    text_test_results[name] = {\n",
    "        \"scalar_reward\": scalar_reward,\n",
    "        \"token_rewards\": token_rewards,\n",
    "        \"description\": description,\n",
    "    }\n",
    "\n",
    "    print(f\"{name:15s}: {scalar_reward:.4f} - {description}\")\n",
    "\n",
    "# Show variance\n",
    "rewards_list = [r[\"scalar_reward\"] for r in text_test_results.values()]\n",
    "reward_std = np.std(rewards_list)\n",
    "reward_range = max(rewards_list) - min(rewards_list)\n",
    "\n",
    "print(f\"\\nReward statistics:\")\n",
    "print(f\"  Range: {reward_range:.4f} (max - min)\")\n",
    "print(f\"  Std:   {reward_std:.4f}\")\n",
    "\n",
    "if reward_range > 0.1:\n",
    "    print(f\"  → Text conditioning DOES affect rewards significantly!\")\n",
    "else:\n",
    "    print(f\"  → Text conditioning has minimal effect on rewards\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Visualize text conditioning effect\n",
    "fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n",
    "\n",
    "# Left: Bar chart of absolute rewards\n",
    "names = [r[\"description\"] for r in text_test_results.values()]\n",
    "rewards = [r[\"scalar_reward\"] for r in text_test_results.values()]\n",
    "colors = [\"gray\", \"green\", \"yellow\", \"blue\", \"orange\", \"red\"]\n",
    "\n",
    "bars = ax1.bar(range(len(names)), rewards, color=colors, alpha=0.7, edgecolor=\"black\", linewidth=1.5)\n",
    "ax1.set_xticks(range(len(names)))\n",
    "ax1.set_xticklabels(names, rotation=45, ha=\"right\", fontsize=10)\n",
    "ax1.set_ylabel(\"Reward Score\", fontsize=12)\n",
    "ax1.set_title(\"Same Audio Tokens, Different Text Conditioning\", fontsize=13, fontweight=\"bold\")\n",
    "ax1.grid(True, alpha=0.3, axis=\"y\")\n",
    "ax1.axhline(\n",
    "    y=text_test_results[\"no_text\"][\"scalar_reward\"],\n",
    "    color=\"k\",\n",
    "    linestyle=\"--\",\n",
    "    linewidth=2,\n",
    "    alpha=0.5,\n",
    "    label=\"Baseline (no text)\",\n",
    ")\n",
    "ax1.legend(fontsize=10)\n",
    "\n",
    "# Right: Reward difference from baseline\n",
    "baseline_reward = text_test_results[\"no_text\"][\"scalar_reward\"]\n",
    "reward_diffs = [r - baseline_reward for r in rewards]\n",
    "\n",
    "bars2 = ax2.bar(\n",
    "    range(len(names)), reward_diffs, color=colors, alpha=0.7, edgecolor=\"black\", linewidth=1.5\n",
    ")\n",
    "ax2.set_xticks(range(len(names)))\n",
    "ax2.set_xticklabels(names, rotation=45, ha=\"right\", fontsize=10)\n",
    "ax2.set_ylabel(\"Reward Difference from Baseline\", fontsize=12)\n",
    "ax2.set_title(\"Effect of Text Conditioning (Δ from No Text)\", fontsize=13, fontweight=\"bold\")\n",
    "ax2.axhline(y=0, color=\"k\", linestyle=\"-\", linewidth=2, alpha=0.5)\n",
    "ax2.grid(True, alpha=0.3, axis=\"y\")\n",
    "\n",
    "# Add value labels on bars\n",
    "for i, (bar, val) in enumerate(zip(bars2, reward_diffs)):\n",
    "    height = bar.get_height()\n",
    "    ax2.text(\n",
    "        bar.get_x() + bar.get_width() / 2.0,\n",
    "        height,\n",
    "        f\"{val:+.3f}\",\n",
    "        ha=\"center\",\n",
    "        va=\"bottom\" if val > 0 else \"top\",\n",
    "        fontsize=9,\n",
    "    )\n",
    "\n",
    "plt.tight_layout()\n",
    "# plt.savefig(\"text_conditioning_effect.png\", dpi=150, bbox_inches=\"tight\")\n",
    "plt.show()\n",
    "\n",
    "print(\"✓ Saved to text_conditioning_effect.png\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"=\" * 80)\n",
    "print(\"TEXT CONDITIONING ANALYSIS\")\n",
    "print(\"=\" * 80)\n",
    "print()\n",
    "\n",
    "baseline = text_test_results[\"no_text\"][\"scalar_reward\"]\n",
    "original = text_test_results[\"original\"][\"scalar_reward\"]\n",
    "happy = text_test_results[\"happy\"][\"scalar_reward\"]\n",
    "sad = text_test_results[\"sad\"][\"scalar_reward\"]\n",
    "mismatched = text_test_results[\"mismatched\"][\"scalar_reward\"]\n",
    "\n",
    "print(f\"Baseline (no text):        {baseline:.4f}\")\n",
    "print(f\"Original (matched):        {original:.4f}  ({original - baseline:+.4f})\")\n",
    "print(f\"Happy lyrics:              {happy:.4f}  ({happy - baseline:+.4f})\")\n",
    "print(f\"Sad lyrics:                {sad:.4f}  ({sad - baseline:+.4f})\")\n",
    "print(f\"Mismatched (dark):         {mismatched:.4f}  ({mismatched - baseline:+.4f})\")\n",
    "print()\n",
    "\n",
    "# Interpret results\n",
    "if abs(original - baseline) > 0.05:\n",
    "    print(\"✓ Text conditioning affects rewards\")\n",
    "    print(f\"  Original prompt changes reward by {original - baseline:+.4f}\")\n",
    "\n",
    "    if happy > sad and happy > mismatched:\n",
    "        print(\"✓ Happy text → higher reward (model prefers positive content)\")\n",
    "    elif sad > happy:\n",
    "        print(\"⚠ Sad text → higher reward (unexpected - audio was upbeat)\")\n",
    "\n",
    "    if mismatched < baseline:\n",
    "        print(\"✓ Mismatched text → lower reward (model detects misalignment)\")\n",
    "\n",
    "    print(\"\\n→ Reward model learned text-audio alignment!\")\n",
    "else:\n",
    "    print(\"≈ Text conditioning has minimal effect (< 0.05)\")\n",
    "    print(\"→ Reward model focuses primarily on audio quality, largely ignores text\")\n",
    "\n",
    "print(\"=\" * 80)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 12. Final Summary\n",
    "\n",
    "Complete analysis of reward model behavior:\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"=\" * 80)\n",
    "print(\"REWARD MODEL DEMO - COMPLETE SUMMARY\")\n",
    "print(\"=\" * 80)\n",
    "print()\n",
    "\n",
    "print(\"1. CFG COMPARISON (LOW vs HIGH)\")\n",
    "print(\"-\" * 80)\n",
    "print(f\"   LOW CFG (0.5):  {scalar_reward_low:.4f}\")\n",
    "print(f\"   HIGH CFG (2.0): {scalar_reward_high:.4f}\")\n",
    "print(f\"   Difference:     {reward_diff:+.4f}\")\n",
    "if abs(reward_diff) > 0.1:\n",
    "    winner = \"HIGH CFG\" if reward_diff > 0 else \"LOW CFG\"\n",
    "    print(f\"   → {winner} preferred by reward model\")\n",
    "else:\n",
    "    print(f\"   → CFG has minimal effect on reward\")\n",
    "print()\n",
    "\n",
    "print(\"2. TEXT CONDITIONING EFFECT\")\n",
    "print(\"-\" * 80)\n",
    "baseline = text_test_results[\"no_text\"][\"scalar_reward\"]\n",
    "for name, data in text_test_results.items():\n",
    "    diff = data[\"scalar_reward\"] - baseline\n",
    "    print(f\"   {data['description']:30s}: {data['scalar_reward']:.4f} ({diff:+.4f})\")\n",
    "\n",
    "print(f\"\\n   Range: {reward_range:.4f}\")\n",
    "if reward_range > 0.1:\n",
    "    print(f\"   → Text DOES affect rewards (multi-modal model)\")\n",
    "else:\n",
    "    print(f\"   → Text has minimal effect (audio-only model)\")\n",
    "print()\n",
    "\n",
    "print(\"3. GENERATED FILES\")\n",
    "print(\"-\" * 80)\n",
    "print(\"   - reward_comparison_cfg.png: Token-level rewards (LOW vs HIGH CFG)\")\n",
    "print(\"   - progressive_reward_comparison.png: Cumulative reward evolution\")\n",
    "print(\"   - text_conditioning_effect.png: Text conditioning impact\")\n",
    "print()\n",
    "\n",
    "print(\"4. KEY INSIGHTS\")\n",
    "print(\"-\" * 80)\n",
    "print(f\"   • Reward model outputs both rewards AND generation logits\")\n",
    "print(f\"   • Can use same model for generation and evaluation\")\n",
    "print(f\"   • Token-level rewards show temporal quality assessment\")\n",
    "print(f\"   • Progressive rewards reveal how assessment evolves\")\n",
    "\n",
    "if reward_range > 0.1:\n",
    "    print(f\"   • Model is multi-modal (text + audio)\")\n",
    "else:\n",
    "    print(f\"   • Model focuses on audio quality\")\n",
    "\n",
    "print()\n",
    "print(\"=\" * 80)\n",
    "print(\"✓ Demo complete! Review plots and analysis above.\")\n",
    "print(\"=\" * 80)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 13. Bonus: Score Existing Songs from S3\n",
    "\n",
    "Load real production songs and evaluate with reward model.\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def load_song_from_s3_npz(clip_id):\n",
    "    \"\"\"Load song codes from S3 NPZ file.\n",
    "\n",
    "    Args:\n",
    "        clip_id: Clip ID (e.g., \"a17f7ec0-e7a4-437d-92c6-4d1faa6e99f5\")\n",
    "\n",
    "    Returns:\n",
    "        X: Tensor (1, n_streams, seq_len) ready for model\n",
    "        duration_s: Song duration in seconds\n",
    "    \"\"\"\n",
    "    import tempfile\n",
    "    import boto3\n",
    "\n",
    "    # Use boto3 directly\n",
    "    s3 = boto3.client(\"s3\")\n",
    "\n",
    "    with tempfile.TemporaryDirectory() as td:\n",
    "        npz_path = os.path.join(td, f\"{clip_id}.npz\")\n",
    "\n",
    "        # Download from S3\n",
    "        s3.download_file(\"suno-data-uploads\", f\"studio/uploads/{clip_id}.npz\", npz_path)\n",
    "\n",
    "        # Load codes\n",
    "        npz_data = np.load(npz_path)\n",
    "        codes = npz_data[\"fuller_arr\"]  # Shape: (n_codebooks, seq_len)\n",
    "\n",
    "        # fuller_arr is (n_codebooks, seq_len) - need to transpose and add batch dim\n",
    "        # Result should be (1, n_codebooks, seq_len)\n",
    "        X = torch.tensor(codes, dtype=torch.long).unsqueeze(0).to(\"cuda\")\n",
    "\n",
    "        # Actual sequence length is in the second dimension\n",
    "        seq_len = codes.shape[1]\n",
    "        duration_s = seq_len / 25.0\n",
    "\n",
    "        return X, duration_s\n",
    "\n",
    "\n",
    "print(\"✓ S3 loader defined\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Example clip IDs (replace with your own)\n",
    "clip_ids = [\n",
    "    \"fa70e7d6-c28c-40ae-a8ea-578a696606f8\",  # Update with real clip IDs\n",
    "    # Add more...\n",
    "]\n",
    "\n",
    "s3_results = {}\n",
    "\n",
    "print(\"Loading and scoring songs from S3...\")\n",
    "for clip_id in clip_ids:\n",
    "    X, duration_s = load_song_from_s3_npz(clip_id)\n",
    "    print(X.shape, duration_s)\n",
    "    # Get reward\n",
    "    with torch.no_grad():\n",
    "        with torch.amp.autocast(device_type=\"cuda\", dtype=torch.bfloat16):\n",
    "            output = reward_model(X, return_logits=True)\n",
    "\n",
    "    reward_logits = output[\"reward_logits\"][0].float().cpu()\n",
    "\n",
    "    # Create Y from X (shifted by 1) to detect padding properly\n",
    "    # Y = X[:, 1:, 1:] (remove text stream and shift)\n",
    "    Y = X[:, 1:, 1:].clone()  # (1, n_codebooks, seq_len-1)\n",
    "\n",
    "    # Find semantic generation start\n",
    "    semantic_stream = X[0, 1, :]\n",
    "    semantic_infer_token = reward_model.config.semantic_infer_token\n",
    "\n",
    "    if (semantic_stream == semantic_infer_token).any():\n",
    "        semantic_start_pos = (\n",
    "            (semantic_stream == semantic_infer_token).nonzero(as_tuple=True)[0][0].item()\n",
    "        )\n",
    "    else:\n",
    "        semantic_start_pos = 0  # fallback\n",
    "\n",
    "    print(f\"semantic_infer_token={semantic_infer_token}, semantic_start_pos={semantic_start_pos}\")\n",
    "\n",
    "    # Compute end index to exclude padding\n",
    "    from scripts.reward_eval_utils import compute_loss_end_indices\n",
    "\n",
    "    loss_end_index_list = compute_loss_end_indices(Y, len(reward_logits))\n",
    "\n",
    "    # Extract scalar reward with proper masking\n",
    "    scalar_reward = extract_scalar_rewards(\n",
    "        output[\"reward_logits\"], [semantic_start_pos], loss_end_index_list\n",
    "    ).item()\n",
    "\n",
    "    s3_results[clip_id] = {\n",
    "        \"reward\": scalar_reward,\n",
    "        \"token_rewards\": reward_logits,\n",
    "        \"duration_s\": duration_s,\n",
    "    }\n",
    "\n",
    "    print(f\"  {clip_id[:12]}...: {scalar_reward:.4f} ({duration_s:.1f}s)\")\n",
    "\n",
    "print(f\"\\n✓ Scored {len(s3_results)} songs\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "reward_logits[:10], reward_logits[-10:]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "if len(s3_results) > 0:\n",
    "    # Plot rewards\n",
    "    fig, ax = plt.subplots(1, 1, figsize=(12, 5))\n",
    "\n",
    "    clip_names = [cid[:12] + \"...\" for cid in s3_results.keys()]\n",
    "    rewards = [data[\"reward\"] for data in s3_results.values()]\n",
    "\n",
    "    bars = ax.bar(range(len(clip_names)), rewards, alpha=0.7, edgecolor=\"black\")\n",
    "    ax.set_xticks(range(len(clip_names)))\n",
    "    ax.set_xticklabels(clip_names, rotation=45, ha=\"right\")\n",
    "    ax.set_ylabel(\"Reward Score\", fontsize=12)\n",
    "    ax.set_title(\"Real Songs from S3: Reward Scores\", fontsize=13, fontweight=\"bold\")\n",
    "    ax.grid(True, alpha=0.3, axis=\"y\")\n",
    "    ax.axhline(\n",
    "        np.mean(rewards), color=\"r\", linestyle=\"--\", label=f\"Mean: {np.mean(rewards):.3f}\", linewidth=2\n",
    "    )\n",
    "    ax.legend()\n",
    "\n",
    "    plt.tight_layout()\n",
    "    # plt.savefig(\"s3_songs_rewards.png\", dpi=150)\n",
    "    plt.show()\n",
    "\n",
    "    print(\"✓ Saved to s3_songs_rewards.png\")\n",
    "else:\n",
    "    print(\"No songs loaded - update clip_ids with real IDs\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
