Download src/hamiltonzero/model/_custom_lap_primitives.py from Twobombs/HamiltonZero: direct link, hf CLI and curl.
- Browser
- Download file 2.66 kB
-
https://huggingface.co/Twobombs/HamiltonZero/resolve/main/src/hamiltonzero/model/_custom_lap_primitives.py
- Command line
-
hf download hf://Twobombs/HamiltonZero/src/hamiltonzero/model/_custom_lap_primitives.py
-
curl -L -o _custom_lap_primitives.py https://huggingface.co/Twobombs/HamiltonZero/resolve/main/src/hamiltonzero/model/_custom_lap_primitives.py
2.66 kB
| # Copyright (c) 2026 Simulacra Research Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| from __future__ import annotations | |
| import contextvars | |
| import jax | |
| import jax.numpy as jnp | |
| from jax.extend import core | |
| from jax.interpreters import batching | |
| _custom_lap_active_var: contextvars.ContextVar[bool] = contextvars.ContextVar( | |
| "spin_custom_lap_active", | |
| default=False, | |
| ) | |
| def custom_lap_active() -> bool: | |
| return _custom_lap_active_var.get() | |
| def enter_custom_lap(): | |
| return _custom_lap_active_var.set(True) | |
| def restore_custom_lap(token): | |
| _custom_lap_active_var.reset(token) | |
| quadrilinear_merge_p = core.Primitive("quadrilinear_merge_lap") | |
| quadrilinear_merge_p.multiple_results = False | |
| def _quadrilinear_merge_impl(T, u_a, u_b): | |
| groups, rank, _, _ = T.shape | |
| if u_a.ndim == 1: | |
| u_a_2d = u_a.reshape(groups, rank) | |
| u_b_2d = u_b.reshape(groups, rank) | |
| contracted = jnp.einsum("ijkl,ik->ijl", T, u_a_2d) | |
| out_2d = jnp.einsum("ijl,il->ij", contracted, u_b_2d) | |
| return out_2d.reshape(-1) | |
| batch = u_a.shape[0] | |
| u_a_3d = u_a.reshape(batch, groups, rank) | |
| u_b_3d = u_b.reshape(batch, groups, rank) | |
| out_3d = jnp.einsum("ijkl,Bik,Bil->Bij", T, u_a_3d, u_b_3d) | |
| return out_3d.reshape(batch, -1) | |
| def _quadrilinear_merge_abstract_eval(T_aval, u_a_aval, u_b_aval): | |
| del u_b_aval | |
| return jax.core.ShapedArray(u_a_aval.shape, T_aval.dtype) | |
| quadrilinear_merge_p.def_impl(_quadrilinear_merge_impl) | |
| quadrilinear_merge_p.def_abstract_eval(_quadrilinear_merge_abstract_eval) | |
| def _quadrilinear_merge_batched(args, dims): | |
| T, u_a, u_b = args | |
| T_axis, u_a_axis, u_b_axis = dims | |
| if T_axis is not None: | |
| raise ValueError("quadrilinear merge parameters cannot be batched") | |
| batch = None | |
| if u_a_axis is not None: | |
| u_a = jnp.moveaxis(u_a, u_a_axis, 0) | |
| batch = u_a.shape[0] | |
| if u_b_axis is not None: | |
| u_b = jnp.moveaxis(u_b, u_b_axis, 0) | |
| batch = u_b.shape[0] if batch is None else batch | |
| if batch is None: | |
| return quadrilinear_merge_p.bind(T, u_a, u_b), None | |
| if u_a_axis is None: | |
| u_a = jnp.broadcast_to(u_a[None], (batch,) + u_a.shape) | |
| if u_b_axis is None: | |
| u_b = jnp.broadcast_to(u_b[None], (batch,) + u_b.shape) | |
| u_a_flat = u_a.reshape((-1, u_a.shape[-1])) | |
| u_b_flat = u_b.reshape((-1, u_b.shape[-1])) | |
| out_flat = quadrilinear_merge_p.bind(T, u_a_flat, u_b_flat) | |
| out = out_flat.reshape(u_a.shape[:-1] + (out_flat.shape[-1],)) | |
| return out, 0 | |
| batching.primitive_batchers[quadrilinear_merge_p] = _quadrilinear_merge_batched | |
| __all__ = [ | |
| "custom_lap_active", | |
| "quadrilinear_merge_p", | |
| ] | |