{
  "cells": [
    {
      "cell_type": "markdown",
      "id": "sage-00",
      "metadata": {},
      "source": [
        "# Lab 5: Exchangeability, Pólya's urn, and random limits\n",
        "\n",
        "Use the **SageMath kernel**, not the Python kernel used by Labs 1–4.\n",
        "Run the cells in order. This lab accompanies the optional Chapter 10\n",
        "and connects its finite-sampling proof to the birthday bound.\n",
        "The matching script runs with `sage 05_exchangeability_and_polya_sage.sage`.\n",
        "\n",
        "Exact rational calculations come first; seeded simulations follow.\n",
        "Equality in a finite table illustrates a theorem but does not prove the\n",
        "corresponding statement for all lengths or for infinite sequences.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-01",
      "metadata": {},
      "source": [
        "from itertools import product\n",
        "import random\n",
        "\n",
        "biases = [QQ(1)/4, QQ(3)/4]\n",
        "weights = [QQ(1)/2, QQ(1)/2]\n",
        "\n",
        "def mixture_word(word, biases=biases, weights=weights):\n",
        "    s, n = sum(word), len(word)\n",
        "    return sum(w*p^s*(1-p)^(n-s) for p, w in zip(biases, weights))\n",
        "\n",
        "words = list(product([0, 1], repeat=4))\n",
        "assert sum(mixture_word(word) for word in words) == 1\n",
        "assert mixture_word([1, 1, 1, 0]) == mixture_word([1, 0, 1, 1])\n",
        "print(table([(word, mixture_word(word)) for word in words],\n",
        "            header_row=[\"Ordered word\", \"Shared-coin probability\"]))\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-02",
      "metadata": {},
      "source": [
        "## 1. One shared coin, or a fresh coin each time?\n",
        "\n",
        "Selecting a bias once and sharing it across a run creates dependence.\n",
        "Selecting an independent fresh bias before each toss gives fair i.i.d.\n",
        "bits in this example. Both models have the same one-toss marginal.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-03",
      "metadata": {},
      "source": [
        "mean = sum(w*p for p, w in zip(biases, weights))\n",
        "second = sum(w*p^2 for p, w in zip(biases, weights))\n",
        "covariance = second-mean^2\n",
        "assert mean == 1/2 and covariance == 1/16\n",
        "assert mixture_word([1, 1]) == 5/16\n",
        "fresh_HH = mean^2\n",
        "assert fresh_HH == 1/4\n",
        "print(\"Shared bias: P(HH) =\", mixture_word([1, 1]), \"Cov =\", covariance)\n",
        "print(\"Fresh bias:  P(HH) =\", fresh_HH, \"Cov = 0\")\n",
        "\n",
        "def posterior_and_prediction(s, n):\n",
        "    likelihoods = [w*p^s*(1-p)^(n-s) for p, w in zip(biases, weights)]\n",
        "    posterior = [w/sum(likelihoods) for w in likelihoods]\n",
        "    predictive = sum(w*p for p, w in zip(biases, posterior))\n",
        "    return posterior, predictive\n",
        "\n",
        "assert posterior_and_prediction(2, 2) == ([1/10, 9/10], 7/10)\n",
        "assert posterior_and_prediction(1, 2) == ([1/2, 1/2], 1/2)\n",
        "print(\"After HH:\", posterior_and_prediction(2, 2))\n",
        "print(\"After HT (or TH):\", posterior_and_prediction(1, 2))\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-04",
      "metadata": {},
      "source": [
        "**Prove.** Factor the fresh-bias word probability into one-toss\n",
        "marginals. Identify why this factorization fails for one shared bias.\n",
        "Derive the posterior odds $3^{2s-n}$ for the two-coin example.\n",
        "\n",
        "## 2. Reinforcement and a beta mixture give the same word law\n",
        "\n",
        "An urn starts with $a$ red and $b$ blue balls. Replace each drawn ball\n",
        "and add one of the same color. The product along a word depends only\n",
        "on its number of reds, although the factors occur in different orders.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-05",
      "metadata": {},
      "source": [
        "def urn_word(word, a=1, b=1):\n",
        "    a, b, probability = ZZ(a), ZZ(b), QQ(1)\n",
        "    for bit in word:\n",
        "        probability *= QQ(a if bit else b)/(a+b)\n",
        "        a, b = a+bit, b+1-bit\n",
        "    return probability\n",
        "\n",
        "P = PolynomialRing(QQ, 'p')\n",
        "p = P.gen()\n",
        "\n",
        "def beta_integral(a, b):\n",
        "    # Positive integer parameters suffice for the physical urn.\n",
        "    primitive = (p^(a-1)*(1-p)^(b-1)).integral()\n",
        "    return primitive(1)-primitive(0)\n",
        "\n",
        "def beta_word(s, n, a=1, b=1):\n",
        "    return beta_integral(a+s, b+n-s)/beta_integral(a, b)\n",
        "\n",
        "assert urn_word([1, 1, 1, 0]) == urn_word([1, 0, 1, 1]) == 1/20\n",
        "for a, b in [(1, 1), (2, 3), (4, 2)]:\n",
        "    for word in product([0, 1], repeat=6):\n",
        "        assert urn_word(word, a, b) == beta_word(sum(word), len(word), a, b)\n",
        "print(\"Every length-six word agrees with its beta-mixture integral.\")\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-06",
      "metadata": {},
      "source": [
        "**Prove for arbitrary length.** Group the red and blue numerators in\n",
        "the urn product to obtain\n",
        "$a^{\\overline{s}}b^{\\overline{n-s}}/(a+b)^{\\overline n}$.\n",
        "Evaluate the beta integral by integration by parts. Specify why the\n",
        "factors are rising factorials, unlike sampling without replacement.\n",
        "\n",
        "## 3. A word probability is not a count probability\n",
        "\n",
        "Multiply by $\\binom ns$ to get the law of the number of reds. With\n",
        "one red and one blue initially, the count is uniform on $0,\\ldots,n$.\n",
        "Compare with the binomial law for independent fair tosses.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-07",
      "metadata": {},
      "source": [
        "n = 12\n",
        "fair = [binomial(n, s)/2^n for s in range(n+1)]\n",
        "shared = [binomial(n, s)*mixture_word([1]*s+[0]*(n-s)) for s in range(n+1)]\n",
        "urn = [binomial(n, s)*beta_word(s, n) for s in range(n+1)]\n",
        "assert sum(fair) == sum(shared) == sum(urn) == 1\n",
        "assert all(probability == 1/(n+1) for probability in urn)\n",
        "comparison = graphics_array([\n",
        "    bar_chart(fair, color='#5064a2', title='Independent fair tosses'),\n",
        "    bar_chart(shared, color='#c47a37', title='One shared bias: 1/4 or 3/4'),\n",
        "    bar_chart(urn, color='#278679', title='Pólya urn: uniform mixing'),\n",
        "], nrows=1)\n",
        "comparison\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-08",
      "metadata": {},
      "source": [
        "## 4. Sampling noise shrinks; variation between runs persists\n",
        "\n",
        "For conditionally i.i.d. Bernoulli trials,\n",
        "$\\operatorname{Cov}(X_i,X_j)=\\operatorname{Var}(\\Theta)$ and\n",
        "$\\operatorname{Var}(S_n/n)=E[\\Theta(1-\\Theta)]/n+\\operatorname{Var}(\\Theta)$.\n",
        "Compute the variance from the complete count distribution as an\n",
        "independent check of that decomposition.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-09",
      "metadata": {},
      "source": [
        "def variance_of_frequency(probabilities):\n",
        "    n = len(probabilities)-1\n",
        "    mean = sum(QQ(s)/n*w for s, w in enumerate(probabilities))\n",
        "    return sum((QQ(s)/n-mean)^2*w for s, w in enumerate(probabilities))\n",
        "\n",
        "assert variance_of_frequency(fair) == 1/(4*n)\n",
        "assert variance_of_frequency(shared) == 3/(16*n)+1/16\n",
        "assert variance_of_frequency(urn) == 1/(6*n)+1/12\n",
        "print(table([(n, 1/(4*n), 3/(16*n)+1/16, 1/(6*n)+1/12)\n",
        "             for n in [10, 100, 1000]],\n",
        "            header_row=[\"n\", \"Fair variance\", \"Shared-coin variance\", \"Urn variance\"]))\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-10",
      "metadata": {},
      "source": [
        "**Predict before simulating.** Across many independent runs, where\n",
        "should the final red frequencies lie in each model? How does this\n",
        "differ from following just one long run? The following experiment\n",
        "uses a fixed seed for reproducibility. Change `draws` to 5000 for\n",
        "the experiment proposed in the book.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-11",
      "metadata": {},
      "source": [
        "rng = random.Random(int(20260910))\n",
        "draws, runs = 2000, 100\n",
        "\n",
        "def simulate_run(model, draws):\n",
        "    reds = 0\n",
        "    bias_quarters = rng.choice([1, 3]) if model == 'shared' else 2\n",
        "    trace = []\n",
        "    for j in range(draws):\n",
        "        bit = (rng.randrange(j+2) < reds+1 if model == 'urn'\n",
        "               else rng.randrange(4) < bias_quarters)\n",
        "        reds += int(bit)\n",
        "        trace.append(float(reds)/float(j+1))\n",
        "    return trace\n",
        "\n",
        "simulations = {model: [simulate_run(model, draws) for _ in range(runs)]\n",
        "               for model in ['fair', 'shared', 'urn']}\n",
        "assert all(0 <= path[-1] <= 1 for paths in simulations.values() for path in paths)\n",
        "histograms = graphics_array([\n",
        "    histogram([path[-1] for path in simulations[model]], bins=[j/20 for j in range(21)],\n",
        "              title=title, color=color)\n",
        "    for model, title, color in [('fair', 'Independent fair tosses', '#5064a2'),\n",
        "                                ('shared', 'Shared coin', '#c47a37'),\n",
        "                                ('urn', 'Pólya urn', '#278679')]\n",
        "], nrows=1)\n",
        "histograms\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "code",
      "id": "sage-12",
      "metadata": {},
      "source": [
        "trajectories = sum(list_plot(list(enumerate(path, start=1)), plotjoined=True,\n",
        "                             color=color, alpha=0.8)\n",
        "                   for path, color in zip(simulations['urn'][:5],\n",
        "                                          ['#278679', '#5064a2', '#c47a37', '#b45368', '#77649c']))\n",
        "trajectories.axes_labels(['draw', 'red frequency'])\n",
        "trajectories\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-13",
      "metadata": {},
      "source": [
        "## 5. Finite exchangeability can have negative covariance\n",
        "\n",
        "An always-disagreeing pair is exchangeable but cannot be extended to\n",
        "an exchangeable triple. More generally, sampling without replacement\n",
        "from $M$ red and $N-M$ blue balls gives covariance\n",
        "$-p(1-p)/(N-1)$, where $p=M/N$.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-14",
      "metadata": {},
      "source": [
        "N, M = 20, 8\n",
        "p_red = QQ(M)/N\n",
        "p_two_red = QQ(M*(M-1))/(N*(N-1))\n",
        "assert p_two_red-p_red^2 == -p_red*(1-p_red)/(N-1) < 0\n",
        "assert QQ(0)-QQ(1)/4 == -1/4  # the always-disagreeing pair\n",
        "print(\"Without-replacement covariance:\", p_two_red-p_red^2)\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-15",
      "metadata": {},
      "source": [
        "**Prove the obstruction.** An infinite Bernoulli mixture has\n",
        "nonnegative pairwise covariance. Give a second proof for the\n",
        "always-disagreeing pair: can all three pairs of three bits disagree?\n",
        "\n",
        "## 6. The finite bridge and the birthday collision bound\n",
        "\n",
        "Conditional on its total, a finite exchangeable binary vector is a\n",
        "uniformly shuffled population. Compare its first $n$ draws with\n",
        "draws made with replacement. Total variation is half the sum of the\n",
        "absolute differences of the count probabilities: within each count\n",
        "both models assign equal probabilities to all ordered words, so this\n",
        "also equals total variation for the ordered color strings.\n"
      ]
    },
    {
      "cell_type": "code",
      "id": "sage-16",
      "metadata": {},
      "source": [
        "def sampling_comparison(N, M, n):\n",
        "    if not (0 <= M <= N and 1 <= n <= N):\n",
        "        raise ValueError(\"Require 0 <= M <= N and 1 <= n <= N.\")\n",
        "    p = QQ(M)/N\n",
        "    without = [binomial(M, s)*binomial(N-M, n-s)/binomial(N, n) for s in range(n+1)]\n",
        "    with_replacement = [binomial(n, s)*p^s*(1-p)^(n-s) for s in range(n+1)]\n",
        "    tv = sum(abs(a-b) for a, b in zip(without, with_replacement))/2\n",
        "    collision = 1-prod(QQ(N-j)/N for j in range(n))\n",
        "    bound = min(QQ(1), QQ(binomial(n, 2))/N)\n",
        "    assert sum(without) == sum(with_replacement) == 1\n",
        "    assert tv <= collision <= bound\n",
        "    return tv, collision, bound\n",
        "\n",
        "for N, M, n in [(10, 4, 3), (100, 40, 5), (1000, 400, 5), (1000, 400, 100)]:\n",
        "    tv, collision, bound = sampling_comparison(N, M, n)\n",
        "    print(\"N, M, n =\", (N, M, n), \"TV ≈\", RR(tv),\n",
        "          \"collision ≈\", RR(collision), \"union bound =\", bound)\n",
        "assert sampling_comparison(1000, 400, 5)[2] == 1/100\n"
      ],
      "execution_count": null,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "id": "sage-17",
      "metadata": {},
      "source": [
        "**Finish the proof outline.** Explain how the coupling bounds the\n",
        "error for **any event**, not just one word. For a fixed prefix length\n",
        "$n$, why does the bound tend to zero as the population length $N$\n",
        "increases? State the compactness input still needed to construct the\n",
        "mixing measure. A finite computation cannot supply that analytic step.\n",
        "\n",
        "## Further investigations\n",
        "\n",
        "1. Derive the Beta$(a+s,b+n-s)$ posterior and its predictive mean.\n",
        "2. Prove that pairwise independence of an **infinite exchangeable**\n",
        "   Bernoulli sequence implies full independence.\n",
        "3. Use Bernstein polynomials to connect the law of $S_n/n$ with\n",
        "   uniqueness of the mixing measure, following the book's proof.\n",
        "\n",
        "References: course Chapter 10; Werner Kirsch,\n",
        "[An elementary proof of de Finetti's Theorem](https://arxiv.org/abs/1809.00882);\n",
        "Diaconis and Freedman, [Finite Exchangeable Sequences](https://doi.org/10.1214/aop/1176994663).\n"
      ]
    }
  ],
  "metadata": {
    "kernelspec": {
      "display_name": "SageMath",
      "language": "sage",
      "name": "sagemath"
    },
    "language_info": {
      "name": "sage",
      "file_extension": ".sage",
      "mimetype": "text/x-sage"
    }
  },
  "nbformat": 4,
  "nbformat_minor": 5
}
