Skip to content

Instantly share code, notes, and snippets.

@gurusura
Forked from dellaert/jev_minimal_repro.ipynb
Created September 23, 2026 15:45
Show Gist options
  • Select an option

  • Save gurusura/ec491017d8e08d6ef9d1d0a203c9f7ac to your computer and use it in GitHub Desktop.

Select an option

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
Display the source blob
Display the rendered blob
Raw
{
"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