Skip to content
Open
Show file tree
Hide file tree
Changes from 3 commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
109 changes: 109 additions & 0 deletions .github/jupyterlite/pytensor_numba.ipynb
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
{
"cells": [
{
"cell_type": "markdown",
"id": "7fb27b941602401d91542211134fc71a",
"metadata": {},
"source": [
"# PyTensor and Numba in JupyterLite\n",
"\n",
"This example is adapted from PyTensor's introduction notebook and runs with the Numba linker entirely in the browser."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "acae54e37e7d407bbb7b55eff062a284",
"metadata": {},
"outputs": [],
"source": [
"import sys\n",
"from importlib.metadata import version\n",
"\n",
"import numba\n",
"import numpy as np\n",
"\n",
"import pytensor\n",
"import pytensor.tensor as pt\n",
"\n",
"\n",
"assert sys.platform == \"emscripten\"\n",
"{\n",
" \"Platform\": sys.platform,\n",
" \"PyTensor\": version(\"pytensor\"),\n",
" \"Numba\": numba.__version__,\n",
" \"NumPy\": np.__version__,\n",
"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "9a63283cbaf04dbcab1f6479b197f3a8",
"metadata": {},
"outputs": [],
"source": [
"x = pt.vector(\"x\", shape=(None,))\n",
"z = pt.exp(pt.sin(x))\n",
"out = pt.cos((z[None, :] @ z[:, None]).squeeze())\n",
"\n",
"out.dprint()"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "8dd0d8092fe74a7c96281538738b07e2",
"metadata": {},
"outputs": [],
"source": [
"numba_fn = pytensor.function([x], out, mode=\"NUMBA\")\n",
"numba_fn.dprint(print_destroy_map=True)"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "72eea5119410473aa328ad9291626812",
"metadata": {},
"outputs": [],
"source": [
"values = np.array([0.25, -0.5, 1.0])\n",
"result = numba_fn(values)\n",
"expected_z = np.exp(np.sin(values))\n",
"expected = np.cos(expected_z @ expected_z)\n",
"\n",
"np.testing.assert_allclose(result, expected)\n",
"{\"Result\": result, \"Expected\": expected}"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "8edb47106e1a46a883d545849b8ab81b",
"metadata": {},
"outputs": [],
"source": [
"# Reuse the compiled function with a different vector length.\n",
"values = np.linspace(-1.0, 1.0, 5)\n",
"result = numba_fn(values)\n",
"expected_z = np.exp(np.sin(values))\n",
"np.testing.assert_allclose(result, np.cos(expected_z @ expected_z))\n",
"result"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python (XPython)",
"language": "python",
"name": "xpython"
},
"language_info": {
"name": "python",
"version": "3.13"
}
},
"nbformat": 4,
"nbformat_minor": 5
}
79 changes: 79 additions & 0 deletions .github/workflows/deploy-jupyterlite.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
name: Build and deploy JupyterLite

on:
workflow_dispatch:
pull_request:
push:
branches:
- main

permissions:
contents: read
pages: write
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
id-token: write
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed

jobs:
build:
runs-on: ubuntu-latest

steps:
- uses: actions/checkout@v5
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
with:
fetch-depth: 0
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed

- name: Install build environment
uses: mamba-org/setup-micromamba@v2
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
with:
environment-file: environment-wasm-build.yml
environment-name: pytensor-wasm-build
init-shell: bash

- name: Build PyTensor from this checkout
shell: bash -l {0}
run: |
set -euxo pipefail
PYTENSOR_PURE_PYTHON=1 python -m build --wheel --no-isolation

mkdir -p build/pytensor-wheel
wheels=(dist/pytensor-*.whl)
test "${#wheels[@]}" -eq 1
python -m zipfile -e "${wheels[0]}" build/pytensor-wheel

- name: Create the WebAssembly environment
shell: bash -l {0}
run: |
set -euxo pipefail
micromamba create -y \
-f environment-wasm-host.yml \
--platform=emscripten-wasm32

echo "PYTENSOR_WASM_PREFIX=${MAMBA_ROOT_PREFIX}/envs/pytensor-wasm-host" >> "${GITHUB_ENV}"

- name: Build JupyterLite
shell: bash -l {0}
run: |
set -euxo pipefail
jupyter lite build \
--XeusAddon.prefix="${PYTENSOR_WASM_PREFIX}" \
--XeusAddon.mounts="$(pwd)/build/pytensor-wheel:/lib/python3.13/site-packages" \
--contents=.github/jupyterlite/pytensor_numba.ipynb \
--output-dir=dist-jupyterlite

- name: Upload Pages artifact
uses: actions/upload-pages-artifact@v4
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
with:
path: dist-jupyterlite

deploy:
needs: build
if: github.ref == 'refs/heads/main'
runs-on: ubuntu-latest

environment:
name: github-pages
url: ${{ steps.deployment.outputs.page_url }}

steps:
- name: Deploy to GitHub Pages
id: deployment
uses: actions/deploy-pages@v4
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
15 changes: 15 additions & 0 deletions environment-wasm-build.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
name: pytensor-wasm-build

channels:
- conda-forge

dependencies:
- python=3.13
- python-build
- setuptools
- cython
- numpy=2.4.*
- versioneer=0.29
- jupyterlite-core
- jupyterlite-xeus

19 changes: 19 additions & 0 deletions environment-wasm-host.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
name: pytensor-wasm-host

channels:
- https://prefix.dev/emscripten-forge-4x
- https://prefix.dev/conda-forge

dependencies:
- python=3.13
- xeus-python
- numpy=2.4.*
- scipy
- numba=0.66.*
- llvmlite=0.48.*
- setuptools
- filelock
- etuples
- logical-unification
- minikanren
- cons
17 changes: 6 additions & 11 deletions pytensor/link/numba/dispatch/vectorize_codegen.py
Original file line number Diff line number Diff line change
Expand Up @@ -307,7 +307,7 @@ def compute_itershape(
broadcast_pattern: tuple[tuple[bool, ...], ...],
size: list[ir.Instruction] | None,
):
one = ir.IntType(64)(1)
one = ctx.get_constant(types.intp, 1)
batch_ndim = len(broadcast_pattern[0])
shape = [None] * batch_ndim
if size is not None:
Expand Down Expand Up @@ -417,7 +417,7 @@ def make_outputs(
"""
output_arrays = []
output_arry_types = []
one = ir.IntType(64)(1)
one = ctx.get_constant(types.intp, 1)
inplace_dict = dict(inplace)
for i, (core_shape, bc, dtype) in enumerate(
zip(output_core_shapes, out_bc, dtypes, strict=True)
Expand Down Expand Up @@ -448,7 +448,7 @@ def make_outputs(
# reduction identity. A flat scan over every element is valid
# regardless of which axes are reduced, and seeds each kept-axis
# accumulator cell (size-1 reduced axes included).
nitems = ir.IntType(64)(1)
nitems = ctx.get_constant(types.intp, 1)
for dim_len in shape:
nitems = builder.mul(nitems, dim_len)
ident = ctx.get_constant(dtype, reduce_identities[i])
Expand Down Expand Up @@ -522,7 +522,7 @@ def make_loop_call(
]
destroyed_inputs = {in_idx: out_idx for out_idx, in_idx in inplace}

zero = ir.Constant(ir.IntType(64), 0)
zero = context.get_constant(types.intp, 0)

def _wrap_negative_index(idx_val, dim_size, signed):
"""Wrap a negative index by adding the dimension size: idx + size if idx < 0.
Expand Down Expand Up @@ -580,12 +580,7 @@ def _wrap_negative_index(idx_val, dim_size, signed):
val = builder.load(ptr)
val.set_metadata("alias.scope", input_scope_set)
val.set_metadata("noalias", output_scope_set)
i64 = ir.IntType(64)
if val.type != i64:
if idx_arr_type.dtype.signed:
val = builder.sext(val, i64)
else:
val = builder.zext(val, i64)
val = context.cast(builder, val, idx_arr_type.dtype, types.intp)
indirect_idxs.append(val)

# Load values from input arrays
Expand Down Expand Up @@ -1076,7 +1071,7 @@ def codegen(ctx, builder, sig, args):
cgutils.unpack_tuple(builder, idx_arrs[k].shape) for k in range(n_indices)
]

one = ir.IntType(64)(1)
one = ctx.get_constant(types.intp, 1)
iter_shapes = list(in_shapes)
iter_bc = list(input_bc_patterns)

Expand Down
14 changes: 8 additions & 6 deletions setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,12 +13,14 @@

NAME: str = dist.get_name() # type: ignore

# Check if building for Pyodide
is_pyodide = os.getenv("PYODIDE", "0") == "1"

if is_pyodide:
# For pyodide we build a universal wheel that must be pure-python
# so we must omit the cython-version of scan.
# Build without optional compiled extensions. Keep PYODIDE as a compatibility
# alias for existing downstream builds.
is_pure_python = (
os.getenv("PYTENSOR_PURE_PYTHON", "0") == "1" or os.getenv("PYODIDE", "0") == "1"
)

if is_pure_python:
# Omit the optional Cython implementation of scan.
ext_modules = []
else:
ext_modules = [
Expand Down
40 changes: 40 additions & 0 deletions tests/link/numba/test_vectorize_codegen.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
from llvmlite import ir
from numba import types

from pytensor.link.numba.dispatch.vectorize_codegen import compute_itershape


class Mock32BitContext:
"""Minimal Numba context whose ``intp`` LLVM representation is i32."""

class CallConv:
@staticmethod
def return_user_exc(builder, exc, args):
pass

call_conv = CallConv()

@staticmethod
def get_constant(typ, value):
assert typ is types.intp
return ir.Constant(ir.IntType(32), value)


def test_compute_itershape_uses_target_intp_width():
module = ir.Module()
function_type = ir.FunctionType(ir.VoidType(), ())
function = ir.Function(module, function_type, "compute_itershape")
builder = ir.IRBuilder(function.append_basic_block("entry"))

one_i32 = ir.Constant(ir.IntType(32), 1)
shape = compute_itershape(
Mock32BitContext(),
builder,
in_shapes=[[one_i32]],
broadcast_pattern=((True,),),
size=None,
)
builder.ret_void()

assert shape == [one_i32]
assert "icmp ne i32" in str(module)
Loading