{
  "cells": [
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "# Processing large HSI cubes without exhausting memory\n",
        "\n",
        "A synthetic, NumPy-only equivalence demonstration. This notebook tests a finite-radius model, not an HSI-trained classifier. It uses no external data and runs on CPU. The separate benchmark.py performs fresh-process Linux RSS measurements.\n"
      ]
    },
    {
      "cell_type": "code",
      "metadata": {},
      "source": [
        "\"\"\"Finite-context HSI teaching harness. NumPy only; no training or real scene data.\n",
        "\n",
        "Axis contract: input H,W,B float32; output ceil(H/s),ceil(W/s),3 logits.\n",
        "Two radius-2 box means with available-neighbour normalization give radius 4.\n",
        "Output decimation is applied AFTER the two filters; it is not a generic CNN tiler.\n",
        "\"\"\"\n",
        "from __future__ import annotations\n",
        "from dataclasses import dataclass\n",
        "from pathlib import Path\n",
        "import numpy as np\n",
        "\n",
        "RADIUS = 4\n",
        "ATOL = 2e-6\n",
        "RTOL = 1e-5\n",
        "\n",
        "\n",
        "def synthetic_window(y0, y1, x0, x1, bands=32):\n",
        "    \"\"\"Coordinate-defined finite fixture: identical values for any read partition.\"\"\"\n",
        "    y = np.arange(y0, y1, dtype=np.float32)[:, None, None]\n",
        "    x = np.arange(x0, x1, dtype=np.float32)[None, :, None]\n",
        "    b = np.arange(bands, dtype=np.float32)[None, None, :]\n",
        "    z = (0.48 + 0.19*np.sin(y/7+b/9) + 0.14*np.cos(x/11-b/8)\n",
        "         + 0.06*np.sin((x+y)/3+b/4)\n",
        "         + 0.15*((x % 31 < 9) & (y % 29 < 13)))\n",
        "    return z.astype(np.float32)\n",
        "\n",
        "\n",
        "def write_fixture(path, shape=(768, 1024, 32), row_chunk=32):\n",
        "    \"\"\"Bounded construction; do not materialize the complete cube first.\"\"\"\n",
        "    arr = np.lib.format.open_memmap(path, mode='w+', dtype=np.float32, shape=shape)\n",
        "    for y in range(0, shape[0], row_chunk):\n",
        "        arr[y:y+row_chunk] = synthetic_window(y, min(y+row_chunk, shape[0]), 0, shape[1], shape[2])\n",
        "    arr.flush()\n",
        "    del arr\n",
        "\n",
        "\n",
        "def box_mean_available(a, radius=2):\n",
        "    \"\"\"Per-layer true-edge rule: mean over available in-domain neighbours.\n",
        "\n",
        "    Artificial tile edges use this same rule but are trimmed away with halo >= 4.\n",
        "    This function is not a nodata/masked-data policy and rejects NaNs upstream.\n",
        "    \"\"\"\n",
        "    h,w,_ = a.shape\n",
        "    p = np.pad(a, ((radius,radius),(radius,radius),(0,0)), mode='constant')\n",
        "    total = np.zeros_like(a)\n",
        "    for dy in range(2*radius+1):\n",
        "        for dx in range(2*radius+1):\n",
        "            total += p[dy:dy+h, dx:dx+w]\n",
        "    ys = np.arange(h); xs = np.arange(w)\n",
        "    cy = np.minimum(h, ys+radius+1) - np.maximum(0, ys-radius)\n",
        "    cx = np.minimum(w, xs+radius+1) - np.maximum(0, xs-radius)\n",
        "    total /= (cy[:,None]*cx[None,:]).astype(np.float32)[...,None]\n",
        "    return total\n",
        "\n",
        "\n",
        "def dense_model(cube):\n",
        "    \"\"\"Fixed coefficients; 4 hidden channels, 3 output logits; no fitted weights.\"\"\"\n",
        "    a = np.asarray(cube, dtype=np.float32)\n",
        "    if a.ndim != 3 or min(a.shape) < 1 or not np.isfinite(a).all():\n",
        "        raise ValueError('Expected a nonempty finite H,W,B cube')\n",
        "    bands = a.shape[-1]\n",
        "    b = np.arange(bands, dtype=np.float32)\n",
        "    # Fixed synthetic constants, not statistics estimated from a tile or scene.\n",
        "    center = 0.45 + 0.025*np.sin(b/5)\n",
        "    scale = 0.25 + 0.015*np.cos(b/7)\n",
        "    z = (a-center)/scale\n",
        "    w = np.stack([np.cos(b/3), np.sin(b/4), np.cos(b/7+1), np.ones_like(b)], axis=1)/bands\n",
        "    hidden = np.einsum('hwb,bk->hwk', z, w, optimize=False, dtype=np.float32)\n",
        "    hidden = np.tanh(box_mean_available(hidden))\n",
        "    hidden = np.tanh(box_mean_available(hidden))\n",
        "    head = np.array([[1.1,-0.6,0.3],[-0.2,0.9,0.5],[0.4,0.2,-1.0],[0.6,-0.2,0.4]], dtype=np.float32)\n",
        "    return np.einsum('hwk,kc->hwc', hidden, head, optimize=False, dtype=np.float32)\n",
        "\n",
        "\n",
        "def eager(cube, stride=2):\n",
        "    if stride < 1: raise ValueError('stride must be positive')\n",
        "    return dense_model(cube)[::stride,::stride].copy()\n",
        "\n",
        "\n",
        "@dataclass(frozen=True)\n",
        "class Tile:\n",
        "    y0: int; y1: int; x0: int; x1: int\n",
        "    ry0: int; ry1: int; rx0: int; rx1: int\n",
        "\n",
        "\n",
        "def tile_plan(height, width, core=96, halo=4, stride=2, align=True):\n",
        "    if min(height,width,core,stride) < 1 or halo < 0:\n",
        "        raise ValueError('Invalid geometry')\n",
        "    if core % stride:\n",
        "        raise ValueError('core must be a multiple of output stride')\n",
        "    for y0 in range(0,height,core):\n",
        "        for x0 in range(0,width,core):\n",
        "            y1=min(y0+core,height); x1=min(x0+core,width)\n",
        "            ry0=max(0,y0-halo); rx0=max(0,x0-halo)\n",
        "            if align:\n",
        "                ry0=(ry0//stride)*stride; rx0=(rx0//stride)*stride\n",
        "            yield Tile(y0,y1,x0,x1,ry0,min(height,y1+halo),rx0,min(width,x1+halo))\n",
        "\n",
        "\n",
        "def tiled(reader, writer, shape, core=96, halo=4, stride=2, align=True):\n",
        "    \"\"\"reader(yslice,xslice) -> H,W,B; writer(yslice,xslice, logits).\n",
        "\n",
        "    align=False deliberately demonstrates a bad local-output phase: it labels\n",
        "    local stride samples as though they were anchored to the global origin.\n",
        "    Do not use align=False for inference.\n",
        "    \"\"\"\n",
        "    stats={'tiles':0,'read_pixels':0,'max_read_shape':[0,0,shape[2]],'max_read_bytes':0}\n",
        "    for t in tile_plan(*shape[:2],core,halo,stride,align):\n",
        "        patch=reader(slice(t.ry0,t.ry1),slice(t.rx0,t.rx1))\n",
        "        local=eager(patch,stride)\n",
        "        # Ceil ensures the first selected local sample is at/after the core start.\n",
        "        offset = stride-1 if align else 0\n",
        "        ly=(t.y0-t.ry0+offset)//stride; lx=(t.x0-t.rx0+offset)//stride\n",
        "        ny=(t.y1-t.y0+stride-1)//stride; nx=(t.x1-t.x0+stride-1)//stride\n",
        "        kept=local[ly:ly+ny,lx:lx+nx]\n",
        "        if kept.shape != (ny,nx,3): raise ValueError('Tile output shape mismatch')\n",
        "        writer(slice(t.y0//stride,t.y0//stride+ny),slice(t.x0//stride,t.x0//stride+nx),kept)\n",
        "        stats['tiles']+=1; stats['read_pixels']+=patch.shape[0]*patch.shape[1]\n",
        "        if patch.nbytes > stats['max_read_bytes']:\n",
        "            stats['max_read_bytes']=int(patch.nbytes); stats['max_read_shape']=list(patch.shape)\n",
        "    stats['read_amplification']=stats['read_pixels']/(shape[0]*shape[1])\n",
        "    return stats\n",
        "\n",
        "\n",
        "def tiled_array(cube, **kwargs):\n",
        "    s=kwargs.get('stride',2)\n",
        "    out=np.empty(((cube.shape[0]+s-1)//s,(cube.shape[1]+s-1)//s,3),np.float32)\n",
        "    tiled(lambda y,x:cube[y,x,:],lambda y,x,v:out.__setitem__((y,x,slice(None)),v),cube.shape,**kwargs)\n",
        "    return out\n",
        "\n",
        "\n",
        "def compare(reference, actual):\n",
        "    err=np.abs(actual-reference)\n",
        "    return {'max_abs_logit_error':float(err.max()),'mean_abs_logit_error':float(err.mean()),\n",
        "            'label_disagreement_fraction':float(np.mean(actual.argmax(-1)!=reference.argmax(-1))),\n",
        "            'allclose':bool(np.allclose(actual,reference,rtol=RTOL,atol=ATOL)),\n",
        "            'atol':ATOL,'rtol':RTOL}\n"
      ],
      "execution_count": 1,
      "outputs": []
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## Compare halo sizes\n",
        "\n",
        "The nominal dense receptive radius is four. The stride-two alignment can expand a requested left/top halo; do not infer a universally minimal halo from one decimated fixture."
      ]
    },
    {
      "cell_type": "code",
      "metadata": {},
      "source": [
        "cube = synthetic_window(0, 67, 0, 83, 16)\n",
        "reference = eager(cube)\n",
        "for halo in [0, 1, 2, 4, 5, 8]:\n",
        "    report = compare(reference, tiled_array(cube, core=16, halo=halo))\n",
        "    print(halo, report)\n",
        "assert compare(reference, tiled_array(cube, core=16, halo=4))['allclose']\n",
        "assert not compare(reference, tiled_array(cube, core=16, halo=2))['allclose']\n"
      ],
      "execution_count": 2,
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "0 {'max_abs_logit_error': 0.20316123962402344, 'mean_abs_logit_error': 0.01648474670946598, 'label_disagreement_fraction': 0.03431372549019608, 'allclose': False, 'atol': 2e-06, 'rtol': 1e-05}\n",
            "1 {'max_abs_logit_error': 0.09822648763656616, 'mean_abs_logit_error': 0.0040090144611895084, 'label_disagreement_fraction': 0.007703081232492998, 'allclose': False, 'atol': 2e-06, 'rtol': 1e-05}\n",
            "2 {'max_abs_logit_error': 0.09822648763656616, 'mean_abs_logit_error': 0.0028485713992267847, 'label_disagreement_fraction': 0.0063025210084033615, 'allclose': False, 'atol': 2e-06, 'rtol': 1e-05}\n",
            "4 {'max_abs_logit_error': 0.0, 'mean_abs_logit_error': 0.0, 'label_disagreement_fraction': 0.0, 'allclose': True, 'atol': 2e-06, 'rtol': 1e-05}\n",
            "5 {'max_abs_logit_error': 0.0, 'mean_abs_logit_error': 0.0, 'label_disagreement_fraction': 0.0, 'allclose': True, 'atol': 2e-06, 'rtol': 1e-05}\n",
            "8 {'max_abs_logit_error': 0.0, 'mean_abs_logit_error': 0.0, 'label_disagreement_fraction': 0.0, 'allclose': True, 'atol': 2e-06, 'rtol': 1e-05}\n"
          ]
        }
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## Break the output-grid phase\n",
        "\n",
        "Resetting the stride origin at an unaligned read window changes which pixels the model samples."
      ]
    },
    {
      "cell_type": "code",
      "metadata": {},
      "source": [
        "bad = tiled_array(cube, core=16, halo=5, align=False)\n",
        "print(compare(reference, bad))\n",
        "assert not compare(reference, bad)['allclose']\n"
      ],
      "execution_count": 3,
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "{'max_abs_logit_error': 0.20237118005752563, 'mean_abs_logit_error': 0.03298979252576828, 'label_disagreement_fraction': 0.06862745098039216, 'allclose': False, 'atol': 2e-06, 'rtol': 1e-05}\n"
          ]
        }
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## Use file-backed input and output\n",
        "\n",
        "The operation allocates tile-sized work arrays. Mapped pages still use RAM, and this small example is not a memory benchmark."
      ]
    },
    {
      "cell_type": "code",
      "metadata": {},
      "source": [
        "import tempfile\n",
        "with tempfile.TemporaryDirectory() as folder:\n",
        "    source = Path(folder) / 'cube.npy'\n",
        "    write_fixture(source, shape=(67,83,16), row_chunk=11)\n",
        "    mapped = np.load(source, mmap_mode='r')\n",
        "    output = np.lib.format.open_memmap(Path(folder) / 'logits.npy', mode='w+', dtype=np.float32, shape=(34,42,3))\n",
        "    stats = tiled(lambda y,x: mapped[y,x,:], lambda y,x,v: output.__setitem__((y,x,slice(None)),v), mapped.shape, core=16, halo=4, stride=2)\n",
        "    output.flush()\n",
        "    print(stats)\n",
        "    np.testing.assert_allclose(output, reference, rtol=RTOL, atol=ATOL)\n",
        "    del mapped, output\n"
      ],
      "execution_count": 4,
      "outputs": [
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "{'tiles': 30, 'read_pixels': 11956, 'max_read_shape': [24, 24, 16], 'max_read_bytes': 36864, 'read_amplification': 2.1499730264340946}\n"
          ]
        }
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {},
      "source": [
        "## Scope\n",
        "\n",
        "Global attention, global/spatial normalization, recursive state, nodata and geospatial output alignment require their own contracts. Finite overlap alone cannot generally reproduce them. See the article and SOURCES.md for primary documentation.\n"
      ]
    }
  ],
  "metadata": {
    "kernelspec": {
      "display_name": "Python 3",
      "language": "python",
      "name": "python3"
    },
    "language_info": {
      "name": "python",
      "version": "3.12.14"
    }
  },
  "nbformat": 4,
  "nbformat_minor": 5
}
