{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 第5章 学生Notebook：美食语义匹配器\n",
    "\n",
    "> 同样的字不等于同样的意思\n",
    "\n",
    "本Notebook完成以下任务：\n",
    "1. 用共同字基线匹配菜品\n",
    "2. 用嵌入向量 + 余弦相似度匹配菜品\n",
    "3. 对比两种方法，分析各自的失败情况\n",
    "4. 修改一个变量（阈值或数据集），观察变化\n",
    "\n",
    "**准备**：确认 `data/candidates.json` 存在。"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 1. 加载数据\n",
    "\n",
    "读取菜品和查询数据。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import numpy as np\n",
    "from pathlib import Path\n",
    "\n",
    "# 加载数据\n",
    "data_path = Path('../data/candidates.json')\n",
    "with open(data_path, 'r', encoding='utf-8') as f:\n",
    "    data = json.load(f)\n",
    "\n",
    "dishes = data['dishes']\n",
    "queries = data['queries']\n",
    "\n",
    "print(f'菜品数量: {len(dishes)}')\n",
    "print(f'查询数量: {len(queries)}')\n",
    "print()\n",
    "print('前3道菜品:')\n",
    "for d in dishes[:3]:\n",
    "    print(f'  {d[\"id\"]} {d[\"name\"]}: {d[\"description\"]}')\n",
    "print()\n",
    "print('第1条查询:')\n",
    "q = queries[0]\n",
    "print(f'  {q[\"id\"]}: {q[\"text\"]}')\n",
    "print(f'  相关: {q[\"relevant\"]}')\n",
    "print(f'  不相关: {q[\"irrelevant\"]}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 2. 方法A：共同字基线\n",
    "\n",
    "数查询和每道菜品名称+描述中有多少相同的字符。共同字越多，相似度越高。\n",
    "\n",
    "**运行**：观察每种查询的排序结果。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def char_overlap_score(query, text):\n",
    "    \"\"\"计算两段文字的共同字数量\"\"\"\n",
    "    query_chars = set(query)\n",
    "    text_chars = set(text)\n",
    "    overlap = query_chars & text_chars\n",
    "    # 去掉标点空格\n",
    "    overlap = {c for c in overlap if c.strip()}\n",
    "    return len(overlap)\n",
    "\n",
    "def char_overlap_ranking(query_text, dishes):\n",
    "    \"\"\"对菜品按共同字数量排序\"\"\"\n",
    "    scores = []\n",
    "    for d in dishes:\n",
    "        combined = d['name'] + d['description']\n",
    "        score = char_overlap_score(query_text, combined)\n",
    "        scores.append((d['id'], d['name'], score))\n",
    "    scores.sort(key=lambda x: x[2], reverse=True)\n",
    "    return scores\n",
    "\n",
    "# 对第1条查询运行\n",
    "q = queries[0]\n",
    "print(f'查询: {q[\"text\"]}')\n",
    "print(f'预期相关: {q[\"relevant\"]}')\n",
    "print()\n",
    "print('共同字排序 Top-5:')\n",
    "ranking = char_overlap_ranking(q['text'], dishes)\n",
    "for i, (did, name, score) in enumerate(ranking[:5]):\n",
    "    marker = '✓' if did in q['relevant'] else ' '\n",
    "    print(f'  {i+1}. [{marker}] {did} {name} (共同字: {score})')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "**思考**：对Q02\"想吃点酸的开胃菜\"运行共同字方法，观察结果。\n",
    "\n",
    "在下方写下你的预测：共同字方法会排前3的是什么？"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# 运行 Q02\n",
    "q2 = queries[1]\n",
    "print(f'查询: {q2[\"text\"]}')\n",
    "print(f'预期相关: {q2[\"relevant\"]}')\n",
    "print()\n",
    "ranking2 = char_overlap_ranking(q2['text'], dishes)\n",
    "print('共同字排序 Top-5:')\n",
    "for i, (did, name, score) in enumerate(ranking2[:5]):\n",
    "    marker = '✓' if did in q2['relevant'] else ' '\n",
    "    print(f'  {i+1}. [{marker}] {did} {name} (共同字: {score})')\n",
    "\n",
    "# 在下方写下你的观察\n",
    "# 我的观察："
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 3. 方法B：嵌入向量 + 余弦相似度\n",
    "\n",
    "### 余弦相似度的计算\n",
    "\n",
    "两个向量的余弦相似度公式：\n",
    "\n",
    "$$\\cos(\\theta) = \\frac{\\vec{a} \\cdot \\vec{b}}{|\\vec{a}| \\times |\\vec{b}|}$$\n",
    "\n",
    "- 值域：-1 到 1\n",
    "- 1 = 方向完全相同\n",
    "- 0 = 方向垂直（无关）\n",
    "\n",
    "**运行**：先看完整代码，再观察结果。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def cosine_similarity(a, b):\n",
    "    \"\"\"计算两个向量的余弦相似度\"\"\"\n",
    "    a = np.array(a, dtype=float)\n",
    "    b = np.array(b, dtype=float)\n",
    "    dot_product = np.dot(a, b)\n",
    "    norm_a = np.linalg.norm(a)\n",
    "    norm_b = np.linalg.norm(b)\n",
    "    if norm_a == 0 or norm_b == 0:\n",
    "        return 0.0\n",
    "    return dot_product / (norm_a * norm_b)\n",
    "\n",
    "# 演示：两个已知向量的余弦相似度\n",
    "v1 = [1, 0, 0]  # 指向x轴\n",
    "v2 = [1, 0, 0]  # 相同方向\n",
    "v3 = [0, 1, 0]  # 垂直方向\n",
    "\n",
    "print(f'相同方向: cos = {cosine_similarity(v1, v2):.4f}')  # 应接近 1.0\n",
    "print(f'垂直方向: cos = {cosine_similarity(v1, v3):.4f}')  # 应为 0.0"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 获取嵌入向量\n",
    "\n",
    "三种路径（按优先级）：\n",
    "1. 本地嵌入模型（Ollama + nomic-embed-text）\n",
    "2. 课程云端嵌入API\n",
    "3. 预计算向量（离线兜底）\n",
    "\n",
    "下面的代码自动检测可用路径。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import hashlib\n",
    "\n",
    "EMBED_DIM = 64  # 预计算向量的维度（简化演示用）\n",
    "\n",
    "def get_embedding_local_model(text):\n",
    "    \"\"\"尝试通过本地 Ollama 获取嵌入向量\"\"\"\n",
    "    try:\n",
    "        import urllib.request\n",
    "        url = 'http://localhost:11434/api/embed'\n",
    "        payload = json.dumps({'model': 'nomic-embed-text', 'input': text}).encode()\n",
    "        req = urllib.request.Request(url, data=payload, headers={'Content-Type': 'application/json'})\n",
    "        with urllib.request.urlopen(req, timeout=10) as resp:\n",
    "            result = json.loads(resp.read())\n",
    "            return result['embeddings'][0]\n",
    "    except Exception:\n",
    "        return None\n",
    "\n",
    "def get_embedding_precomputed(text):\n",
    "    \"\"\"离线兜底：用文本哈希生成确定性伪向量\"\"\"\n",
    "    h = hashlib.sha256(text.encode('utf-8')).digest()\n",
    "    rng = np.random.RandomState(int.from_bytes(h[:4], 'big'))\n",
    "    vec = rng.randn(EMBED_DIM).astype(float)\n",
    "    vec = vec / np.linalg.norm(vec)  # 归一化\n",
    "    return vec.tolist()\n",
    "\n",
    "def get_embedding(text):\n",
    "    \"\"\"自动选择嵌入路径\"\"\"\n",
    "    vec = get_embedding_local_model(text)\n",
    "    if vec is not None:\n",
    "        return vec, 'local'\n",
    "    return get_embedding_precomputed(text), 'precomputed'\n",
    "\n",
    "# 测试嵌入路径\n",
    "test_vec, source = get_embedding('测试文本')\n",
    "print(f'嵌入路径: {source}')\n",
    "print(f'向量维度: {len(test_vec)}')\n",
    "print(f'前5个值: {test_vec[:5]}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 计算所有菜品的嵌入并排序\n",
    "\n",
    "**运行**：对每条查询，用余弦相似度排序菜品。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def embedding_ranking(query_text, dishes):\n",
    "    \"\"\"对菜品按嵌入余弦相似度排序\"\"\"\n",
    "    q_vec, source = get_embedding(query_text)\n",
    "    scores = []\n",
    "    for d in dishes:\n",
    "        combined = d['name'] + ' ' + d['description']\n",
    "        d_vec, _ = get_embedding(combined)\n",
    "        sim = cosine_similarity(q_vec, d_vec)\n",
    "        scores.append((d['id'], d['name'], sim))\n",
    "    scores.sort(key=lambda x: x[2], reverse=True)\n",
    "    return scores, source\n",
    "\n",
    "# 对 Q01 运行\n",
    "q = queries[0]\n",
    "ranking_emb, source = embedding_ranking(q['text'], dishes)\n",
    "print(f'查询: {q[\"text\"]}')\n",
    "print(f'嵌入路径: {source}')\n",
    "print(f'预期相关: {q[\"relevant\"]}')\n",
    "print()\n",
    "print('嵌入排序 Top-5:')\n",
    "for i, (did, name, sim) in enumerate(ranking_emb[:5]):\n",
    "    marker = '✓' if did in q['relevant'] else ' '\n",
    "    print(f'  {i+1}. [{marker}] {did} {name} (相似度: {sim:.4f})')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 4. 两种方法全面对照\n",
    "\n",
    "**运行**：对全部6条查询，对比两种方法。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def evaluate_ranking(ranking, relevant_ids, top_k=5):\n",
    "    \"\"\"计算Top-K中命中相关菜品的数量\"\"\"\n",
    "    top_ids = [r[0] for r in ranking[:top_k]]\n",
    "    hits = sum(1 for did in top_ids if did in relevant_ids)\n",
    "    return hits\n",
    "\n",
    "print(f'{\"查询\":<20} {\"共同字命中\":<10} {\"嵌入命中\":<10} {\"差异\"}')\n",
    "print('-' * 60)\n",
    "\n",
    "for q in queries:\n",
    "    char_rank = char_overlap_ranking(q['text'], dishes)\n",
    "    emb_rank, _ = embedding_ranking(q['text'], dishes)\n",
    "    \n",
    "    char_hits = evaluate_ranking(char_rank, q['relevant'])\n",
    "    emb_hits = evaluate_ranking(emb_rank, q['relevant'])\n",
    "    diff = emb_hits - char_hits\n",
    "    diff_str = f'+{diff}' if diff > 0 else str(diff)\n",
    "    \n",
    "    print(f'{q[\"text\"]:<20} {char_hits:<10} {emb_hits:<10} {diff_str}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "**思考**：\n",
    "\n",
    "1. 哪些查询嵌入方法明显优于共同字？为什么？\n",
    "2. 有没有共同字方法更好或持平的查询？为什么？\n",
    "3. 有没有两种方法都失败的查询？为什么？\n",
    "\n",
    "在下方写下你的分析。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# 我的分析：\n",
    "# \n",
    "# 嵌入明显优于共同字的查询：\n",
    "# \n",
    "# 两种方法都失败的查询：\n",
    "# "
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 5. 可视化：排序对照\n",
    "\n",
    "**运行**：生成两种方法的排序对照图。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "try:\n",
    "    import matplotlib\n",
    "    matplotlib.use('Agg')\n",
    "    import matplotlib.pyplot as plt\n",
    "    HAS_MPL = True\n",
    "except ImportError:\n",
    "    HAS_MPL = False\n",
    "    print('Matplotlib 未安装，跳过可视化。核心计算不受影响。')\n",
    "\n",
    "if HAS_MPL:\n",
    "    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n",
    "    fig.suptitle('共同字 vs 嵌入向量：菜品排序对照', fontsize=14)\n",
    "    \n",
    "    for idx, q in enumerate(queries):\n",
    "        ax = axes[idx // 3][idx % 3]\n",
    "        \n",
    "        char_rank = char_overlap_ranking(q['text'], dishes)\n",
    "        emb_rank, _ = embedding_ranking(q['text'], dishes)\n",
    "        \n",
    "        top5_char = char_rank[:5]\n",
    "        top5_emb = emb_rank[:5]\n",
    "        \n",
    "        names = [d['name'] for d in dishes]\n",
    "        char_scores = [0] * len(dishes)\n",
    "        emb_scores = [0] * len(dishes)\n",
    "        \n",
    "        for rank, (did, _, score) in enumerate(char_rank):\n",
    "            i = int(did[1:]) - 1\n",
    "            char_scores[i] = len(dishes) - rank\n",
    "        for rank, (did, _, score) in enumerate(emb_rank):\n",
    "            i = int(did[1:]) - 1\n",
    "            emb_scores[i] = len(dishes) - rank\n",
    "        \n",
    "        colors = ['green' if did in q['relevant'] else 'gray' for did, _, _ in char_rank]\n",
    "        \n",
    "        ax.barh(range(min(8, len(dishes))), \n",
    "                [char_scores[i] for i in range(min(8, len(dishes)))],\n",
    "                alpha=0.5, label='共同字', color='steelblue')\n",
    "        \n",
    "        ax.set_title(f'{q[\"id\"]}: {q[\"text\"][:12]}...', fontsize=10)\n",
    "        ax.set_yticks(range(min(8, len(dishes))))\n",
    "        ax.set_yticklabels([dishes[i]['name'] for i in range(min(8, len(dishes)))], fontsize=8)\n",
    "        ax.invert_yaxis()\n",
    "    \n",
    "    plt.tight_layout()\n",
    "    plt.savefig('../outputs/ch05_comparison.png', dpi=100, bbox_inches='tight')\n",
    "    print('对照图已保存到 outputs/ch05_comparison.png')\n",
    "    plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 6. B档修改：修改一个变量\n",
    "\n",
    "从以下两项中选择**一项**修改，修改前先写预测。\n",
    "\n",
    "**选项A**：新增3道菜品到数据集\n",
    "- 至少一道能被\"想吃点酸的\"匹配到\n",
    "- 至少一道是\"不要辣的\"应该排除的\n",
    "\n",
    "**选项B**：添加阈值判断\n",
    "- 当最高相似度低于阈值时，输出\"请人工确认\"\n",
    "- 测试不同阈值（0.3, 0.5, 0.7）的影响"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# 我选择：选项___\n",
    "# \n",
    "# 修改前预测：\n",
    "# \n",
    "\n",
    "# === 选项A：新增菜品 ===\n",
    "# 取消下面的注释并修改\n",
    "# new_dishes = [\n",
    "#     {\"id\": \"D21\", \"name\": \"___\", \"description\": \"___\", \"category\": \"___\"},\n",
    "#     {\"id\": \"D22\", \"name\": \"___\", \"description\": \"___\", \"category\": \"___\"},\n",
    "#     {\"id\": \"D23\", \"name\": \"___\", \"description\": \"___\", \"category\": \"___\"},\n",
    "# ]\n",
    "# dishes_extended = dishes + new_dishes\n",
    "\n",
    "# === 选项B：阈值判断 ===\n",
    "# 取消下面的注释并修改\n",
    "# THRESHOLD = 0.5  # 修改这个值\n",
    "# for q in queries:\n",
    "#     emb_rank, _ = embedding_ranking(q['text'], dishes)\n",
    "#     top_sim = emb_rank[0][2]\n",
    "#     status = '推荐' if top_sim >= THRESHOLD else '请人工确认'\n",
    "#     print(f'{q[\"text\"]}: 最高相似度={top_sim:.4f} → {status}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 7. 验收\n",
    "\n",
    "运行下面的检查，确认你完成了所有必做项。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "checks = {\n",
    "    '运行了共同字基线': True,  # 你已经运行了第2节\n",
    "    '运行了嵌入排序': True,    # 你已经运行了第3-4节\n",
    "    '写了对比分析': False,      # 检查第4节的分析单元格是否有内容\n",
    "    '完成了B档修改': False,      # 检查第6节是否有修改\n",
    "}\n",
    "\n",
    "print('验收清单:')\n",
    "for check, status in checks.items():\n",
    "    print(f'  [{\"✓\" if status else \" \"}] {check}')\n",
    "\n",
    "print()\n",
    "print('请手动把 False 改为 True，确认你完成了对应项目。')\n",
    "print()\n",
    "print('关键概念检查:')\n",
    "print('  1. 余弦相似度的值域是 ___ 到 ___')\n",
    "print('  2. 嵌入向量的作用是 ___')\n",
    "print('  3. 共同字方法的局限是 ___')\n",
    "print('  4. 嵌入方法也可能失败的情况是 ___')"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.10.0"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
