-
-
Save gurusura/ec491017d8e08d6ef9d1d0a203c9f7ac to your computer and use it in GitHub Desktop.
Jev + GTSAM: minimal standalone reproduction of CPT generation, evidence updates, and Bayesian-network structure recovery
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| { | |
| "cells": [ | |
| { | |
| "cell_type": "markdown", | |
| "id": "bc2c7379", | |
| "metadata": {}, | |
| "source": [ | |
| "# Jev + GTSAM: minimal reproduction\n", | |
| "\n", | |
| "A standalone reproduction of the article’s CPTs, three evidence updates, and structure recovery for the classic Asia/Chest Clinic network." | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "id": "dd6ca096", | |
| "metadata": {}, | |
| "source": [ | |
| "The setup cell installs missing dependencies into the active kernel. It installs or upgrades `gtsam` only when neither `gtsam` nor `gtsam-develop` is newer than `4.3a0`; any newer installed version is left unchanged." | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": null, | |
| "id": "861aa1ee", | |
| "metadata": {}, | |
| "outputs": [], | |
| "source": [ | |
| "import importlib.util\n", | |
| "import subprocess\n", | |
| "import sys\n", | |
| "from importlib.metadata import version, PackageNotFoundError\n", | |
| "\n", | |
| "requirements = {\n", | |
| " \"numpy\": \"numpy==1.26.4\",\n", | |
| " \"typesafe_sdk\": \"typesafe-sdk==0.7.0\",\n", | |
| " \"packaging\": \"packaging\",\n", | |
| "}\n", | |
| "missing_packages = [package for module, package in requirements.items()\n", | |
| " if importlib.util.find_spec(module) is None]\n", | |
| "if missing_packages:\n", | |
| " subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", *missing_packages])\n", | |
| "\n", | |
| "from packaging.version import Version\n", | |
| "\n", | |
| "gtsam_versions = []\n", | |
| "for package in (\"gtsam\", \"gtsam-develop\"):\n", | |
| " try:\n", | |
| " gtsam_versions.append(Version(version(package)))\n", | |
| " except PackageNotFoundError:\n", | |
| " pass\n", | |
| "if any(v > Version(\"4.3a0\") for v in gtsam_versions):\n", | |
| " print(\"GTSAM newer than 4.3a0 is already installed; leaving it unchanged.\")\n", | |
| "else:\n", | |
| " subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", \"--pre\", \"gtsam>4.3a0\"])\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "id": "0a141676", | |
| "metadata": {}, | |
| "source": [ | |
| "You can run this notebook even if you do not have a JEV API key. `LIVE = False` uses the original `jev-1.13.0` probabilities embedded below. Change it to `True` to make **two Jev requests** with the original prompts; enter your key in the hidden prompt or set `TYPESAFE_API_KEY`." | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": null, | |
| "id": "2291852a", | |
| "metadata": {}, | |
| "outputs": [], | |
| "source": [ | |
| "import itertools as it\n", | |
| "import os\n", | |
| "from getpass import getpass\n", | |
| "\n", | |
| "import numpy as np\n", | |
| "import gtsam\n", | |
| "\n", | |
| "LIVE = False\n", | |
| "MODEL = \"jev-1.13.0\"\n", | |
| "\n", | |
| "\n", | |
| "def ask(state, questions, recorded, labels):\n", | |
| " if not LIVE:\n", | |
| " return np.array(recorded, dtype=float)\n", | |
| " from typesafe_sdk import Choice, TypeSafeClient\n", | |
| " if not os.getenv(\"TYPESAFE_API_KEY\", \"\").strip():\n", | |
| " os.environ[\"TYPESAFE_API_KEY\"] = getpass(\"TypeSafe API key: \").strip()\n", | |
| " with TypeSafeClient() as client:\n", | |
| " response = client.system_one(\n", | |
| " model=MODEL, state=state,\n", | |
| " questions={qid: Choice(**spec) for qid, spec in questions.items()},\n", | |
| " ).model_dump(mode=\"json\")\n", | |
| " return np.array([[response[\"answers\"][qid][\"probabilities\"][label]\n", | |
| " for label in labels] for qid in questions], dtype=float)\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "id": "300ea619", | |
| "metadata": {}, | |
| "source": [ | |
| "## 1. Jev supplies 18 CPT rows\n", | |
| "\n", | |
| "Variables are binary (`0 = false`, `1 = true`). Parent assignments are in lexicographic order, with the last parent changing fastest; each GTSAM row is `P(false)/P(true)`.\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": null, | |
| "id": "eb69d9f0", | |
| "metadata": {}, | |
| "outputs": [], | |
| "source": [ | |
| "NODES = {\n", | |
| " 'A': ('the patient recently visited Asia', ()),\n", | |
| " 'S': ('the patient is a smoker', ()),\n", | |
| " 'T': ('the patient has tuberculosis', ('A',)),\n", | |
| " 'L': ('the patient has lung cancer', ('S',)),\n", | |
| " 'B': ('the patient has bronchitis', ('S',)),\n", | |
| " 'E': ('the patient has tuberculosis or lung cancer', ('T', 'L')),\n", | |
| " 'X': (\"the patient's chest X-ray is abnormal\", ('E',)),\n", | |
| " 'D': ('the patient has shortness of breath (dyspnea)', ('E', 'B')),\n", | |
| "}\n", | |
| "\n", | |
| "CPT_STATE = {'task': 'Parameterize a small educational Bayesian network about respiratory disease.',\n", | |
| " 'population': 'Imagine a randomly selected adult patient in a general clinical population. '\n", | |
| " 'Use ordinary real-world medical knowledge, not the memorized textbook '\n", | |
| " 'Asia-network numbers.',\n", | |
| " 'interpretation': 'For each question, return a probability distribution for the child '\n", | |
| " 'variable under exactly the stated parent assignment. Treat unspecified '\n", | |
| " 'variables as unknown; do not assume them false.'}\n", | |
| "\n", | |
| "rows = [(child, values) for child, (_, parents) in NODES.items()\n", | |
| " for values in it.product((0, 1), repeat=len(parents))]\n", | |
| "cpt_questions = {}\n", | |
| "for child, values in rows:\n", | |
| " meaning, parents = NODES[child]\n", | |
| " suffix = \"_\".join(f\"{p}{v}\" for p, v in zip(parents, values)) or \"root\"\n", | |
| " condition = \" \".join(\n", | |
| " f\"It is {'true' if v else 'false'} that {NODES[p][0]}.\"\n", | |
| " for p, v in zip(parents, values)\n", | |
| " ) or \"No other facts about this randomly selected patient are given.\"\n", | |
| " cpt_questions[f\"cpt_{child}_{suffix}\"] = {\n", | |
| " \"instructions\": f\"Given the stated condition, which truth value best describes whether {meaning}? \"\n", | |
| " f\"Condition: {condition}\",\n", | |
| " \"criteria\": {\"false\": f\"It is false that {meaning}.\",\n", | |
| " \"true\": f\"It is true that {meaning}.\"},\n", | |
| " }\n", | |
| "\n", | |
| "# Original Jev probabilities, in the same order as rows.\n", | |
| "recorded_cpts = [\n", | |
| " [0.96, 0.04], # A | {}\n", | |
| " [0.8, 0.2], # S | {}\n", | |
| " [0.97, 0.03], # T | {'A': 0}\n", | |
| " [0.68, 0.32], # T | {'A': 1}\n", | |
| " [0.99, 0.01], # L | {'S': 0}\n", | |
| " [0.92, 0.08], # L | {'S': 1}\n", | |
| " [0.95, 0.05], # B | {'S': 0}\n", | |
| " [0.31, 0.69], # B | {'S': 1}\n", | |
| " [0.99, 0.01], # E | {'T': 0, 'L': 0}\n", | |
| " [0.0, 1.0], # E | {'T': 0, 'L': 1}\n", | |
| " [0.0, 1.0], # E | {'T': 1, 'L': 0}\n", | |
| " [0.0, 1.0], # E | {'T': 1, 'L': 1}\n", | |
| " [0.89, 0.11], # X | {'E': 0}\n", | |
| " [0.01, 0.99], # X | {'E': 1}\n", | |
| " [0.93, 0.07], # D | {'E': 0, 'B': 0}\n", | |
| " [0.15, 0.85], # D | {'E': 0, 'B': 1}\n", | |
| " [0.05, 0.95], # D | {'E': 1, 'B': 0}\n", | |
| " [0.03, 0.97], # D | {'E': 1, 'B': 1}\n", | |
| "]\n", | |
| "cpt_probabilities = ask(CPT_STATE, cpt_questions, recorded_cpts, (\"false\", \"true\"))\n", | |
| "cpt_probabilities /= cpt_probabilities.sum(axis=1, keepdims=True)\n", | |
| "for (child, values), (p0, p1) in zip(rows, cpt_probabilities):\n", | |
| " print(f\"{child} | {str(dict(zip(NODES[child][1], values))):20} {p0:.2f} / {p1:.2f}\")\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "id": "ab748802", | |
| "metadata": {}, | |
| "source": [ | |
| "## 2. GTSAM computes the evidence updates\n", | |
| "\n", | |
| "Add observations as unary factors, then marginalize. These queries make no Jev calls.\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": null, | |
| "id": "e484f169", | |
| "metadata": {}, | |
| "outputs": [], | |
| "source": [ | |
| "keys = {name: (i, 2) for i, name in enumerate(NODES)}\n", | |
| "asia = gtsam.DiscreteBayesNet()\n", | |
| "for child, (_, parents) in NODES.items():\n", | |
| " table = \" \".join(f\"{p0:.12g}/{p1:.12g}\"\n", | |
| " for (name, _), (p0, p1) in zip(rows, cpt_probabilities) if name == child)\n", | |
| " if parents:\n", | |
| " parent_keys = gtsam.DiscreteKeys()\n", | |
| " for parent in parents:\n", | |
| " parent_keys.push_back(keys[parent])\n", | |
| " asia.add(keys[child], parent_keys, table)\n", | |
| " else:\n", | |
| " asia.add(keys[child], table)\n", | |
| "\n", | |
| "\n", | |
| "def posterior(evidence):\n", | |
| " factors = gtsam.DiscreteFactorGraph(asia)\n", | |
| " for name, value in evidence.items():\n", | |
| " factors.add(keys[name], \"0 1\" if value else \"1 0\")\n", | |
| " marginals = gtsam.DiscreteMarginals(factors)\n", | |
| " return [marginals.marginalProbabilities(keys[name])[1] for name in (\"T\", \"L\", \"B\")]\n", | |
| "\n", | |
| "\n", | |
| "scenarios = [(\"No Asia visit\", {\"A\": 0}),\n", | |
| " (\"+ abnormal X-ray\", {\"A\": 0, \"X\": 1}),\n", | |
| " (\"+ shortness of breath\", {\"A\": 0, \"X\": 1, \"D\": 1})]\n", | |
| "evidence_results = np.array([posterior(evidence) for _, evidence in scenarios])\n", | |
| "print(f\"{'Observations':25} {'Tuberculosis':>14} {'Lung cancer':>14} {'Bronchitis':>14}\")\n", | |
| "for (label, _), probabilities in zip(scenarios, evidence_results):\n", | |
| " print(f\"{label:25}\" + \"\".join(f\"{p:14.2%}\" for p in probabilities))\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "id": "9546cfce", | |
| "metadata": {}, | |
| "source": [ | |
| "## 3. Jev proposes structure from variable meanings\n", | |
| "\n", | |
| "Use the original shuffled identifiers (`V1`–`V8`), and ask about all 28 pairs. Neither the reference edges nor CPTs are sent. Each answer is `[P(no edge), P(u → v), P(v → u)]`.\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": null, | |
| "id": "781a003c", | |
| "metadata": {}, | |
| "outputs": [], | |
| "source": [ | |
| "id_to_variable = dict(zip((f\"V{i}\" for i in range(1, 9)), \"LSABETDX\"))\n", | |
| "ids = tuple(id_to_variable)\n", | |
| "pairs = list(it.combinations(ids, 2))\n", | |
| "STRUCTURE_STATE = {'interpretation': 'Consider all listed variables when judging each pair. Prefer direct '\n", | |
| " 'generative or definitional dependencies. Association through a shared '\n", | |
| " 'cause or a path through other listed variables alone does not require a '\n", | |
| " 'direct edge. An observed test result can inform beliefs about a disease '\n", | |
| " 'without being a generative cause of it. Deterministic definitions can '\n", | |
| " 'have incoming edges. If neither direct direction belongs in the model, '\n", | |
| " 'choose no_edge. These questions ask about model structure, not whether a '\n", | |
| " 'variable is true for a particular patient.',\n", | |
| " 'population': 'Adult patients in a general clinical population; no individual patient '\n", | |
| " 'observations are supplied.',\n", | |
| " 'task': 'Specify direct dependencies in a sparse generative Bayesian network over the listed '\n", | |
| " 'binary variables.'}\n", | |
| "STRUCTURE_STATE[\"variables\"] = [\n", | |
| " {\"id\": ident, \"meaning\": NODES[variable][0]} for ident, variable in id_to_variable.items()\n", | |
| "]\n", | |
| "structure_questions = {\n", | |
| " f\"pair_{i:02d}\": {\n", | |
| " \"instructions\": f\"Which direct relationship should connect {u} and {v} in the proposed model, \"\n", | |
| " \"accounting for possible mediation by the other listed variables?\",\n", | |
| " \"criteria\": {\n", | |
| " \"no_edge\": f\"No direct edge between {u} and {v}; an indirect association may still exist.\",\n", | |
| " \"u_to_v\": f\"A direct edge {u} -> {v}: {u} is a parent of {v}.\",\n", | |
| " \"v_to_u\": f\"A direct edge {v} -> {u}: {v} is a parent of {u}.\",\n", | |
| " },\n", | |
| " } for i, (u, v) in enumerate(pairs, 1)\n", | |
| "}\n", | |
| "recorded_pairs = [\n", | |
| " [0.02, 0.0, 0.98], # pair_01\n", | |
| " [0.96, 0.0, 0.04], # pair_02\n", | |
| " [0.85, 0.14, 0.01], # pair_03\n", | |
| " [0.0, 1.0, 0.0], # pair_04\n", | |
| " [0.99, 0.0, 0.01], # pair_05\n", | |
| " [0.28, 0.72, 0.0], # pair_06\n", | |
| " [0.17, 0.83, 0.0], # pair_07\n", | |
| " [1.0, 0.0, 0.0], # pair_08\n", | |
| " [0.1, 0.9, 0.0], # pair_09\n", | |
| " [0.98, 0.02, 0.0], # pair_10\n", | |
| " [0.98, 0.02, 0.0], # pair_11\n", | |
| " [0.98, 0.02, 0.0], # pair_12\n", | |
| " [0.99, 0.01, 0.0], # pair_13\n", | |
| " [0.71, 0.29, 0.0], # pair_14\n", | |
| " [0.37, 0.63, 0.0], # pair_15\n", | |
| " [0.07, 0.93, 0.0], # pair_16\n", | |
| " [1.0, 0.0, 0.0], # pair_17\n", | |
| " [0.99, 0.01, 0.0], # pair_18\n", | |
| " [0.97, 0.01, 0.02], # pair_19\n", | |
| " [0.53, 0.0, 0.47], # pair_20\n", | |
| " [0.11, 0.89, 0.0], # pair_21\n", | |
| " [0.39, 0.61, 0.0], # pair_22\n", | |
| " [0.01, 0.0, 0.99], # pair_23\n", | |
| " [0.22, 0.78, 0.0], # pair_24\n", | |
| " [0.25, 0.75, 0.0], # pair_25\n", | |
| " [0.23, 0.77, 0.0], # pair_26\n", | |
| " [0.21, 0.79, 0.0], # pair_27\n", | |
| " [0.97, 0.01, 0.02], # pair_28\n", | |
| "]\n", | |
| "pair_probabilities = ask(STRUCTURE_STATE, structure_questions, recorded_pairs,\n", | |
| " (\"no_edge\", \"u_to_v\", \"v_to_u\"))\n", | |
| "pair_probabilities /= pair_probabilities.sum(axis=1, keepdims=True)\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "id": "9d59e7aa", | |
| "metadata": {}, | |
| "source": [ | |
| "## 4. Select the graph and compare edges\n", | |
| "\n", | |
| "Maximize `sum(log P(pair choice)) − number of edges` (edge penalty **1**), searching all **8! = 40,320** orders to enforce acyclicity. Ties prefer fewer edges, then the first order. The recorded result is **8 correct directions, 3 extra edges, 0 missing**, with **72.7% precision and 100% recall**.\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": null, | |
| "id": "02429728", | |
| "metadata": {}, | |
| "outputs": [], | |
| "source": [ | |
| "EDGE_PENALTY = 1.0\n", | |
| "log_p = np.log(np.maximum(pair_probabilities, 1e-12))\n", | |
| "orders = np.array(list(it.permutations(range(len(ids)))))\n", | |
| "ranks = np.argsort(orders, axis=1)\n", | |
| "pair_indices = np.array(list(it.combinations(range(len(ids)), 2)))\n", | |
| "forward = ranks[:, pair_indices[:, 0]] < ranks[:, pair_indices[:, 1]]\n", | |
| "gain = np.where(forward, log_p[:, 1] - log_p[:, 0], log_p[:, 2] - log_p[:, 0]) - EDGE_PENALTY\n", | |
| "include = gain > 0\n", | |
| "scores = log_p[:, 0].sum() + np.maximum(gain, 0).sum(axis=1)\n", | |
| "best = np.lexsort((np.arange(len(orders)), include.sum(axis=1), -scores))[0]\n", | |
| "recovered = set()\n", | |
| "for j, (u, v) in enumerate(pairs):\n", | |
| " if include[best, j]:\n", | |
| " u, v = (u, v) if forward[best, j] else (v, u)\n", | |
| " recovered.add((id_to_variable[u], id_to_variable[v]))\n", | |
| "\n", | |
| "reference = {(parent, child) for child, (_, parents) in NODES.items() for parent in parents}\n", | |
| "correct = recovered & reference\n", | |
| "extra = recovered - reference\n", | |
| "missing = reference - recovered\n", | |
| "print(\"Recovered:\", \", \".join(f\"{u} → {v}\" for u, v in sorted(recovered)))\n", | |
| "print(f\"Correct directions: {len(correct)}/{len(reference)}; extra: {len(extra)}; missing: {len(missing)}\")\n", | |
| "print(f\"Precision: {len(correct) / len(recovered):.1%}; recall: {len(correct) / len(reference):.1%}\")\n", | |
| "print(\"Extra edges:\", \", \".join(f\"{u} → {v}\" for u, v in sorted(extra)))\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "id": "274423a1", | |
| "metadata": {}, | |
| "source": [ | |
| "The graph is the classic [Asia/Chest Clinic benchmark (Lauritzen & Spiegelhalter, 1988)](https://doi.org/10.1111/j.2517-6161.1988.tb01721.x); the probabilities are Jev’s estimates. The edge penalty was selected for the article’s comparison. These are modeling demonstrations, not clinical estimates.\n" | |
| ] | |
| } | |
| ], | |
| "metadata": { | |
| "kernelspec": { | |
| "display_name": "Python 3", | |
| "language": "python", | |
| "name": "python3" | |
| }, | |
| "language_info": { | |
| "name": "python" | |
| } | |
| }, | |
| "nbformat": 4, | |
| "nbformat_minor": 5 | |
| } |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment