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.
@@ -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,10 @@
1
+ .venv/
2
+ .probe/
3
+ __pycache__/
4
+ *.pyc
5
+ build/
6
+ dist/
7
+ *.egg-info/
8
+ .pytest_cache/
9
+ .ruff_cache/
10
+ site/
@@ -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,10 @@
1
+ project(
2
+ 'entangle-jax',
3
+ 'cpp', 'cython',
4
+ version: '0.1.0',
5
+ default_options: ['cpp_std=c++17', 'buildtype=release'],
6
+ )
7
+
8
+ python = import('python').find_installation(pure: false)
9
+
10
+ subdir('src/entangle_jax')
@@ -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,3 @@
1
+ #pragma once
2
+
3
+ extern "C" void* entangle_handler_address();
@@ -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
+ }
@@ -0,0 +1,7 @@
1
+ #pragma once
2
+
3
+ extern "C" {
4
+ void* write_handler_address();
5
+ void* read_handler_address();
6
+ void reset_slots();
7
+ }