{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 09｜从 Token ID 到 Next-token Loss\n",
    "\n",
    "目标：构造错位的 input/target，追踪 `[B,T] → [B,T,C] → [B,T,V]`，并验证 `-100` 位置不参与平均 loss。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import torch.nn.functional as F\n",
    "\n",
    "torch.manual_seed(42)\n",
    "vocab_size, width = 12, 8\n",
    "sequence = torch.tensor([2, 5, 1, 9, 4])\n",
    "input_ids = sequence[:-1].unsqueeze(0)\n",
    "labels = sequence[1:].unsqueeze(0)\n",
    "print('input :', input_ids.tolist())\n",
    "print('target:', labels.tolist())\n",
    "assert input_ids.shape == labels.shape == (1, 4)\n",
    "assert torch.equal(input_ids[0, 1:], labels[0, :-1])"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 追踪形状并计算逐位置损失\n",
    "\n",
    "这里故意不加入 Attention；只验证语言模型输入、输出头和交叉熵接口。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "embedding = torch.nn.Embedding(vocab_size, width)\n",
    "lm_head = torch.nn.Linear(width, vocab_size, bias=False)\n",
    "hidden = embedding(input_ids)\n",
    "logits = lm_head(hidden)\n",
    "per_token_loss = F.cross_entropy(\n",
    "    logits.reshape(-1, vocab_size), labels.reshape(-1), reduction='none'\n",
    ").reshape_as(labels)\n",
    "print('hidden:', tuple(hidden.shape))\n",
    "print('logits:', tuple(logits.shape))\n",
    "print('per-token NLL:', per_token_loss.tolist())\n",
    "assert hidden.shape == (1, 4, width)\n",
    "assert logits.shape == (1, 4, vocab_size)\n",
    "assert torch.isfinite(per_token_loss).all()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 验证 Loss Mask\n",
    "\n",
    "把最后一个标签设为 `-100`，总 loss 应等于前三个位置 NLL 的平均。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "masked_labels = labels.clone()\n",
    "masked_labels[:, -1] = -100\n",
    "masked_loss = F.cross_entropy(\n",
    "    logits.reshape(-1, vocab_size), masked_labels.reshape(-1), ignore_index=-100\n",
    ")\n",
    "expected = per_token_loss[:, :-1].mean()\n",
    "print(f'masked loss={masked_loss.item():.6f}, expected={expected.item():.6f}')\n",
    "assert torch.allclose(masked_loss, expected)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 练习与验收\n",
    "\n",
    "1. 把 batch 扩成两条不同序列，重新标注形状。\n",
    "2. mask 前两个位置，确认平均分母只包含有效 token。\n",
    "3. 给 hidden 加一个能读取历史的模块，解释为什么当前独立 embedding 无法根据上下文预测。\n",
    "4. 能脱离代码画出 ID、embedding、logits、labels 和 scalar loss 的关系。"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {"display_name": "Python 3", "language": "python", "name": "python3"},
  "language_info": {"name": "python", "version": "3.12"}
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
