entangle-jax 0.1.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- entangle_jax-0.1.0/.github/workflows/publish.yml +44 -0
- entangle_jax-0.1.0/.github/workflows/tests.yml +27 -0
- entangle_jax-0.1.0/.gitignore +10 -0
- entangle_jax-0.1.0/LICENSE +21 -0
- entangle_jax-0.1.0/PKG-INFO +108 -0
- entangle_jax-0.1.0/README.md +96 -0
- entangle_jax-0.1.0/meson.build +10 -0
- entangle_jax-0.1.0/pyproject.toml +48 -0
- entangle_jax-0.1.0/src/entangle_jax/__init__.py +15 -0
- entangle_jax-0.1.0/src/entangle_jax/_entangle_ffi.cc +30 -0
- entangle_jax-0.1.0/src/entangle_jax/_entangle_ffi.h +3 -0
- entangle_jax-0.1.0/src/entangle_jax/_ffi.pyi +6 -0
- entangle_jax-0.1.0/src/entangle_jax/_ffi.pyx +47 -0
- entangle_jax-0.1.0/src/entangle_jax/_primitive.py +91 -0
- entangle_jax-0.1.0/src/entangle_jax/_testing.py +33 -0
- entangle_jax-0.1.0/src/entangle_jax/_testing_ffi.cc +54 -0
- entangle_jax-0.1.0/src/entangle_jax/_testing_ffi.h +7 -0
- entangle_jax-0.1.0/src/entangle_jax/meson.build +33 -0
- entangle_jax-0.1.0/src/entangle_jax/py.typed +0 -0
- entangle_jax-0.1.0/tests/test_entangle.py +133 -0
- entangle_jax-0.1.0/uv.lock +780 -0
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
name: Publish
|
|
2
|
+
on:
|
|
3
|
+
release:
|
|
4
|
+
types: [published]
|
|
5
|
+
jobs:
|
|
6
|
+
build_wheels:
|
|
7
|
+
strategy:
|
|
8
|
+
fail-fast: false
|
|
9
|
+
matrix:
|
|
10
|
+
os: [ubuntu-latest, windows-latest, macos-latest]
|
|
11
|
+
runs-on: ${{ matrix.os }}
|
|
12
|
+
steps:
|
|
13
|
+
- uses: actions/checkout@v7
|
|
14
|
+
- uses: pypa/cibuildwheel@v4.1.0
|
|
15
|
+
- uses: actions/upload-artifact@v4
|
|
16
|
+
with:
|
|
17
|
+
name: wheels-${{ matrix.os }}
|
|
18
|
+
path: wheelhouse/*.whl
|
|
19
|
+
|
|
20
|
+
build_sdist:
|
|
21
|
+
runs-on: ubuntu-latest
|
|
22
|
+
steps:
|
|
23
|
+
- uses: actions/checkout@v7
|
|
24
|
+
- uses: astral-sh/setup-uv@v7
|
|
25
|
+
- run: uv sync --group dev
|
|
26
|
+
- run: uv build --sdist
|
|
27
|
+
- uses: actions/upload-artifact@v4
|
|
28
|
+
with:
|
|
29
|
+
name: sdist
|
|
30
|
+
path: dist/*.tar.gz
|
|
31
|
+
|
|
32
|
+
publish:
|
|
33
|
+
needs: [build_wheels, build_sdist]
|
|
34
|
+
runs-on: ubuntu-latest
|
|
35
|
+
environment: pypi
|
|
36
|
+
permissions:
|
|
37
|
+
id-token: write
|
|
38
|
+
steps:
|
|
39
|
+
- uses: actions/download-artifact@v4
|
|
40
|
+
with:
|
|
41
|
+
pattern: "*"
|
|
42
|
+
path: dist/
|
|
43
|
+
merge-multiple: true
|
|
44
|
+
- uses: pypa/gh-action-pypi-publish@release/v1
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
name: Tests
|
|
2
|
+
on:
|
|
3
|
+
pull_request:
|
|
4
|
+
push:
|
|
5
|
+
branches:
|
|
6
|
+
- main
|
|
7
|
+
jobs:
|
|
8
|
+
test:
|
|
9
|
+
runs-on: ubuntu-latest
|
|
10
|
+
steps:
|
|
11
|
+
- uses: actions/checkout@v7
|
|
12
|
+
- uses: astral-sh/setup-uv@v7
|
|
13
|
+
- run: uv sync --group dev
|
|
14
|
+
- run: uv run pytest
|
|
15
|
+
- run: uv run ruff check .
|
|
16
|
+
- run: uv run ruff format --check .
|
|
17
|
+
- run: uv run ty check
|
|
18
|
+
|
|
19
|
+
build_wheels:
|
|
20
|
+
strategy:
|
|
21
|
+
fail-fast: false
|
|
22
|
+
matrix:
|
|
23
|
+
os: [ubuntu-latest, windows-latest, macos-latest]
|
|
24
|
+
runs-on: ${{ matrix.os }}
|
|
25
|
+
steps:
|
|
26
|
+
- uses: actions/checkout@v7
|
|
27
|
+
- uses: pypa/cibuildwheel@v4.1.0
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Nardi Lam
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
|
@@ -0,0 +1,108 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: entangle-jax
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: A generic JAX primitive that is opaque to XLA reordering optimizations
|
|
5
|
+
Author-Email: Nardi Lam <mail@nardilam.nl>
|
|
6
|
+
License-Expression: MIT
|
|
7
|
+
Project-URL: Homepage, https://github.com/nardi/entangle-jax
|
|
8
|
+
Requires-Python: >=3.11
|
|
9
|
+
Requires-Dist: jax>=0.10
|
|
10
|
+
Requires-Dist: jaxtyping>=0.2
|
|
11
|
+
Description-Content-Type: text/markdown
|
|
12
|
+
|
|
13
|
+
# entangle-jax
|
|
14
|
+
|
|
15
|
+
A generic JAX primitive that is opaque to XLA reordering optimizations.
|
|
16
|
+
|
|
17
|
+
`entangle(payload, *witnesses)` returns `payload` unchanged, but ordered under
|
|
18
|
+
`jit` after every witness's producer. Useful for sequencing native
|
|
19
|
+
side-effecting calls that XLA would otherwise be free to reorder, since they
|
|
20
|
+
share no ordinary data dependency.
|
|
21
|
+
|
|
22
|
+
Install with `uv add entangle-jax` or `pip install entangle-jax`.
|
|
23
|
+
|
|
24
|
+
## Why
|
|
25
|
+
|
|
26
|
+
When interacting with external native libraries, it is sometimes necessary to
|
|
27
|
+
interact with objects they manage through opaque pointers. This conflicts with
|
|
28
|
+
the XLA memory model where everything is a static-sized buffer that can be
|
|
29
|
+
arbitrarily copied or reused when needed. One way to achieve this interaction is
|
|
30
|
+
by creating a token that XLA can pass around that references this object, but
|
|
31
|
+
XLA is not aware that e.g. multiple copies of this token in fact reference the
|
|
32
|
+
same memory, and so will perform optimizations that are at worst unsafe and at
|
|
33
|
+
best inefficient.
|
|
34
|
+
|
|
35
|
+
Unfortunately, there is no way to tell XLA that a buffer will be modified
|
|
36
|
+
in-place and that it should maintain consistent order of operations involving
|
|
37
|
+
this buffer. Marking a primitive as side-effecting will stop XLA from removing
|
|
38
|
+
it as dead code, but it will still assume that the input and output buffers are
|
|
39
|
+
distinct, independent objects. At most, it is possible to create an in-place
|
|
40
|
+
modifying primitive by using [input-output aliasing](https://openxla.org/xla/aliasing), in which case the input and
|
|
41
|
+
output buffers are the same, but then XLA will simply copy the input buffer if
|
|
42
|
+
it is used by another function to ensure both calls are safe.
|
|
43
|
+
|
|
44
|
+
To work around this, we can create an artificial data dependency with
|
|
45
|
+
`entangle`. The contents of a custom call are not visible to the XLA optimizer, so
|
|
46
|
+
it has to treat the output `payload` as a new distinct variable from the input
|
|
47
|
+
`payload` that depends on all `witness` values. A standard-library alternative
|
|
48
|
+
would be `jax.lax.optimization_barrier`, but this is currently
|
|
49
|
+
[unreliable on CPU without a non-default XLA flag](https://github.com/openxla/xla/issues/20440).
|
|
50
|
+
|
|
51
|
+
The actual runtime behavior of the primitive is a no-op: it uses aliasing to tell XLA that the input and output `payload` should be the same buffer, so it does literally nothing. This also means it is supported on every platform. The usual pattern is to overwrite the `payload` variable so that its buffer can always be donated, but if it cannot be (for example because it is owned by a caller outside of the JIT context) XLA will insert a copy of the `payload` buffer.
|
|
52
|
+
|
|
53
|
+
## Example
|
|
54
|
+
|
|
55
|
+
```python
|
|
56
|
+
import jax.numpy as jnp
|
|
57
|
+
|
|
58
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
59
|
+
x = some_native_read(token1)
|
|
60
|
+
token2 = modify_token(token1)
|
|
61
|
+
y = some_native_read(token2)
|
|
62
|
+
```
|
|
63
|
+
|
|
64
|
+
Here `token*` is a value that has some library object associated with it, and `modify_token` changes something about this backing object. Even if `modify_token` is marked as side-effecting, `x = some_native_read(token1)` and `token2 = modify_token(token1)` read the same token variable. This means XLA could choose to reorder them, which would lead to the following order:
|
|
65
|
+
|
|
66
|
+
```python
|
|
67
|
+
# Equivalent reordering:
|
|
68
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
69
|
+
token2 = modify_token(token1)
|
|
70
|
+
x = some_native_read(token1)
|
|
71
|
+
y = some_native_read(token2)
|
|
72
|
+
```
|
|
73
|
+
|
|
74
|
+
In this case, both `x` and `y` read the same value. To enforce ordering, we can entangle `x` and the `token` used in its computation, before modifying it.
|
|
75
|
+
|
|
76
|
+
```python
|
|
77
|
+
from entangle_jax import entangle
|
|
78
|
+
|
|
79
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
80
|
+
x = some_native_read(token1)
|
|
81
|
+
token2 = entangle(token1, x)
|
|
82
|
+
token3 = modify_token(token2)
|
|
83
|
+
y = some_native_read(token3)
|
|
84
|
+
```
|
|
85
|
+
|
|
86
|
+
This means that `modify_token` cannot be reordered to be before the line that determines `x`, since it uses a value that depends on `x`.
|
|
87
|
+
|
|
88
|
+
```python
|
|
89
|
+
# Invalid reordering:
|
|
90
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
91
|
+
token2 = entangle(token1, x) # x doesn't exist yet!
|
|
92
|
+
token3 = modify_token(token2)
|
|
93
|
+
x = some_native_read(token1)
|
|
94
|
+
y = some_native_read(token3)
|
|
95
|
+
```
|
|
96
|
+
|
|
97
|
+
## Development
|
|
98
|
+
|
|
99
|
+
```bash
|
|
100
|
+
uv sync
|
|
101
|
+
uv run pytest
|
|
102
|
+
uv run ruff check .
|
|
103
|
+
uv run ruff format --check .
|
|
104
|
+
uv run ty check
|
|
105
|
+
```
|
|
106
|
+
|
|
107
|
+
The compiled extension rebuilds automatically on import while developing, so the
|
|
108
|
+
default editable install (`uv sync`) is what you want day to day.
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
# entangle-jax
|
|
2
|
+
|
|
3
|
+
A generic JAX primitive that is opaque to XLA reordering optimizations.
|
|
4
|
+
|
|
5
|
+
`entangle(payload, *witnesses)` returns `payload` unchanged, but ordered under
|
|
6
|
+
`jit` after every witness's producer. Useful for sequencing native
|
|
7
|
+
side-effecting calls that XLA would otherwise be free to reorder, since they
|
|
8
|
+
share no ordinary data dependency.
|
|
9
|
+
|
|
10
|
+
Install with `uv add entangle-jax` or `pip install entangle-jax`.
|
|
11
|
+
|
|
12
|
+
## Why
|
|
13
|
+
|
|
14
|
+
When interacting with external native libraries, it is sometimes necessary to
|
|
15
|
+
interact with objects they manage through opaque pointers. This conflicts with
|
|
16
|
+
the XLA memory model where everything is a static-sized buffer that can be
|
|
17
|
+
arbitrarily copied or reused when needed. One way to achieve this interaction is
|
|
18
|
+
by creating a token that XLA can pass around that references this object, but
|
|
19
|
+
XLA is not aware that e.g. multiple copies of this token in fact reference the
|
|
20
|
+
same memory, and so will perform optimizations that are at worst unsafe and at
|
|
21
|
+
best inefficient.
|
|
22
|
+
|
|
23
|
+
Unfortunately, there is no way to tell XLA that a buffer will be modified
|
|
24
|
+
in-place and that it should maintain consistent order of operations involving
|
|
25
|
+
this buffer. Marking a primitive as side-effecting will stop XLA from removing
|
|
26
|
+
it as dead code, but it will still assume that the input and output buffers are
|
|
27
|
+
distinct, independent objects. At most, it is possible to create an in-place
|
|
28
|
+
modifying primitive by using [input-output aliasing](https://openxla.org/xla/aliasing), in which case the input and
|
|
29
|
+
output buffers are the same, but then XLA will simply copy the input buffer if
|
|
30
|
+
it is used by another function to ensure both calls are safe.
|
|
31
|
+
|
|
32
|
+
To work around this, we can create an artificial data dependency with
|
|
33
|
+
`entangle`. The contents of a custom call are not visible to the XLA optimizer, so
|
|
34
|
+
it has to treat the output `payload` as a new distinct variable from the input
|
|
35
|
+
`payload` that depends on all `witness` values. A standard-library alternative
|
|
36
|
+
would be `jax.lax.optimization_barrier`, but this is currently
|
|
37
|
+
[unreliable on CPU without a non-default XLA flag](https://github.com/openxla/xla/issues/20440).
|
|
38
|
+
|
|
39
|
+
The actual runtime behavior of the primitive is a no-op: it uses aliasing to tell XLA that the input and output `payload` should be the same buffer, so it does literally nothing. This also means it is supported on every platform. The usual pattern is to overwrite the `payload` variable so that its buffer can always be donated, but if it cannot be (for example because it is owned by a caller outside of the JIT context) XLA will insert a copy of the `payload` buffer.
|
|
40
|
+
|
|
41
|
+
## Example
|
|
42
|
+
|
|
43
|
+
```python
|
|
44
|
+
import jax.numpy as jnp
|
|
45
|
+
|
|
46
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
47
|
+
x = some_native_read(token1)
|
|
48
|
+
token2 = modify_token(token1)
|
|
49
|
+
y = some_native_read(token2)
|
|
50
|
+
```
|
|
51
|
+
|
|
52
|
+
Here `token*` is a value that has some library object associated with it, and `modify_token` changes something about this backing object. Even if `modify_token` is marked as side-effecting, `x = some_native_read(token1)` and `token2 = modify_token(token1)` read the same token variable. This means XLA could choose to reorder them, which would lead to the following order:
|
|
53
|
+
|
|
54
|
+
```python
|
|
55
|
+
# Equivalent reordering:
|
|
56
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
57
|
+
token2 = modify_token(token1)
|
|
58
|
+
x = some_native_read(token1)
|
|
59
|
+
y = some_native_read(token2)
|
|
60
|
+
```
|
|
61
|
+
|
|
62
|
+
In this case, both `x` and `y` read the same value. To enforce ordering, we can entangle `x` and the `token` used in its computation, before modifying it.
|
|
63
|
+
|
|
64
|
+
```python
|
|
65
|
+
from entangle_jax import entangle
|
|
66
|
+
|
|
67
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
68
|
+
x = some_native_read(token1)
|
|
69
|
+
token2 = entangle(token1, x)
|
|
70
|
+
token3 = modify_token(token2)
|
|
71
|
+
y = some_native_read(token3)
|
|
72
|
+
```
|
|
73
|
+
|
|
74
|
+
This means that `modify_token` cannot be reordered to be before the line that determines `x`, since it uses a value that depends on `x`.
|
|
75
|
+
|
|
76
|
+
```python
|
|
77
|
+
# Invalid reordering:
|
|
78
|
+
token1 = jnp.asarray(0, jnp.int32)
|
|
79
|
+
token2 = entangle(token1, x) # x doesn't exist yet!
|
|
80
|
+
token3 = modify_token(token2)
|
|
81
|
+
x = some_native_read(token1)
|
|
82
|
+
y = some_native_read(token3)
|
|
83
|
+
```
|
|
84
|
+
|
|
85
|
+
## Development
|
|
86
|
+
|
|
87
|
+
```bash
|
|
88
|
+
uv sync
|
|
89
|
+
uv run pytest
|
|
90
|
+
uv run ruff check .
|
|
91
|
+
uv run ruff format --check .
|
|
92
|
+
uv run ty check
|
|
93
|
+
```
|
|
94
|
+
|
|
95
|
+
The compiled extension rebuilds automatically on import while developing, so the
|
|
96
|
+
default editable install (`uv sync`) is what you want day to day.
|
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
build-backend = "mesonpy"
|
|
3
|
+
requires = ["meson-python>=0.16", "Cython>=3.0", "jax>=0.10"]
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "entangle-jax"
|
|
7
|
+
version = "0.1.0"
|
|
8
|
+
description = "A generic JAX primitive that is opaque to XLA reordering optimizations"
|
|
9
|
+
readme = "README.md"
|
|
10
|
+
license = "MIT"
|
|
11
|
+
requires-python = ">=3.11"
|
|
12
|
+
authors = [{ name = "Nardi Lam", email = "mail@nardilam.nl" }]
|
|
13
|
+
dependencies = ["jax>=0.10", "jaxtyping>=0.2"]
|
|
14
|
+
|
|
15
|
+
[project.urls]
|
|
16
|
+
Homepage = "https://github.com/nardi/entangle-jax"
|
|
17
|
+
|
|
18
|
+
[dependency-groups]
|
|
19
|
+
dev = [
|
|
20
|
+
"pytest>=8.0",
|
|
21
|
+
"ruff>=0.6",
|
|
22
|
+
"meson>=1.3",
|
|
23
|
+
"ninja>=1.11",
|
|
24
|
+
"meson-python>=0.16",
|
|
25
|
+
"Cython>=3.0",
|
|
26
|
+
"ty>=0.0.1a1",
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
[tool.uv]
|
|
30
|
+
no-build-isolation-package = ["entangle-jax"]
|
|
31
|
+
|
|
32
|
+
[tool.meson-python.args]
|
|
33
|
+
setup = ["--vsenv"]
|
|
34
|
+
|
|
35
|
+
[tool.cibuildwheel]
|
|
36
|
+
build = "{cp311-manylinux_x86_64,cp312-manylinux_x86_64,cp313-manylinux_x86_64,cp311-win_amd64,cp312-win_amd64,cp313-win_amd64,cp311-macosx_x86_64,cp312-macosx_x86_64,cp313-macosx_x86_64,cp311-macosx_arm64,cp312-macosx_arm64,cp313-macosx_arm64}"
|
|
37
|
+
test-requires = "pytest"
|
|
38
|
+
test-command = "pytest {project}/tests"
|
|
39
|
+
|
|
40
|
+
[tool.ruff]
|
|
41
|
+
line-length = 100
|
|
42
|
+
target-version = "py311"
|
|
43
|
+
|
|
44
|
+
[tool.ruff.lint]
|
|
45
|
+
select = ["E", "F", "I", "UP", "B"]
|
|
46
|
+
|
|
47
|
+
[tool.pytest.ini_options]
|
|
48
|
+
testpaths = ["tests"]
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
"""A generic ordering primitive for JAX, opaque to XLA's algebraic simplifier.
|
|
2
|
+
|
|
3
|
+
`entangle(payload, *witnesses)` returns `payload` unchanged, but ordered under `jit`
|
|
4
|
+
after every witness's producer. Any XLA custom call is opaque to the algebraic
|
|
5
|
+
simplifier by construction, so this holds without needing `has_side_effect=True`.
|
|
6
|
+
It is more reliable than `jax.lax.optimization_barrier` which needs a non-default flag on
|
|
7
|
+
CPU.
|
|
8
|
+
|
|
9
|
+
Intended for sequencing native side-effecting custom calls that touch the same
|
|
10
|
+
resource but share no ordinary data dependency.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from entangle_jax._primitive import entangle
|
|
14
|
+
|
|
15
|
+
__all__ = ["entangle"]
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
#include "xla/ffi/api/c_api.h"
|
|
2
|
+
#include "xla/ffi/api/ffi.h"
|
|
3
|
+
|
|
4
|
+
#include "_entangle_ffi.h"
|
|
5
|
+
|
|
6
|
+
namespace ffi = xla::ffi;
|
|
7
|
+
|
|
8
|
+
namespace {
|
|
9
|
+
|
|
10
|
+
ffi::Error entangle_impl(const ffi::AnyBuffer payload, const ffi::AnyBuffer witness,
|
|
11
|
+
ffi::Result<ffi::AnyBuffer> out) {
|
|
12
|
+
(void)witness;
|
|
13
|
+
(void)payload;
|
|
14
|
+
(void)out;
|
|
15
|
+
return ffi::Error::Success();
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
XLA_FFI_DEFINE_HANDLER_SYMBOL(
|
|
19
|
+
entangle_handler, entangle_impl,
|
|
20
|
+
ffi::Ffi::Bind()
|
|
21
|
+
.Arg<ffi::AnyBuffer>()
|
|
22
|
+
.Arg<ffi::AnyBuffer>()
|
|
23
|
+
.Ret<ffi::AnyBuffer>()
|
|
24
|
+
);
|
|
25
|
+
|
|
26
|
+
} // namespace
|
|
27
|
+
|
|
28
|
+
extern "C" void* entangle_handler_address() {
|
|
29
|
+
return reinterpret_cast<void*>(entangle_handler);
|
|
30
|
+
}
|
|
@@ -0,0 +1,6 @@
|
|
|
1
|
+
# Type stub for the compiled _ffi extension (built from _ffi.pyx).
|
|
2
|
+
#
|
|
3
|
+
# Importing this module registers the entangle_jax* XLA FFI targets as a side effect.
|
|
4
|
+
# The only name meant for use from Python is the test-only slot reset below.
|
|
5
|
+
|
|
6
|
+
def testing_reset_slots() -> None: ...
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
# distutils: language = c++
|
|
2
|
+
# cython: language_level=3
|
|
3
|
+
"""Registers the compiled entangle FFI handlers as JAX custom call targets.
|
|
4
|
+
|
|
5
|
+
The handler addresses come from _entangle_ffi.cc and _testing_ffi.cc, both compiled
|
|
6
|
+
directly into this extension. Each address is wrapped in a PyCapsule, the format JAX's
|
|
7
|
+
FFI registration expects, and handed to jax.ffi.register_ffi_target. The _testing_*
|
|
8
|
+
handlers back entangle_jax._testing, a private module used only by this package's own
|
|
9
|
+
hazard regression test and not part of the public API.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from cpython.pycapsule cimport PyCapsule_New
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
cdef extern from "_entangle_ffi.h":
|
|
16
|
+
void* entangle_handler_address()
|
|
17
|
+
|
|
18
|
+
cdef extern from "_testing_ffi.h":
|
|
19
|
+
void* write_handler_address()
|
|
20
|
+
void* read_handler_address()
|
|
21
|
+
void reset_slots()
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
cdef object _capsule(void* address):
|
|
25
|
+
return PyCapsule_New(address, NULL, NULL)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _register_targets():
|
|
29
|
+
import jax
|
|
30
|
+
|
|
31
|
+
for platform in ("cpu", "cuda", "rocm", "tpu"):
|
|
32
|
+
jax.ffi.register_ffi_target(
|
|
33
|
+
"entangle_jax", _capsule(entangle_handler_address()), platform=platform
|
|
34
|
+
)
|
|
35
|
+
jax.ffi.register_ffi_target(
|
|
36
|
+
"entangle_jax_testing_write", _capsule(write_handler_address()), platform=platform
|
|
37
|
+
)
|
|
38
|
+
jax.ffi.register_ffi_target(
|
|
39
|
+
"entangle_jax_testing_read", _capsule(read_handler_address()), platform=platform
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
_register_targets()
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def testing_reset_slots():
|
|
47
|
+
reset_slots()
|
|
@@ -0,0 +1,91 @@
|
|
|
1
|
+
"""Primitive definition for the entangle ordering operation."""
|
|
2
|
+
|
|
3
|
+
from typing import TypeVar
|
|
4
|
+
|
|
5
|
+
import jax.custom_batching
|
|
6
|
+
import jax.extend.core
|
|
7
|
+
import jax.interpreters.ad as ad
|
|
8
|
+
import jax.interpreters.batching
|
|
9
|
+
import jax.interpreters.mlir as mlir
|
|
10
|
+
import jax.tree_util as jtu
|
|
11
|
+
from jaxtyping import Array, PyTree
|
|
12
|
+
|
|
13
|
+
from entangle_jax import _ffi # noqa: F401
|
|
14
|
+
|
|
15
|
+
_Payload = TypeVar("_Payload", bound=PyTree[Array])
|
|
16
|
+
"""The pytree of arrays passed as `entangle`'s payload, returned unchanged so a caller
|
|
17
|
+
keeps the exact type it passed in."""
|
|
18
|
+
|
|
19
|
+
entangle_p = jax.extend.core.Primitive("entangle_jax")
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@entangle_p.def_impl
|
|
23
|
+
def _entangle_impl(payload: Array, witness: Array) -> Array:
|
|
24
|
+
return jax.ffi.ffi_call(
|
|
25
|
+
"entangle_jax",
|
|
26
|
+
jax.ShapeDtypeStruct(payload.shape, payload.dtype),
|
|
27
|
+
input_output_aliases={0: 0},
|
|
28
|
+
)(payload, witness)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
@entangle_p.def_abstract_eval
|
|
32
|
+
def _entangle_abstract_eval(payload, witness):
|
|
33
|
+
del witness
|
|
34
|
+
return payload
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
mlir.register_lowering(entangle_p, mlir.lower_fun(_entangle_impl, multiple_results=False))
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def _entangle_p_vmap(vector_arg_values, batch_axes):
|
|
41
|
+
# entangle is a cheap passthrough, so a sequential map over the batch keeps
|
|
42
|
+
# it correct without a bespoke batching rule.
|
|
43
|
+
out = jax.vmap(jax.custom_batching.sequential_vmap(entangle_p.bind), in_axes=batch_axes)(
|
|
44
|
+
*vector_arg_values
|
|
45
|
+
)
|
|
46
|
+
return out, 0
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
jax.interpreters.batching.primitive_batchers[entangle_p] = _entangle_p_vmap
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _entangle_jvp(primals, tangents):
|
|
53
|
+
payload, witness = primals
|
|
54
|
+
t_payload, t_witness = tangents
|
|
55
|
+
if not isinstance(t_witness, ad.Zero):
|
|
56
|
+
raise NotImplementedError(
|
|
57
|
+
"entangle: cannot differentiate w.r.t. a witness argument. Witnesses only "
|
|
58
|
+
"establish an execution-order dependency. You should stop-gradient anything passed as "
|
|
59
|
+
"a witness if it originates from a differentiated computation."
|
|
60
|
+
)
|
|
61
|
+
out = entangle_p.bind(payload, witness)
|
|
62
|
+
if isinstance(t_payload, ad.Zero):
|
|
63
|
+
return out, ad.Zero(out.aval.to_tangent_aval())
|
|
64
|
+
return out, entangle_p.bind(t_payload, witness)
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
ad.primitive_jvps[entangle_p] = _entangle_jvp
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _entangle_transpose(cts_out, payload, witness):
|
|
71
|
+
if ad.is_undefined_primal(witness):
|
|
72
|
+
raise NotImplementedError("entangle: cannot differentiate w.r.t. a witness")
|
|
73
|
+
if ad.is_undefined_primal(payload):
|
|
74
|
+
return cts_out, None
|
|
75
|
+
return None, None
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
ad.primitive_transposes[entangle_p] = _entangle_transpose
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def entangle(payload: _Payload, *witnesses: PyTree[Array]) -> _Payload:
|
|
82
|
+
"""Return `payload` unchanged, ordered under `jit` after every witness's producer."""
|
|
83
|
+
payload_leaves, treedef = jtu.tree_flatten(payload)
|
|
84
|
+
witness_leaves = [leaf for w in witnesses for leaf in jtu.tree_leaves(w)]
|
|
85
|
+
if not witness_leaves:
|
|
86
|
+
return payload
|
|
87
|
+
signal, *rest = witness_leaves
|
|
88
|
+
for w in rest:
|
|
89
|
+
signal = entangle_p.bind(signal, w)
|
|
90
|
+
entangled_leaves = [entangle_p.bind(leaf, signal) for leaf in payload_leaves]
|
|
91
|
+
return jtu.tree_unflatten(treedef, entangled_leaves)
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""Private test-only stand-ins for a mutable native resource.
|
|
2
|
+
|
|
3
|
+
Used only by this package's own hazard regression test, to reproduce the
|
|
4
|
+
sibling-calls reordering hazard entangle exists to fix, without needing a real
|
|
5
|
+
consumer package. Not part of the public API. Do not import from outside this
|
|
6
|
+
package's own test suite.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import jax
|
|
10
|
+
import jax.numpy as jnp
|
|
11
|
+
|
|
12
|
+
from entangle_jax import _ffi # noqa: F401
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def reset_slots():
|
|
16
|
+
_ffi.testing_reset_slots()
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def write(slot_id, value):
|
|
20
|
+
"""Mutate slot `slot_id` to `value`, in place. Has side effect, mirrors refactor."""
|
|
21
|
+
return jax.ffi.ffi_call(
|
|
22
|
+
"entangle_jax_testing_write",
|
|
23
|
+
jax.ShapeDtypeStruct((), jnp.int32),
|
|
24
|
+
has_side_effect=True,
|
|
25
|
+
)(jnp.asarray(slot_id, jnp.int32), jnp.asarray(value, jnp.float32))
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def read(slot_id):
|
|
29
|
+
"""Read slot `slot_id`. No side effect declared, mirrors solve."""
|
|
30
|
+
return jax.ffi.ffi_call(
|
|
31
|
+
"entangle_jax_testing_read",
|
|
32
|
+
jax.ShapeDtypeStruct((), jnp.float32),
|
|
33
|
+
)(jnp.asarray(slot_id, jnp.int32))
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
#include <cstdint>
|
|
2
|
+
#include <cstring>
|
|
3
|
+
|
|
4
|
+
#include "xla/ffi/api/c_api.h"
|
|
5
|
+
#include "xla/ffi/api/ffi.h"
|
|
6
|
+
|
|
7
|
+
#include "_testing_ffi.h"
|
|
8
|
+
|
|
9
|
+
namespace ffi = xla::ffi;
|
|
10
|
+
|
|
11
|
+
static double g_slots[16] = {0};
|
|
12
|
+
|
|
13
|
+
ffi::Error write_impl(const ffi::Buffer<ffi::DataType::S32> id,
|
|
14
|
+
const ffi::Buffer<ffi::DataType::F32> value,
|
|
15
|
+
ffi::Result<ffi::Buffer<ffi::DataType::S32>> out_id) {
|
|
16
|
+
int32_t slot = *id.typed_data();
|
|
17
|
+
g_slots[slot] = *value.typed_data();
|
|
18
|
+
*out_id->typed_data() = slot;
|
|
19
|
+
return ffi::Error::Success();
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
XLA_FFI_DEFINE_HANDLER_SYMBOL(
|
|
23
|
+
write_handler, write_impl,
|
|
24
|
+
ffi::Ffi::Bind()
|
|
25
|
+
.Arg<ffi::Buffer<ffi::DataType::S32>>()
|
|
26
|
+
.Arg<ffi::Buffer<ffi::DataType::F32>>()
|
|
27
|
+
.Ret<ffi::Buffer<ffi::DataType::S32>>());
|
|
28
|
+
|
|
29
|
+
ffi::Error read_impl(const ffi::Buffer<ffi::DataType::S32> id,
|
|
30
|
+
ffi::Result<ffi::Buffer<ffi::DataType::F32>> out_value) {
|
|
31
|
+
int32_t slot = *id.typed_data();
|
|
32
|
+
*out_value->typed_data() = g_slots[slot];
|
|
33
|
+
return ffi::Error::Success();
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
XLA_FFI_DEFINE_HANDLER_SYMBOL(
|
|
37
|
+
read_handler, read_impl,
|
|
38
|
+
ffi::Ffi::Bind()
|
|
39
|
+
.Arg<ffi::Buffer<ffi::DataType::S32>>()
|
|
40
|
+
.Ret<ffi::Buffer<ffi::DataType::F32>>());
|
|
41
|
+
|
|
42
|
+
extern "C" {
|
|
43
|
+
void* write_handler_address() {
|
|
44
|
+
return reinterpret_cast<void*>(write_handler);
|
|
45
|
+
}
|
|
46
|
+
|
|
47
|
+
void* read_handler_address() {
|
|
48
|
+
return reinterpret_cast<void*>(read_handler);
|
|
49
|
+
}
|
|
50
|
+
|
|
51
|
+
void reset_slots() {
|
|
52
|
+
std::memset(g_slots, 0, sizeof(g_slots));
|
|
53
|
+
}
|
|
54
|
+
}
|