diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000000000000000000000000000000000000..d6902233fdb399f746106445e3aa25e21765689d --- /dev/null +++ b/.gitattributes @@ -0,0 +1,38 @@ +*.7z filter=lfs diff=lfs merge=lfs -text +*.arrow filter=lfs diff=lfs merge=lfs -text +*.bin filter=lfs diff=lfs merge=lfs -text +*.bz2 filter=lfs diff=lfs merge=lfs -text +*.ckpt filter=lfs diff=lfs merge=lfs -text +*.ftz filter=lfs diff=lfs merge=lfs -text +*.gz filter=lfs diff=lfs merge=lfs -text +*.h5 filter=lfs diff=lfs merge=lfs -text +*.joblib filter=lfs diff=lfs merge=lfs -text +*.lfs.* filter=lfs diff=lfs merge=lfs -text +*.mlmodel filter=lfs diff=lfs merge=lfs -text +*.model filter=lfs diff=lfs merge=lfs -text +*.msgpack filter=lfs diff=lfs merge=lfs -text +*.npy filter=lfs diff=lfs merge=lfs -text +*.npz filter=lfs diff=lfs merge=lfs -text +*.onnx filter=lfs diff=lfs merge=lfs -text +*.ot filter=lfs diff=lfs merge=lfs -text +*.parquet filter=lfs diff=lfs merge=lfs -text +*.pb filter=lfs diff=lfs merge=lfs -text +*.pickle filter=lfs diff=lfs merge=lfs -text +*.pkl filter=lfs diff=lfs merge=lfs -text +*.pt filter=lfs diff=lfs merge=lfs -text +*.pth filter=lfs diff=lfs merge=lfs -text +*.rar filter=lfs diff=lfs merge=lfs -text +*.safetensors filter=lfs diff=lfs merge=lfs -text +saved_model/**/* filter=lfs diff=lfs merge=lfs -text +*.tar.* filter=lfs diff=lfs merge=lfs -text +*.tar filter=lfs diff=lfs merge=lfs -text +*.tflite filter=lfs diff=lfs merge=lfs -text +*.tgz filter=lfs diff=lfs merge=lfs -text +*.wasm filter=lfs diff=lfs merge=lfs -text +*.xz filter=lfs diff=lfs merge=lfs -text +*.zip filter=lfs diff=lfs merge=lfs -text +*.zst filter=lfs diff=lfs merge=lfs -text +*tfevents* filter=lfs diff=lfs merge=lfs -text +datasets/train/*.jsonl filter=lfs diff=lfs merge=lfs -text +datasets/eval/*.json filter=lfs diff=lfs merge=lfs -text +weights/*.eqx filter=lfs diff=lfs merge=lfs -text diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..fe0bcc9d094c66b139496d78873a160d37e75f1d --- /dev/null +++ b/.gitignore @@ -0,0 +1,13 @@ +__pycache__/ +*.py[cod] +*.egg-info/ +.venv/ +build/ +dist/ +outputs/ +weights/ +*.eqx +*.hzc +*.zst +!weights/ +!weights/hamiltonzero_v1.eqx diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..d645695673349e3947e8e5ae42332d0ac3164cd7 --- /dev/null +++ b/LICENSE @@ -0,0 +1,202 @@ + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/NOTICE b/NOTICE new file mode 100644 index 0000000000000000000000000000000000000000..6ce2203aa87f771ceb26107ba7db841a7fbf4cdc --- /dev/null +++ b/NOTICE @@ -0,0 +1,5 @@ +HamiltonZero +Copyright 2026 Simulacra Research Inc. + +This product includes third-party software. See THIRD_PARTY_NOTICES.md for +attribution and license information. diff --git a/README.md b/README.md new file mode 100644 index 0000000000000000000000000000000000000000..a080080a926d44fcf6de42c2d2529e73d1eadd5d --- /dev/null +++ b/README.md @@ -0,0 +1,304 @@ +--- +license: apache-2.0 +--- + +# HamiltonZero + +HamiltonZero is Simulacra Research's research release for compiled neural +wavefunctions of quantum spin Hamiltonians. It exposes three workflows: + +- learned-router multisystem training; +- compiled single-system fine-tuning; +- compiled single-system evaluation, with optional router contest or large-N + execution. + +## Installation + +HamiltonZero requires Python 3.12 and JAX-compatible accelerator drivers. + +```bash +python -m pip install . +``` + +The package pins the Python `jax` package to a +[`TakeOver/jax` commit](https://github.com/TakeOver/jax/commit/79f82535b15a444516d4a5e2beb71d283665b2ff), +also published as `hamiltonzero-jax-v0.11.0-spin.1`, and pins +`jaxlib==0.11.0`. The fork contains the symbolic-zero JVP support used by the +tuned Pallas attention kernel; stock Python JAX 0.11.0 is not sufficient for +that pathway. Install the accelerator plugin appropriate for the host using +the standard JAX instructions. + +Learned-router training uses eight visible accelerators and requires an MCMC +batch size divisible by eight. Fine-tuning uses all visible accelerators and +requires its MCMC batch size to be divisible by their count. Evaluation chooses +a visible-device subset compatible with its walker batch. + +## Foundation checkpoint + +This Hugging Face repository stores the directly loadable HamiltonZero v1 +foundation checkpoint at `weights/hamiltonzero_v1.eqx`. To download the +checkpoint without cloning the repository: + +```bash +hf download simulacra-research/HamiltonZero \ + weights/hamiltonzero_v1.eqx \ + --local-dir . +``` + +The checkpoint contains the complete foundation wavefunction and its learned +router. `router` is the checkpoint kind, not a router-only parameter file. + +To load the model directly, construct an architecture template and deserialize +its array leaves: + +```python +import jax + +from hamiltonzero.checkpoint import load_model +from hamiltonzero.config import ModelConfig +from hamiltonzero.model import build_model + +template = build_model( + ModelConfig(), + jax.random.PRNGKey(0), + n_max=64, +) +model = load_model("weights/hamiltonzero_v1.eqx", template) +``` + +The template key initializes placeholder values only; deserialization replaces +all serialized array leaves. Set `n_max` to the padded width of the system when +constructing a template for direct model use. The command-line evaluation path +does this from the input system automatically. + +## Hamiltonians and NetworkX + +The public API follows the textbook convention + +\[ +H = \sum_{i.eqx` files. For a one-system training panel it may be a single +file. + +## Fine-tune + +Fine-tuning selects and freezes a route from a router checkpoint, compiles the +single-system wavefunction, and optimizes that compiled model: + +```bash +hamiltonzero finetune examples/finetune.json +``` + +The example fine-tunes on the 256-spin PPP-Ohno system and writes +`outputs/ppp_ohno_n256.eqx`. A neighboring `.eqx.json` sidecar records the +compiled-fine-tune kind, frozen model width, and configured ranks. A compatible +single post-burn-in state can also be supplied: + +```bash +hamiltonzero finetune examples/finetune.json --reuse-mcmc path/to/state.eqx +``` + +## Evaluate + +Compiled evaluation uses the route selected by a router checkpoint, or the +embedded frozen route in a compiled fine-tune checkpoint: + +```bash +hamiltonzero eval examples/eval.json +``` + +Use router contest to compare candidate routes before evaluating the winner: + +```bash +hamiltonzero eval examples/eval.json --contest +``` + +Use the sequence-sharded large-N implementation for the large systems: + +```bash +hamiltonzero eval examples/eval_large_n.json --large-n +``` + +Each evaluation writes `eval.json` and `eval.metrics.jsonl` inside its +configured output directory. + +Training and fine-tuning metrics are written beside the final checkpoint as +`.metrics.jsonl`. Evaluation writes the same per-measurement fields +to `eval.metrics.jsonl`. These JSONL rows contain step, energy, energy standard +deviation, step wall time, and total wall time. Final `eval.json` additionally +reports exchange/field channels and lag-one autocorrelation when available. + +## Configuration + +Every command accepts one JSON configuration. The files in `examples/` are +minimal runnable configurations; omitted parameters use the defaults in +`hamiltonzero.config`. + +The KFAC-JAX fork is vendored under `src/kfac_jax`. + +## License + +HamiltonZero first-party source, datasets, and released model weights are +licensed under Apache-2.0, copyright Simulacra Research Inc. The vendored +KFAC-JAX fork and JAX-derived large-N attention kernel remain under +Apache-2.0. The Microsoft-Folx-derived attention forward and reverse-mode +kernels remain under MIT. See +[`THIRD_PARTY_NOTICES.md`](THIRD_PARTY_NOTICES.md). diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md new file mode 100644 index 0000000000000000000000000000000000000000..854a8dc726cc406236a784a16877a943a02fb6ee --- /dev/null +++ b/THIRD_PARTY_NOTICES.md @@ -0,0 +1,13 @@ +# Third-party notices + +The KFAC-JAX fork under `src/kfac_jax`, including modifications copyright +2026 Simulacra Research Inc., is licensed under Apache-2.0. Its license is in +`third_party/kfac_jax/LICENSE`. + +The attention forward and reverse-mode kernels derived from Microsoft Folx, +with modifications copyright 2026 Simulacra Research Inc., are licensed +under MIT. Its license is in `third_party/folx/LICENSE`. + +The large-N Pallas attention kernel derived from JAX, with modifications +copyright 2026 Simulacra Research Inc., is licensed under Apache-2.0. Its +license is in `third_party/jax/LICENSE`. diff --git a/datasets/README.md b/datasets/README.md new file mode 100644 index 0000000000000000000000000000000000000000..ef300fa461dcbd2708ac073622832ba438bf6b34 --- /dev/null +++ b/datasets/README.md @@ -0,0 +1,38 @@ +# Datasets + +All files use the public textbook convention + +\[ +H=\sum_{i4} {site:>5} {system.nodes[site]}\")" + ] + }, + { + "cell_type": "markdown", + "id": "permutation-convention", + "metadata": {}, + "source": [ + "## Permutation convention\n", + "\n", + "`leaf_to_input[leaf]` gives the index in `system.nodes` placed at that compiled leaf. `input_to_leaf[site]` is the inverse used below to color each public lattice site by its merge-tree cell. The NetworkX loader preserves the supplied node order. HamiltonZero applies the learned permutation once to the context and walkers, then compiles an identity-routed tree, so **no extra bit reversal is needed**. For padded systems the arrays also contain virtual leaves, which must remain when determining merge groups even though they are not drawn as physical sites." + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "id": "visualize", + "metadata": {}, + "outputs": [ + { + "data": { + "image/png": "iVBORw0KGgoAAAANSUhEUgAAB0QAAAGYCAYAAADSo1HRAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAASdAAAEnQB3mYfeAAAhqxJREFUeJzs3Xd0VOX69vErnRZKgFASem8aQIqgotKlK1WpUgQsCCqi6EEEURQVRTgqKCqiHJBmpEiRDhJACAJSJZAAgQABQhJSn/cP38zPIYUJBHZm+H7Wcp3M3s/sueZhnTs7+97FzRhjBNwFzpw5oypVqig2Nlb33HOP+vTpo3Llyuny5cs6duyYfvnlFyUlJengwYN273Nzc1P9+vW1c+dOi5JnbtasWRo8eLDy5Mmjl19+Wffee69CQ0M1ZcoUXbt2Td9884369etndUwAuCWNGjVSSEiI/P39NXDgQNWqVUuSdPz4cW3YsEHr1q1TWFiYAgMDbe9p0aKF1q5dq+joaBUuXNii5BmbNWuWhg0bpvvvv19t27bVhx9+qAsXLuTKrABws8aOHatJkybJx8dHTz75pB544AH5+vrq5MmT2r17t5YuXapp06apf//+tve89957eu2117RgwQJ17drVuvAZeOCBB3TmzBl17NhRVatWVbFixXTgwAFNmzZNFy5cUJ06dRQaGio3NzerowLATVu3bp0effRRSf/sT3fp0kUlSpTQ2bNndfDgQS1evFjNmjXT999/b3vPzp071aBBAz3zzDP6/PPPrYqeocjISJUvX15t27ZVo0aNVLFiRcXGxmrt2rX64YcfZIzR3Llz9eSTT1odFQBu2qVLl1SpUiVdvHhR1apV04ABA1ShQgXFxsbq2LFjWrFihU6fPq0zZ87Yva9w4cIqVqyYjh49alFyxxw7dky1a9dWgwYNtGnTJrVu3VorV660OhaclKfVAYA7ZeHChYqNjVWjRo20ceNGeXt7261/9913FRYWZk24m/TOO+9Ikr755hv16NFDktS1a1fVqlVLvXr10sSJE2mIAnBqBw8eVEhIiHx9fRUSEqJy5crZrX/jjTcUGRmpIkWKWJQw+9q3b69u3bqpUKFCkqTp06dbnAgAct63334rSfruu+/UvXv3dOvj4uJ09erVOx3rpn3//fcqX758uuU9e/ZUvXr19Oeff2rPnj2qW7funQ8HADkkrXYPHjxYX375Zbr1n3zyicLDw+90rJvm5+eniIgIFStWzG75gAEDVL58eb3zzjv6+eefaYgCcGq//PKLLl68qNq1a2v79u3Kly+f3fqJEyc63THvf3v++edVokQJvfHGG2rdurXVceDkaIjirnH27FlJUrt27dI1Q9P8+yDHli1b9O6770qSjh49qvbt29vWPfnkk3Y7zCkpKVq8eLHWrl2rs2fPqlChQnrooYf05JNPysfHxzbu2rVrtobl5MmTtXjxYv3yyy+6cuWKatasqSFDhiggIMCh7xMaGqqwsDBVrFgx3UGmHj166PXXX9fRo0e1f/9+29VUAOBs0mp33bp10zVD05QsWdL2c2xsrHr06KHQ0FBJ/xyo9vT8Z3cnrfb+29atW7VkyRIdP35cnp6eCgoKUr9+/ey2KUkvv/yy7az4P//8U999950iIiJUunRp9e7dWw0bNnT4O12/bQBwRWfPnpWbm5s6d+6c4fp8+fLZHawZOXKk7Uzv9957T998841t3eLFi+Xl5WV7HRYWph9++EH79+9XYmKiKlasqJ49e6ZrRn7//feaN2+exo0bp8DAQM2cOVP79u1T/vz59dhjj6lr164OX9GZUTNUkqpVq6bGjRvrt99+U2RkpEPbAoDcKm3fu0uXLhmud3NzU9myZW2vv/32W82aNUuStGLFCrvjJuPGjVODBg1sr69evaoffvhB27Zt0+XLl1WyZEm1adNGHTp0sKvFf/75p1577TV16tRJffv21TfffKPNmzcrKSlJjRo10pAhQ5Q/f36Hvo+3t3e6Zmia++67T5Lk4eHh0LYAILdKq92tWrVK1wxN8+992dDQUI0dO1axsbFKSkqyq92dOnXS4MGDba+NMfr555+1evVqnTlzRgUKFFCTJk3Up0+fdJ/VqVMnBQYGavr06VqxYoUWL15su2p10KBBqlChQra/26JFi7RixQotWbIk0+8GZIsB7hJff/21kWTatGnj0PgFCxYYSRn+N27cONu48PBwExQUlOG46tWrm5MnT9rGxsTEGEmmadOmZuDAgenGFyxY0GzZssWhfLNnzzaSTP/+/TNc369fPyPJzJ0716HtAUBudPLkSSPJFC9e3Fy8ePGG46OjozOt3U2bNrWNS0xMNH369MlwnK+vr/n111/tttuoUSMjycycOdN4eHjYjXdzczMfffTRTX/HgIAAI8lER0ff9DYAILepWrWqkWTWr1/v0Pj69etnWr/j4+Nt46ZPn268vb3TjXFzczNvv/223TbHjh1rJJnJkyeb4sWLp3tPly5dTFJS0i1/13r16hlJ5siRI7e8LQCw0vDhw40k88Ybbzg0/tVXX820dgcHB9vGbdu2zZQsWTLDca1bt7ar8+vWrTOSzDPPPGPbB//3f5UrVzbh4eG39D0vX75sWrRoYSSZefPm3dK2AMBqixYtMpJMkyZNTEpKyg3Hr169OtPaPWLECNu4s2fPmiZNmmQ4rkKFCubw4cN22/Xw8DC1atUyo0aNSjc+X758ZtWqVdn6XlevXjVlypQx7du3N8YYs2nTJtvvDeBmcYUo7hpdunTR2LFjtXLlSjVu3FjdunVTw4YNFRQUJF9f33TjH3jgAQUHB6tDhw6qXLmyPv74Y9u6qlWrSvrnytCOHTtqz549atmypTp37qwSJUro/PnzmjdvntavX68+ffpo/fr1dtvetWuXtmzZov79+6t58+a6dOmSZs+erT/++EPdu3fXoUOHbnjGY0REhCRlenZN2pk/znQ7GwC4XpkyZdShQwcFBwcrKChIvXv3VpMmTVS/fv0Mr7QsUKCAgoOD9cYbbyg0NFTz5s2z1VM/Pz/buFdffVVz5sxRjRo11LdvX1WsWFHXrl3T+vXrbbd3PHbsmIoWLWq3/eeff15NmzZVr169lDdvXi1btkwLFizQyy+/rCZNmqhRo0a3d0IAwEk8++yzGjFihNq1a6ennnpKDz/8sOrXr68qVapkeFXm1KlTNWvWLH377bcaM2aMmjZtaluXdneX5cuX69lnn1XRokXVv39/BQUFydvbW6Ghofr888/1n//8Rw0aNFCbNm3stj1u3DgVL15cH3zwgQIDA7Vr1y59+umnWrx4saZMmaIxY8bc9Pf87bff9Mcff+jRRx9V5cqVb3o7AJAbDBkyRDNnztQ777yj0NBQtWvXTvfdd5/q1KmT4Z22+vfvr5IlS2rkyJFq27athg8fbluXdgeVM2fOqH379rp8+bKeeuopPfLIIypcuLBOnjypmTNn6tdff9Ubb7yhKVOm2G077fa9o0ePVt26dXXy5El98sknOnr0qPr37681a9Y4/L1SU1PVsWNHSVJ0dLRCQ0OVnJyscePG2R4/BADOqk2bNqpUqZK2bt2q+vXrq2fPnmrYsKHq1atne1TPvwUFBSk4OFg9e/ZUwYIF7W6RXrFiRdvP3bp109atW/XQQw+pa9euKl26tKKjo7Vw4UKtXLlSPXr00K5du+z27dPuVtizZ0+1adNGcXFx+v7777V161b16tVLhw4dSnecJTNvv/22zp8/r2nTpt3C7ADXsbojC9xJu3fvNvfdd5/dGSoeHh6mQYMG5osvvjDJycnp3iPJ1K9fP8PtzZ8/30gyL774Yrp1qamp5rHHHjOSzF9//WWM+b8rRCWZKVOm2I1PSEgwDRo0MJLMrFmzbvhd0s7E/PDDDzNcP2XKlGyd2QkAudWFCxdMjx49jLu7u139rlixonn99dfNhQsX0r2nefPmmV51GRkZaby9vU2DBg3szkZP8+GHHxpJZtq0abZlaWenP/bYY+nOuHzppZeMJNOzZ8+b+n5cIQrAFaWkpJjx48eb/Pnz29VuPz8/07dvX7N///5073n33XeNJLNgwYIMt1mvXj2TL18+8/fff6dbFxoaaiSZTp062ZalXSHq7++f7nfFihUrbOscOZM+IydOnDAlS5Y0vr6+XB0KwGUsWbLElC1b1q5258mTx7Rq1cruqs80O3bssF3RmZHRo0cbSeb7779Pty4mJsZUqFDBFCxY0HY8Ju0KUUlm3bp1duPPnj1r/Pz8jCSzZ88eh79TUlJSuquVOnfubEJDQx3eBgDkZgcPHjRNmza1q3Pu7u6mbt265tNPPzUJCQnp3lOoUCFTqVKlDLe3atUqI8kMGDAgw/W9evUyksy2bdtsy9LupnX9sejk5GTbMZrrj4dnZv/+/cbLy8tMnDjRtowrRJETuEIUd5WgoCDt2LFDe/bssZ3NvXnzZu3YsUM7duzQokWLFBwcbPeMoqz8+uuvkv654rNz584yxkiSjDEyxuj48eOSpH379ql69eq29/n6+uqFF16w25a3t7dGjx6tbt26aePGjRo4cGCWn532TLzk5OQM1yclJdmNAwBn5efnp3nz5mnKlClatWqVQkJCFBISoj179mjSpEmaM2eONm3alOkzRq+3bt06JSYmKiYmRj179pT0f3Vbki5duiTpn9p9vddff13u7u52y8aOHasPP/xQGzduvIVvCQCuxd3dXf/5z380atQorVq1Stu2bdMff/yhrVu36rvvvtP//vc/LViwQB06dHBoe+fPn9cff/yhokWLauTIkZLsa7cxRl5eXhnW7uHDh9vdJUD650z6+vXra9euXTp06JBq1KiRre8XGRmpVq1aKTo6WsHBwVwdCsBldOrUSe3bt9fmzZu1ceNG23GTVatWadWqVRo9erQmT57s8PbSjpuk1f3rj5tcu3ZNV65cUXh4uN0z7h544AE9/PDDdtvy9/fX4MGDNXnyZG3cuFH33nuvQxk8PDwUHBwsY4zOnj2rDRs2aN68efr111+1YsUKNWvWzOHvAwC5UbVq1bR582bt379fa9eutd2dcPfu3dq9e7f+97//afXq1cqbN69D20ur3X/99Ve6Y96SFBYWJumf4yaNGze2vc/DwyPd3VfSlq1du1YbN27USy+9dMPPf/bZZ1WhQgW98sorDuUFHEWnBHeloKAgBQUF2V5v3LhR3bt316+//qpvv/1WgwYNcmg7abet3bRpU5bjYmNj7V5XqlQpw6ZrWtP0zJkzN/zstFseXLhwIcP1Fy9elCQVLlz4htsCAGcQGBiop59+Wk8//bSkf3bA+/Tpo82bN+vVV1/VvHnzHNpOWu0+ePCgDh48mOm462u3JLuTW9IUKVJEJUqUUGRkpEOfDwB3kwIFCujxxx/X448/Lkm6evWqXnvtNX322WcaPHiwIiIiHDqBL612X7hwQUuXLs10nKO1O235rl27dObMmWw1RE+fPq1HH31Uf//9txYtWqSWLVs6/F4AcAYeHh5q1qyZrVFojNGcOXM0cOBAvf/+++rZs6fq1q3r0LbS6ndwcHCW466v31nVbsmx4yZp3Nzc1L59e9vrQYMGqU2bNurdu7deeukl7dy50+FtAUBuVqtWLdWqVcv2OiQkRN27d9eWLVs0ffp0vfzyyw5tJ612//7771mOu752lylTJsPHwGWndn///fdav369Vq9eneHt2oFbQUMUkPTQQw/plVde0csvv6w1a9Y43BBNO3gzZcoUVatWLdNx/26+SlJMTEyG49KW+/j43PCz055jun///gzXpy2vUqXKDbcFAM6ofPny+uyzzxQUFJStZwil1e4+ffqoe/fumY4LDAxMtywmJibD511cvXrV4bsLAMDdrECBAvr00081f/58nT17Vvv27Uu3r5yRtNpdp04dTZo0KdNxefLkSbcsJ/a905w8eVLNmzfXiRMnNH/+fLsD7ADgqtzc3NS3b1+tWLFC8+bN09q1ax1uiHp6esrd3V0//fRTlvvLZcuWtXudk7U7Iz179tTgwYO1Z88eJSUlsS8PwCU1bNhQb775pgYNGqQ1a9Y43BBN2/ceP3686tWrl+m4fzdfpZyp3Z988ony5s2rqVOnaurUqbbl0dHRkqQ//vhD7du3V8WKFfXpp5/ecHvAv9EQBf6/1NRUSVJiYqLdcjc3N6WkpGT4nlq1amnZsmWKi4vL1sGQsLAwhYWF2d0ORpLWr18v6f+anVlp0qSJ3N3dtX79el26dMnuStBLly5p/fr18vT0tLttAQC4msxqd9ptbTOq32k77CdPnsz2gez169erf//+dst27Nih2NhY3XPPPdnaFgDczdLq87/rd1a1u1KlSvLx8dHx48fVtGlTFSlSxOHPWr9+fboTHuPj4xUSEiLJ8RMI//77bz366KM6deqU5s+fr86dOzucAQBcQUb73lnVbumffe/ffvtNefLkUdu2bR3+rK1btyoxMTHd1UHZOW6SlYsXLyo+Pl6enp7y8PC4pW0BQG6W1XGTrGq3JF2+fDlbx00uXLigffv2qXbt2nbLs1O7k5KSFB8fr2XLlmW4PioqSsuWLXP4tunAv7nfeAjgGoKDg/XVV1/p/Pnz6dZt27ZNU6ZMkaR0xbRQoUI6efKkEhIS0r2vd+/ecnNz04QJE/Tdd9+l+yVy7NgxjR8/Pt37UlJSNHDgQNttbSVp8+bNev/99yXJoYMrxYsXV5s2bRQbG6thw4bZnhmalJSkoUOH2pq01z8vCQCcyZEjR/Tuu+/q0KFD6dadOnXK9jzmjGq3JB0+fDjd+5o1a6ayZctqw4YNGjFihK5cuWK3/tKlS5oxY0aGn/nGG28oNDTULsOwYcMkOVa7AeBuMXbsWG3evNl2ACZNbGysRo0apQsXLsjb29vulohZ1e68efOqa9euunr1qjp37qxjx47ZrU9OTtaSJUsyvCXjjz/+qB9++MH2OiEhQc8//7wiIyPVpEkT+fv73/D7HD58WA899JBOnTqlefPmqUuXLjd8DwA4m88//1yLFi1SXFyc3fLU1FT98MMPWrx4sST7fe+sarf0z11ZJOnpp5+2HRD/t5CQEH322WfploeHh+vFF1+0HeuQpG+++UaLFi1Snjx51Lp16xt+nyVLlmR4J5nw8HA9+eSTkqQHH3zQ1tQFAGe0Zs0aff755xk+xmfPnj2aMGGCpIyPm0RGRqY7JiL9cxW9l5eXpk6dqv/+979KTk62W3/y5Em9/fbbdjU6zeDBg+2y7Nq1y3Z83JHjJtOmTVNwcHC6/9KeX12vXj0FBwdr2rRpN9wWcD03k/YkXMDFTZw4UW+++abc3d0VGBiowMBA5cmTRydOnLAdUClVqpT27t2rYsWK2d7XunVrrVq1SmXKlFHNmjXl6empJ5980rbz/NZbb9mKur+/v6pWrSo3NzeFhYUpIiJC3t7eunbtmqR/bqno6+ur8uXL6/z583Jzc9M999yjK1euaP/+/UpNTVWXLl20aNEih77T/v371ahRI8XGxqp06dKqWbOmDhw4oNOnT8vX11chISGZPncDAJzB77//rvvvv1/SPzU2MDBQRYsW1blz57Rv3z6lpKTI09NTK1euVPPmzW3v++CDDzR69GgVLFhQ9913n/LmzatatWrZdqDXrl2rxx57TImJiSpQoICqVaumQoUKKTw8XGFhYUpKStKmTZv0wAMPSJIaN26s7du3q06dOvrrr79Uu3Zt5cmTR3v37lVcXJzKlCmjP//803ZAKCsXLlxQv379bK/Xrl2ra9euqXXr1rbb0rz99ttZ3pYGAHK7AgUKKDY2Vr6+vipbtqxKlSqlmJgYHThwwHbLrDfffFNvv/227T27d+9WvXr15OHhoYYNG9pO7Fu8eLG8vLwUGRmpxo0b68SJE/Lw8FD16tVVqlQpnTt3Tn///beuXr2qcePG6a233pL0z0ks77zzjurUqaM///xTlStXVkBAgPbv36/z58/L09NTGzZsUJMmTW74fYKCghQaGqpSpUplWp9feOEFtWrV6hZnDgCs07lzZy1dulTe3t4qU6aMAgMDlZqaqqNHj9qe+9a0aVNt3LjR1kRMTU1VyZIlFRUVpZo1a6pcuXJyd3fXuHHj1KBBA6Wmpqpz5862E1bKly+vChUqKDY2VsePH1dUVJSaNm2qzZs3S/rnKqJHHnlEtWvX1r59++Tv76+aNWvq5MmT+vvvvyVJkyZN0muvvXbD7zNmzBhNnjxZfn5+qlChgooUKaLIyEj99ddfSklJUf78+bV+/Xrdd999t2M6AeCO+Oyzz/T888/Lzc1NAQEBCgwMVP78+RUREWE70dvPz0979uxRmTJlbO/r0aOH5s+fr9KlS6t27dry8vJSp06dNHjwYEnS1KlTNXLkSElSsWLFVLVqVXl5eSksLEzh4eFKTU1VfHy87ZEVnp6e8vf3V2JiouLj43XvvfcqLi5O+/fvV3Jyspo3b67Vq1fLzc3tpr7n5s2b9eCDD6p169ZauXLlrUwZ7mYGuEuEhoaaQYMGmRIlShhJdv95eXmZrl27mrCwsHTv27VrlyldurTd+HHjxtmN+eabb0z58uXTbbdq1armvffes42LiYkxkkzTpk3NmjVrTMmSJW1j3dzczJNPPmni4uKy9b02b95sqlSpku5zt27delPzBAC5yYULF8xrr71matSoka7GSjINGzY0v/32W7r3Xb582TRp0sRubNOmTe3GbN++3TRt2jTdNv38/Mxzzz1noqKibGMbNWpkJJlDhw6Zhg0b2o2vX7++OXz4sMPfKTw8PMPv8u//VqxYcfOTBgC5wKeffmqaNWtmPDw80tW4smXLms8++yzD9z3//PPG3d3dbnx8fLxtfWRkpHnyySeNp6en3Rhvb2/z+OOPm127dtnGjh071kgyCxcuNP3797fbbokSJczPP//s8PepVKnSDWv3zJkzb37CACAXWLlypenataspUKBAuhpXsGBB8/zzz5srV66ke9+8efNM/vz57cYHBwfb1iclJZm33nrLFClSJN12GzdubP73v//Zxq5bt85IMiNGjDDTp083+fLls6v11x+Pycr27dtNp06djJeXl91nuru7mxYtWpg9e/bc0nwBQG5w8OBBM2zYMBMQEJCuxnp4eJj27dubQ4cOpXvfX3/9le549ogRI+zGLFiwwFStWjXdditUqGDGjRtnUlJSbGM9PDxMrVq1zNatW03ZsmXtxnfu3Nlcvnz5lr7npk2bjCTTunXrW9oO7m5cIYq7jjFGERERCg8P19WrV+Xn56caNWoof/78mb4nKSlJe/fu1blz55SSkqKqVatmeM/zQ4cOKTw8XPny5VO5cuUUEBBgtz7tCtG0sx+Tk5P1xx9/6MqVK6pRo0a68dn5Tn/++afOnj2rkiVLqk6dOje1HQDIzc6fP6+TJ08qKipK+fPnV7Vq1VS8ePEs33Po0CGdOHFCiYmJ8vPzy/AqoMjISB06dEipqakqW7asypUrZ7tSM03aFaJpZz8eOHBAp06dUunSpW3P1nDUtWvXMrx11781atToht8NAJxBbGysTpw4odOnT8vDw0Nly5ZVxYoVszwz/OzZs9q/f7/i4+NljNFjjz2W7naGV69e1Z9//qmrV6+qVKlSqlChQrr9+bQrRIODg9W+fXudPn1aBw8eVL58+VS/fn15eXk5/D3WrVun2NjYLMfce++9dmfdA4CzSkpK0smTJxUREaHExESVLFlS1atXz7JuXr16VXv37tWlS5eUmpqqhg0bprsleXJysvbt26dz587ZrtosWrSo3Zi0K0RHjBihqVOnKiYmRnv27FFycrLq1q2rwoULZ/v7xMbG6ujRo4qMjJSvr69q1qx5U9sBgNzMGKPTp08rPDxcly9fVpEiRVSjRg35+vpm+p6UlBTt3btXZ8+eVXJysipWrKiaNWumG3f06FGdOHFCPj4+KleunAIDA9Ptz3t6eqp69eq2u3nt2bNHFy9eVNWqVVWuXLlb/n7R0dHasmWL/P391bBhw1veHu5ONESBO+j6higAwDlc3xAFAOR+1zdEAQC53/UNUQCAc/h3QxTIrXhqOAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMvyvPEQADklb968Cg4Olp+fn9VRAADZ8OGHHyo6Olre3t5WRwEAOKhPnz5q3LgxzxgCACdSp04dBQcHq2LFilZHAQBkw88//6wCBQpYHQPIEs8QBQAAAAAAAAAAAOCyuGUuAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBleVodIKdcunRJGzZsUJkyZeTj42N1HABwCgkJCQoPD1ezZs1UuHBhSzJQvwEge6jdAOB8qN0A4Jysrt/UbgDIvsxqt8s0RDds2KDOnTtbHQMAnNKSJUvUqVMnSz6b+g0AN4faDQDOh9oNAM7JqvpN7QaAm3d97XaZhmiZMmUkSQWfmSrP4mUtTuP8Ln0+QpJUeOgnFidxDcxnzmEuc1b0Z8NlLkXaaqgVqN/IrS59PkKpSUnSE+OsjgLYm/+GdPUCtRsAnAj73QDgnKyu39RuAMi+zGq3yzRE024Z4Fm8rDxLV7Y4jfNz8/KWJOYyhzCfOYe5zFlunl4ykqW3XaF+I7dy8/KWjJtUvJzVUQB7Hl6SqN0A4EzY7wYA52R1/aZ2A0D2ZVa73a2JAwAAAAAAAAAAAAC3Hw1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABcFg1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABcFg1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABcFg1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABclqfVATZv3qz169crOTlZ9913n9q1ayc3NzerYwEAsnDs2DEtXbpUFy9eVOXKlfXEE0/I19fX6lgAgCxcunRJCxcu1PHjx1W8eHF17txZ5cqVszoWACALKSkpCg4O1u7du+Xj46PmzZurUaNGVscCANxASEiI1qxZo4SEBAUFBaljx47y8PCwOhYA3NUsu0LUGKOBAweqTZs2On/+vJKSkjR48GC1atVK165dsyoWAOAGZs+erZo1ayokJEQ+Pj6aNm2aatWqpSNHjlgdDQCQidDQUFWvXl0zZ85Unjx5tGnTJlWvXl0//fST1dEAAJmIiYnRQw89pOeee06pqak6c+aMHnnkEb344otWRwMAZGHkyJFq1qyZTp8+LWOMRowYoQcffFAxMTFWRwOAu5plV4h+/fXX+vrrr7Vy5Uq1bt1aktSvXz/dc889Gj9+vN59912rogEAMnH06FE988wzGjFihD744ANJ0ujRo1W/fn099dRTCgkJsTghAOB6KSkp6tmzp8qWLavNmzfL0/OfPwGGDRum/v37q0mTJipdurTFKQEA1xs9erRCQ0O1f/9+2xX9Dz/8sLp27aoHH3xQTzzxhMUJAQDXW7hwoaZOnap58+apR48ekqQhQ4aoZs2aGj16tP773/9anBAA7l6WXSE6Y8YM1axZ09YMlaSqVauqXbt2+uKLL5SUlGRVNABAJmbOnKmkpCSNHDnStszHx0fPPvusduzYoe3bt1uYDgCQkd9++00HDx7U888/b2uGStKoUaMUGxur2bNnW5gOAJCR2NhYffvtt3riiSfsbm+e9vqzzz6zMB0AIDMzZsxQQECAunfvblsWGBiobt266dtvv9XVq1ctTAcAdzdLGqKxsbHavXt3hs+9uP/++xUdHa39+/dbkAwAkJVNmzapTJky6a4kuv/++23rAQC5S1ptvn7fu0qVKipatCi1GwByoV27dik+Pj7D4yaNGzfWtm3blJKSYkEyAEBmUlNTtW3bNjVs2FBubm526+6//37Fx8dr586dFqUDAFhyy9zw8HAZYzK8NVfashMnTigoKCjD9587d05RUVF2y44ePZrjOQEA9k6ePKnAwMB0y/9du7NC/QaAO+/kyZOSlOm+N7UbAHKfG9XuhIQERUZGKiAgIMP3U7sB4M47d+6c4uPjb3jMO6v3U7sB4PaxpCEaHx8v6Z/bLF4vT548dmMyMmPGDI0fP/72hAMAZCo+Pv6ma7dE/QYAK9xo3/vKlStZvp/aDQB3HsdNAMD5ULsBIHezpCGaP39+SRn/AkhbljYmI8OHD1e3bt3slh09elSdO3fOuZAAgHTy589/07Vbon4DgBX+ve/t5eVlty4+Pp7aDQC5EMdNAMD5ULsBIHezpCFatmxZubu7KyIiIt268PBwSVKFChUyfb+/v7/8/f1vWz4AQMbKly+f4e1aHKndEvUbAKxQvnx5SVJERIRq1qxpty4iIkJNmzbN8v3UbgC48/5du68XHh6ufPnyqUSJEpm+n9oNAHde8eLFlT9/fo55A0Au5W7Fh+bJk0eNGjXS1q1b063bvHmz/P39VaNGDQuSAQCy8vDDD+vMmTM6fvy43fLNmzfb1gMAcpe02rxlyxa75fv27dOlS5eo3QCQC9WrV0++vr7pandqaqq2bdumBx98UO7ulhzSAQBkws3NTQ899JC2b9+ulJQUu3WbN29WgQIFVL9+fYvSAQAs23t+4YUXdOzYMc2bN8+2bNeuXVq5cqWef/55eXh4WBUNAJCJwYMHK2/evHrnnXdsy65cuaJp06apWbNmCgoKsi4cACBDDz74oOrWraupU6cqLi7OtnzSpEkqVKiQ+vfvb104AECG8uTJoyFDhmjp0qXav3+/bfns2bN1+vRpvfjii9aFAwBk6oUXXtC5c+c0c+ZM27KDBw9q4cKFeuaZZ5Q3b14L0wHA3c2SW+ZKUs+ePbVnzx7169dPS5cuVb58+bRgwQJ169ZNY8aMsSoWACALZcqU0Q8//KA+ffro2LFjqlWrllasWKG8efNq7ty5VscDAGTAzc1N8+fPV9u2bVW3bl21bNlSoaGh2rdvnxYsWKBixYpZHREAkIGJEyfq0KFDeuCBB9S1a1ddvnxZS5cu1cSJE9WmTRur4wEAMtCmTRu98847GjFihNasWSM/Pz/99NNPat68uSZOnGh1PAC4q1nWEJWk9957TwMGDNDGjRuVnJysoUOHqkGDBlZGAgDcQOfOnXX8+HH9+uuvunjxotq1a6cWLVrIy8vL6mgAgExUrlxZ+/bt0+rVq3X8+HE1bdpUbdq0UZEiRayOBgDIRJ48eRQcHKxt27Zp9+7d8vb21rvvvqtKlSpZHQ0AkIXXX39dPXv21Lp165SQkKD+/furSZMmVscCgLuepQ1RSapWrZqqVatmdQwAQDYUK1ZMTz31lNUxAADZ4OPjo/bt21sdAwCQTffff7/uv/9+q2MAALKhYsWKqlixotUxAAD/YtkzRAEAAAAAAAAAAADgdqMhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy/K0OkBOu/T5CLl5eVsdw+mlnDkmSbowoYvFSVwD85lzmMuclRp9xuoINtTvW5dy8YzcJLn7lbI6iktIqzeaNdTaIK7iSpTkJsm3uNVJnF/MeasT2FC7cwb1O+cwlzmL+cw5uWm/GwAAALgbuVxDFADgnFKTkiTjZnUM55ZqZNyklMQkq5O4FA9vL6sjuIQUd7d/Dqozn7csJReVSmp3DqF+5xzmMmcxnznHWB0AAAAAuLu5XEO08NBP5Fm6stUxnF7a1XdF31xscRLXwHzmHOYyZ53/Tzulng+3OsY/nhgnFS9ndQrg/8waKg9vL+oNch1qNwA4oc+fli5HWp0CAAAAuGvxDFEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXZWlDNCQkRC+//LKqV6+uwMBAHThwwMo4AAAHnDhxQh988IEaNWqkwMBAffHFF1ZHAgDcwOXLl/Xdd9+pffv2KlOmjPr06WN1JADADSQnJ+vXX3/VwIEDVaFCBZUvX97qSAAAB+zZs0evv/66atWqpcDAQG3bts3qSAAAWdgQ/fTTT/Xcc8+pZMmSevTRR3Xq1CklJiZaFQcA4IB9+/bp4Ycf1rlz5zRgwACdOnVKMTExVscCANxAo0aNtGbNGg0dOlQXL15UVFSU1ZEAADcwZMgQffTRR2ratKkqVaqkiIgIqyMBAG7g22+/Vf/+/eXr66t27drp1KlTSkhIsDoWAECSp1UfPHz4cL3wwguSpIkTJ1oVAwCQDTVq1NDx48clSb///rvFaQAAjtq3b588Pf/Z9Xdzc7M4DQDAEV9++aWtdv/8888WpwEAOOKpp55Sv379JEmfffaZxWkAAP9m2RWiaTv1AADn4eHhYXUEAMBNYN8bAJwPtRsAnA+1GwByL4caok2bNtVbb72lLVu2KDk5+XZnAgAAAAAAAAAAAIAc4dApKx4eHpo0aZLGjx8vX19fPfLII2rZsqVatmypatWq3e6M6Zw7dy7dc4+OHj16x3MAALKH+g0AzofaDQDOh9oNAM6H2g0At5dDDdGNGzfq6tWrWr9+vVavXq3Vq1fbnl9RpkwZtWzZUl999dVtDfpvM2bM0Pjx4+/Y5wEAcgb1GwCcD7UbAJwPtRsAnA+1GwBuL4dval6gQAG1b99e7du3lyTt2bNHY8aM0a+//qqvv/76jjZEhw8frm7dutktO3r0qDp37nzHMgAAso/6DQDOh9oNAM6H2g0AzofaDQC3l8MN0cTERG3ZskWrVq3S6tWr9ccff8jb21vNmzdXy5Ytb2fGdPz9/eXv739HPxMAcOuo3wDgfKjdAOB8qN0A4Hyo3QBweznUEH3ssce0YcMGxcXFqU6dOmrZsqXeeecdPfTQQ8qbN+/tzggAAAAAAAAAAAAAN8WhhuiKFSuUN29evfbaa+rTp49q1Khxu3MBAAAAAAAAAAAAwC1zd2TQt99+qyeeeEKzZ89WzZo1VaZMGT399NP68ccfFRUVdVMfvHv3bgUGBiowMFAffPCBJKl169YKDAzUPffcc1PbBADcftWqVVNgYKA6duwoSZo4caKtnh84cMDidACAjAwaNMhWq+Pi4rR+/Xrb6y+++MLqeACADHzzzTe2Wv3rr78qJSXF9rpPnz5WxwMAZODYsWO2Wv3mm29Kkrp166bAwECVL1/e2nAAcJdz6ArRvn37qm/fvpKkP//8U6tXr9bq1as1aNAgxcfHKygoSH/88Ue2PrhWrVr6/fffM1zn4eGRrW0BAO6cdevWKTU1NcN1JUqUuMNpAACOmDx5st56660M1xUuXPiOZgEAOKZbt25q0aJFhuvy5Mlzh9MAABxRrly5TI95u7m53eE0AIB/c6gh+m+1a9dWamqqUlJSFBcXp40bN2r37t3Z/mBvb28FBgZm+30AAGuVLl3a6ggAgGwqWrSo1REAANmUP39+5c+f3+oYAIBs8PT05Jg3AORSDjVEIyMjtXr1aq1atUqrV6/W2bNn5ebmpnvuuUcvv/xypmcsAgAAAAAAAAAAAICVHGqIli5dWsYYBQYGqm3btmrZsqVatGghf3//250PAAAAAAAAAAAAAG6aQw3RqVOnqmXLlqpRo8btzgMAAAAAAAAAAAAAOcahhugLL7xg+/ns2bM6f/68ihUrphIlSty2YAAAAAAAAAAAAABwq9wdHbhy5UrVrl1bJUuWtPvfVatW3c58AAAAAAAAAAAAAHDTHLpCdP369Wrfvr0qVKigl156SSVLltTZs2e1ZMkStWvXTr/99psefPDB250VAAAAAAAAAAAAALLFoYbo22+/rV69eumbb76Rh4eHbfl7772n/v3766233tLatWtvW0gAAAAAAAAAAAAAuBkONURDQkK0f/9+u2aoJHl4eGjChAmqU6fObQkHAAAAAAAAAAAAALfCoWeIpqSkKE+ePBmuy5s3r1JSUnI0FAAAAAAAAAAAAADkBIcaorVq1dK0adMyXDd9+nTVqlUrR0MBAAAAAAAAAAAAQE5w6Ja5I0eOVO/evbV79249/vjjKlmypM6ePavFixfrl19+0Y8//ni7cwIAAAAAAAAAAABAtjnUEH3qqad07tw5jRs3TsuXL7ct9/X11SeffKKePXvetoAAAAAAAAAAAAAAcLMcaoheunRJTz/9tIYMGaKQkBBdvHhRRYsWVYMGDZQ/f/7bnREAAAAAAAAAAAAAbopDDVE/Pz/NmzdP3bt31yOPPHK7MwEAAAAAAAAAAABAjnB3ZFCpUqXUvHnz250FAAAAAAAAAAAAAHKUQw3RAQMGaN68ebc7CwAAAAAAAAAAAADkKIdumdurVy+NHz9e+/fvV8eOHRUQECAvLy+7MdWrV78tAQEAAAAAAAAAAADgZjnUEK1du7bt5//+978ZjjHG5EwiAAAAAAAAAAAAAMghDjVEZ8+efbtzAAAAAAAAAAAAAECOc6gh2r9//9scAwAAAAAAAAAAAABynrvVAQAAAAAAAAAAAADgdqEhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgshxuiqamp+uWXX/TKK6/o6aefti1fu3atkpOTb0s4AAAAAAAAAAAAALgVno4Munr1qtq3b68NGzYob968io+P19dffy1J+vrrr3XhwgV17979tgYFAAAAAAAAAAAAgOxy6ArRN954Q2fOnNHmzZsVGxtrt27w4MH68ssvb0s4AAAAAAAAAAAAALgVDl0humDBAv3000+6//77062rVauWduzYkePBAAAAAAAAAAAAAOBWOXSF6Llz53TPPffYXru5udl+9vb2VkJCQs4nAwAAAAAAAAAAAIBb5FBDtGTJkgoNDbW9/ndDNCQkROXKlcv5ZAAAAAAAAAAAAABwixxqiHbs2FHPPvusTpw4Ybf89OnTevnll9W5c+fbkQ0AAAAAAAAAAAAAbolDDdFx48bp8uXLqlatmh588EEZY9S2bVtVq1ZN8fHxeu211253TgAAAAAAAAAAAADINk9HBvn7+2v79u2aMGGCgoOD5ePjo3379qlPnz4aP368ChcufJtjOu7S5yPk5uVtdQynl3LmmCTpwoQuFidxDcxnzmEuc1Zq9BmrI/yfheMlT+r3LbkSJblJ8i1udRLXcOGkUkS9ySkpF8/ITZK7Xymrozg9arcLon7nHOYyZzGfOSfmvNUJbDhukjPYt8k5zGXOYj5zVq7a9wYA3BKHGqKSVLx4cX366af69NNPb2ceAMBdyt3LS25eXlbHcGop7m7//OHrzTzmhJS0/01MsjSHy0g1Mm7MZ44wVgf4P9TunEH9zjnMZc5iPnNOipvVCf5PalKSZHJRIGfFvk3OYS5zFvOZs3LRvjcA4NY41BCNi4vTsmXL1K1bt3TrFixYoHbt2ilfvnw5Hu5mFB76iTxLV7Y6htNLuxqm6JuLLU7iGpjPnMNc5qzz/2mn1PPhVseQRP1G7nNhQpd/DiIM+tzqKIC9z5+WLkdanUIStRsAHJWb9rv1xDipeDmrUwCAc8hF+94AgFvj0DNEP/jgA/31118Zrjtw4ICmTJmSo6EAAAAAAAAAAAAAICc41BCdO3eunnzyyQzXPfnkk/rhhx9yNBQAAAAAAAAAAAAA5ASHGqInTpxQiRIlMlxXokQJhYWF5WQmAAAAAAAAAAAAAMgRDjVES5UqpZCQkAzXhYSEyN/fP0dDAQAAAAAAAAAAAEBOcKgh2qFDBz333HM6duyY3fJjx47pueeeU4cOHW5LOAAAAAAAAAAAAAC4FZ6ODHrzzTcVHBysGjVqqFGjRgoICNCpU6e0fft2lS5dWuPGjbvdOQEAAAAAAAAAAAAg2xy6QtTf31/bt2/XwIEDdfLkSS1dulQnTpzQoEGDtH37dm6ZCwAAAAAAAAAAACBXcugKUUkqUaKE/vvf/97OLAAAAAAAAAAAAACQoxy6QhQAAAAAAAAAAAAAnFGGV4ju3LlTknTffffZvc5K2lgAAAAAAAAAAAAAyC0ybIg2aNBAkmSMsXudlbSxAAAAAAAAAAAAAJBbZNgQXbx4cZavAQAAAAAAAAAAAMAZZNgQ7dy5c5avAQAAAAAAAAAAAMAZuDsyaN68ebe0HgAAAAAAAAAAAACs4FBDtFevXre0HgAAAAAAAAAAAACs4FBDNCsJCQny8PDIiSwAAAAAAAAAAAAAkKMyfIaoJO3ZsyfL19I/zdDly5crMDAwp3MBAAAAAAAAAAAAwC3LtCFat27dLF+ncXd315QpU3I2FQAAAAAAAAAAAADkgEwboj/++KPt5169etm9TpM/f37VrFlTlSpVuj3pAAAAAAAAAAAAAOAWZNoQ7dmzp+3nsLAwu9cAAAAAAAAAAAAA4AzcHRlUvnz5LNfPmzcvJ7IAAAAAAAAAAAAAQI5yqCHaq1evW1oPAAAAAAAAAAAAAFZwqCGalYSEBHl4eOREFgAAAAAAAAAAAADIUZk+Q3TPnj1Zvpb+aYYuX75cgYGBOZ0LAAAAAAAAAAAAAG5Zpg3RunXrZvk6jbu7u6ZMmZKzqQAAAAAAAAAAAAAgB2TaEP3xxx9tP/fq1cvudZr8+fOrZs2aqlSp0u1JBwAAAAAAAAAAAAC3INOGaM+ePW0/h4WF2b0GAAAAAAAAAAAAAGfg7sigMWPG3O4cAAAAAAAAAAAAAJDjHGqIStKsWbPUoEEDFSlSRAUKFEj3HwAAAAAAAAAAAADkNg41RD/77DMNHz5clStX1qVLl9SrVy81btxYCQkJeuihh9S/f//bHBMAAAAAAAAAAAAAss+hhujMmTP12Wef6ccff7S9XrNmjY4cOaK4uDj17t37toYEAAAAAAAAAAAAgJvhUEP00KFDeuKJJ2yvU1JSJEnly5fXtGnT9PLLL9+edAAAAAAAAAAAAABwCxxqiCYkJKho0aKSpHz58uns2bO2dZUqVdLu3btvTzoAAAAAAAAAAAAAuAUONUT/rWbNmvrpp59sr3/55Rf5+fnddIDIyEiFhIQoIiJCxpib3g4A4M6JiYnRrl27dOTIESUlJVkdBwDggMTERB04cEChoaG6evWq1XEAAA4wxigsLEw7duxQVFSU1XEAAA6KiopSSEiITp48qdTUVKvjAAB0Ew3RAQMGaOTIkWrVqpXat2+vp556Sr169cr2B3/00UeqUaOGqlWrpueee05BQUGqXr26li1blu1tAQDujOXLl6t58+YqVqyYhgwZolatWsnf31/vv/8+J7UAQC71119/qW/fvipcuLC6dOmiPn36qGjRoho0aJAuXrxodTwAQAYuX76s1157TaVKldL999+v4cOHq0KFCmrWrJlCQ0OtjgcAyMTnn3+ue++9VxUrVtRzzz2nBg0aqFKlSlqwYIHV0QDgrudQQ3Tbtm22n4cNG6Z33nlHERER+vvvv/Xyyy9rwoQJ2f7g0aNHq1u3bjp9+rTtCtFGjRqpY8eO2rRpU7a3BwC4/T766CMVKFBAx44d065du3T8+HG9//77evXVV2/qdwEA4PZbunSpduzYoTVr1ujQoUPau3evtm3bpp9++kmdO3e2Oh4AIANHjhzRjBkzNGnSJEVERGjHjh06fvy4EhIS9Mgjj9g9yggAkHu8/vrratGihU6dOmU75t2uXTt1796dC4EAwGIONUQbN25s+9nNzU1jxozRgQMHdODAAb377rvy8fHJ9gcvXLhQb7/9tvLnzy9JypMnjz7++GOlpqbq66+/zvb2AAC33wsvvKAlS5YoMDDQtmzw4MEKCgrSzJkzLUwGAMhM48aN9fvvv6tJkya2ZfXq1dPgwYO1adMmHTp0yMJ0AICM+Pn5afPmzXr66afl4eEhSSpevLjGjx+v6Ohou0cZAQByj9mzZ+vDDz9UwYIFJUleXl6aMmWKvLy8NGvWLIvTAcDdzdOqD+7UqVO6Zd7e3nJzc9OlS5fufCAAwA117Ngxw+U+Pj7UbgDIpR5++OEMl6ed1Ej9BoDcp2LFihkup3YDQO6W0TFvDw8PeXp6UrsBwGIZNkR37tyZ7Q3dd999txzmq6++kjFGDzzwQJbjzp07p6ioKLtlR48eveXPBwBk3969exUSEqLWrVvfcCz1GwByh2vXrmnu3LkqWLCg7rnnnizHUrsBIPdIu7qI4yYA4Dzmzp2r+Ph4ajcAWCzDhmiDBg2yvSFjzC0F2b9/v9544w1VqFBBzzzzTJZjZ8yYofHjx9/S5wEAbl1cXJz69OkjT09Ph54hSv0GgNzhxRdfVFhYmD788EPlzZs3y7HUbgDIHX766SfNnTtXbdu2VbNmzbIcS+0GgNzh+PHjGjVqlEqWLKkXX3wxy7HUbgC4vTJsiC5evPiOhjh58qQee+wxeXp6auHChSpQoECW44cPH65u3brZLTt69Kg6d+58G1MCAP4tKSlJ3bp10969e/XFF184dKcA6jcAWG/y5Mn64osv1KNHD40cOfKG46ndAGC9TZs2qW/fvqpWrZq+/fbbG46ndgOA9c6dO6c2bdooISFBP//8s4oWLZrleGo3ANxeGTZE72SRPXPmjFq0aKHo6GitWrVKdevWveF7/P395e/vfwfSAQAykpKSol69emn58uWaOnWqhgwZ4tD7qN8AYK1p06ZpzJgx6tKli77//nu5ubnd8D3UbgCwVkhIiNq1a6eAgACtXbtWxYsXv+F7qN0AYK2LFy+qZcuWOnnypH755Zcb3i5XonYDwO2WYUP0TomKilLz5s115swZ/frrr2rcuLGVcQAADkhNTVXfvn21cOFCffTRRxoxYoTVkQAADpg5c6ZGjBihTp066X//+588PS39UwAA4IDdu3erdevWKl68uNatW6eAgACrIwEAbuDy5ctq1aqVDh8+rKVLl6p58+ZWRwIASHK36oOjo6PVokULhYeHa8WKFWrSpIlVUQAADjLGaPDgwfrhhx80ZcoUh261CACw3pw5czR06FB16NBB8+fPl5eXl9WRAAA3sH//frVq1Up+fn5av369AgMDrY4EALiBq1evqm3bttq/f7+WLFmiVq1aWR0JAPD/WXJaeFJSktq0aaO9e/dqwoQJSk5O1vr1623rCxYsqHr16lkRDQCQhdGjR+vrr79Wly5dVL9+fbvaLUkPPvigPDw8rAkHAMjQsmXLNGDAAFWsWFHPPvustm7dare+du3aKlasmEXpAAAZiYiIUMuWLRUbG6uPPvpIx44d07Fjx2zrAwICVKVKFQsTAgCuZ4xRx44dtW3bNo0ZM0Y+Pj52x03y5cunhg0bWhcQAO5yljRE4+LilDdvXjVr1kxr1qzRmjVr7NbXrFlTM2bMsCIaACAL0dHRatasmS5evKi33nor3foVK1Yob968dz4YACBTJ0+etD2zaNKkSenWT5w40aFnGgEA7pzTp0+ratWqkqSvvvoq3fouXbrw6AoAyGVSU1OVmpqqZs2aadu2bdq2bZvd+rJly+q7776zKB0AwJKGaKFChdJdVQQAyP1mzZpldQQAQDYNGzZMw4YNszoGACAbGjZsyHETAHAyHh4e1G4AyMUse4YoAAAAAAAAAAAAANxuNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABcFg1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABcFg1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABcFg1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABcFg1RAAAAAAAAAAAAAC6LhigAAAAAAAAAAAAAl0VDFAAAAAAAAAAAAIDLoiEKAAAAAAAAAAAAwGXREAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZnlYHyGmXPh8hNy9vq2M4vZQzxyRJFyZ0sTiJa2A+cw5zmbNSo89YHcGG+n3rUi6ekZskd79SVkdxCWn1RrOGWhvEVVyJktwk+Ra3OonzizlvdQIbanfOoH7nHOYyZzGfOSc37XcDAAAAdyOXa4gCAJxTalKSZNysjuHcUo2Mm5SSmGR1Epfi4e1ldQSXkOLu9s9BdebzlqXkolJJ7c4h1O+cw1zmLOYz5xirAwAAAAB3N5driBYe+ok8S1e2OobTS7v6ruibiy1O4hqYz5zDXOas8/9pp9Tz4VbH+McT46Ti5axOAfyfWUPl4e1FvUGuQ+0GACf0+dPS5UirUwAAAAB3LZ4hCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LJoiAIAAAAAAAAAAABwWTREAQAAAAAAAAAAALgsGqIAAAAAAAAAAAAAXBYNUQAAAAAAAAAAAAAui4YoAAAAAAAAAAAAAJdFQxQAAAAAAAAAAACAy6IhCgAAAAAAAAAAAMBl0RAFAAAAAAAAAAAA4LI8rfzwM2fOKDg4WIcPH1aBAgVUp04dderUSZ6elsYCAGQhPj5ewcHB2rt3r5KTk1WxYkV169ZNRYoUsToaACATxhitW7dO27dvV1RUlAIDA9W+fXtVrVrV6mgAgCzs379fa9asUVhYmIoXL677779fjzzyiNWxAABZiIqK0s8//6xDhw4pT548qlmzph5//HF5e3tbHQ0A7mqWXSE6adIkPfroowoNDVVAQIDi4uI0fPhwVa1aVUeOHLEqFgAgC2vXrlXt2rUVHBwsX19f5c+fXzNmzFCZMmW0cOFCq+MBADJw4cIF1alTRx9++KGuXbumUqVKaf369apRo4Zeeuklq+MBADLRvXt3DRgwQOHh4SpbtqxOnDihdu3a6aGHHlJMTIzV8QAAGfjss8/UtGlT7dixQyVLllRycrJefvllVaxYUXv37rU6HgDc1Sy7FLNLly567bXX5ObmZls2dOhQVatWTa+++qoWLVpkVTQAQCYqV66svXv3Kn/+/LZlr776qurXr6/BgwdzlT8A5EI+Pj5avny5ypYta1v2yiuv6Nlnn9VHH32kHj16qGHDhhYmBABk5K233lLNmjXtlrVp00aPP/64pk2bptdff92iZACAzLRs2VLDhg2Th4eHbdnzzz+vKlWq6MUXX9Rvv/1mYToAuLtZdoVojRo17JqhklSxYkWVLl1aJ0+etCgVACAr5cqVs2uGSpK3t7caNmyo6OhozlQHgFyoQIECds3QNA888IAkse8NALnU9c1QidoNALldtWrV7JqhklSqVClVqlSJ2g0AFrOsIZqRnTt3KiIiQm3btrU6CgDAQZcvX9a6devUuHFjniMKAE7CGKOlS5eqQIECtoPrAIDcb8mSJZLEcRMAcCJ//fWXDh06RO0GAItZfl/D1157TRcuXNCpU6e0a9cuTZw4UaNHj87yPefOnVNUVJTdsqNHj97OmACAf/n222+1bds2RUdHa/369WrXrp3ef//9G76P+g0A1gkJCdHXX3+tuLg47dixQyVKlNCGDRtUsmTJLN9H7QYA61y5ckWjR49WcnKyDh8+rNOnT2vu3Lnq1KlTlu+jdgOAtSZMmKCIiAhFRkZq27ZtGj16tN54440s30PtBoDby/KGaK1atXT58mUVKlRIO3fu1M8//6wePXqoYsWKmb5nxowZGj9+/B1MCQD4t7Jlyyo+Pl6RkZE6dOiQVq1apW7dut3wbEfqNwBYp0iRIgoKCtKVK1d06dIlrVu3TsuWLVO9evWyfB+1GwCs4+XlpaCgICUkJMjDw0O7d+/WkiVL1L59exUsWDDT91G7AcBa1apVU9GiRVWkSBHt2rVLwcHB6tmzZ4a3Q09D7QaA28vyhmjv3r1tP7/yyiu655571L17d+3cuTPT9wwfPlzdunWzW3b06FF17tz5dsUEAPzLI488okceeUSS9J///Edt27ZV165ddeTIEZUuXTrT91G/AcA6VapUUZUqVSRJo0eP1uTJkzVmzBjVqFFDXbt2zfR91G4AsE7evHk1dOhQ2+uBAweqSZMm8vPz0+eff57p+6jdAGCt7t27235+7bXXFBQUpC5duuivv/6Su3vGT7GjdgPA7WV5Q/Tf/P399dhjj2n27Nm6ePGi/Pz8Mh3n7+9/h9MBADLi7u6ufv36adWqVdq8ebPdTv/1qN8AkHsMGDBAY8aM0apVq7JsiFK7ASD3aNiwoWrWrKlVq1ZlOY7aDQC5R6FChdSlSxd9/PHHCgsLy/TOiNRuALi9Mj4dxUIxMTHy8PCQj4+P1VEAAA6KiYmRJOXPn9/iJAAAR1G7AcA5xcTEULsBwMmk7Xvny5fP4iQAcPeypCGanJysBQsWyBhjt3zDhg0KDg5Wly5d2LkHgFxo6dKlunr1qt2yM2fOaMqUKQoICLDdRhcAkHts3rxZJ06csFuWkJCgsWPHyt3dXU8++aRFyQAAmQkLC9OWLVvSLZ8+fbrCwsLsHj8EAMg95s2bp5SUFLtlO3fu1I8//qgWLVqoZMmSFiUDAFhyy1w3NzctXbpUr7/+umrUqKEiRYro6NGj+v3339W1a1fNmjXLilgAgBs4fvy4goKCVL58eQUGBurs2bNav369atSooSVLlnCmIwDkQomJierQoYMKFCigSpUqKTY2Vtu2bVNqaqr+97//qUGDBlZHBABcx8vLS2+99ZYiIyNVvXp1eXl5ae/evTp69Khee+01vfLKK1ZHBABk4LffftMbb7yh6tWrq1ixYjp+/Lg2b96sxx57TN98843V8QDgrmZJQ9TDw0Pff/+9zp07p127dun06dPq3LmzGjZsqICAACsiAQAc8OKLL+qZZ55RSEiIjh8/rjx58uiDDz5Q7dq1rY4GAMjEo48+qtDQUO3evVsHDx5UcnKyRo0apUaNGsnLy8vqeACADAQEBGj16tX6+++/FRoaqujoaPXp00f333+/ChcubHU8AEAmvvzyS124cEE7d+5URESEOnTooO+++07lypWzOhoA3PUsaYim8ff3V9u2ba2MAADIprx586pZs2Zq1qyZ1VEAAA5yc3NTvXr1VK9ePaujAACyoWLFiqpYsaLVMQAA2VC0aFG1bt3a6hgAgOtY8gxRAAAAAAAAAAAAALgTaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyPK0OkFMSEhIkSdGfDZebp5fFaZxfavQZSdL5/7SzOIlrYD5zDnOZs1Ivnpb0fzXUCrbPnv+G5EH9Ri4Sc14pbtQb5D7UbgBwQleiJFG7AcDpWFy/OeYNANmX2XETl2mIhoeHS5LMpUgZi7O4ktTz4VZHcCnMZ85hLnNWeHi46tWrZ9lnS5KuXrDk84Ebod4gt6J2A4DzoXYDgHOyqn5zzBsAbt71tdvNGOMStfTSpUvasGGDypQpIx8fH6vjZOro0aPq3LmzlixZosqVK1sdx+kxnzmHucxZzjKfCQkJCg8PV7NmzVS4cGFLMjhD/XaWf09nwXzmLOYz5zjLXFK7Hecs/6bOgLnMWcxnznKG+aR2O84Z/j2dCfOZc5jLnOUs82l1/aZ2352Yz5zFfOYcZ5nLzGq3y1whWrhwYXXq1MnqGA6rXLmyatWqZXUMl8F85hzmMmc5w3xadYZ6Gmeq387w7+lMmM+cxXzmHGeYS2p39jjDv6mzYC5zFvOZs3L7fFK7sye3/3s6G+Yz5zCXOcsZ5tPK+k3tvrsxnzmL+cw5zjCXGdVudwtyAAAAAAAAAAAAAMAdQUMUAAAAAAAAAAAAgMuiIQoAAAAAAAAAAADAZdEQvcOKFy+ucePGqXjx4lZHcQnMZ85hLnMW8+la+PfMWcxnzmI+cw5z6Xr4N805zGXOYj5zFvPpWvj3zFnMZ85hLnMW8+la+PfMWcxnzmI+c46zz6WbMcZYHQIAAAAAAAAAAAAAbgeuEAUAAAAAAAAAAADgsmiIAgAAAAAAAAAAAHBZNEQBAAAAAAAAAAAAuCwaogAAAAAAAAAAAABclqfVAe4WsbGx+uSTT7RhwwalpKSoYcOGGjVqlIoVK2Z1NKdjjNHvv/+u+fPna/v27apQoYLmzp1rdSynZIzR+vXrtWjRIh06dEgFChTQvffeq+eee05Fixa1Op7TSU1N1YoVK/Tzzz/r6NGjKlSokIKCgvTMM8+oRIkSVsfDTZo7d64WLlyoixcvqmrVqnr22Wd17733Wh3LKf3999+aP3++Vq5cqWvXrmn58uXy8/OzOpZTOnTokL7//nvt3btXiYmJqlKlioYMGaLatWtbHc0p7d27V3PmzNG+ffvk5uamypUra+DAgfx/3Ylt27ZNX3zxhY4fP64SJUroySefVOfOna2O5ZSio6O1ZMkSLVq0SFFRURo/frxat25tdSyndOHCBX3//ffavn27zp49q3LlyqlLly7q0KGD1dGc0tmzZ/XNN99ox44dio6OVtmyZfX444+rffv2cnNzszoebkJERIQ+/vhj7d69W3ny5FHLli317LPPytvb2+poTicpKUmrV6/W/PnzdfDgQbVr105vvvmm1bGcUmJiopYsWaJVq1bp+PHj8vf3V9OmTTV48GD5+PhYHc/pXLt2TfPmzdO6det08uRJlShRQg8++KAGDBigfPnyWR0PN+HatWuaPn261qxZo8TERNWtW1cvvfSSSpUqZXU0p7Rr1y7Nnz9fmzdvVtGiRfXzzz9bHclpbdmyRT/99JP++usv+fj4qHbt2nr++edVsmRJq6M5pTVr1mjRokU6fPiwChQooDp16uiZZ55RYGCg1dGyhStE74CYmBg1adJE33//vYYMGaJRo0Zpw4YNql+/vs6cOWN1PKfTsmVLjRo1SmXLltWpU6cUGhpqdSSn1adPH7Vq1UrGGL388svq1auXli9frkqVKikkJMTqeE6nb9++mjNnjho1aqQ33nhDnTt31oIFC1SlShXt2rXL6ni4CYMHD9awYcPUokULjRs3TsYYNWjQQKtXr7Y6mtOZNGmSWrVqpUuXLsnX11fbt29XYmKi1bGc0uzZs1W9enXt3LlT/fr10wsvvKDLly/r3nvv1bRp06yO53Rmzpyp4cOHKyAgQKNGjdIzzzyjkydPKigoSDNmzLA6Hm7Cjz/+qAceeEB+fn4aP368GjdurJ49e3Ig+CZs2bJF1atX18aNG1W9enVt375dUVFRVsdySseOHVNAQIC++OILPfLII3r99ddVoUIF9ejRQ127dpUxxuqITuXQoUNq1qyZ4uPj1bdvX7366qsqUaKEunTpoq5du1odDzfh8OHDCgoK0p9//qlXX31Vffv21aeffqqWLVsqKSnJ6nhOJSUlRWXKlNFnn32mxo0ba/v27Tp27JjVsZzWPffco6FDh6pcuXIaO3asWrRooQ8++EB16tRRZGSk1fGcTpMmTbRjxw61bt1a//nPf9SkSRNNmDBBderU0fnz562Oh2y6du2aHn30UU2fPl39+/fX6NGjFRoaqqCgIP39999Wx3M6TzzxhIYMGaKiRYsqOjpaf/zxh9WRnNZzzz2nhx56SDExMXrxxRfVv39/bdmyRZUrV9a6deusjud0hg8frhkzZqhu3boaO3asunfvrpUrV6pq1arasGGD1fGyx+C2e/31142Hh4c5evSobVl0dLQpXLiw6du3r4XJnNOFCxdsP1erVs3UqlXLwjTObdCgQSYkJMRu2ZUrV0zJkiVNw4YNLUrlvGJiYtIti4yMNO7u7uaJJ56wIBFuxW+//WYkmZkzZ9otb9GihSlbtqxJSkqyKJlzunjxou3nZ5991kgyZ86csTCR8/rss8/Ml19+mW55p06djLe3t4mKirIglfPKqHanpqaaoKAgU7hwYQsS4VZcvnzZFClSxPTo0cNu+bvvvmvc3d3Nn3/+aVEy5xQTE2MSExONMcYsXrzYSDJz5syxOJVz+vPPP83AgQPNtWvX7JZPmzbNSDKLFi2yKJlzio+Pz3BfbNSoUUaS2bFjhwWpcCtatWplAgMDTVxcnG3Zjh07jCTz6aefWpjMOaXte8fExBhJpl+/ftYGcmLt2rUzERERdssOHDhgJJnhw4dblMp5ZbTvvXz5ciPJvPvuuxYkwq14//33jSSze/du27LY2FhTunRp07FjR+uCOal/H/Nu1KiRCQgIsDCNc3vhhRfM+vXr7ZZdu3bNVKpUyVStWtWiVM4ro9p96dIlkzdvXtO8eXMLEt08rhC9A7777js98MADqlSpkm1Z4cKF1blzZ82fP19xcXEWpnM+3F4x50yfPl0NGjSwW+br66u6detyFtJNKFCgQLplJUqUUP78+XX16lULEuFWfPfdd/Ly8tKTTz5pt7x///46efIkZ5RlU5EiRayO4DIGDx6swYMHp1v+0EMPKTExUfv27bMglfPKqHa7ubmpfPnyiouLU0pKigWpcLOCg4MVHR2t/v372y0fMGCAUlNTNWfOHGuCOakCBQrIy8vL6hguoWrVqpo1a1a62ys+9NBDksTdRLIpT5488vRM/wSgihUrShL73k4mMjJSq1evVo8ePZQ3b17b8vvuu0+1a9fWN998Y104J8W+d85ZtGiRAgIC7JbVqFFDxYsXp3bfhIz2vandzuu7777Tvffeq6CgINuyfPnyqXv37lq2bBlX/WYTx7xzzgcffKBmzZrZLfPx8VHDhg11+PBh6k02ZVS7CxUqpKJFizrdXNIQvc3Onj2riIgI3XPPPenWBQUF6dq1a9q/f78FyQBl+CyW5ORk7d27V/7+/hYkci3GGH322We6evWqnn76aavjIJt27typypUrp3uOSdqO/s6dOy1IBWRcu6X/O5jOM4tvXUhIiNasWaMBAwbIw8PD6jjIhrTafP2+d4kSJVSyZElqNyxD7b79zp07p1mzZqlmzZq6//77rY6DbNi1a5eMMZkeN0l7ZjpghYzq94kTJ3ThwgVqdw5ISEjQhx9+qHz58qU7GRm5W3x8vA4cOJBp7U5JSdHu3bstSAZkXLuNMdq9e7d8fX15ZnEO+O677xQREaGBAwdaHSVb0p9SiRyV9jyB4sWLp1tXrFgxSeI5oshVJk+erFOnTmncuHFWR3Fajz76qK5evarw8HClpqZq2bJlatu2rdWxkE2RkZGqWbNmuuXUbuRG27Zt07x589S0aVPVqFHD6jhO6dVXX9WGDRt04cIFnThxQuPHj9err75qdSxk0432vandyE2io6M1btw4+fr6qkePHlbHcUo//fSTpkyZotjYWB05ckQdO3bUl19+me5KXORuN6rdycnJOn/+vEqXLn2nowEZGjlypFJTUzO8awtu7MSJE+rRo4eSkpL0999/q0yZMtq+fXuGf38j9zp37pxSU1M55g2nMWPGDB08eFAjRoyQuzvXCd6Mjh07KjIyUqdPn1ZcXJwWLFigrl27Wh0rW2iI3mZJSUmSlOHVBWm3+EkbA1htxYoVGjdunOrXr6/XXnvN6jhOa/LkyUpISNCxY8f08ccfa+DAgQoODlb9+vWtjoZsSEpKonbDKYSHh6tbt24qUKCAvv76a6vjOK2BAweqc+fOOnPmjH788UeNGzdOxYsX16BBg6yOhmy40b73tWvX7nQkIENJSUnq1auXwsPDNWfOHO7OcpMefPBBBQYG6tKlS1q3bp0+/fRTvfjii5o9e7bc3NysjgcHcdwEzmTSpElavHix+vfvr/bt21sdxymVKFFCU6dO1bVr17R3715NnjxZffv21fLly1WyZEmr48FB1G44k02bNumll15S9erVNXHiRKvjOK3x48crLi5OYWFhmjZtmoYOHaoSJUrowQcftDqaw2iI3maFChWSlPF98GNiYuzGAFbauHGjnnjiCVWtWlXLly/nrOpbkPZc1gceeECPP/64atSooaefflqhoaEWJ0N2FCpUiNqNXO/s2bNq0aKFLl26pJUrV6pq1apWR3Ja/567xx9/XE888YSeffZZtW3bNt1zo5B7/Xvfu2DBgnbrYmJiVLhwYQtSAfZSUlLUu3dv/frrr3r//ffVu3dvqyM5rRIlSthuWdmmTRuVK1dOzz77rFq3bq1evXpZnA6O4rgJnMV///tfjR07Vu3bt9eXX35pdRynlSdPHjVu3FiS9PDDD6tVq1aqVauWxo4dq6+++sridHAUtRvOYteuXWrfvr1Kly6tVatWZfg8TDimbt26kqSmTZvqiSee0L333qs+ffooLCzM2mDZwLXBt1n58uXl7e2tv//+O926tGXVqlW707EAO7///rvatWunsmXL6rfffuMM9Rzk6+urpk2bau/evUpISLA6DrKhWrVq1G7kaufPn1fz5s0VERGhZcuW6YEHHrA6kktp06aNEhMTOZnFyaTV5uvrd1JSkiIiIqjdsFxqaqoGDBig+fPna9KkSXrllVesjuRS0h5TsWPHDouTIDsyq91py/z9/TmhBZabPXu27WS5n376SV5eXlZHchnVq1dX+fLlqd1Opnjx4vLz8+O4CXK1vXv3qlWrVipSpIh+++03lSlTxupILiNPnjx6+OGHdeLECZ07d87qOA6jIXqbeXl5qUWLFlq/fr2Sk5Pt1q1atUq1a9fm/4iw1K5du9SmTRuVLl1a69at4/Ykt8GJEyfk5+fHVbdOpm3btrpw4YJ2795tt3zVqlVyd3dX69atLUoGSJcuXVKrVq10/Phx/fLLL2rWrJnVkVzOiRMnJEmlSpWyOAmyI60Zsnr1arvlGzZsUEJCgh577DErYgGSJGOMhg4dqjlz5mjixIk8ouI2oHY7p6CgIJUsWTJd7b569aq2bt1K7YblfvjhBw0aNEitW7fW4sWL+ds+hyUkJCgyMpLa7YTatGmjrVu3Ki4uzm75qlWrVKZMGdWuXduiZID0119/qUWLFipQoIDWrVun8uXLWx3J5Zw4cUJ58+Z1qqvBaYjeAWPHjtX58+c1btw427Ivv/xSu3bt0ltvvWVdMNz19u3bp9atW6tEiRJat24dO5+34OLFi5owYYIuXbpkW5aYmKh33nlH27dv18iRI60Lh5syaNAgBQQEaNSoUYqNjZX0z5llM2bM0DPPPKPSpUtbnBB3q6tXr6pNmzY6dOiQfvnlFz3yyCNWR3JqkyZN0pEjR+yWrVixQlOnTtXDDz9suyUMnEOdOnXUpUsXTZkyRUePHpX0zwkEr776qmrWrKnu3btbnBB3s5EjR2rmzJmaOHGixo4da3Ucp/b999/rt99+s1t28OBBvfDCCypatKj69OljUTLcDHd3d7355ptavXq1Fi5cKOmfq6lfeeUVpaamasyYMRYnxN1s8eLF6tevn1q1aqUlS5bQDL0FO3bs0FdffaXExETbsosXL2rgwIGKj4/XiBEjLEyHmzFmzBglJCRo9OjRSk1NlSTNnz9fa9eu1VtvvcXzvGGZY8eOqUWLFsqbN6/Wr1+vChUqWB3JaSUmJtr6W2mSk5P16aef6tdff9Xw4cOd6nejmzHGWB3ibrBw4UI9++yz8vHxkY+Pj86dO6d3331Xw4YNszqa0/nwww+1YMECSVJoaKjc3Nx0zz33SJJ69uypF1980cJ0zuWhhx7Spk2bVLlyZRUtWjTd+l9//dWpzvCwUlJSkt5//319/vnn8vDwUMGCBRUWFqaCBQvqlVdeYcfeSR04cEB9+vTR8ePHVaZMGR0+fFi9e/fW9OnT5e3tbXU8p7Jx40aNHj1akhQWFqazZ8+qXr168vLyUqVKlTR37lyLEzqPt99+W+PGjVPx4sVVsWLFdOsnTpyoFi1aWJDMOS1YsEDvvfeewsPDVbp0aZ05c0ZxcXHq06ePJk+eLF9fX6sjIptiYmI0aNAg/fzzz6pWrZqOHz+ue++9V3PmzFG5cuWsjudU4uPjbSddREdH6/Dhw6pUqZKKFSsmSVq+fLn8/PysjOg0tm7dqqZNm8rHx0dBQUHp1nfq1IkrRrNh9+7dGj9+vNavX6+yZcvqypUrOnXqlFq0aKGPP/5Y1atXtzoibsI777yjd999V2XLltXly5fl5eWlWbNmsV9zEwYNGqR9+/YpNTVVO3bsULFixVSpUiVJ0vjx47nbTTbky5dP8fHxCgoKSnfAt1SpUlq8eLFFyZxPVFSUJkyYoLlz56p48eJyc3PT8ePHVaVKFU2aNEkdOnSwOiJuwooVK/TMM88oNTVVvr6+ioiI0H/+8x8eC3ATvvjiC82ePVuStH//fiUmJtpO0G3Xrp3efPNNK+M5lY4dOyo4OFgVKlTI8NFwP/30kwIDAy1I5nyMMfrwww81ffp0paSkqEiRIjpx4oTy5MmjESNG6NVXX5W7u/Ncd0lD9A5KSUnRoUOHlJKSoqpVqzpV5zw3CQsLU2RkZIbrSpUqxYGubNi/f7/tQecZue++++Tp6XkHE7mGiIgInT17VsWLF1fZsmWtjoMcEBYWpujoaFWoUIHnF92k6OhoHTp0KMN1+fLls53YghuLiIhQREREpuurVKmS4UkuyFp0dLTCwsKUL18+VapUid9/LuD8+fM6efKk/P39+WP3JqWmpiokJCTT9fXr1+cZag66cuWKDhw4kOl6f3//DE9yQdauXbumo0ePyhijihUrKn/+/FZHwi2Kj4/XkSNH5OPjo6pVq3J10U3K6m/9ypUr205swY1t375dmR069fHx4W4iNyElJUXHjx9XTEyMAgICMmxWwLmkpqbq8OHDSkxMVJUqVZQ3b16rIzmlrP7WL168uO3EFtzYwYMH7e7id726devSm7kJZ86c0enTp+Xn56fy5cs75X4aDVEAAAAAAAAAAAAALst5rmUFAAAAAAAAAAAAgGyiIQoAAAAAAAAAAADAZdEQBQAAAAAAAAAAAOCyaIgCAAAAAAAAAAAAcFk0RAEAAAAAAAAAAAC4LBqiAAAAAAAAAAAAAFwWDVEAAAAAAAAAAAAALouGKAAAAAAAAAAAAACXRUMUlklISFBERIQSEhJc+jNzkrPnB+AaTp06pStXrrj8Z+YkZ88PwPmdP39e586dc/nPzEnOnh+A87t69aoiIiJkjHHpz8xJzp4fgPNLSkpSRESE4uPjXfozc5Kz54fzoCEKy2zbtk1lypTRpk2bsv3erBqDWa27lc/MDZw9PwDXUK5cOU2aNOmm3hsREZFpYzCrdbfymbmBs+cH4Px69+6txx577KbeGxUVpaioqGyvu5XPzA2cPT8A5/f555+rTJkyunz5crbfGxMTk2ljMKt1t/KZuYGz5wfg/Pbv368yZcpo2bJl2X5vYmKiIiIidO3atWytu5XPzA2cPT+cBw1ROKV169apTJky2rZtW7bW+fj4KCAgQHny5LkTMQEA/5KcnKwyZcro/fffz9Y6SQoMDFShQoVud0QAQAa6deumDh06ZHtd8eLFVaJEidsZDQCQiU8++URlypRRbGxsttb5+voqICBA7u4cMgSAO+2PP/5QmTJltHLlymyt8/b2VkBAgPLly3cnYgJOy9PqAMCddP/99ysiIsLqGACAbAoLC7M6AgAgm+bMmWN1BABANj3zzDN65plnrI4BAMiGmjVrcswbcACneyFXSU5OVkREhO2/8+fPp7uFS2xsrM6fPy/pn1t0pY2NjY3Ncp1042dwpqSk6Pz580pNTc0y47lz527qOZ6ObP/atWuKiopSSkqKQ9uMjY1VREREuvFpcxkXF2e37YiICCUmJkr651YLGd2eMikpSRcuXMg037+3kZSUZJtzAHevyMhIW809e/askpKS7NanpKTo1KlTkv7vFl0RERG6ePFiluvS3OgZnNHR0RneNiaNMUYXLlxQTEzMTX2/G20/u78bMqrRaSIjIxUdHW17bYxRRESELXtqaqrd3Px7XGa/Y67fRtrY5ORkh/ICcE3R0dG2mnv69OlMa1JCQoLtFl1pY2+0TrrxMzhjY2NveEvDmJiYTPdLb+RG27+Z3w3X1+g00dHROnPmjN2y628nHB0dneE+fnR0dKa/PzLaBs9WAu5ucXFxtpp76tQpXbp0Kd2Y6Oho277z6dOnbeOTkpKyXCfd+BmciYmJNzwGcO3aNZ07d+6m9jUd2X7a7wZHnxOaUY2W/m8u/53zypUrdk2F2NjYDOtuXFxcpn+fXL+NuLg4buEL3OVSU1PtjnlHRUWl+9s9Pj7ett934cIF29iYmJgs10k3fgZnamqqzp8/n+Xx5pSUFJ07d+6m9jUd2X5CQkK2fjdkVKPTPisiIkJXr1612/a/j/knJSVlWHeTk5Mz/dvi+m0kJydn2JuAkzOARdatW2ckmdWrV9uWHT161AQEBNj+y5s3rylcuLAZMWKESUhIMMYYM2fOHFOsWDEjyRQrVsw2ds6cOVmuy+wzjTHm5MmTpkePHiZ//vzG19fX+Pr6mt69e5tTp07ZjenWrZvJmzevKVSokPHy8jIdO3Y04eHhN/yujmz/jz/+MM2aNTMeHh7G19fXFChQwAwePNhcuXIlyzmbNm2akZQux19//WUkmdmzZ9uWLV682EgyGzduNMOGDTNFihQx7u7uplmzZiYqKsokJiaaoUOH2pYHBQWZo0eP2m03bRubN282o0aNMkWKFDGenp6mQoUKZt26dTecCwDOz8PDw7z66qt2y+677z5bzS1cuLDx9vY27du3NydOnDDGGHP8+HETEBBgJBlfX1/b2CFDhmS5LqvPvHbtmnnttddMiRIljLe3t/H19TVNmzY1v//+u92YMWPGmGLFipl8+fIZLy8vU7duXbNhw4Ybfk9Htn/hwgXTp08fky9fPuPr62s8PT1Ny5YtzV9//ZXlnGVUo9MEBASYp556yvY6JibGSDITJkwwX375pSlZsqTx8fExFSpUMNu2bTPGGPP555+bkiVLmjx58hh/f3/z888/223z39uYM2eOKV26tMmXL58pWLCg+eijj244FwCcX+vWrU39+vXtlo0cOdJWc/39/Y2Hh4epW7eu+e2332xjGjdubLy9vY23t7dtbM2aNW+4LrPPNMaY7777ztSsWdO4u7ubQoUKmYoVK5pvvvnGbsy3335rqlWrZjw9PU2BAgVMqVKlzPTp0x36rjfafkpKipk4caLx9/c3+fLlM56enqZmzZpm6dKlN5yz62t0mn79+pkSJUrYLWvUqJFp1qyZ2b59u6levbopUKCAKVCggPnvf/9rjDFm69atpmbNmsbX19fkyZPHvPnmm+m2m7aNPXv2mDp16piCBQsaT09P07t3bxMfH+/QfABwXh988IGRZKKjo23LFi1aZKu5pUuXNl5eXiYwMNBMnTrVNuaVV14xBQsWNJJM6dKlbeNDQ0OzXJfZZxpjzPbt282jjz5qPD09TaFChYyfn58ZM2aMiY2NtRvz8MMP28bkz5/fDBs2zMTFxd3wuzqy/SVLlphatWoZDw8Pky9fPuPv728mTpxoUlJSspyzjGq0Mf8cX5Jk/vzzT9uysWPHGknmzJkzpm3btqZgwYLG3d3dPP300yYpKcmcO3fOtGnTxhQsWNC4ubmZDh06mJiYGLvtpm0jKirKdOnSxRQqVMi4ubmZhg0bmmPHjt1wLgA4t927dxtJZsGCBbZl586dszvmnXaMeMiQIebq1avGGGOWLl1qihcvbiSZokWL2sZOnz49y3WZfaYxxkRGRpq+ffsaX19f22c+8cQT5vjx47YxZ8+eNb179zb58+c3BQsWNF5eXqZVq1bpjglnxJHtHzhwwLRq1cp4enoaX19fkzdvXtOnTx9z4cKFLOcsoxptjDFnzpwxkszHH39sW5Z2zHzZsmXm5ZdfNn5+fsbd3d00atTIREREmJSUFPPSSy8ZPz8/4+HhYWrUqGH27dtnt920bSxfvty8+eabpmjRosbLy8sEBASYX3755YZzAedAQxSWyaw5eb3Vq1ebIkWKmDFjxtiWrVixwkjKsAGX1bqMPvP06dOmdOnSJigoyOzZs8cYY0xcXJyZO3eumT9/vjHmn+IeEBBgGjZsaA4dOmSM+af4NmvWzFSpUsVuB/16jmz/yJEjxtfX1zzyyCPmzJkzxhhjNmzYYIoXL24efPBBk5qammn+m2mItm3b1ixatMikpqaa8PBwU758edO1a1czcuRIs2DBApOSkmJOnTplKlasaFq3bm233bRtdOjQwcybN8+kpKSYy5cvm2bNmplSpUpxYAa4C2TUnLzekSNHTIMGDUyjRo1sByiSkpKMJDN27Nh047Nal9FnpqammlatWhk/Pz+zaNEik5SUZFJSUsyWLVvMxIkTbWMee+wxU6xYMbNy5UqTmppqrl27ZkaNGmV8fHzM7t27M83vyPYTExNNvXr1TEBAgNm+fbsx5p8TYBo1amSKFi1qIiIiMs1/Mw3Rhx9+2Lz++uvm2rVr5tq1a6Zjx46mZMmSZu7cueaVV14x165dMwkJCebxxx83RYoUsTuhJm0bzZs3Ny+//LKJi4szycnJZsyYMUaS2blzZ6ZzAcA1ZNac/Lfo6GgzePBgU7BgQXPy5Enb8mbNmplGjRpl+J6s1mX0me+9955xc3MzEyZMsB38+fvvv83QoUNtY6ZMmWLc3NzMhx9+aBISEkxqaqqZP3++8fb2Np9++mmW38GR7b/44ovGy8vLzJ492yQnJ5vY2FgzbNgw4+bmZndCSU40RGvXrm169Ohhzp07Z1JTU23fbdGiRaZTp07m7NmzJjU11UydOjXDv2EaNWpk6tSpY3r06GFOnz5tjDFm+fLlxt3d3UyePDnLuQDg/DJrTv5bUlKS+eKLL4yHh4dZuHChbfmECROMpHTNuhuty+gzt2zZYry9vU3nzp1t+7jR0dHm3XffNX/88Ycx5p+Gpo+Pj+nRo4eJiooyxhgTGhpqKlSoYDp27Jjl93Rk+8HBwcbNzc0MHTrUXL161SQnJ5tvvvnGeHl5mREjRmSZ/2Yaon369DG7du0yxhizbds24+PjYyZNmmSeeOIJs2PHDmOMMSEhISZv3rzp/oZJ20a/fv1sJzD+/fffpmzZsqZly5ZZzgUA55dZc/J6mzdvNiVLlrTbT922bZuRZBYvXpxufFbrMvrMixcvmooVK5rq1aub7du3246JLFy40Hay4KVLl0zlypVNnTp1zN69e40xxpw/f960bdvWBAYGmosXL2aa35HtR0REmKJFi5pGjRrZ/r4ICQkxAQEBpm7duiYxMTHT/DfTEG3ZsqX5/vvvTUpKijl79qypUaOGadGihXnzzTfNt99+a5KTk01UVJSpXbt2ur9h0rbRpk0b89VXX5mkpCQTGxtr2rVrZwoXLmwuXbqU6VzAedAQhWVu1BC9fPmyiYiIMOHh4aZfv36mfPnytnU52RB9/vnnjZeXlzly5EimWUeOHGk8PT1NWFiY3fLjx48bNzc38+WXX2b6Xke2P2jQIOPt7W138NwYY7744gvb2S2Z5b+ZhujIkSPtxk6aNMm4u7vb/RFhjDGTJ0+2ndV4/TauH7t27VojiTNmgLtAVg3R+Ph4c/r0aRMeHm5mzZplJNnOKszJhujSpUuNJDNr1qxMcy5btsxIMt9++63d8pSUFFOtWjXTvXv3TN/ryPbnzZtnJJkff/zRbvnhw4eNh4eHXa3NiYboPffcYztBxph//oiQZO6991675bt27UqXK20btWvXthsb9//aO/OoJq73jT+BQBIQJLKJEVIUUWupUmgbkcVWQRB6EETlWLqoRRRFLVgFFbuoHBDEBUQFtWKVio2KgpyyHDfEglpNUbZDtVVSwBYUqCJLZH5/eJJfhiQDFfi2tffzX973znNn5hwe7rz3zp22NkpfX59atWqVxuskEAgvB0wTol1dXdSDBw+o2tpahT/t3btXkR+oCdHGxkaKy+VSgYGBGs/z0aNHlJ6eHvXBBx+o5JYuXUqZmJjQ3gRSpi/69fX1FJvNphYtWkSLy2QyRTFI0/lT1F+fENXR0aGN07u6uqhhw4ZRXC6X9mwhk8koExMTlfOSayhPUFMURbm7u1OvvfaaxuskEAgvB71NiDY1NSnqJvb29lRAQIAiN5ATopMnT6YsLS0ZF0C7uLhQr7zyCtXe3k6LnzhxggKgmNhUR1/0J02aRI0aNYqSyWS0eHBwMKWtra3YgWugJkSPHj1Ka+vv70+x2WyVZ4s5c+ZQQqGQFpNryHcrkyOvsTQ0NGi8TgKB8O+ntwnR1tZW6rfffqNqa2upsLAwysjISJEbyAnR6OhoisViMS4G//LLLykAKm9LPnjwgOJwONTWrVs1HtsX/dWrV1NaWlqKF4zkZGZm0rx2oCZEFyxYQGubkpKiNi6vuSvX6+UaH330Ea2tvMZy5MgRjddJ+PdAviFK+EfR3d2NzZs3w8rKCsbGxnBwcIBIJMLJkydx7949xm9vviiFhYWYMGECbGxsNLbJz8/H2LFjweFw0NDQgPr6etTV1UFHRwfm5ua4cuVKv/SLiopgZ2cHgUBAi3t5eSnyA8nUqVNpv0ePHo3u7m64urrS4vJzvnfvnorGu+++S/s9duxYAMCvv/46cCdKIBD+NZw+fRr29vbQ19eHnZ0dRCIRIiMjAQB3794d8P4KCwsBALNmzdLYJj8/HwDg6OhI8+6GhgbY2dn16t296cu92dPTkxYfM2YMRo8ePeDe7ebmBhaLpfg9evRoAICTkxMtzuTd77zzDq0tj8eDpaUl8W4C4T9KZWUlZs6cCX19fYwZMwZvv/02pk+fDmBwvLu4uBjt7e2M3nr58mW0tbXBycmJ5t11dXUYN24cGhsbUVNT88L6JSUlkMlkKt6tra0NDw8P3Lp1a0C/82ZjY4ORI0cqfrPZbAiFQlhaWkIoFNL6t7a2VuvdNjY2sLS0pMXGjh1LvJtA+I/S2tqKJUuWgM/nQyAQwNHRESKRCBUVFYPi3Y8fP0ZJSQlmzJgBLperts2TJ09QXFwMJycnxfc65f5tbW0NABrH3n3R//PPPyGRSODu7g5tbW1azsvLC8+ePcMPP/zQj6tURV3dRCaTqcRtbGwglUrVfjeP1E0IBIIyiYmJGDVqFPh8Puzt7SESiXD48GE0Nzer/UZ9fyksLIRQKMSkSZM0tsnPz8fIkSNhbGxMG3t3dXVBKBT2WjfpTb+oqAjW1tawtbWlxWfOnKnIDyTqvBsAnJ2daXFS8/7vwv67T4BAUCYuLg5ffPEFDh48iMDAQOjq6gIAwsPDsX37dnR3d0NLa2Dn8R89eoSJEycytmlqakJzczMcHR1Vctra2mCzNf8p9UW/paVF8ZCgjKmpqSKvCeXCtjJMH7E2MzOj/dbT02OMyz/QzaShr6+vsS2BQHi5uXr1Kvz8/BAaGorz58/DyMgIAJCbmwtvb29GP3pRHj16BDabjWHDhmls09TUBBaLBQ8PD7X5IUOG9Eu/paUFbDZbcb3KmJqa4vfff9d4rCbvBjT792B4N/Dcv4l3Ewj/Pdra2jBt2jQIhUJUVVUpxqIdHR3g8XiD5t0AYG5urrFNU1MTAGDjxo3YvHmzSl4gEKCtre2F9eXjavk4WxnlsffQoUPVHv9Xx97qfFdPT0/h1T3jf8W7Hz9+rLZPAoHwcrNgwQJcuHABZ86cgbOzs8KXXF1d0draOuD9tbS0gKIoRm9tbm5Gd3c3cnJycPHiRZW8QCDQuMC9L/ry6+rNuzUx2HWTZ8+e4enTpyrPF6RuQiAQ5OzZswcRERHYvXs3Fi5cqFgA8tVXX+Hzzz8ftLE3k7cCz8fef/zxh9qaNwDo6Oj0S7+lpUWtdw8ZMgQcDofUvAn/c8iEKOEfxalTpyASifDhhx/S4j///POg9SkQCNSuBlHGwsICI0aMwI8//jgo+ubm5pBKpSpxeWz48OEaj+Xz+QCe/4NRXn2uTo9AIBAGg6ysLLBYLMTFxSkGisDge7dMJkNdXZ3K2/VyLCwsQFEUysrKGCc2X1Tf3NwcMpkMDx48UHkIkEqlsLKy0qiv7N3KdHV1MU6kEggEwkBRUlKC+vp6pKam0hbm3blzBxRFDUqfcj9lGhtbWFgAAHbv3o05c+YMuL7crzWNvbW1tWFiYqLxeD6fr7ZwQ8beBALhf8GzZ89w5swZLFu2DC4uLrTcnTt31Bad+4uJiQl0dXUZvdXY2Bg6OjoIDAzEvn37BkWfzWb3q25CvJtAIPydnDp1ChMmTEBoaCgtPth1k/LycsY2FhYW6OzsxJ07dwZF39zcXK12Y2MjOjo6+lzzVoZ4N6E/kC1zCf8ouFwuOjs7abH79+8rtj2UY2hoCABob29X0WDKqcPf3x9VVVW4dOmSSk6+gjEgIAA//fQTbt68qVaDaWVKX/S9vLxQXl6OsrIyWv7IkSOKvCbkWw70nKw9fvy4xmMIBAJhIOFyuaAoCl1dXYpYd3c3Dhw4QGvHZrOhp6en1p+Zcurw9/cHAKSmpqrklL0bAA4dOqRWozfv7k1f7s0ZGRm0fHFxMe7du8fo3WZmZjAyMlLx7u+++27QJiIIBAJBGfmq9J5jb3W+Z2hoqNGfmXI9cXZ2hpmZGaO3uri4wMzM7IW8uy/6U6ZMgYGBgYp3P3nyBFlZWXBzc1P79qYcW1tbSCQS2nnU1taipKRE4zEEAoEwUGhpaamtm+Tk5KC+vp4WG6i6CYfDgbe3N06fPq124V53dze4XC58fHxw5swZPHz4UK2OpjdE+6o/depUZGdn48mTJ7T8kSNHYGBgoLIdojK2trZob29HRUWFIkZRFMRiscZjCAQCYSBR592NjY04deoULTbQNe+GhgZkZ2er5JTrJnfv3lVbtwZ6r5v0pu/l5QWpVKqyNS6peRP+LsiEKOEfxfz583Ht2jVs2bIFVVVVyMnJgb+/P2bPnk1rN27cOOjp6SEjIwM1NTWQSqWKQTFTTh2ffvopnJyc4Ofnh9TUVJSXl6OoqAgrVqxQFFPCw8Ph4uKCmTNnYt++fbh16xbKysqQmZkJT09PZGVl9Ut/zZo1GDVqFHx9fZGdnY2KigokJCQgJiYGwcHBePPNNzXqOzo6YvLkyYiOjkZBQQHKysoQFRUFAwODvt52AoFA6BezZ8+Gjo4OFi1ahLKyMhQXF8PPzw9OTk4qbd944w3k5eVBIpFAKpXSCiZMuZ689dZb+Oyzz7BlyxZERkbi2rVrkEgk2L59O4KDgxVt1q9fj6ioKERFRaG0tBTV1dXIzc1FSEgIIiIi+qU/bdo0zJ07F+vWrUNycjIqKipw4sQJzJkzB3Z2dlixYgXjfQsLC0NmZiZSU1NRWVmJgwcPoqCgACNGjGA8jkAgEAYCBwcH2NraYv369bhw4QLKysqwYcMGPH36VOX7bA4ODqisrERBQQFqa2tRV1fXp1xPuFwuDhw4gOvXr+O9997D+fPnUVVVhWPHjin+Z/B4PKSnp+PcuXPw9/fHuXPnUFNTgwsXLmDTpk1wc3Prl76BgQESEhKQm5uLZcuWQSKRoKioCF5eXujs7ERiYiLjfQsLC4NUKsXKlStRXl6OvLw8hISEYMaMGb3ecwKBQOgvLBYL8+bNQ3p6Oo4ePYrq6mocOnQIMTExcHV1pbV1cHAAAOzfvx+//PILpFKpYgEjU04dO3bswNChQ+Hi4oKTJ0+iuroa+fn58PPzQ2lpKQBg165diolLsViM6upqXL16FWlpabC3t2d8A7Qv+gkJCejs7ISXlxcuXboEiUSC5cuX4+zZs4iPj1dMFKjj/fffh7GxMRYvXoxr166htLQUQUFBEIlEfbjrBAKB0H/mz5+PmpoaREVFoaqqCnl5efD29lapeVtbW4PP5yMzMxPV1dWQSqWKrVqZcupYvHgxPDw8EBQUhKSkJNy6dQtXrlzB2rVrkZCQoGjj4+MDf39/7Nq1CxKJBLdv34ZYLIavry/S09P7pb98+XLY2dlh3rx5EIvFqKysxO7duxEVFYW5c+fC3d1do76NjQ28vLwQExODnJwc3L59G5s2bSKLyAn9gmyZS/jb4HA4EAgEitXpABASEgIWi4X09HSkp6fDzs4OX3/9Nb7//ntcvHhRsXf4sGHDcPToUSQmJsLd3R0ymQyxsbEICgpizKnrk8fj4dy5c0hOTsbhw4cRFxcHKysrBAQE4JNPPgHwvLhSUFCA/fv3QywWIz4+HgYGBnj11VexevVqTJ8+XeN19kWfz+ejtLQUcXFx2LhxI1paWmBlZYWUlBQsWrSI8Z4BwIkTJ7B27VqEhobCyMgIoaGhcHV1hVgspm1fyePxIBAIwOFwVM5RXZzL5arENbXV0tKCQCBgfAghEAgvByNHjqR9W23ChAkoKChATEwMZs+ejREjRmDlypWwsLBAdnY2zbPS0tIQGRmJgIAAdHR0KBaa9Jbr2ScAbN26FSKRCAcPHoRYLAafz8fUqVMRHx+vaLN582a4uroiLS0NH3/8MYDng+pZs2YhKCiI8Tr7op+RkYE9e/YgMzMTiYmJ4PP5WLhwIdasWUPzX3Xnv3HjRgBAUlISduzYAR8fH+zduxfOzs60LX41+SuLxWKMK/fH5NHm5ubEuwmE/wCmpqa0iU4Oh4P8/HxER0dj6dKl4PF4CAgIQEpKCvLy8mjfR46IiEBDQwNWrVqF1tZWGBoaKrbHYsr17BMAfHx8cPXqVSQmJmLZsmVgsViwt7en7Srg6ekJiUSCnTt3Ijw8HI8fP4ZQKISrqytOnjzJeJ190V+8eDGEQiGSk5MREBAAXV1diEQi7N+/X7ESXdP5u7m54dtvv8XOnTvh6+sLR0dH7NmzBzt37sT9+/dpbc3MzFTGzPK4OkxNTdHR0dEnjaFDh2rc0p1AILw8GBgYQCAQQEvr/99nSEpKwvDhw7Ft2zY8ffoUzs7OyMrKQkREBM0vpkyZgm3btiEjIwMpKSno7u5Gbm4uXn/9dcacuj6trKxw48YNbNu2DVu2bEFraytsbW0REhKCyZMnA3g+3r158yZ27dqF7du3o76+Hubm5rC3t8fhw4dp27P3pC/6EydOxPXr1xEbG4slS5ags7MT48ePR15eHjw8PBjvmYGBAQoKCrBhwwYEBgZCKBQiOjoazc3NyM7Opn0jT+6vPb9dZ2hoyBhX7k+ThrzG0rOmQyAQXi50dXUhEAhou44EBgais7MTaWlpEIvFGD9+PJKTk3Hjxg0UFhYqxpw8Hg/Hjh1DbGwsPD090dXVhXXr1iE0NJQxp65PNpuNs2fPYu/evRCLxdixYwcEAgF8fX0RFhamaHP69Gmkp6cjMzMTSUlJ0NPTw7hx4xAcHAxvb2+N19kXfX19fVy+fBlbt25FbGwsHj58CIFAgPj4eCxdupTxngHAN998g8jISISHh2PIkCFYuHAhlixZguPHj9NeBtJUM5fHeTxer3FNGgBUaiyEfy8sikypEwgEAoFAIBAIBAKBQCAQCAQCgUAgEAiElxSyZS6BQCAQCAQCgUAgEAgEAoFAIBAIBAKBQHhp+T+pTqK/SWoV+gAAAABJRU5ErkJggg==", + "text/plain": [ + "
" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "palette = ListedColormap([\n", + " \"#1696d2\", \"#0077b6\", \"#00a676\", \"#f59f00\",\n", + " \"#7b2cbf\", \"#e76f51\", \"#577590\", \"#90be6d\",\n", + "])\n", + "\n", + "def color_regions(labels):\n", + " regions = sorted(np.unique(labels))\n", + " neighbors = {region: set() for region in regions}\n", + " for row in range(L):\n", + " for column in range(L):\n", + " here = labels[row, column]\n", + " if row + 1 < L:\n", + " there = labels[row + 1, column]\n", + " if here != there:\n", + " neighbors[here].add(there)\n", + " neighbors[there].add(here)\n", + " if column + 1 < L:\n", + " there = labels[row, column + 1]\n", + " if here != there:\n", + " neighbors[here].add(there)\n", + " neighbors[there].add(here)\n", + " assigned = {}\n", + " for region in regions:\n", + " used = {assigned[value] for value in neighbors[region] if value in assigned}\n", + " assigned[region] = next(color for color in range(palette.N) if color not in used)\n", + " return np.vectorize(assigned.__getitem__)(labels)\n", + "\n", + "if leaf_to_input.shape != (L * L,) or not np.array_equal(np.sort(leaf_to_input), np.arange(L * L)):\n", + " raise ValueError(\"the returned order is not a permutation of the 4x4 public sites\")\n", + "leaf_position = input_to_leaf\n", + "levels = int(np.log2(leaf_to_input.size)) + 1\n", + "fig, axes = plt.subplots(1, levels, figsize=(3.1 * levels, 3.4), constrained_layout=True)\n", + "for level, axis in enumerate(axes):\n", + " labels = (leaf_position // (2**level)).reshape(L, L)\n", + " colors = np.zeros_like(labels) if level == 0 else color_regions(labels)\n", + " axis.imshow(colors, cmap=palette, vmin=0, vmax=palette.N - 1, origin=\"upper\")\n", + " for boundary in range(L + 1):\n", + " axis.plot([-0.5, L - 0.5], [boundary - 0.5, boundary - 0.5], color=\"black\", linewidth=1.0) if boundary in (0, L) else None\n", + " axis.plot([boundary - 0.5, boundary - 0.5], [-0.5, L - 0.5], color=\"black\", linewidth=1.0) if boundary in (0, L) else None\n", + " for row in range(L):\n", + " for column in range(L - 1):\n", + " if labels[row, column] != labels[row, column + 1]:\n", + " axis.plot([column + 0.5, column + 0.5], [row - 0.5, row + 0.5], color=\"black\", linewidth=1.0)\n", + " for row in range(L - 1):\n", + " for column in range(L):\n", + " if labels[row, column] != labels[row + 1, column]:\n", + " axis.plot([column - 0.5, column + 0.5], [row + 0.5, row + 0.5], color=\"black\", linewidth=1.0)\n", + " axis.set_title(f\"Step {level}\")\n", + " axis.set_xticks(range(L))\n", + " axis.set_yticks(range(L))\n", + " axis.set_xlabel(\"lattice column\")\n", + " axis.set_ylabel(\"lattice row\" if level == 0 else \"\")\n", + "plt.show()" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python", + "version": "3.12" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/networkx_system.py b/examples/networkx_system.py new file mode 100644 index 0000000000000000000000000000000000000000..0d1d3f2e669eb1c6b01e8544753f73893d913c8e --- /dev/null +++ b/examples/networkx_system.py @@ -0,0 +1,16 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from pathlib import Path + +import networkx as nx + +from hamiltonzero import SpinHamiltonian +from hamiltonzero.data import save_system + + +graph = nx.path_graph(8) +nx.set_edge_attributes(graph, 1.0, "J") +nx.set_node_attributes(graph, 0.0, "h") +system = SpinHamiltonian.from_networkx(graph) +save_system(Path("outputs/systems/chain_8.json"), system) diff --git a/examples/train.json b/examples/train.json new file mode 100644 index 0000000000000000000000000000000000000000..f5cd157210ea22d46c5449a0de64453a8d1a7e97 --- /dev/null +++ b/examples/train.json @@ -0,0 +1,6 @@ +{ + "systems": "datasets/train/foundation_5000.jsonl", + "output": "outputs/foundation.eqx", + "steps": 1000, + "n_max": 64 +} diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000000000000000000000000000000000000..0975c664462ea9f3244d1d689d8d8201c4b18aac --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,47 @@ +[build-system] +requires = ["setuptools>=77", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "hamiltonzero" +version = "0.1.0" +description = "Compiled neural wavefunctions for quantum spin systems" +readme = "README.md" +requires-python = ">=3.12" +license = "Apache-2.0 AND MIT" +license-files = [ + "LICENSE", + "NOTICE", + "THIRD_PARTY_NOTICES.md", + "third_party/folx/LICENSE", + "third_party/jax/LICENSE", + "third_party/kfac_jax/LICENSE", +] +dependencies = [ + "absl-py>=2.5", + "distrax>=0.1.9", + "dm-tree>=0.1.10", + "equinox==0.13.6", + "immutabledict>=4.3", + "jax @ git+https://github.com/TakeOver/jax.git@79f82535b15a444516d4a5e2beb71d283665b2ff", + "jaxlib==0.11.0", + "jaxtyping==0.3.9", + "networkx>=3.5", + "numpy>=2.3", + "optax>=0.2.8", + "packaging>=26", + "typing-extensions>=4.15", +] + +[project.optional-dependencies] +notebooks = [ + "jupyterlab>=4.4", + "matplotlib>=3.10", +] + +[project.scripts] +hamiltonzero = "hamiltonzero.cli:main" + +[tool.setuptools.packages.find] +where = ["src"] +include = ["hamiltonzero*", "kfac_jax*"] diff --git a/src/hamiltonzero/__init__.py b/src/hamiltonzero/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..ec64f348fed5e96b3d59da75d1cff2c6cb04f8a6 --- /dev/null +++ b/src/hamiltonzero/__init__.py @@ -0,0 +1,40 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .hamiltonian import SpinHamiltonian +from .inference import ( + CompiledOrder, + EnergySamples, + PreparedInference, + burn_in, + energy, + prepare, + spin, + step, +) +from .renyi2 import ( + BasisSamplerState, + Renyi2Result, + burn_in_basis, + measure_renyi2, + renyi2_purity, + step_basis, +) + +__all__ = [ + "BasisSamplerState", + "CompiledOrder", + "EnergySamples", + "PreparedInference", + "Renyi2Result", + "SpinHamiltonian", + "burn_in", + "burn_in_basis", + "energy", + "measure_renyi2", + "prepare", + "renyi2_purity", + "spin", + "step", + "step_basis", +] diff --git a/src/hamiltonzero/checkpoint.py b/src/hamiltonzero/checkpoint.py new file mode 100644 index 0000000000000000000000000000000000000000..01b5ddce5ee6a2d7782dedd5ce7dceb046baf749 --- /dev/null +++ b/src/hamiltonzero/checkpoint.py @@ -0,0 +1,65 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any, Literal, TypeVar + +import equinox as eqx + + +T = TypeVar("T") +CheckpointKind = Literal["router", "compiled_finetune"] + + +def _metadata_path(path: str | Path) -> Path: + source = Path(path) + return source.with_name(source.name + ".json") + + +def save_model( + path: str | Path, + model: object, + *, + kind: CheckpointKind | None = None, + metadata: dict[str, Any] | None = None, +) -> None: + destination = Path(path) + destination.parent.mkdir(parents=True, exist_ok=True) + eqx.tree_serialise_leaves(destination, model) + if kind is not None: + payload = {"kind": kind, **(metadata or {})} + _metadata_path(destination).write_text( + json.dumps(payload, indent=2, sort_keys=True) + "\n" + ) + + +def load_model(path: str | Path, template: T) -> T: + return eqx.tree_deserialise_leaves(Path(path), template) + + +def load_model_metadata(path: str | Path) -> dict[str, Any] | None: + source = _metadata_path(path) + return json.loads(source.read_text()) if source.exists() else None + + +def save_mcmc(path: str | Path, state: object) -> None: + destination = Path(path) + destination.parent.mkdir(parents=True, exist_ok=True) + eqx.tree_serialise_leaves(destination, state) + + +def load_mcmc(path: str | Path, template: T) -> T: + return eqx.tree_deserialise_leaves(Path(path), template) + + +__all__ = [ + "CheckpointKind", + "load_mcmc", + "load_model", + "load_model_metadata", + "save_mcmc", + "save_model", +] diff --git a/src/hamiltonzero/cli.py b/src/hamiltonzero/cli.py new file mode 100644 index 0000000000000000000000000000000000000000..35c459e9c75f9cef1fc6ce32e58ddde17af0be22 --- /dev/null +++ b/src/hamiltonzero/cli.py @@ -0,0 +1,69 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import argparse +import dataclasses +import json +from pathlib import Path + +from .config import load_config + + +def _write_metric(path: Path, metric) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("a", encoding="utf-8") as stream: + stream.write( + json.dumps(dataclasses.asdict(metric), separators=(",", ":")) + "\n" + ) + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(prog="hamiltonzero") + commands = parser.add_subparsers(dest="mode", required=True) + for mode in ("train", "finetune"): + command = commands.add_parser(mode) + command.add_argument("config", type=Path) + command.add_argument("--reuse-mcmc", type=Path) + evaluate = commands.add_parser("eval") + evaluate.add_argument("config", type=Path) + pathway = evaluate.add_mutually_exclusive_group() + pathway.add_argument("--contest", action="store_true") + pathway.add_argument("--large-n", action="store_true") + return parser + + +def main(argv: list[str] | None = None) -> None: + args = _parser().parse_args(argv) + config = load_config(args.config, args.mode) + if args.mode == "eval": + from .modes.eval import run + + if args.contest or args.large_n: + config = dataclasses.replace( + config, + contest=bool(args.contest), + large_n=bool(args.large_n), + ) + run(config) + return + if args.reuse_mcmc is not None: + config = dataclasses.replace( + config, + mcmc=dataclasses.replace(config.mcmc, reuse_mcmc=args.reuse_mcmc), + ) + metrics_path = config.output.with_name(config.output.name + ".metrics.jsonl") + sink = lambda metric: _write_metric(metrics_path, metric) + if args.mode == "train": + from .modes.train import run_train + + run_train(config, metric_sink=sink) + else: + from .modes.finetune import run_finetune + + run_finetune(config, metric_sink=sink) + + +if __name__ == "__main__": + main() diff --git a/src/hamiltonzero/compiled/__init__.py b/src/hamiltonzero/compiled/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..73bee8e197a6b25cc9314d020272060e398a5d38 --- /dev/null +++ b/src/hamiltonzero/compiled/__init__.py @@ -0,0 +1,2 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 diff --git a/src/hamiltonzero/compiled/api.py b/src/hamiltonzero/compiled/api.py new file mode 100644 index 0000000000000000000000000000000000000000..2d50faff9f67d8602966c1072adb11684f979d9d --- /dev/null +++ b/src/hamiltonzero/compiled/api.py @@ -0,0 +1,34 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import jax +import jax.numpy as jnp + +from .tree import compile_physical_tree_reference +from .trunk import bind_shared_kernel, compile_shared_trunk +from .types import CompiledWaveFunction, CompiledWaveFunctions + + +def compile_wavefunction(model, context, perm) -> CompiledWaveFunction: + trunk = compile_shared_trunk(model, context) + tree = compile_physical_tree_reference(model, trunk, perm) + return CompiledWaveFunction(kernel=bind_shared_kernel(model), tree=tree) + + +def compile_wavefunctions(model, context, perms) -> CompiledWaveFunctions: + trunk = compile_shared_trunk(model, context) + perms = jnp.asarray(perms, dtype=jnp.int32) + trees = jax.vmap(lambda perm: compile_physical_tree_reference(model, trunk, perm))( + perms + ) + return CompiledWaveFunctions(kernel=bind_shared_kernel(model), trees=trees) + + +def select_compiled_wavefunction( + candidates: CompiledWaveFunctions, + winner, +) -> CompiledWaveFunction: + tree = jax.tree_util.tree_map(lambda value: value[winner], candidates.trees) + return CompiledWaveFunction(kernel=candidates.kernel, tree=tree) diff --git a/src/hamiltonzero/compiled/execute.py b/src/hamiltonzero/compiled/execute.py new file mode 100644 index 0000000000000000000000000000000000000000..eff747d76688c35dafc1f62443c9c08632ecf30c --- /dev/null +++ b/src/hamiltonzero/compiled/execute.py @@ -0,0 +1,118 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +import jax +import jax.numpy as jnp + +from hamiltonzero.model import normalize_leaf_carriers + +from .types import CARRY_LEFT, CARRY_RIGHT, EMPTY, MERGE + + +def _single(values: tuple[Any, ...], name: str) -> Any: + if len(values) != 1: + raise ValueError( + f"compiled HT executor requires exactly one {name}; got {len(values)}" + ) + return values[0] + + +def _factorized_apply(factor: Any, h: jax.Array, x: jax.Array) -> jax.Array: + odd_dtype = jnp.float32 + x_compute = x if x.dtype == odd_dtype else x.astype(odd_dtype) + V = factor.V if factor.V.dtype == odd_dtype else factor.V.astype(odd_dtype) + U = factor.U if factor.U.dtype == odd_dtype else factor.U.astype(odd_dtype) + mixed = (x_compute @ V) * h + mixed = mixed if mixed.dtype == odd_dtype else mixed.astype(odd_dtype) + return mixed @ U + + +def _compiled_quadrilinear_merge( + T: jax.Array, + u_a: jax.Array, + u_b: jax.Array, +) -> jax.Array: + from hamiltonzero.energy import custom_lap_active, quadrilinear_merge_p + + odd_dtype = jnp.float32 + T = T if T.dtype == odd_dtype else T.astype(odd_dtype) + u_a = u_a if u_a.dtype == odd_dtype else u_a.astype(odd_dtype) + u_b = u_b if u_b.dtype == odd_dtype else u_b.astype(odd_dtype) + if custom_lap_active(): + return quadrilinear_merge_p.bind(T, u_a, u_b) + G, d_r, _, _ = T.shape + leading = u_a.shape[:-1] + u_a_flat = u_a.reshape((-1, G, d_r)) + u_b_flat = u_b.reshape((-1, G, d_r)) + out_flat = jnp.einsum("ijkl,Bik,Bil->Bij", T, u_a_flat, u_b_flat) + return out_flat.reshape((*leading, G * d_r)) + + +def _opcode_gates(opcodes: jax.Array, dtype: jnp.dtype) -> tuple[jax.Array, ...]: + both = (opcodes == MERGE).astype(dtype) + left = (opcodes == CARRY_LEFT).astype(dtype) + right = (opcodes == CARRY_RIGHT).astype(dtype) + return both, left, right + + +def _gate_reference( + candidate: jax.Array, + left_value: jax.Array, + right_value: jax.Array, + opcodes: jax.Array, + *, + feature_axis: bool, +) -> jax.Array: + both, left, right = _opcode_gates(opcodes, candidate.dtype) + if feature_axis: + both, left, right = both[..., None], left[..., None], right[..., None] + pad = candidate.ndim - both.ndim + shape = (1,) * pad + both.shape + both, left, right = both.reshape(shape), left.reshape(shape), right.reshape(shape) + return both * candidate + left * left_value + right * right_value + + +def execute_wavefunction(kernel: Any, tree: Any, q_routed: jax.Array): + if len(tree.leaf_combiner_h) != 0 or len(tree.readout_combiner_h) != 0: + raise ValueError("single-head compiled HT executor does not accept combiners") + if len(tree.merge_h) != len(tree.opcodes): + raise ValueError("merge_h and opcodes must have one entry per tree level") + q_weight = kernel.q_to_odd.weight + leaf_factor = _single(kernel.leaf_factors, "leaf factor") + merge_factor = _single(kernel.merge_factors, "merge factor") + readout_factor = _single(kernel.readout_factors, "readout factor") + leaf_h = _single(tree.leaf_h, "leaf conditioner") + readout_h = _single(tree.readout_h, "readout conditioner") + odd_dtype = jnp.float32 + q_compute = q_routed if q_routed.dtype == odd_dtype else q_routed.astype(odd_dtype) + q_weight = q_weight if q_weight.dtype == odd_dtype else q_weight.astype(odd_dtype) + z = q_compute @ q_weight + u_raw = _factorized_apply(leaf_factor, leaf_h, z) + u, log_rms = normalize_leaf_carriers(u_raw) + s = jnp.zeros(u.shape[:-1], dtype=odd_dtype) + s = s + log_rms.astype(s.dtype) + for h_level, opcodes in zip(tree.merge_h, tree.opcodes, strict=True): + if u.shape[-2] != 2 * h_level.shape[-2]: + raise ValueError("compiled merge level has incompatible shrinking shape") + u_left, u_right = u[..., 0::2, :], u[..., 1::2, :] + s_left, s_right = s[..., 0::2], s[..., 1::2] + raw = _compiled_quadrilinear_merge(kernel.merge_T, u_left, u_right) + out = raw + _factorized_apply(merge_factor, h_level, raw) + scale = jnp.sqrt(jnp.mean(out * out, axis=-1) + kernel.merge_eps) + candidate_u = out / scale[..., None] + candidate_s = s_left + s_right + jnp.log(scale) + u = _gate_reference(candidate_u, u_left, u_right, opcodes, feature_axis=True) + s = _gate_reference(candidate_s, s_left, s_right, opcodes, feature_axis=False) + if u.shape[-2] != 1: + raise ValueError("compiled tree did not reduce to one root") + u_root = u[..., 0, :] + s_root = s[..., 0] + psi = _factorized_apply(readout_factor, readout_h, u_root) + psi_re, psi_im = psi[..., 0], psi[..., 1] + log_abs = 0.5 * jnp.log(psi_re * psi_re + psi_im * psi_im) + s_root + phase = jnp.arctan2(psi_im, psi_re) + return log_abs, phase diff --git a/src/hamiltonzero/compiled/model.py b/src/hamiltonzero/compiled/model.py new file mode 100644 index 0000000000000000000000000000000000000000..bd0efd6c82e3c9abeae4ee69e8cd2be81a76d121 --- /dev/null +++ b/src/hamiltonzero/compiled/model.py @@ -0,0 +1,456 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array + +from hamiltonzero.model import ( + normalize_leaf_carriers, + quadrilinear_merge, + tagged_dense_no_bias, +) +from hamiltonzero.optim import register_scale_and_shift + +from .execute import ( + _compiled_quadrilinear_merge, + _factorized_apply, + _gate_reference, + _single, +) +from .types import EMPTY, MERGE, OPCODE_DTYPE, CompiledTree, SharedKernel, level_widths + + +def _scale_tag(y_flat, x_flat, scale_param, *, tag_id: str): + return register_scale_and_shift( + y_flat, + x_flat, + scale=scale_param, + tag_id=tag_id, + ) + + +class CompiledFinetuneWaveFunction(eqx.Module): + kernel: SharedKernel + leaf_h: Array + merge_h: Array + readout_h: Array + perm: Array + inv_perm: Array + leaf_real: Array + opcodes: Array + n_sites: int = eqx.field(static=True) + r_leaf: int = eqx.field(static=True) + r_merge: int = eqx.field(static=True) + + def leaf_h_rows(self) -> Array: + return self.leaf_h.reshape(self.n_sites, self.r_leaf) + + def merge_h_level(self, level: int, width: int) -> Array: + return self.merge_h[level, : width * self.r_merge].reshape(width, self.r_merge) + + @classmethod + def from_compiled(cls, kernel: SharedKernel, tree: CompiledTree): + leaf_h = _single(tree.leaf_h, "leaf conditioner") + readout_h = _single(tree.readout_h, "readout conditioner") + n_sites = int(tree.perm.shape[-1]) + widths = level_widths(n_sites) + w_max = widths[0] + r_leaf = int(leaf_h.shape[-1]) + r_merge = int(tree.merge_h[0].shape[-1]) + merge_rows = [] + opcode_rows = [] + for width, h_l, ops_l in zip(widths, tree.merge_h, tree.opcodes, strict=True): + if h_l.shape[-2] != width or ops_l.shape[-1] != width: + raise ValueError( + f"level width mismatch: expected {width}, got " + f"{h_l.shape[-2]}/{ops_l.shape[-1]}" + ) + pad = w_max - width + merge_rows.append( + jnp.pad(h_l, ((0, pad), (0, 0))).reshape(-1) if pad else h_l.reshape(-1) + ) + opcode_rows.append( + jnp.pad(ops_l, (0, pad), constant_values=EMPTY) if pad else ops_l + ) + return cls( + kernel=kernel, + leaf_h=leaf_h.reshape(-1), + merge_h=jnp.stack(merge_rows), + readout_h=readout_h, + perm=tree.perm, + inv_perm=tree.inv_perm, + leaf_real=tree.leaf_real, + opcodes=jnp.stack(opcode_rows).astype(OPCODE_DTYPE), + n_sites=n_sites, + r_leaf=r_leaf, + r_merge=r_merge, + ) + + def as_compiled_tree(self) -> CompiledTree: + widths = level_widths(self.n_sites) + return CompiledTree( + perm=self.perm, + inv_perm=self.inv_perm, + leaf_real=self.leaf_real, + leaf_h=(self.leaf_h_rows(),), + leaf_combiner_h=(), + merge_h=tuple( + self.merge_h_level(i, width) for i, width in enumerate(widths) + ), + opcodes=tuple(self.opcodes[i, :width] for i, width in enumerate(widths)), + readout_h=(self.readout_h,), + readout_combiner_h=(), + ) + + def route_q(self, q: Array) -> Array: + return jnp.take(q, self.perm, axis=-2) + + def param_counts(self) -> dict: + kernel_leaves = { + "q_to_odd": self.kernel.q_to_odd.weight.size, + "leaf_V": self.kernel.leaf_factors[0].V.size, + "leaf_U": self.kernel.leaf_factors[0].U.size, + "merge_T": self.kernel.merge_T.size, + "merge_V": self.kernel.merge_factors[0].V.size, + "merge_U": self.kernel.merge_factors[0].U.size, + "readout_V": self.kernel.readout_factors[0].V.size, + "readout_U": self.kernel.readout_factors[0].U.size, + } + tree_leaves = { + "leaf_h": self.leaf_h.size, + "merge_h": self.merge_h.size, + "readout_h": self.readout_h.size, + } + return { + "kernel": kernel_leaves, + "tree": tree_leaves, + "kernel_total": sum(kernel_leaves.values()), + "tree_total": sum(tree_leaves.values()), + "total": sum(kernel_leaves.values()) + sum(tree_leaves.values()), + } + + def __call__(self, q: Array, ctx: Any = None, t: Any = 0.0): + del ctx, t + return self._forward_plain(q) + + def _forward_plain(self, q: Array): + kernel = self.kernel + odd_dtype = jnp.float32 + q_c = q if q.dtype == odd_dtype else q.astype(odd_dtype) + weight = kernel.q_to_odd.weight + weight = weight if weight.dtype == odd_dtype else weight.astype(odd_dtype) + z = q_c @ weight + u_raw = _factorized_apply(kernel.leaf_factors[0], self.leaf_h_rows(), z) + u, log_rms = normalize_leaf_carriers(u_raw) + s = jnp.zeros(u.shape[:-1], dtype=odd_dtype) + log_rms.astype(odd_dtype) + widths = level_widths(self.n_sites) + for level, width in enumerate(widths): + h_level = self.merge_h_level(level, width) + opcodes = self.opcodes[level, :width] + u_left, u_right = u[..., 0::2, :], u[..., 1::2, :] + s_left, s_right = s[..., 0::2], s[..., 1::2] + raw = _compiled_quadrilinear_merge(kernel.merge_T, u_left, u_right) + out = raw + _factorized_apply(kernel.merge_factors[0], h_level, raw) + scale = jnp.sqrt(jnp.mean(out * out, axis=-1) + kernel.merge_eps) + candidate_u = out / scale[..., None] + candidate_s = s_left + s_right + jnp.log(scale) + u = _gate_reference( + candidate_u, u_left, u_right, opcodes, feature_axis=True + ) + s = _gate_reference( + candidate_s, s_left, s_right, opcodes, feature_axis=False + ) + return self._readout( + _factorized_apply(kernel.readout_factors[0], self.readout_h, u[..., 0, :]), + s[..., 0], + ) + + def _readout(self, psi: Array, s_root: Array): + psi_re, psi_im = psi[..., 0], psi[..., 1] + log_abs = 0.5 * jnp.log(psi_re * psi_re + psi_im * psi_im) + s_root + return log_abs, jnp.arctan2(psi_im, psi_re) + + def call_tagged(self, q: Array, ctx: Any = None, t: Any = 0.0): + del ctx, t + if q.ndim != 2: + raise ValueError( + f"call_tagged is per-walker: expected q [P, 4], got {q.shape}" + ) + kernel = self.kernel + odd_dtype = jnp.float32 + q_c = q if q.dtype == odd_dtype else q.astype(odd_dtype) + real = self.leaf_real + z = tagged_dense_no_bias( + kernel.q_to_odd.weight, + q_c, + tag_id="compiled.q_to_odd", + pathway="odd", + kfac_structural_mask=real, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + ) + leaf = kernel.leaf_factors[0] + vz = tagged_dense_no_bias( + leaf.V, + z, + tag_id="compiled.leaf.V", + pathway="odd", + kfac_structural_mask=real, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + ) + vz_flat = vz.reshape(-1) + mixed = _scale_tag( + vz_flat * self.leaf_h, + vz_flat, + self.leaf_h, + tag_id="compiled.leaf.h", + ).reshape(vz.shape) + u_raw = tagged_dense_no_bias( + leaf.U, + mixed, + tag_id="compiled.leaf.U", + pathway="odd", + kfac_structural_mask=real, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + ) + u, log_rms = normalize_leaf_carriers(u_raw) + s = jnp.zeros(u.shape[:-1], dtype=odd_dtype) + log_rms.astype(odd_dtype) + merge = kernel.merge_factors[0] + merge_T = kernel.merge_T + merge_eps = kernel.merge_eps + + def level_body(carry, xs): + u_buffer, s_buffer = carry + h_level, opcodes = xs + u_left, u_right = u_buffer[0::2], u_buffer[1::2] + s_left, s_right = s_buffer[0::2], s_buffer[1::2] + merge_rows = opcodes == MERGE + raw = quadrilinear_merge( + merge_T, + u_left, + u_right, + tag_id="compiled.merge.T", + pathway="odd", + kfac_structural_mask=merge_rows, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + vx = tagged_dense_no_bias( + merge.V, + raw, + tag_id="compiled.merge.V", + pathway="odd", + kfac_structural_mask=merge_rows, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + vx_flat = vx.reshape(-1) + mixed_level = _scale_tag( + vx_flat * h_level, + vx_flat, + h_level, + tag_id="compiled.merge.h", + ).reshape(vx.shape) + correction = tagged_dense_no_bias( + merge.U, + mixed_level, + tag_id="compiled.merge.U", + pathway="odd", + kfac_structural_mask=merge_rows, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + out = raw + correction + scale = jnp.sqrt(jnp.mean(out * out, axis=-1) + merge_eps) + candidate_u = out / scale[..., None] + candidate_s = s_left + s_right + jnp.log(scale) + u_next = _gate_reference( + candidate_u, u_left, u_right, opcodes, feature_axis=True + ) + s_next = _gate_reference( + candidate_s, s_left, s_right, opcodes, feature_axis=False + ) + return ( + jnp.concatenate([u_next, jnp.zeros_like(u_next)], axis=0), + jnp.concatenate([s_next, jnp.zeros_like(s_next)], axis=0), + ), None + + (u_buffer, s_buffer), _ = jax.lax.scan( + level_body, + (u, s), + (self.merge_h, self.opcodes), + ) + readout = kernel.readout_factors[0] + vr = tagged_dense_no_bias( + readout.V, + u_buffer[0], + tag_id="compiled.readout.V", + pathway="odd", + ) + mixed_readout = _scale_tag( + vr * self.readout_h, + vr, + self.readout_h, + tag_id="compiled.readout.h", + ) + psi = tagged_dense_no_bias( + readout.U, + mixed_readout, + tag_id="compiled.readout.U", + pathway="odd", + ) + return self._readout(psi, s_buffer[0]) + + +def _expand_stage(V, U, h_2d, new_rank: int, key): + old_rank = V.shape[-1] + if new_rank < old_rank: + raise ValueError(f"cannot shrink rank {old_rank} -> {new_rank}") + if new_rank == old_rank: + return V, U, h_2d + extra = new_rank - old_rank + key_u, key_h = jax.random.split(key) + v_new = jnp.zeros((*V.shape[:-1], extra), dtype=V.dtype) + u_new = jnp.std(U) * jax.random.normal(key_u, (extra, *U.shape[1:]), dtype=U.dtype) + h_new = jnp.std(h_2d) * jax.random.normal( + key_h, (*h_2d.shape[:-1], extra), dtype=h_2d.dtype + ) + return ( + jnp.concatenate([V, v_new], axis=-1), + jnp.concatenate([U, u_new], axis=0), + jnp.concatenate([h_2d, h_new], axis=-1), + ) + + +def expand_rank( + model: CompiledFinetuneWaveFunction, + *, + leaf_rank: int, + merge_rank: int, + key, +) -> CompiledFinetuneWaveFunction: + key_leaf, key_merge = jax.random.split(jnp.asarray(key), 3)[:2] + kernel = model.kernel + leaf = kernel.leaf_factors[0] + merge = kernel.merge_factors[0] + n_sites = model.n_sites + max_width = n_sites // 2 + n_levels = model.merge_h.shape[0] + leaf_h = model.leaf_h.reshape(n_sites, model.r_leaf) + merge_h = model.merge_h.reshape(n_levels, max_width, model.r_merge) + readout_h = model.readout_h + r_leaf, r_merge = model.r_leaf, model.r_merge + V, U, leaf_h = _expand_stage(leaf.V, leaf.U, leaf_h, int(leaf_rank), key_leaf) + leaf = eqx.tree_at(lambda factor: (factor.V, factor.U), leaf, (V, U)) + r_leaf = int(leaf_rank) + V, U, merge_h = _expand_stage(merge.V, merge.U, merge_h, int(merge_rank), key_merge) + merge = eqx.tree_at(lambda factor: (factor.V, factor.U), merge, (V, U)) + r_merge = int(merge_rank) + kernel = eqx.tree_at( + lambda value: ( + value.leaf_factors, + value.merge_factors, + ), + kernel, + ((leaf,), (merge,)), + ) + return CompiledFinetuneWaveFunction( + kernel=kernel, + leaf_h=leaf_h.reshape(-1), + merge_h=merge_h.reshape(n_levels, -1), + readout_h=readout_h, + perm=model.perm, + inv_perm=model.inv_perm, + leaf_real=model.leaf_real, + opcodes=model.opcodes, + n_sites=n_sites, + r_leaf=r_leaf, + r_merge=r_merge, + ) + + +def compile_finetune_model( + eager_model, + ctx_row, + *, + leaf_rank: int, + merge_rank: int, + physical_perm, + key, +) -> CompiledFinetuneWaveFunction: + from .tree import compile_physical_tree_reference + from .trunk import bind_shared_kernel, compile_shared_trunk + + shared_trunk = compile_shared_trunk(eager_model, ctx_row) + n_sites = int(shared_trunk.real_mask.shape[-1]) + identity = jnp.arange(n_sites, dtype=jnp.int32) + tree = compile_physical_tree_reference(eager_model, shared_trunk, identity) + model = CompiledFinetuneWaveFunction.from_compiled( + bind_shared_kernel(eager_model), tree + ) + model = expand_rank( + model, + leaf_rank=leaf_rank, + merge_rank=merge_rank, + key=key, + ) + physical_perm = jnp.asarray(physical_perm, dtype=jnp.int32) + if physical_perm.shape != (n_sites,): + raise ValueError( + f"physical_perm must have shape {(n_sites,)}, got {physical_perm.shape}" + ) + model = eqx.tree_at( + lambda value: (value.perm, value.inv_perm), + model, + ( + physical_perm, + jnp.argsort(physical_perm).astype(jnp.int32), + ), + ) + return model + + +def build_finetune_template_model( + eager_model, + n_sites: int, + *, + leaf_rank: int, + merge_rank: int, +) -> CompiledFinetuneWaveFunction: + from .trunk import bind_shared_kernel + + kernel = bind_shared_kernel(eager_model) + widths = level_widths(int(n_sites)) + r_leaf = int(kernel.leaf_factors[0].V.shape[-1]) + r_merge = int(kernel.merge_factors[0].V.shape[-1]) + r_readout = int(kernel.readout_factors[0].V.shape[-1]) + tree = CompiledTree( + perm=jnp.arange(n_sites, dtype=jnp.int32), + inv_perm=jnp.arange(n_sites, dtype=jnp.int32), + leaf_real=jnp.ones((n_sites,), dtype=jnp.bool_), + leaf_h=(jnp.ones((n_sites, r_leaf), dtype=jnp.float32),), + leaf_combiner_h=(), + merge_h=tuple( + jnp.ones((width, r_merge), dtype=jnp.float32) for width in widths + ), + opcodes=tuple( + jnp.full((width,), MERGE, dtype=OPCODE_DTYPE) for width in widths + ), + readout_h=(jnp.ones((r_readout,), dtype=jnp.float32),), + readout_combiner_h=(), + ) + model = CompiledFinetuneWaveFunction.from_compiled(kernel, tree) + return expand_rank( + model, + leaf_rank=leaf_rank, + merge_rank=merge_rank, + key=jax.random.PRNGKey(0), + ) diff --git a/src/hamiltonzero/compiled/tree.py b/src/hamiltonzero/compiled/tree.py new file mode 100644 index 0000000000000000000000000000000000000000..d525038d03922967e93f7d64654916d819f90ac2 --- /dev/null +++ b/src/hamiltonzero/compiled/tree.py @@ -0,0 +1,582 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +import equinox as eqx +import jax +import jax.numpy as jnp + +from hamiltonzero.model import ( + edge_merge_masked, + tagged_dense, + tagged_rms_eqx_style, + tree_active_clock_depth, + tree_depth_count_features, + tree_sphere, +) + +from .types import ( + CARRY_LEFT, + CARRY_RIGHT, + EMPTY, + MERGE, + CompiledTree, + LadderProjectionKernel, + PhysicalCompilerKernel, +) + + +def bind_physical_compiler_kernel(model: Any) -> PhysicalCompilerKernel: + leaf = eqx.tree_at(lambda x: (x.P_u.V, x.P_u.U), model.leaf, (None, None)) + merge = eqx.tree_at( + lambda x: (x.T, x.output_hypernet.V, x.output_hypernet.U), + model.merge, + (None, None, None), + ) + readout = eqx.tree_at( + lambda x: (x.output_hypernet.V, x.output_hypernet.U), + model.readout, + (None, None), + ) + return PhysicalCompilerKernel( + contextualizer=model.readout_leaf_context, + global_fork=model.gladder_fork_phys, + leaf=leaf, + merge=merge, + readout=readout, + leaf_projection=LadderProjectionKernel( + model.gladder_to_gemb_w, + model.gladder_to_gemb_b, + model.gladder_gemb_ln_s, + ), + tree_pool=model.gladder_tree_pool, + tree_update=model.gladder_tree_update, + tree_projection_weight=model.gladder_tree_proj_w, + tree_projection_bias=model.gladder_tree_proj_b, + root_projection=LadderProjectionKernel( + model.gladder_root_proj_w, + model.gladder_root_proj_b, + model.gladder_root_ln_s, + ), + ) + + +def _project_global( + projection: LadderProjectionKernel, + value, + *, + dense_tag: str, + norm_tag: str, +): + structural_active = jnp.asarray(True) + out = tagged_dense( + projection.weight, + projection.bias, + value, + tag_id=dense_tag, + pathway="even", + kfac_structural_mask=structural_active, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + return tagged_rms_eqx_style( + projection.norm_scale, + out, + tag_id=norm_tag, + pathway="even", + kfac_structural_mask=structural_active, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + + +def compile_context_only_reduction( + *, + merge, + c_leaf, + leaf_real, + g_emb, + edges, + structural_mask, + gladder, + n_total=None, + clock_depth=None, + initial_counts=None, + level_offset: int = 0, + feature_n_levels=None, +): + c = jnp.asarray(c_leaf) + m = jnp.asarray(leaf_real, dtype=c.dtype) + if c.shape[0] != m.shape[0]: + raise ValueError("c_leaf and leaf_real widths differ") + n = c.shape[0] + if n == 0 or n & (n - 1): + raise ValueError(f"context-only width must be a power of two, got {n}") + k = jnp.asarray(structural_mask, dtype=c.dtype) + if k.shape != m.shape: + raise ValueError("structural_mask and leaf_real widths differ") + n_total = jnp.sum(m) if n_total is None else n_total + feature_n_levels = ( + tree_active_clock_depth(m) if feature_n_levels is None else feature_n_levels + ) + clock_depth = tree_active_clock_depth(k) if clock_depth is None else clock_depth + counts = m if initial_counts is None else jnp.asarray(initial_counts, dtype=m.dtype) + if counts.shape != m.shape: + raise ValueError("initial_counts and leaf_real widths differ") + candidates = [] + carried = [] + depth_levels = [] + opcode_levels = [] + g_curr = g_emb + e_curr = edges + e_curr = tree_sphere(e_curr) + level = int(level_offset) + while c.shape[0] > 1: + c_a, c_b = c[0::2], c[1::2] + m_a, m_b = m[0::2], m[1::2] + k_a, k_b = k[0::2], k[1::2] + pair_count = c_a.shape[0] + pair_idx = jnp.arange(pair_count, dtype=jnp.int32) + both_struct = k_a * k_b + pair_base = jnp.maximum( + jnp.sum((k_a + k_b - k_a * k_b).astype(jnp.int32)), + jnp.asarray(2, dtype=jnp.int32), + ) + cnt_a, cnt_b = counts[0::2], counts[1::2] + depth = tree_depth_count_features( + cnt_a, cnt_b, n_total, level, feature_n_levels, c.dtype + ) + counts = cnt_a + cnt_b + depth_levels.append(depth) + d_edge = e_curr.shape[-1] + e_pairs = e_curr.reshape(pair_count, 2, pair_count, 2, d_edge) + sibling_lr = e_pairs[pair_idx, 0, pair_idx, 1] + sibling_rl = e_pairs[pair_idx, 1, pair_idx, 0] + level_active = jnp.any(both_struct.astype(bool)) + g_level = g_curr + g_level = tagged_dense( + gladder[2], + gladder[3], + g_curr, + tag_id="gladder.tree.proj", + pathway="even", + kfac_structural_mask=level_active, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + + def candidate_one(ca, cb, elr, erl, dep, pidx, struct_active): + return merge.context_candidate( + ca, + cb, + g_level, + sibling_edge_lr=elr, + sibling_edge_rl=erl, + level_idx=jnp.int32(level), + pair_idx=pidx, + pair_base=pair_base, + clock_depth=clock_depth, + depth_feats=dep, + kfac_structural_mask=struct_active, + kfac_g_structural_mask=level_active, + kfac_scan_shared=False, + ) + + candidate = jax.vmap(candidate_one)( + c_a, + c_b, + sibling_lr, + sibling_rl, + depth, + pair_idx, + both_struct, + ) + candidates.append(candidate) + gate_m_a, gate_m_b = k_a, k_b + gate_both = gate_m_a * gate_m_b + gate_a = gate_m_a * (1.0 - gate_m_b) + gate_b = (1.0 - gate_m_a) * gate_m_b + c = ( + gate_both[:, None] * candidate + + gate_a[:, None] * c_a + + gate_b[:, None] * c_b + ) + m = m_a + m_b - m_a * m_b + k = k_a + k_b - k_a * k_b + opcode_levels.append( + jnp.where( + m_a.astype(jnp.bool_), + jnp.where(m_b.astype(jnp.bool_), MERGE, CARRY_LEFT), + jnp.where(m_b.astype(jnp.bool_), CARRY_RIGHT, EMPTY), + ).astype(jnp.uint8) + ) + attn_mask = both_struct + d_edge = e_curr.shape[-1] + e_blocks = e_curr.reshape(pair_count, 2, pair_count, 2, d_edge) + e00, e01 = e_blocks[:, 0, :, 0], e_blocks[:, 0, :, 1] + e10, e11 = e_blocks[:, 1, :, 0], e_blocks[:, 1, :, 1] + + def edge_row(e0, e1, e2, e3, ma, mb, ka, kb, ca, cb): + return jax.vmap( + lambda x0, x1, x2, x3, mqa, mqb, kqa, kqb, cqa, cqb: edge_merge_masked( + x0, + x1, + x2, + x3, + ma, + mb, + mqa, + mqb, + ca, + cb, + cqa, + cqb, + merge.edge_merge, + k_2i=ka, + k_2i1=kb, + k_2j=kqa, + k_2j1=kqb, + kfac_scan_shared=False, + )[0] + )(e0, e1, e2, e3, m_a, m_b, k_a, k_b, c_a, c_b) + + e_new = jax.vmap(edge_row)(e00, e01, e10, e11, m_a, m_b, k_a, k_b, c_a, c_b) + e_new = merge.tree_edge_fwl.apply_residual( + e_new, c, attn_mask, kfac_scan_shared=False + ) + edge_keep = (both_struct[:, None] * both_struct[None, :]).astype(bool) + e_curr = jnp.where(edge_keep[..., None], e_new, e00) + e_curr = jnp.where(edge_keep[..., None], tree_sphere(e_curr), e00) + c_skip = c + c = merge.level_edge_attn( + c, + e_curr, + attn_mask, + level_idx=jnp.int32(level), + kfac_scan_shared=False, + ) + c = jnp.where(attn_mask.astype(bool)[:, None], tree_sphere(c), c_skip) + carried.append(c) + level_mask = k + update_active = jnp.any(attn_mask.astype(bool)) + pool_structural_mask = level_mask.astype(c.dtype) * update_active.astype( + c.dtype + ) + pooled = gladder[0]( + g_curr, + c, + level_mask.astype(c.dtype), + kfac_structural_mask=pool_structural_mask, + kfac_update_mask=update_active, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + ) + g_curr = gladder[1]( + g_curr, + pooled, + update_mask=update_active, + kfac_structural_mask=update_active, + kfac_scan_shared=False, + ) + level += 1 + e_root = e_curr[0, 0] + merge_h = tuple( + compile_merge_h( + merge, + candidate, + depth_levels[i], + ) + for i, candidate in enumerate(candidates) + ) + return { + "c_candidate": tuple(candidates), + "c_carried": tuple(carried), + "depth_features": tuple(depth_levels), + "merge_h": merge_h, + "opcodes": tuple(opcode_levels), + "c_root": c[0], + "e_root": e_root, + "g_final": g_curr, + } + + +def project_conditioner(context, hypernet): + return jnp.matmul(context, hypernet.W_h) + + +def leaf_context(leaf_builder, e_leaf, g_emb): + g_broadcast = jnp.broadcast_to(g_emb, e_leaf.shape[:-1] + g_emb.shape) + return jnp.concatenate((e_leaf, g_broadcast), axis=-1) + + +def compile_target_leaf_h(leaf_builder, e_leaf, g_emb): + context = leaf_context(leaf_builder, e_leaf, g_emb) + return (project_conditioner(context, leaf_builder.P_u),) + + +def merge_context(merge, c_p_candidate, depth_features): + return jnp.concatenate( + (c_p_candidate, depth_features.astype(c_p_candidate.dtype)), axis=-1 + ) + + +def compile_merge_h(merge, c_p_candidate, depth_features): + return project_conditioner( + merge_context(merge, c_p_candidate, depth_features), merge.output_hypernet + ) + + +def readout_context(readout, e_root, c_root, g_emb): + e_norm = readout.ln_e(e_root, pathway="even") + return jnp.concatenate( + (e_norm, c_root.astype(e_norm.dtype), g_emb.astype(e_norm.dtype)), axis=-1 + ) + + +def compile_target_readout_h(readout, e_root, c_root, g_emb): + context = readout_context(readout, e_root, c_root, g_emb) + return (project_conditioner(context, readout.output_hypernet),) + + +def classify_merge_opcodes(leaf_real): + active = jnp.asarray(leaf_real, dtype=jnp.bool_) + n = active.shape[0] + if n == 0 or n & (n - 1): + raise ValueError(f"leaf_real width must be a nonzero power of two, got {n}") + levels = [] + while active.shape[0] > 1: + left = active[0::2] + right = active[1::2] + opcode = jnp.where( + left, + jnp.where(right, MERGE, CARRY_LEFT), + jnp.where(right, CARRY_RIGHT, EMPTY), + ).astype(jnp.uint8) + levels.append(opcode) + active = left | right + return tuple(levels) + + +def assemble_compiled_tree(*, perm, leaf_real, boundaries) -> CompiledTree: + perm = jnp.asarray(perm, dtype=jnp.int32) + if perm.ndim != 1: + raise ValueError(f"perm must be rank one, got shape {perm.shape}") + leaf_real = jnp.asarray(leaf_real, dtype=jnp.bool_) + if leaf_real.shape != perm.shape: + raise ValueError( + f"leaf_real shape {leaf_real.shape} must match perm {perm.shape}" + ) + inv_perm = jnp.argsort(perm).astype(jnp.int32) + return CompiledTree( + perm=perm, + inv_perm=inv_perm, + leaf_real=leaf_real, + leaf_h=tuple(boundaries["leaf_h"]), + leaf_combiner_h=tuple(boundaries["leaf_combiner_h"]), + merge_h=tuple(boundaries["merge_h"]), + opcodes=tuple(boundaries["opcodes"]), + readout_h=tuple(boundaries["readout_h"]), + readout_combiner_h=tuple(boundaries["readout_combiner_h"]), + ) + + +def compile_physical_tree_from_reduced_state( + kernel: PhysicalCompilerKernel, + *, + perm, + leaf_real, + leaf_h, + c_reduced, + edge_reduced, + real_reduced, + structural_reduced, + counts_reduced, + g_reduced, + early_merge_h=(), + early_opcodes=(), + full_structural_mask=None, +) -> CompiledTree: + perm = jnp.asarray(perm, dtype=jnp.int32) + leaf_real = jnp.asarray(leaf_real) + if perm.ndim != 1 or leaf_real.shape != perm.shape: + raise ValueError("perm and leaf_real must be matching rank-one arrays") + early_merge_h = tuple(early_merge_h) + early_opcodes = tuple(early_opcodes) + if len(early_merge_h) != len(early_opcodes): + raise ValueError("early merge_h/opcode level counts differ") + if full_structural_mask is None: + full_structural_mask = leaf_real + level_offset = len(early_merge_h) + reduced = compile_context_only_reduction( + merge=kernel.merge, + c_leaf=c_reduced, + leaf_real=real_reduced, + g_emb=g_reduced, + edges=edge_reduced, + structural_mask=structural_reduced, + n_total=jnp.sum(leaf_real), + clock_depth=tree_active_clock_depth(jnp.asarray(full_structural_mask)), + gladder=( + kernel.tree_pool, + kernel.tree_update, + kernel.tree_projection_weight, + kernel.tree_projection_bias, + ), + initial_counts=counts_reduced, + level_offset=level_offset, + feature_n_levels=tree_active_clock_depth(leaf_real), + ) + readout_g_emb = _project_global( + kernel.root_projection, + reduced["g_final"], + dense_tag="gladder.root_proj", + norm_tag="gladder.root_ln", + ) + boundaries = { + "leaf_h": tuple(leaf_h), + "leaf_combiner_h": (), + "merge_h": early_merge_h + tuple(reduced["merge_h"]), + "opcodes": early_opcodes + tuple(reduced["opcodes"]), + "readout_h": compile_target_readout_h( + kernel.readout, + reduced["e_root"], + reduced["c_root"], + readout_g_emb, + ), + "readout_combiner_h": (), + } + return assemble_compiled_tree( + perm=perm, + leaf_real=leaf_real, + boundaries=boundaries, + ) + + +def compile_physical_tree_from_shared_trunk( + kernel: PhysicalCompilerKernel, + shared_trunk, + perm, +) -> CompiledTree: + perm = jnp.asarray(perm, dtype=jnp.int32) + if perm.ndim != 1 or perm.shape != shared_trunk.real_mask.shape: + raise ValueError("perm must be rank one and match the shared trunk site width") + node = shared_trunk.node_raw[perm] + edge = shared_trunk.edge_raw[perm][:, perm] + leaf_real = shared_trunk.real_mask[perm] + structural_mask = shared_trunk.balanced_mask + e_leaf, edge_leaf, g_stream = kernel.contextualizer.with_edge( + node, + edge, + leaf_real, + structural_mask, + g=shared_trunk.global_stream, + ) + g_stream = kernel.global_fork(g_stream, edge_leaf, structural_mask) + leaf_g_emb = _project_global( + kernel.leaf_projection, + g_stream, + dense_tag="gladder.to_gemb", + norm_tag="gladder.gemb_ln", + ) + c_leaf = tree_sphere(kernel.leaf.P_c(e_leaf, pathway="even")) + reduced = compile_context_only_reduction( + merge=kernel.merge, + c_leaf=c_leaf, + leaf_real=leaf_real, + g_emb=g_stream, + edges=edge_leaf, + structural_mask=structural_mask, + gladder=( + kernel.tree_pool, + kernel.tree_update, + kernel.tree_projection_weight, + kernel.tree_projection_bias, + ), + ) + readout_g_emb = _project_global( + kernel.root_projection, + reduced["g_final"], + dense_tag="gladder.root_proj", + norm_tag="gladder.root_ln", + ) + boundaries = { + "leaf_h": compile_target_leaf_h( + kernel.leaf, + e_leaf, + leaf_g_emb, + ), + "leaf_combiner_h": (), + "merge_h": reduced["merge_h"], + "opcodes": reduced["opcodes"], + "readout_h": compile_target_readout_h( + kernel.readout, + reduced["e_root"], + reduced["c_root"], + readout_g_emb, + ), + "readout_combiner_h": (), + } + return assemble_compiled_tree( + perm=perm, + leaf_real=leaf_real, + boundaries=boundaries, + ) + + +def compile_physical_tree_reference( + model, + shared_trunk, + perm, +) -> CompiledTree: + if model.gladder_post is None or model.gladder_fork_phys is None: + raise ValueError( + "reference physical compiler requires the target global ladder" + ) + perm = jnp.asarray(perm, dtype=jnp.int32) + if perm.ndim != 1 or perm.shape != shared_trunk.real_mask.shape: + raise ValueError("perm must be rank one and match the shared trunk site width") + node = shared_trunk.node_raw[perm] + edge = shared_trunk.edge_raw[perm][:, perm] + leaf_real = shared_trunk.real_mask[perm] + structural_mask = shared_trunk.balanced_mask + e_leaf, edge_leaf, g_stream = model._contextualize_leaf_even_with_edge_g( + node, + edge, + leaf_real, + structural_mask, + shared_trunk.global_stream, + ) + g_stream = model.gladder_fork_phys(g_stream, edge_leaf, structural_mask) + leaf_g_emb = model._gladder_project(g_stream) + c_leaf = tree_sphere(model.leaf.P_c(e_leaf, pathway="even")) + reduced = compile_context_only_reduction( + merge=model.merge, + c_leaf=c_leaf, + leaf_real=leaf_real, + g_emb=g_stream, + edges=edge_leaf, + structural_mask=structural_mask, + gladder=model._gladder_tree_refs(), + ) + readout_g_emb = model._gladder_root_project(reduced["g_final"]) + boundaries = { + "leaf_h": compile_target_leaf_h(model.leaf, e_leaf, leaf_g_emb), + "leaf_combiner_h": (), + "merge_h": reduced["merge_h"], + "opcodes": reduced["opcodes"], + "readout_h": compile_target_readout_h( + model.readout, + reduced["e_root"], + reduced["c_root"], + readout_g_emb, + ), + "readout_combiner_h": (), + } + return assemble_compiled_tree( + perm=perm, + leaf_real=leaf_real, + boundaries=boundaries, + ) diff --git a/src/hamiltonzero/compiled/trunk.py b/src/hamiltonzero/compiled/trunk.py new file mode 100644 index 0000000000000000000000000000000000000000..c20df4a3054cce8aec762f0268655586c9d446ba --- /dev/null +++ b/src/hamiltonzero/compiled/trunk.py @@ -0,0 +1,129 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +import jax +import jax.numpy as jnp + +from hamiltonzero.model import tree_sphere + +from .types import ( + CanonicalHamiltonian, + FactorizedQSide, + QSideLeafInput, + SharedKernel, + SharedTrunk, + TrunkCompilerKernel, +) + + +def _qside(hypernet) -> FactorizedQSide: + return FactorizedQSide(V=hypernet.V, U=hypernet.U) + + +def bind_shared_kernel(model) -> SharedKernel: + return SharedKernel( + q_to_odd=QSideLeafInput(weight=model.q_to_odd.weight), + leaf_factors=(_qside(model.leaf.P_u),), + leaf_combiner_factors=(), + merge_T=model.merge.T, + merge_factors=(_qside(model.merge.output_hypernet),), + readout_factors=(_qside(model.readout.output_hypernet),), + merge_eps=float(model.merge.eps), + ) + + +def bind_trunk_compiler_kernel(model) -> TrunkCompilerKernel: + return TrunkCompilerKernel( + featurizer=model.featurizer, + trunk=model.trunk, + shared_global=model.gladder_post, + ) + + +class _TrunkMaskContext(eqx.Module): + mask: jax.Array + + +def compile_canonical_shared_trunk( + kernel: TrunkCompilerKernel, + canonical: CanonicalHamiltonian, +) -> SharedTrunk: + if canonical.node_mask.ndim != 2 or canonical.node_mask.shape[0] != 1: + raise ValueError("compiled shared trunk requires exact physical P=1") + if canonical.balanced_mask.shape != canonical.node_mask.shape: + raise ValueError("balanced_mask must match node_mask shape") + graph = canonical.graph_inputs + if graph.node.shape[:2] != canonical.node_mask.shape: + raise ValueError("canonical graph node width must match node_mask") + if graph.edge.shape[:3] != ( + 1, + canonical.node_mask.shape[1], + canonical.node_mask.shape[1], + ): + raise ValueError("canonical graph edge width must match node_mask") + + def one(edge_input, node_input, real_mask, balanced_mask): + edge_feat, local_feat, global_feat = kernel.featurizer( + edge_input, + real_mask, + node_input, + ) + g_seed = tree_sphere(global_feat.astype(local_feat.dtype)) + node_raw, edge_raw, g_seed = kernel.trunk( + _TrunkMaskContext(real_mask), + edge_feat, + local_feat, + g_seed, + ) + global_stream = kernel.shared_global( + g_seed.astype(edge_raw.dtype), edge_raw, real_mask + ) + return SharedTrunk( + node_raw=node_raw, + edge_raw=edge_raw, + global_raw=global_feat, + global_stream=global_stream, + real_mask=real_mask, + balanced_mask=balanced_mask, + ) + + return jax.vmap(one)( + graph.edge, + graph.node, + canonical.node_mask, + canonical.balanced_mask, + ) + + +def select_single_physical_trunk(trunk: SharedTrunk) -> SharedTrunk: + leaves = jax.tree_util.tree_leaves(trunk) + if not leaves or any(x.ndim < 1 or x.shape[0] != 1 for x in leaves): + raise ValueError("production SharedTrunk must have exact leading P=1") + return jax.tree_util.tree_map(lambda x: x[0], trunk) + + +def compile_shared_trunk(model, ctx) -> SharedTrunk: + edge_feat, local_feat, global_feat = model.featurizer( + ctx.J_double_prime, + ctx.mask, + ctx.h_prime, + ) + g_seed = tree_sphere(global_feat.astype(local_feat.dtype)) + node_raw, edge_raw, g_seed = model.trunk( + ctx, + edge_feat, + local_feat, + g_seed, + ) + global_stream = model._gladder_g_stream(edge_raw, ctx.mask, g_seed) + return SharedTrunk( + node_raw=node_raw, + edge_raw=edge_raw, + global_raw=global_feat, + global_stream=global_stream, + real_mask=ctx.mask, + balanced_mask=ctx.bmask, + ) diff --git a/src/hamiltonzero/compiled/types.py b/src/hamiltonzero/compiled/types.py new file mode 100644 index 0000000000000000000000000000000000000000..e9d22aafd8c69c3002a7ece5dc56b586597c3f94 --- /dev/null +++ b/src/hamiltonzero/compiled/types.py @@ -0,0 +1,169 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from enum import IntEnum +from typing import Any + +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array + + +class MergeOpcode(IntEnum): + MERGE = 0 + CARRY_LEFT = 1 + CARRY_RIGHT = 2 + EMPTY = 3 + + +MERGE = int(MergeOpcode.MERGE) +CARRY_LEFT = int(MergeOpcode.CARRY_LEFT) +CARRY_RIGHT = int(MergeOpcode.CARRY_RIGHT) +EMPTY = int(MergeOpcode.EMPTY) +OPCODE_DTYPE = jnp.uint8 + + +class QSideLeafInput(eqx.Module): + weight: Array + + +class FactorizedQSide(eqx.Module): + V: Array + U: Array + + +class SharedKernel(eqx.Module): + q_to_odd: QSideLeafInput + leaf_factors: tuple[FactorizedQSide, ...] + leaf_combiner_factors: tuple[FactorizedQSide, ...] + merge_T: Array + merge_factors: tuple[FactorizedQSide, ...] + readout_factors: tuple[FactorizedQSide, ...] + merge_eps: float = eqx.field(static=True) + + +class TrunkCompilerKernel(eqx.Module): + featurizer: Any + trunk: Any + shared_global: Any + + +class LadderProjectionKernel(eqx.Module): + weight: Array + bias: Array + norm_scale: Array + + +class PhysicalCompilerKernel(eqx.Module): + contextualizer: Any + global_fork: Any + leaf: Any + merge: Any + readout: Any + leaf_projection: LadderProjectionKernel + tree_pool: Any + tree_update: Any + tree_projection_weight: Array + tree_projection_bias: Array + root_projection: LadderProjectionKernel + + +class ModelHamiltonianArrays(eqx.Module): + coupling: Array + full_coupling: Array + field: Array + + +class GraphInputs(eqx.Module): + node: Array + edge: Array + + +class QuotientInputs(eqx.Module): + node_key: Array + edge_key: Array + + +class EnergyInputs(eqx.Module): + custom_lap_J_eff: Array + custom_lap_radial_const: Array + one_body_fields: tuple[Array, ...] + + +class EnergyMasks(eqx.Module): + real: Array + balanced: Array + + +class EnergyFrame(eqx.Module): + custom_lap_J_eff: Array + w_levels: tuple[Array, ...] + custom_lap_radial_const: Array + one_body_fields: tuple[Array, ...] + masks: EnergyMasks + + +EnergyFrameBatch = EnergyFrame + + +class CanonicalHamiltonian(eqx.Module): + model_coupling_fields: ModelHamiltonianArrays + node_mask: Array + balanced_mask: Array + graph_inputs: GraphInputs + quotient_inputs: QuotientInputs + energy_inputs: EnergyInputs + system_identity: Array + + +class SharedTrunk(eqx.Module): + node_raw: Array + edge_raw: Array + global_raw: Array + global_stream: Array + real_mask: Array + balanced_mask: Array + + +class CompiledTree(eqx.Module): + perm: Array + inv_perm: Array + leaf_real: Array + leaf_h: tuple[Array, ...] + leaf_combiner_h: tuple[Array, ...] + merge_h: tuple[Array, ...] + opcodes: tuple[Array, ...] + readout_h: tuple[Array, ...] + readout_combiner_h: tuple[Array, ...] + + +CompiledTreeBatch = CompiledTree + + +class CompiledWaveFunction(eqx.Module): + kernel: SharedKernel + tree: CompiledTree + + def __call__(self, q_routed, _ctx=None, _t=0.0): + from .execute import execute_wavefunction + + return execute_wavefunction(self.kernel, self.tree, q_routed) + + +class CompiledWaveFunctions(eqx.Module): + kernel: SharedKernel + trees: CompiledTreeBatch + + def __call__(self, q_routed, _ctx=None, _t=0.0): + from .execute import execute_wavefunction + + return execute_wavefunction(self.kernel, self.trees, q_routed) + + +def level_widths(n_sites: int) -> tuple[int, ...]: + if n_sites <= 0 or n_sites & (n_sites - 1): + raise ValueError(f"n_sites must be a positive power of two, got {n_sites}") + return tuple(n_sites >> level for level in range(1, n_sites.bit_length())) diff --git a/src/hamiltonzero/config.py b/src/hamiltonzero/config.py new file mode 100644 index 0000000000000000000000000000000000000000..c96d54b29f40f1878e5278695d7903902e4afaf4 --- /dev/null +++ b/src/hamiltonzero/config.py @@ -0,0 +1,301 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import dataclasses +import json +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Literal, TypeVar + + +AttentionImplementation = Literal["tuned", "einsum"] + + +@dataclass(frozen=True, slots=True) +class ModelConfig: + d_e: int = 512 + d_o: int = 32 + d_c: int = 512 + d_r: int = 32 + n_heads: int = 16 + n_layers: int = 8 + rank: int = 32 + edge_channels: int = 96 + attention_qk_dim: int = 256 + attention_v_dim: int = 256 + merge_dim: int = 1536 + trunk_edge_node_context_dim: int = 256 + trunk_edge_hidden_dim: int = 128 + trunk_attention_bias_hidden_dim: int = 64 + trunk_ffn_hidden_dim: int = 2048 + trunk_two_hop_hidden_dim: int = 128 + tree_edge_node_context_dim: int = 256 + global_dim: int = 256 + merge_hypernet_rank: int = 256 + featurizer_bond_dim: int = 128 + featurizer_heads: int = 8 + featurizer_head_dim: int = 64 + featurizer_global_queries: int = 4 + featurizer_edge_hidden_dim: int = 192 + featurizer_zeeman_hidden_dim: int = 2048 + featurizer_global_hidden_dim: int = 16384 + featurizer_combine_hidden_dim: int = 12288 + featurizer_token_initial_scale: float = 0.02 + polar_group_norm_tau: float = 0.001 + polar_bond_hidden_dim: int = 256 + polar_bond_groups: int = 16 + polar_bond_group_dim: int = 16 + polar_zeeman_groups: int = 16 + polar_zeeman_group_dim: int = 16 + router_max_n: int = 128 + router_model_dim: int = 512 + router_heads: int = 16 + router_attention_dim: int = 128 + router_score_dim: int = 512 + router_candidate_dim: int = 1024 + router_summary_dim: int = 1024 + router_ffn_dim: int = 1024 + router_score_initial_scale: float = 0.0 + router_rope_base: float = 10000.0 + router_rope_scaling: float = 1.0 + router_tree_prefix_layers: int = 4 + router_tree_candidate_layers: int = 2 + router_tree_merge_dim: int = 1024 + router_tree_post_layers: int = 4 + router_context_layers: int = 2 + router_context_heads: int = 4 + router_context_attention_dim: int = 256 + router_context_edge_node_dim: int = 256 + level_edge_heads: int = 8 + level_edge_mlp_dim: int = 384 + level_edge_mlp_blocks: int = 3 + level_edge_ffn_dim: int = 1024 + level_edge_rope_base: float = 10000.0 + level_edge_rope_scaling: float = 1.0 + root_readout_edge_rank: int = 128 + ngpt_alpha_initial: float = 0.25 + ngpt_alpha_initial_fraction: float = 0.25 + ngpt_alpha_maximum: float = 0.8 + global_ladder_tap_dim: int = 256 + level_edge_bias_mlp_dim: int = 128 + level_edge_bias_mlp_blocks: int = 1 + merge_context_mlp_dim: int = 1024 + readout_context_layers: int = 2 + readout_context_heads: int = 8 + readout_context_attention_dim: int = 128 + readout_context_edge_node_dim: int = 256 + readout_context_summary_dim: int = 1024 + readout_context_mlp_dim: int = 2048 + readout_context_bias_dim: int = 32 + readout_context_edge_ffn_dim: int = 384 + readout_context_rope_base: float = 10000.0 + readout_context_rope_scaling: float = 1.0 + two_hop_channels: int = 64 + tree_fwl_channels: int = 128 + attention: AttentionImplementation = "tuned" + + +@dataclass(frozen=True, slots=True) +class MCMCConfig: + batch_size: int = 512 + replicas: int = 8 + steps: int = 32 + burn_in: int = 256 + burn_in_replica_steps: int = 2 + walker_chunk_size: int | None = None + initial_sigma: float = 0.3 + initial_haar_sites: int = 1 + sigma_scale: float = 1.1 + langevin_target_acceptance: float = 0.574 + haar_target_acceptance: float = 0.234 + beta_history_weight: float = 0.9 + adapt_every: int = 1 + reuse_mcmc: Path | None = None + + +@dataclass(frozen=True, slots=True) +class KFACConfig: + learning_rate_numerator: float = 0.05 + learning_rate_offset: float = 5.0 + learning_rate_decay_steps: float = 5000.0 + curvature_ema: float = 0.995 + curvature_update_period: int = 2 + inverse_update_period: int = 2 + damping: float = 0.001 + minimum_damping: float = 0.0001 + norm_constraint: float = 0.001 + mad_clip_width: float = 5.0 + momentum: float = 0.0 + l2_regularization: float = 0.0 + + +@dataclass(frozen=True, slots=True) +class RouterConfig: + temperature: float = 1.0 + loss_weight: float = 1.0 + + +@dataclass(frozen=True, slots=True) +class EnergyConfig: + mu: float | None = None + eps: float = 0.1 + chunk_size: int = 512 + + +@dataclass(frozen=True, slots=True) +class TrainConfig: + systems: Path + output: Path + steps: int + seed: int = 777 + n_max: int = 64 + model: ModelConfig = field(default_factory=ModelConfig) + router: RouterConfig = field(default_factory=RouterConfig) + mcmc: MCMCConfig = field(default_factory=MCMCConfig) + kfac: KFACConfig = field(default_factory=KFACConfig) + energy: EnergyConfig = field(default_factory=EnergyConfig) + + +def _finetune_mcmc() -> MCMCConfig: + return MCMCConfig(batch_size=256, replicas=8, steps=2, burn_in=256) + + +def _finetune_kfac() -> KFACConfig: + return KFACConfig( + learning_rate_numerator=0.002, + learning_rate_offset=1.0, + learning_rate_decay_steps=10000.0, + curvature_ema=0.99, + curvature_update_period=2, + inverse_update_period=4, + damping=0.001, + ) + + +def _finetune_energy() -> EnergyConfig: + return EnergyConfig(mu=2.86) + + +@dataclass(frozen=True, slots=True) +class FineTuneConfig: + system: Path + checkpoint: Path + output: Path + steps: int = 10000 + seed: int = 777 + leaf_rank: int = 1536 + merge_rank: int = 1024 + route_temperature: float = 1.0 + model: ModelConfig = field(default_factory=lambda: ModelConfig(attention="einsum")) + mcmc: MCMCConfig = field(default_factory=_finetune_mcmc) + kfac: KFACConfig = field(default_factory=_finetune_kfac) + energy: EnergyConfig = field(default_factory=_finetune_energy) + + +def _eval_mcmc() -> EvalMCMCConfig: + return EvalMCMCConfig( + batch_size=256, + replicas=8, + steps=24, + burn_in=1024, + walker_chunk_size=16, + ) + + +@dataclass(frozen=True, slots=True) +class EvalMCMCConfig: + batch_size: int = 256 + replicas: int = 8 + steps: int = 24 + burn_in: int = 1024 + burn_in_replica_steps: int = 2 + walker_chunk_size: int = 16 + initial_sigma: float = 0.3 + initial_haar_sites: int = 1 + sigma_scale: float = 1.1 + langevin_target_acceptance: float = 0.574 + haar_target_acceptance: float = 0.234 + beta_history_weight: float = 0.9 + + +@dataclass(frozen=True, slots=True) +class EvalConfig: + system: Path + checkpoint: Path + output: Path + seed: int = 777 + contest: bool = False + large_n: bool = False + measurements: int = 256 + contest_candidates: int = 8 + contest_beam_width: int = 8 + contest_preburn: int = 128 + contest_measurements: int = 128 + contest_se_multiplier: float = 2.0 + route_temperature: float = 4.0 + large_n_sequence_shards: int = 0 + large_n_pair_tile_size: int = 128 + contextualizer_attention: AttentionImplementation | None = None + model: ModelConfig = field(default_factory=ModelConfig) + mcmc: EvalMCMCConfig = field(default_factory=_eval_mcmc) + energy: EnergyConfig = field(default_factory=EnergyConfig) + + def __post_init__(self) -> None: + if self.contest and self.large_n: + raise ValueError("contest and large_n are mutually exclusive") + + +Config = TrainConfig | FineTuneConfig | EvalConfig +T = TypeVar("T") + + +def _coerce(cls: type[T], values: dict[str, Any]) -> T: + nested = { + "model": ModelConfig, + "router": RouterConfig, + "mcmc": MCMCConfig, + "kfac": KFACConfig, + "energy": EnergyConfig, + } + if cls is EvalConfig: + nested["mcmc"] = EvalMCMCConfig + data = dict(values) + fields_by_name = {item.name: item for item in dataclasses.fields(cls)} + for name, nested_cls in nested.items(): + if name in data and isinstance(data[name], dict): + item = fields_by_name.get(name) + defaults: dict[str, Any] = {} + if item is not None and item.default_factory is not dataclasses.MISSING: + defaults = dataclasses.asdict(item.default_factory()) + nested_values = {**defaults, **data[name]} + if name == "mcmc" and nested_values.get("reuse_mcmc") is not None: + nested_values["reuse_mcmc"] = Path(nested_values["reuse_mcmc"]) + data[name] = nested_cls(**nested_values) + path_fields = {"systems", "system", "checkpoint", "output"} + for item in dataclasses.fields(cls): + if item.name in path_fields and item.name in data: + data[item.name] = Path(data[item.name]) + return cls(**data) + + +def load_config(path: str | Path, mode: Literal["train", "finetune", "eval"]) -> Config: + values = json.loads(Path(path).read_text()) + cls = {"train": TrainConfig, "finetune": FineTuneConfig, "eval": EvalConfig}[mode] + return _coerce(cls, values) + + +__all__ = [ + "AttentionImplementation", + "EnergyConfig", + "EvalConfig", + "EvalMCMCConfig", + "FineTuneConfig", + "KFACConfig", + "MCMCConfig", + "ModelConfig", + "RouterConfig", + "TrainConfig", + "load_config", +] diff --git a/src/hamiltonzero/data/__init__.py b/src/hamiltonzero/data/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..cb79748562b8423a2b16489bd3244c0f6d755441 --- /dev/null +++ b/src/hamiltonzero/data/__init__.py @@ -0,0 +1,22 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .systems import ( + build_context, + build_context_and_energy, + build_multi_context, + load_system, + load_systems, + padded_model_arrays, + save_system, +) + +__all__ = [ + "build_context", + "build_context_and_energy", + "build_multi_context", + "load_system", + "load_systems", + "padded_model_arrays", + "save_system", +] diff --git a/src/hamiltonzero/data/systems.py b/src/hamiltonzero/data/systems.py new file mode 100644 index 0000000000000000000000000000000000000000..3fca2167d438234b1f0c121071c249460cf16dc4 --- /dev/null +++ b/src/hamiltonzero/data/systems.py @@ -0,0 +1,264 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import json +from dataclasses import replace +from pathlib import Path +from typing import Any, Iterable + +import numpy as np + +from hamiltonzero.hamiltonian import SpinHamiltonian, _exchange_matrix + + +def _load_payload(path: str | Path) -> dict[str, Any]: + source = Path(path) + if source.suffix == ".jsonl": + records = [] + with source.open(encoding="utf-8") as stream: + for line_number, line in enumerate(stream, 1): + if not line.strip(): + continue + record = json.loads(line) + if not isinstance(record, dict): + raise ValueError(f"JSONL record {line_number} must be an object") + records.append(record) + return {"systems": records} + payload = json.loads(source.read_text()) + if not isinstance(payload, dict): + raise ValueError("Hamiltonian dataset must be a JSON object") + return payload + + +def _from_sparse(spec: dict[str, Any]) -> SpinHamiltonian: + n_sites = int(spec.get("n_sites", spec.get("n_spins", 0))) + if n_sites <= 0: + raise ValueError("sparse Hamiltonian requires a positive n_sites") + exchange = np.zeros((n_sites, n_sites, 3, 3), dtype=np.float32) + coupling = np.zeros((n_sites, n_sites), dtype=np.float32) + seen: set[tuple[int, int]] = set() + for offset, term in enumerate(spec["exchange"]): + if not isinstance(term, list) or len(term) != 3: + raise ValueError(f"exchange term {offset} must be [i,j,J]") + left, right, value = term + if ( + not isinstance(left, int) + or isinstance(left, bool) + or not isinstance(right, int) + or isinstance(right, bool) + or not 0 <= left < right < n_sites + ): + raise ValueError( + f"exchange term {offset} must satisfy 0 <= i < j < n_sites" + ) + pair = (left, right) + if pair in seen: + raise ValueError(f"duplicate exchange term for sites {pair}") + seen.add(pair) + matrix = _exchange_matrix(value) + exchange[left, right] = matrix + exchange[right, left] = matrix.T + coupling[left, right] = coupling[right, left] = 1.0 + field = spec.get("field", spec.get("h", spec.get("h_field", 0.0))) + return SpinHamiltonian.from_arrays( + exchange, + field, + coupling=coupling, + nodes=spec.get("nodes"), + mu=spec.get("mu"), + ) + + +def _next_power_of_two(value: int) -> int: + return 1 if value <= 1 else 1 << (value - 1).bit_length() + + +def _from_record( + record: dict[str, Any], + *, + needs_fwl2: bool | None = None, +) -> SpinHamiltonian: + outer = record + spec = record.get("spec", record) + convention = spec.get("convention", "textbook") + if convention != "textbook": + raise ValueError("public Hamiltonian JSON must use convention='textbook'") + if "exchange" in spec: + system = _from_sparse(spec) + else: + system = SpinHamiltonian.from_arrays( + spec["J"], + spec.get("h", spec.get("h_field", 0.0)), + coupling=spec.get("coupling"), + nodes=spec.get("nodes"), + mu=spec.get("mu"), + ) + if needs_fwl2 is None: + needs_fwl2 = outer.get("needs_fwl2", spec.get("needs_fwl2")) + metadata = { + name: outer.get(name, spec.get(name)) + for name in ("category", "tag", "topology_class", "j_class") + } + return replace( + system, + _needs_fwl2=needs_fwl2, + _category=metadata["category"], + _tag=metadata["tag"], + _topology_class=metadata["topology_class"], + _j_class=metadata["j_class"], + ) + + +def load_system(path: str | Path) -> SpinHamiltonian: + payload = _load_payload(path) + if "systems" in payload: + systems = payload["systems"] + if len(systems) != 1: + raise ValueError("load_system requires exactly one system") + dispatch = payload.get("needs_fwl2", payload.get("dataset_needs_fwl2")) + if dispatch is not None: + if not isinstance(dispatch, list) or len(dispatch) != 1: + raise ValueError("needs_fwl2 sidecar must align with systems") + return _from_record(systems[0], needs_fwl2=bool(dispatch[0])) + return _from_record(systems[0]) + return _from_record(payload) + + +def load_systems(path: str | Path) -> list[SpinHamiltonian]: + payload = _load_payload(path) + records = payload.get("systems", [payload]) + dispatch = ( + payload.get("needs_fwl2", payload.get("dataset_needs_fwl2")) + if "systems" in payload + else None + ) + if dispatch is None and isinstance(payload.get("per_system"), list): + derived = payload["per_system"] + if len(derived) == len(records) and all( + isinstance(value, dict) and "needs_fwl2" in value for value in derived + ): + dispatch = [value["needs_fwl2"] for value in derived] + if dispatch is not None: + if not isinstance(dispatch, list) or len(dispatch) != len(records): + raise ValueError("needs_fwl2 sidecar must align with systems") + return [ + _from_record(record, needs_fwl2=bool(value)) + for record, value in zip(records, dispatch, strict=True) + ] + return [_from_record(record) for record in records] + + +def save_system(path: str | Path, system: SpinHamiltonian) -> None: + destination = Path(path) + destination.parent.mkdir(parents=True, exist_ok=True) + destination.write_text(json.dumps(system.to_dict(), indent=2) + "\n") + + +def padded_model_arrays( + system: SpinHamiltonian, + n_max: int | None = None, +) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: + width = _next_power_of_two(system.n_spins) if n_max is None else int(n_max) + if width < system.n_spins: + raise ValueError("n_max cannot be smaller than the system") + if width <= 0 or width & (width - 1): + raise ValueError("n_max must be a positive power of two") + coupling, exchange, field = system.model_arrays() + padding = width - system.n_spins + coupling = np.pad(coupling, ((0, padding), (0, padding))) + exchange = np.pad(exchange, ((0, padding), (0, padding), (0, 0), (0, 0))) + field = np.pad(field, ((0, padding), (0, 0))) + mask = np.zeros((width,), dtype=np.int32) + mask[: system.n_spins] = 1 + return coupling, exchange, field, mask + + +def _context_arrays(system: SpinHamiltonian, n_max: int | None): + import jax.numpy as jnp + + from hamiltonzero.model.route_quotient import system_needs_fwl2 + + _coupling, exchange, field, mask = padded_model_arrays(system, n_max) + needs_fwl2 = system._needs_fwl2 + if needs_fwl2 is None: + _, physical_exchange, physical_field = system.model_arrays() + needs_fwl2 = system_needs_fwl2( + physical_exchange, + physical_field, + system.n_spins, + category=system._category, + tag=system._tag, + topology_class=system._topology_class, + j_class=system._j_class, + ) + return ( + jnp.asarray(exchange), + jnp.asarray(field), + jnp.asarray(mask), + needs_fwl2, + ) + + +def build_context( + system: SpinHamiltonian, + n_max: int | None = None, +): + from hamiltonzero.model import SpinContext + + exchange, field, mask, needs_fwl2 = _context_arrays(system, n_max) + return SpinContext( + J_full=exchange, + h=field, + mask=mask, + needs_fwl2=needs_fwl2, + ) + + +def build_context_and_energy( + system: SpinHamiltonian, + n_max: int | None = None, + mu: float | None = None, + eps: float = 0.1, +): + from hamiltonzero.energy.frame import build_energy_inputs + from hamiltonzero.model import SpinContext + + exchange, field, mask, needs_fwl2 = _context_arrays(system, n_max) + context = SpinContext( + J_full=exchange, + h=field, + mask=mask, + needs_fwl2=needs_fwl2, + ) + energy_inputs = build_energy_inputs( + exchange, + field, + mask, + system.mu if mu is None else mu, + eps, + ) + return context, energy_inputs + + +def build_multi_context( + systems: Iterable[SpinHamiltonian], + n_max: int, +): + from hamiltonzero.model import MultiSystemContext + + system_list = list(systems) + contexts = [build_context(system, n_max=n_max) for system in system_list] + return MultiSystemContext.stack(contexts) + + +__all__ = [ + "build_context", + "build_context_and_energy", + "build_multi_context", + "load_system", + "load_systems", + "padded_model_arrays", + "save_system", +] diff --git a/src/hamiltonzero/energy/__init__.py b/src/hamiltonzero/energy/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..546f284a75fa4831cc3eca555e72fc3edcd74024 --- /dev/null +++ b/src/hamiltonzero/energy/__init__.py @@ -0,0 +1,22 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from hamiltonzero.model._custom_lap_primitives import ( + custom_lap_active, + quadrilinear_merge_p, +) + +from .compiled import ( + vmc_energy_custom_lap_compiled, + vmc_energy_custom_lap_finetune, +) + + +__all__ = [ + "custom_lap_active", + "quadrilinear_merge_p", + "vmc_energy_custom_lap_compiled", + "vmc_energy_custom_lap_finetune", +] diff --git a/src/hamiltonzero/energy/compiled.py b/src/hamiltonzero/energy/compiled.py new file mode 100644 index 0000000000000000000000000000000000000000..9e0ddbe9e07381d104b571be44083e0f2aa1c721 --- /dev/null +++ b/src/hamiltonzero/energy/compiled.py @@ -0,0 +1,57 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +from typing import Any + +import jax + +from hamiltonzero.energy.kernel import ( + _vmc_energy_custom_lap_finetune, + _vmc_energy_custom_lap_prebuilt, +) + + +def vmc_energy_custom_lap_compiled( + kernel: Any, + tree: Any, + energy_frame: Any, + q_routed: jax.Array, + *, + chunk_size: int | None = 512, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array]: + return _vmc_energy_custom_lap_prebuilt( + kernel, + tree, + energy_frame, + q_routed, + chunk_size=chunk_size, + ) + + +def vmc_energy_custom_lap_finetune( + model: Any, + energy_frame: Any, + q_routed: jax.Array, + *, + chunk_size: int | None = 512, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array]: + from hamiltonzero.compiled.model import CompiledFinetuneWaveFunction + + if not isinstance(model, CompiledFinetuneWaveFunction): + raise TypeError("fine-tune energy requires CompiledFinetuneWaveFunction") + return _vmc_energy_custom_lap_finetune( + model, + energy_frame, + q_routed, + 0.0, + chunk_size=chunk_size, + ) + + +__all__ = [ + "vmc_energy_custom_lap_compiled", + "vmc_energy_custom_lap_finetune", +] diff --git a/src/hamiltonzero/energy/custom_lap.py b/src/hamiltonzero/energy/custom_lap.py new file mode 100644 index 0000000000000000000000000000000000000000..5d87a82d8ef9f66af696a7b0016d1e709d5a6f2b --- /dev/null +++ b/src/hamiltonzero/energy/custom_lap.py @@ -0,0 +1,1491 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +from typing import Any, Callable, NamedTuple + +import jax +import jax.numpy as jnp +from jax import lax +from jax.extend import core +from jax.extend.core import Literal + + +class JLP(NamedTuple): + value: Any + jac: Any + lap: Any + level: int + has_chunk_axis: bool + chunk_idx: int = -1 + + +def _is_jlp(x) -> bool: + return isinstance(x, JLP) + + +from hamiltonzero.model._custom_lap_primitives import ( + custom_lap_active, + enter_custom_lap, + restore_custom_lap, + quadrilinear_merge_p, +) + + +class use_custom_lap: + def __enter__(self): + self._cla_token = enter_custom_lap() + return self + + def __exit__(self, *exc): + restore_custom_lap(self._cla_token) + + +def build_W_levels(W_full, N: int) -> list: + + assert W_full.shape == (3 * N, 3 * N), ( + f"W_full must be [3N, 3N]; got {W_full.shape}" + ) + assert (N & (N - 1)) == 0, f"N must be a power of 2; got {N}" + levels = [] + k = 1 + while k <= N: + n_chunks = N // k + W_reshaped = W_full.reshape(n_chunks, 3 * k, n_chunks, 3 * k) + idx = jnp.arange(n_chunks) + W_k = W_reshaped[idx, :, idx, :] + levels.append(W_k) + k *= 2 + return levels + + +def _level_idx(k: int) -> int: + + assert k > 0 and (k & (k - 1)) == 0, f"k must be a power of 2; got {k}" + return k.bit_length() - 1 + + +def _W_at_level(W_levels, k: int): + return W_levels[_level_idx(k)] + + +_RULE_REGISTRY: dict[core.Primitive, Callable] = {} + + +def _params_except(params, *drop, **defaults): + + out = {k: params[k] for k in params if k not in drop} + for k, v in defaults.items(): + out.setdefault(k, v) + return out + + +def _shape_bind_params(params, *drop): + + out = {k: params[k] for k in params if k not in drop and k != "out_sharding"} + out.setdefault("sharding", params.get("out_sharding", None)) + return out + + +def _select_W_for_jlp( + level: int, chunk_idx: int, has_chunk_axis: bool, W_levels, M_jac: int +): + + W_k = _W_at_level(W_levels, level) + if has_chunk_axis: + assert W_k.shape[0] == M_jac, ( + f"has_chunk_axis: W_k chunks {W_k.shape[0]} must equal jac M {M_jac}" + ) + return W_k + if chunk_idx >= 0: + assert M_jac == 1 + return W_k[chunk_idx : chunk_idx + 1] + + assert W_k.shape[0] == M_jac, ( + f"multi-chunk: W_k chunks {W_k.shape[0]} must equal jac M {M_jac}" + ) + return W_k + + +def _jac_self_quad_form( + jac, level: int, chunk_idx: int, has_chunk_axis: bool, W_levels +): + + M = jac.shape[1] + W_used = _select_W_for_jlp(level, chunk_idx, has_chunk_axis, W_levels, M) + n_trailing = jac.ndim - 2 + if n_trailing == 0: + out = jnp.einsum("mc,cmn,nc->c", jac, W_used, jac) + elif n_trailing == 1: + out = jnp.einsum("mca,cmn,nca->ca", jac, W_used, jac) + elif n_trailing == 2: + out = jnp.einsum("mcab,cmn,ncab->cab", jac, W_used, jac) + elif n_trailing == 3: + out = jnp.einsum("mcabd,cmn,ncabd->cabd", jac, W_used, jac) + else: + raise NotImplementedError( + f"_jac_self_quad_form: trailing rank {n_trailing} not supported" + ) + if not has_chunk_axis: + out = out.sum(axis=0) + return out + + +def _make_unary_rule(prim, f_prime_fn, f_dprime_fn): + + def rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + v_out = prim.bind(x.value, **params) + fp = f_prime_fn(x.value) + jac_out = fp * x.jac + fpp = f_dprime_fn(x.value) + cross = _jac_self_quad_form( + x.jac, x.level, x.chunk_idx, x.has_chunk_axis, W_levels + ) + lap_out = fp * x.lap + fpp * cross + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=x.has_chunk_axis, + chunk_idx=x.chunk_idx, + ) + + _RULE_REGISTRY[prim] = rule + + +_make_unary_rule(lax.sin_p, jnp.cos, lambda x: -jnp.sin(x)) +_make_unary_rule(lax.cos_p, lambda x: -jnp.sin(x), lambda x: -jnp.cos(x)) +_make_unary_rule( + lax.tanh_p, + lambda x: 1.0 - jnp.tanh(x) ** 2, + lambda x: -2.0 * jnp.tanh(x) * (1.0 - jnp.tanh(x) ** 2), +) +_make_unary_rule(lax.exp_p, jnp.exp, jnp.exp) +_make_unary_rule(lax.log_p, lambda x: 1.0 / x, lambda x: -1.0 / (x * x)) +_make_unary_rule(lax.neg_p, lambda x: -jnp.ones_like(x), lambda x: jnp.zeros_like(x)) +_make_unary_rule(lax.abs_p, lambda x: jnp.sign(x), lambda x: jnp.zeros_like(x)) +_make_unary_rule( + lax.sqrt_p, lambda x: 0.5 / jnp.sqrt(x), lambda x: -0.25 / (x * jnp.sqrt(x)) +) +_make_unary_rule( + lax.rsqrt_p, + lambda x: -0.5 / (x * jnp.sqrt(x)), + lambda x: 0.75 / (x * x * jnp.sqrt(x)), +) + + +def _logistic_prime(x): + s = jax.nn.sigmoid(x) + return s * (1.0 - s) + + +def _logistic_dprime(x): + s = jax.nn.sigmoid(x) + return s * (1.0 - s) * (1.0 - 2.0 * s) + + +_logistic_p = lax.logistic_p +_make_unary_rule(_logistic_p, _logistic_prime, _logistic_dprime) + + +def _integer_pow_rule(invals, params, W_levels): + [x] = invals + y = params["y"] + assert _is_jlp(x) + v_out = lax.integer_pow_p.bind(x.value, **params) + if y == 0: + return JLP( + value=jnp.ones_like(x.value), + jac=jnp.zeros_like(x.jac), + lap=jnp.zeros_like(x.lap), + level=x.level, + has_chunk_axis=x.has_chunk_axis, + chunk_idx=x.chunk_idx, + ) + if y == 1: + return x + fp = float(y) * lax.integer_pow_p.bind(x.value, y=y - 1) + jac_out = fp * x.jac + fpp = float(y * (y - 1)) * lax.integer_pow_p.bind(x.value, y=max(y - 2, 0)) + cross = _jac_self_quad_form(x.jac, x.level, x.chunk_idx, x.has_chunk_axis, W_levels) + lap_out = fp * x.lap + fpp * cross + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=x.has_chunk_axis, + chunk_idx=x.chunk_idx, + ) + + +_RULE_REGISTRY[lax.integer_pow_p] = _integer_pow_rule + + +def _convert_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + new_dtype = params["new_dtype"] + return JLP( + value=x.value.astype(new_dtype), + jac=x.jac.astype(new_dtype), + lap=x.lap.astype(new_dtype), + level=x.level, + has_chunk_axis=x.has_chunk_axis, + chunk_idx=x.chunk_idx, + ) + + +_RULE_REGISTRY[lax.convert_element_type_p] = _convert_rule + + +def _broadcast_jac_lap_for_op(x, target_shape): + + lap_out = jnp.broadcast_to(x.lap, target_shape) + jac_out = _broadcast_jac_to_value_shape(x.jac, target_shape, x.has_chunk_axis) + return jac_out, lap_out + + +def _broadcast_jac_to_value_shape(jac, target_value_shape, has_chunk_axis: bool): + + if has_chunk_axis: + new_shape = (jac.shape[0], target_value_shape[0]) + tuple( + target_value_shape[1:] + ) + else: + new_shape = (jac.shape[0], jac.shape[1]) + tuple(target_value_shape) + while jac.ndim < len(new_shape): + jac = jnp.expand_dims(jac, axis=jac.ndim) + return jnp.broadcast_to(jac, new_shape) + + +def _promote_jlp_one_level(x: JLP) -> JLP: + + k = x.level + new_k = 2 * k + if x.has_chunk_axis: + assert x.chunk_idx in (0, 1), ( + f"_promote_jlp_one_level (chunked): chunk_idx must be 0 or 1; got {x}" + ) + is_left = x.chunk_idx == 0 + zero_shape = (3 * k,) + x.jac.shape[1:] + zeros = jnp.zeros(zero_shape, dtype=x.jac.dtype) + if is_left: + new_jac = jnp.concatenate([x.jac, zeros], axis=0) + else: + new_jac = jnp.concatenate([zeros, x.jac], axis=0) + return JLP( + value=x.value, + jac=new_jac, + lap=x.lap, + level=new_k, + has_chunk_axis=True, + chunk_idx=-1, + ) + + assert x.chunk_idx >= 0, ( + f"_promote_jlp_one_level: per-node requires chunk_idx>=0; got {x}" + ) + is_left = x.chunk_idx % 2 == 0 + new_chunk_idx = x.chunk_idx // 2 + zero_shape = (3 * k,) + x.jac.shape[1:] + zeros = jnp.zeros(zero_shape, dtype=x.jac.dtype) + if is_left: + new_jac = jnp.concatenate([x.jac, zeros], axis=0) + else: + new_jac = jnp.concatenate([zeros, x.jac], axis=0) + return JLP( + value=x.value, + jac=new_jac, + lap=x.lap, + level=new_k, + has_chunk_axis=False, + chunk_idx=new_chunk_idx, + ) + + +def _align_jlp_levels(a: JLP, b: JLP): + + while a.level < b.level: + a = _promote_jlp_one_level(a) + while b.level < a.level: + b = _promote_jlp_one_level(b) + assert a.has_chunk_axis == b.has_chunk_axis, ( + "_align_jlp_levels: chunk-axis mismatch" + ) + if a.has_chunk_axis: + if ( + a.chunk_idx in (0, 1) + and b.chunk_idx in (0, 1) + and a.chunk_idx != b.chunk_idx + ): + a = _promote_jlp_one_level(a) + b = _promote_jlp_one_level(b) + return a, b + + while a.chunk_idx != b.chunk_idx: + a = _promote_jlp_one_level(a) + b = _promote_jlp_one_level(b) + return a, b + + +def _add_or_sub_rule(sign: float): + def rule(invals, params, W_levels): + a, b = invals + if _is_jlp(a) and _is_jlp(b): + need_promote = ( + (a.level != b.level) + or ( + not a.has_chunk_axis + and not b.has_chunk_axis + and a.chunk_idx != b.chunk_idx + ) + or ( + a.has_chunk_axis + and b.has_chunk_axis + and a.chunk_idx in (0, 1) + and b.chunk_idx in (0, 1) + and a.chunk_idx != b.chunk_idx + ) + ) + if need_promote: + a, b = _align_jlp_levels(a, b) + assert a.has_chunk_axis == b.has_chunk_axis, "add/sub: chunk-axis mismatch" + v_out = a.value + sign * b.value + a_jac_b = _broadcast_jac_to_value_shape( + a.jac, v_out.shape, a.has_chunk_axis + ) + b_jac_b = _broadcast_jac_to_value_shape( + b.jac, v_out.shape, b.has_chunk_axis + ) + jac_out = a_jac_b + sign * b_jac_b + a_lap_b = jnp.broadcast_to(a.lap, v_out.shape) + b_lap_b = jnp.broadcast_to(b.lap, v_out.shape) + lap_out = a_lap_b + sign * b_lap_b + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + if _is_jlp(a): + v_out = a.value + sign * b + jac_out, lap_out = _broadcast_jac_lap_for_op(a, v_out.shape) + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + + v_out = a + sign * b.value + jac_out, lap_out = _broadcast_jac_lap_for_op(b, v_out.shape) + return JLP( + value=v_out, + jac=sign * jac_out, + lap=sign * lap_out, + level=b.level, + has_chunk_axis=b.has_chunk_axis, + chunk_idx=b.chunk_idx, + ) + + return rule + + +_RULE_REGISTRY[lax.add_p] = _add_or_sub_rule(+1.0) +_RULE_REGISTRY[lax.sub_p] = _add_or_sub_rule(-1.0) + + +def _mul_rule(invals, params, W_levels): + a, b = invals + if _is_jlp(a) and _is_jlp(b): + need_promote = (a.level != b.level) or ( + not a.has_chunk_axis and not b.has_chunk_axis and a.chunk_idx != b.chunk_idx + ) + if need_promote: + a, b = _align_jlp_levels(a, b) + assert a.has_chunk_axis == b.has_chunk_axis + + v_out = a.value * b.value + a_jac_b = _broadcast_jac_to_value_shape(a.jac, v_out.shape, a.has_chunk_axis) + b_jac_b = _broadcast_jac_to_value_shape(b.jac, v_out.shape, b.has_chunk_axis) + a_val_b = jnp.broadcast_to(a.value, v_out.shape) + b_val_b = jnp.broadcast_to(b.value, v_out.shape) + jac_out = a_val_b * b_jac_b + b_val_b * a_jac_b + a_lap_b = jnp.broadcast_to(a.lap, v_out.shape) + b_lap_b = jnp.broadcast_to(b.lap, v_out.shape) + cross = _jac_cross_quad_form( + a_jac_b, b_jac_b, a.level, a.chunk_idx, a.has_chunk_axis, W_levels + ) + lap_out = a_val_b * b_lap_b + b_val_b * a_lap_b + 2.0 * cross + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + if _is_jlp(a): + v_out = a.value * b + jac_b, lap_b = _broadcast_jac_lap_for_op(a, v_out.shape) + return JLP( + value=v_out, + jac=b * jac_b, + lap=b * lap_b, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + + v_out = a * b.value + jac_b, lap_b = _broadcast_jac_lap_for_op(b, v_out.shape) + return JLP( + value=v_out, + jac=a * jac_b, + lap=a * lap_b, + level=b.level, + has_chunk_axis=b.has_chunk_axis, + chunk_idx=b.chunk_idx, + ) + + +def _jac_cross_quad_form( + jac_a, jac_b, level: int, chunk_idx: int, has_chunk_axis: bool, W_levels +): + + M = jac_a.shape[1] + W_used = _select_W_for_jlp(level, chunk_idx, has_chunk_axis, W_levels, M) + n_trailing = jac_a.ndim - 2 + if n_trailing == 0: + out = jnp.einsum("mc,cmn,nc->c", jac_a, W_used, jac_b) + elif n_trailing == 1: + out = jnp.einsum("mca,cmn,nca->ca", jac_a, W_used, jac_b) + elif n_trailing == 2: + out = jnp.einsum("mcab,cmn,ncab->cab", jac_a, W_used, jac_b) + elif n_trailing == 3: + out = jnp.einsum("mcabd,cmn,ncabd->cabd", jac_a, W_used, jac_b) + else: + raise NotImplementedError( + f"_jac_cross_quad_form: trailing rank {n_trailing} not supported" + ) + if not has_chunk_axis: + out = out.sum(axis=0) + return out + + +_RULE_REGISTRY[lax.mul_p] = _mul_rule + + +def _div_rule(invals, params, W_levels): + a, b = invals + if _is_jlp(a) and _is_jlp(b): + need_promote = (a.level != b.level) or ( + not a.has_chunk_axis and not b.has_chunk_axis and a.chunk_idx != b.chunk_idx + ) + if need_promote: + a, b = _align_jlp_levels(a, b) + assert a.has_chunk_axis == b.has_chunk_axis + v_out = a.value / b.value + a_val_b = jnp.broadcast_to(a.value, v_out.shape) + b_val_b = jnp.broadcast_to(b.value, v_out.shape) + a_jac_b = _broadcast_jac_to_value_shape(a.jac, v_out.shape, a.has_chunk_axis) + b_jac_b = _broadcast_jac_to_value_shape(b.jac, v_out.shape, b.has_chunk_axis) + inv_b = 1.0 / b_val_b + jac_out = inv_b * a_jac_b - (a_val_b * inv_b * inv_b) * b_jac_b + a_lap_b = jnp.broadcast_to(a.lap, v_out.shape) + b_lap_b = jnp.broadcast_to(b.lap, v_out.shape) + lap_first = inv_b * a_lap_b - (a_val_b * inv_b * inv_b) * b_lap_b + cross_bb = _jac_self_quad_form( + b_jac_b, b.level, b.chunk_idx, b.has_chunk_axis, W_levels + ) + cross_ab = _jac_cross_quad_form( + a_jac_b, b_jac_b, a.level, a.chunk_idx, a.has_chunk_axis, W_levels + ) + lap_second = (2.0 * a_val_b * inv_b**3) * cross_bb + ( + -2.0 * inv_b * inv_b + ) * cross_ab + lap_out = lap_first + lap_second + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + if _is_jlp(a): + inv_b = 1.0 / b + v_out = a.value * inv_b + jac_b, lap_b = _broadcast_jac_lap_for_op(a, v_out.shape) + return JLP( + value=v_out, + jac=inv_b * jac_b, + lap=inv_b * lap_b, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + + v_out = a / b.value + inv_b = 1.0 / b.value + fp = -a * inv_b * inv_b + jac_out = fp * b.jac + fpp = 2.0 * a * inv_b**3 + cross = _jac_self_quad_form(b.jac, b.level, b.chunk_idx, b.has_chunk_axis, W_levels) + lap_out = fp * b.lap + fpp * cross + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=b.level, + has_chunk_axis=b.has_chunk_axis, + chunk_idx=b.chunk_idx, + ) + + +_RULE_REGISTRY[lax.div_p] = _div_rule + + +def _atan2_rule(invals, params, W_levels): + a_arg, b_arg = invals + if _is_jlp(a_arg) and _is_jlp(b_arg): + need_promote = (a_arg.level != b_arg.level) or ( + not a_arg.has_chunk_axis + and not b_arg.has_chunk_axis + and a_arg.chunk_idx != b_arg.chunk_idx + ) + if need_promote: + a_arg, b_arg = _align_jlp_levels(a_arg, b_arg) + assert a_arg.has_chunk_axis == b_arg.has_chunk_axis + a, b = a_arg, b_arg + elif _is_jlp(a_arg): + a = a_arg + b = _trivial_jlp_like(b_arg, a_arg) + else: + b = b_arg + a = _trivial_jlp_like(a_arg, b_arg) + v_out = jnp.arctan2(a.value, b.value) + r_sq = a.value**2 + b.value**2 + inv_r2 = 1.0 / r_sq + fa = b.value * inv_r2 + fb = -a.value * inv_r2 + aj = _broadcast_jac_to_value_shape(a.jac, v_out.shape, a.has_chunk_axis) + bj = _broadcast_jac_to_value_shape(b.jac, v_out.shape, b.has_chunk_axis) + jac_out = fa * aj + fb * bj + faa = -2.0 * a.value * b.value * inv_r2 * inv_r2 + fbb = 2.0 * a.value * b.value * inv_r2 * inv_r2 + fab = (a.value**2 - b.value**2) * inv_r2 * inv_r2 + al = jnp.broadcast_to(a.lap, v_out.shape) + bl = jnp.broadcast_to(b.lap, v_out.shape) + cross_aa = _jac_self_quad_form(aj, a.level, a.chunk_idx, a.has_chunk_axis, W_levels) + cross_bb = _jac_self_quad_form(bj, b.level, b.chunk_idx, b.has_chunk_axis, W_levels) + cross_ab = _jac_cross_quad_form( + aj, bj, a.level, a.chunk_idx, a.has_chunk_axis, W_levels + ) + lap_out = fa * al + fb * bl + faa * cross_aa + fbb * cross_bb + 2.0 * fab * cross_ab + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + + +def _trivial_jlp_like(plain_val, template_jlp: JLP) -> JLP: + + val = jnp.asarray(plain_val) + if template_jlp.has_chunk_axis: + jac_shape = (template_jlp.jac.shape[0], template_jlp.jac.shape[1]) + tuple( + val.shape[1:] + ) + else: + jac_shape = (template_jlp.jac.shape[0], template_jlp.jac.shape[1]) + tuple( + val.shape + ) + return JLP( + value=val, + jac=jnp.zeros(jac_shape, dtype=val.dtype), + lap=jnp.zeros_like(val), + level=template_jlp.level, + has_chunk_axis=template_jlp.has_chunk_axis, + chunk_idx=template_jlp.chunk_idx, + ) + + +_RULE_REGISTRY[lax.atan2_p] = _atan2_rule + + +def _dot_general_rule(invals, params, W_levels): + lhs, rhs = invals + dimension_numbers = params["dimension_numbers"] + (lhs_contract, rhs_contract), (lhs_batch, rhs_batch) = dimension_numbers + + if _is_jlp(lhs) and _is_jlp(rhs): + raise NotImplementedError( + "dot_general: JLP × JLP outside quadrilinear_merge_p is not supported. " + "If a model op needs same-level bilinear, route it through the merge " + "primitive or rewrite as elementwise mul + reduce_sum." + ) + if _is_jlp(lhs): + return _dot_general_jlp_plain( + lhs, rhs, dimension_numbers, params, jlp_is_lhs=True + ) + return _dot_general_jlp_plain(rhs, lhs, dimension_numbers, params, jlp_is_lhs=False) + + +def _dot_general_jlp_plain( + jlp_arg, plain_arg, dimension_numbers, params, jlp_is_lhs: bool +): + + (lhs_contract, rhs_contract), (lhs_batch, rhs_batch) = dimension_numbers + dot_kw = { + "precision": params.get("precision", None), + "preferred_element_type": params.get("preferred_element_type", None), + "out_sharding": params.get("out_sharding", None), + } + + if jlp_is_lhs: + v_out = lax.dot_general(jlp_arg.value, plain_arg, dimension_numbers, **dot_kw) + lap_out = lax.dot_general(jlp_arg.lap, plain_arg, dimension_numbers, **dot_kw) + else: + v_out = lax.dot_general(plain_arg, jlp_arg.value, dimension_numbers, **dot_kw) + lap_out = lax.dot_general(plain_arg, jlp_arg.lap, dimension_numbers, **dot_kw) + + jac = jlp_arg.jac + leading_3k = jac.shape[0] + + if jlp_arg.has_chunk_axis: + jac_for_dot = jac + shift = 1 + else: + jac_for_dot = jac.reshape((leading_3k,) + tuple(jlp_arg.value.shape)) + shift = 1 + + if jlp_is_lhs: + new_lhs_contract = tuple(a + shift for a in lhs_contract) + new_rhs_contract = tuple(rhs_contract) + new_lhs_batch = tuple(a + shift for a in lhs_batch) + new_rhs_batch = tuple(rhs_batch) + new_dim_nums = ( + (new_lhs_contract, new_rhs_contract), + (new_lhs_batch, new_rhs_batch), + ) + jac_out_raw = lax.dot_general(jac_for_dot, plain_arg, new_dim_nums, **dot_kw) + + n_batch = len(new_lhs_batch) + pos_3k = n_batch + else: + new_lhs_contract = tuple(lhs_contract) + new_rhs_contract = tuple(a + shift for a in rhs_contract) + new_lhs_batch = tuple(lhs_batch) + new_rhs_batch = tuple(a + shift for a in rhs_batch) + new_dim_nums = ( + (new_lhs_contract, new_rhs_contract), + (new_lhs_batch, new_rhs_batch), + ) + jac_out_raw = lax.dot_general(plain_arg, jac_for_dot, new_dim_nums, **dot_kw) + + n_batch = len(new_lhs_batch) + lhs_ndim = jnp.asarray(plain_arg).ndim + n_lhs_nonbatch = lhs_ndim - n_batch - len(new_lhs_contract) + pos_3k = n_batch + n_lhs_nonbatch + + if pos_3k != 0: + jac_out = jnp.moveaxis(jac_out_raw, pos_3k, 0) + else: + jac_out = jac_out_raw + + if jlp_arg.has_chunk_axis: + new_has_chunk_axis = True + new_chunk_idx = -1 + else: + jac_out = jac_out[:, None] + new_has_chunk_axis = False + new_chunk_idx = jlp_arg.chunk_idx + + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=jlp_arg.level, + has_chunk_axis=new_has_chunk_axis, + chunk_idx=new_chunk_idx, + ) + + +_RULE_REGISTRY[lax.dot_general_p] = _dot_general_rule + + +def _broadcast_in_dim_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + shape = params["shape"] + broadcast_dimensions = params["broadcast_dimensions"] + extra = _shape_bind_params(params, "shape", "broadcast_dimensions") + + v_out = lax.broadcast_in_dim_p.bind( + x.value, shape=shape, broadcast_dimensions=broadcast_dimensions, **extra + ) + lap_out = lax.broadcast_in_dim_p.bind( + x.lap, shape=shape, broadcast_dimensions=broadcast_dimensions, **extra + ) + + if x.has_chunk_axis: + new_shape = (x.jac.shape[0],) + tuple(shape) + new_bd = (0,) + tuple(d + 1 for d in broadcast_dimensions) + jac_out = lax.broadcast_in_dim_p.bind( + x.jac.reshape((x.jac.shape[0],) + x.value.shape), + shape=new_shape, + broadcast_dimensions=new_bd, + **extra, + ) + + new_has_chunk_axis = True + new_chunk_idx = -1 + else: + new_shape = (x.jac.shape[0], 1) + tuple(shape) + new_bd = (0, 1) + tuple(d + 2 for d in broadcast_dimensions) + jac_out = lax.broadcast_in_dim_p.bind( + x.jac, + shape=new_shape, + broadcast_dimensions=new_bd, + **extra, + ) + new_has_chunk_axis = False + new_chunk_idx = x.chunk_idx + + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=new_has_chunk_axis, + chunk_idx=new_chunk_idx, + ) + + +_RULE_REGISTRY[lax.broadcast_in_dim_p] = _broadcast_in_dim_rule + + +def _reduce_sum_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + axes = tuple(params["axes"]) + extra = _params_except(params, "axes", out_sharding=None) + v_out = lax.reduce_sum_p.bind(x.value, axes=axes, **extra) + lap_out = lax.reduce_sum_p.bind(x.lap, axes=axes, **extra) + if x.has_chunk_axis: + if 0 in axes: + non_chunk_axes = tuple(a for a in axes if a != 0) + jac_reduce_axes = tuple(a + 1 for a in non_chunk_axes) + if jac_reduce_axes: + jac_out = lax.reduce_sum_p.bind(x.jac, axes=jac_reduce_axes, **extra) + else: + jac_out = x.jac + new_has_chunk_axis = False + new_chunk_idx = -1 + else: + jac_reduce_axes = tuple(a + 1 for a in axes) + jac_out = lax.reduce_sum_p.bind(x.jac, axes=jac_reduce_axes, **extra) + new_has_chunk_axis = True + new_chunk_idx = -1 + else: + jac_reduce_axes = tuple(a + 2 for a in axes) + jac_out = lax.reduce_sum_p.bind(x.jac, axes=jac_reduce_axes, **extra) + new_has_chunk_axis = False + new_chunk_idx = x.chunk_idx + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=new_has_chunk_axis, + chunk_idx=new_chunk_idx, + ) + + +_RULE_REGISTRY[lax.reduce_sum_p] = _reduce_sum_rule + + +def _reduce_max_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + axes = tuple(params["axes"]) + extra = _params_except(params, "axes", out_sharding=None) + v_out = lax.reduce_max_p.bind(x.value, axes=axes, **extra) + + keep_shape = list(x.value.shape) + for a in axes: + keep_shape[a] = 1 + v_max_kept = lax.reduce_max_p.bind(x.value, axes=axes, **extra).reshape( + tuple(keep_shape) + ) + mask = (x.value == v_max_kept).astype(x.value.dtype) + + mask_sum_axes = lax.reduce_sum_p.bind( + mask, + axes=axes, + **extra, + ).reshape(tuple(keep_shape)) + mask = mask / (mask_sum_axes + 1e-30) + + def _gated_reduce(jac_arr, jac_axes): + + n_lead = jac_arr.ndim - x.value.ndim + mask_b = mask.reshape((1,) * n_lead + mask.shape) + + return lax.reduce_sum_p.bind( + jac_arr * mask_b, + axes=jac_axes, + **extra, + ) + + if x.has_chunk_axis: + if 0 in axes: + non_chunk_axes = tuple(a for a in axes if a != 0) + jac_reduce_axes = tuple(a + 1 for a in non_chunk_axes) + (1,) + jac_out = _gated_reduce(x.jac, jac_reduce_axes) + new_has_chunk_axis = False + new_chunk_idx = -1 + else: + jac_reduce_axes = tuple(a + 1 for a in axes) + jac_out = _gated_reduce(x.jac, jac_reduce_axes) + new_has_chunk_axis = True + new_chunk_idx = -1 + else: + jac_reduce_axes = tuple(a + 2 for a in axes) + jac_out = _gated_reduce(x.jac, jac_reduce_axes) + new_has_chunk_axis = False + new_chunk_idx = x.chunk_idx + lap_out = jnp.zeros_like(v_out) + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=new_has_chunk_axis, + chunk_idx=new_chunk_idx, + ) + + +_RULE_REGISTRY[lax.reduce_max_p] = _reduce_max_rule + + +def _make_minmax_rule(sign_for_first): + + def rule(invals, params, W_levels): + a, b = invals + if _is_jlp(a) and _is_jlp(b): + need_promote = ( + (a.level != b.level) + or ( + not a.has_chunk_axis + and not b.has_chunk_axis + and a.chunk_idx != b.chunk_idx + ) + or ( + a.has_chunk_axis + and b.has_chunk_axis + and a.chunk_idx in (0, 1) + and b.chunk_idx in (0, 1) + and a.chunk_idx != b.chunk_idx + ) + ) + if need_promote: + a, b = _align_jlp_levels(a, b) + cmp = (a.value - b.value) * sign_for_first + mask_a = (cmp > 0).astype(a.value.dtype) + mask_b = 1.0 - mask_a + v_out = mask_a * a.value + mask_b * b.value + + jac_a_b = _broadcast_jac_to_value_shape( + a.jac, v_out.shape, a.has_chunk_axis + ) + jac_b_b = _broadcast_jac_to_value_shape( + b.jac, v_out.shape, b.has_chunk_axis + ) + n_lead_a = jac_a_b.ndim - mask_a.ndim + n_lead_b = jac_b_b.ndim - mask_b.ndim + mask_a_lead = mask_a.reshape((1,) * n_lead_a + mask_a.shape) + mask_b_lead = mask_b.reshape((1,) * n_lead_b + mask_b.shape) + jac_out = mask_a_lead * jac_a_b + mask_b_lead * jac_b_b + lap_out = mask_a * jnp.broadcast_to( + a.lap, v_out.shape + ) + mask_b * jnp.broadcast_to(b.lap, v_out.shape) + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + if _is_jlp(a): + cmp = (a.value - b) * sign_for_first + mask_a = (cmp > 0).astype(a.value.dtype) + v_out = mask_a * a.value + (1 - mask_a) * b + n_lead = a.jac.ndim - mask_a.ndim + mask_a_lead = mask_a.reshape((1,) * n_lead + mask_a.shape) + jac_out = mask_a_lead * a.jac + lap_out = mask_a * a.lap + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=a.level, + has_chunk_axis=a.has_chunk_axis, + chunk_idx=a.chunk_idx, + ) + + cmp = (a - b.value) * sign_for_first + mask_a = (cmp > 0).astype(b.value.dtype) + v_out = mask_a * a + (1 - mask_a) * b.value + mask_b = 1 - mask_a + n_lead = b.jac.ndim - mask_b.ndim + mask_b_lead = mask_b.reshape((1,) * n_lead + mask_b.shape) + jac_out = mask_b_lead * b.jac + lap_out = mask_b * b.lap + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=b.level, + has_chunk_axis=b.has_chunk_axis, + chunk_idx=b.chunk_idx, + ) + + return rule + + +_RULE_REGISTRY[lax.max_p] = _make_minmax_rule(+1.0) +_RULE_REGISTRY[lax.min_p] = _make_minmax_rule(-1.0) + + +def _reshape_jac_for_value(x: JLP, new_value_shape: tuple, new_dimensions): + + leading_3k = x.jac.shape[0] + if x.has_chunk_axis: + M = x.value.shape[0] + old_trailing = x.value.shape[1:] + + if len(new_value_shape) >= 1 and new_value_shape[0] == M: + jac_axes_perm = None + if new_dimensions is not None: + jac_axes_perm = (0,) + tuple(d + 1 for d in new_dimensions) + jac_pre = jnp.transpose(x.jac, jac_axes_perm) + else: + jac_pre = x.jac + new_jac_shape = (leading_3k,) + tuple(new_value_shape) + jac_new = jnp.reshape(jac_pre, new_jac_shape) + + return jac_new, True, -1 + + if M == 1: + assert new_dimensions is None, ( + "reshape with M=1-squeeze + transpose not yet supported" + ) + + new_jac_shape = (leading_3k, 1) + tuple(new_value_shape) + + jac_new = jnp.reshape(x.jac, new_jac_shape) + return jac_new, False, 0 + raise NotImplementedError( + f"reshape: cannot reshape JLP value {x.value.shape} (has_chunk_axis, M={M}) " + f"to {new_value_shape} — chunk axis would be fused/lost." + ) + + assert new_dimensions is None or all(d >= 0 for d in new_dimensions) + jac_pre = x.jac + if new_dimensions is not None: + jac_axes_perm = (0, 1) + tuple(d + 2 for d in new_dimensions) + jac_pre = jnp.transpose(jac_pre, jac_axes_perm) + new_jac_shape = (leading_3k, 1) + tuple(new_value_shape) + jac_new = jnp.reshape(jac_pre, new_jac_shape) + return jac_new, False, x.chunk_idx + + +def _reshape_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + new_sizes = tuple(params["new_sizes"]) + dimensions = params.get("dimensions") + + extra = _shape_bind_params(params, "new_sizes", "dimensions") + v_out = lax.reshape_p.bind( + x.value, new_sizes=new_sizes, dimensions=dimensions, **extra + ) + lap_out = lax.reshape_p.bind( + x.lap, new_sizes=new_sizes, dimensions=dimensions, **extra + ) + jac_new, new_has_chunk, new_idx = _reshape_jac_for_value(x, new_sizes, dimensions) + return JLP( + value=v_out, + jac=jac_new, + lap=lap_out, + level=x.level, + has_chunk_axis=new_has_chunk, + chunk_idx=new_idx, + ) + + +_RULE_REGISTRY[lax.reshape_p] = _reshape_rule + + +def _transpose_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + perm = tuple(params["permutation"]) + extra = {k: params[k] for k in params if k != "permutation"} + v_out = lax.transpose_p.bind(x.value, permutation=perm, **extra) + lap_out = lax.transpose_p.bind(x.lap, permutation=perm, **extra) + + if x.has_chunk_axis: + if perm[0] != 0: + raise NotImplementedError( + "transpose: chunk axis (value axis 0) must remain at position 0; " + f"got permutation {perm}." + ) + + jac_perm = (0, 1) + tuple(p + 1 for p in perm[1:]) + else: + jac_perm = (0, 1) + tuple(p + 2 for p in perm) + jac_new = lax.transpose_p.bind(x.jac, permutation=jac_perm, **extra) + return JLP( + value=v_out, + jac=jac_new, + lap=lap_out, + level=x.level, + has_chunk_axis=x.has_chunk_axis, + chunk_idx=x.chunk_idx, + ) + + +_RULE_REGISTRY[lax.transpose_p] = _transpose_rule + + +def _slice_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + start = tuple(params["start_indices"]) + limit = tuple(params["limit_indices"]) + strides = params.get("strides") + extra = { + k: params[k] + for k in params + if k not in ("start_indices", "limit_indices", "strides") + } + v_out = lax.slice_p.bind( + x.value, start_indices=start, limit_indices=limit, strides=strides, **extra + ) + lap_out = lax.slice_p.bind( + x.lap, start_indices=start, limit_indices=limit, strides=strides, **extra + ) + + if x.has_chunk_axis: + new_chunk_idx = x.chunk_idx + old_M = x.value.shape[0] + stride0 = strides[0] if strides is not None else 1 + v_out_M = v_out.shape[0] + chunk_axis_touched = start[0] != 0 or limit[0] != old_M or stride0 != 1 + if chunk_axis_touched: + if v_out_M == 1: + new_chunk_idx = start[0] + elif stride0 > 1 and v_out_M * stride0 == old_M and start[0] in (0, 1): + new_chunk_idx = start[0] + elif v_out_M != old_M: + raise NotImplementedError( + f"slice on chunk axis: unsupported sub-range " + f"start={start[0]}, limit={limit[0]}, stride={stride0}, " + f"old_M={old_M}, v_out_M={v_out_M}" + ) + jac_start = (0, start[0]) + tuple(start[1:]) + jac_limit = (x.jac.shape[0], limit[0]) + tuple(limit[1:]) + jac_strides = None if strides is None else (1, strides[0]) + tuple(strides[1:]) + jac_out = lax.slice_p.bind( + x.jac, + start_indices=jac_start, + limit_indices=jac_limit, + strides=jac_strides, + **extra, + ) + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=True, + chunk_idx=new_chunk_idx, + ) + + jac_start = (0, 0) + tuple(start) + jac_limit = (x.jac.shape[0], x.jac.shape[1]) + tuple(limit) + jac_strides = None if strides is None else (1, 1) + tuple(strides) + jac_out = lax.slice_p.bind( + x.jac, + start_indices=jac_start, + limit_indices=jac_limit, + strides=jac_strides, + **extra, + ) + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=False, + chunk_idx=x.chunk_idx, + ) + + +_RULE_REGISTRY[lax.slice_p] = _slice_rule + + +def _squeeze_rule(invals, params, W_levels): + [x] = invals + assert _is_jlp(x) + dims = tuple(params["dimensions"]) + extra = {k: params[k] for k in params if k != "dimensions"} + v_out = lax.squeeze_p.bind(x.value, dimensions=dims, **extra) + lap_out = lax.squeeze_p.bind(x.lap, dimensions=dims, **extra) + if x.has_chunk_axis and 0 in dims: + non_chunk_dims = tuple(d for d in dims if d != 0) + jac_dims = tuple(d + 1 for d in non_chunk_dims) + + if jac_dims: + jac_out = lax.squeeze_p.bind(x.jac, dimensions=jac_dims, **extra) + else: + jac_out = x.jac + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=False, + chunk_idx=x.chunk_idx, + ) + if x.has_chunk_axis: + jac_dims = tuple(d + 1 for d in dims) + else: + jac_dims = tuple(d + 2 for d in dims) + if jac_dims: + jac_out = lax.squeeze_p.bind(x.jac, dimensions=jac_dims, **extra) + else: + jac_out = x.jac + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=x.level, + has_chunk_axis=x.has_chunk_axis, + chunk_idx=x.chunk_idx, + ) + + +_RULE_REGISTRY[lax.squeeze_p] = _squeeze_rule + + +_stack_p = lax.stack_p + + +def _stack_rule(invals, params, W_levels): + del W_levels + axis = int(params["axis"]) + extra = {k: params[k] for k in params if k != "axis"} + + jlp_inputs = [v for v in invals if _is_jlp(v)] + levels = {v.level for v in jlp_inputs} + assert len(levels) == 1, f"stack: mixed JLP levels {levels}" + level = next(iter(levels)) + has_chunk = jlp_inputs[0].has_chunk_axis + chunk_idx = jlp_inputs[0].chunk_idx + for v in jlp_inputs[1:]: + assert v.has_chunk_axis == has_chunk and v.chunk_idx == chunk_idx, ( + "stack: inconsistent chunk metadata" + ) + + if has_chunk and axis == 0: + raise NotImplementedError( + "stack: inserting an axis before the JLP chunk axis is not supported" + ) + + values = [v.value if _is_jlp(v) else v for v in invals] + value_out = _stack_p.bind(*values, axis=axis, **extra) + + laps = [v.lap if _is_jlp(v) else jnp.zeros_like(v) for v in invals] + lap_out = _stack_p.bind(*laps, axis=axis, **extra) + + jac_axis = axis + (1 if has_chunk else 2) + ref_jac = jlp_inputs[0].jac + jacs = [v.jac if _is_jlp(v) else jnp.zeros_like(ref_jac) for v in invals] + jac_out = _stack_p.bind(*jacs, axis=jac_axis, **extra) + return JLP( + value=value_out, + jac=jac_out, + lap=lap_out, + level=level, + has_chunk_axis=has_chunk, + chunk_idx=chunk_idx, + ) + + +_RULE_REGISTRY[_stack_p] = _stack_rule + + +def _concatenate_rule(invals, params, W_levels): + dim = params["dimension"] + extra = {k: params[k] for k in params if k != "dimension"} + + jlp_inputs = [v for v in invals if _is_jlp(v)] + levels = set(v.level for v in jlp_inputs) + assert len(levels) == 1, f"concatenate: mixed JLP levels {levels}" + level = next(iter(levels)) + has_chunk = jlp_inputs[0].has_chunk_axis + chunk_idx = jlp_inputs[0].chunk_idx + for v in jlp_inputs[1:]: + assert v.has_chunk_axis == has_chunk and v.chunk_idx == chunk_idx, ( + "concatenate: inconsistent chunk metadata" + ) + if has_chunk and dim == 0: + raise NotImplementedError("concatenate on chunk axis not supported") + + values = [v.value if _is_jlp(v) else v for v in invals] + v_out = lax.concatenate_p.bind(*values, dimension=dim, **extra) + laps = [v.lap if _is_jlp(v) else jnp.zeros_like(v) for v in invals] + lap_out = lax.concatenate_p.bind(*laps, dimension=dim, **extra) + + jac_dim = dim + 2 if not has_chunk else dim + 1 + jacs = [] + for v in invals: + if _is_jlp(v): + jacs.append(v.jac) + else: + shape_v = v.shape if hasattr(v, "shape") else jnp.asarray(v).shape + if has_chunk: + M = jlp_inputs[0].jac.shape[1] + jac_shape = (jlp_inputs[0].jac.shape[0], M) + tuple(shape_v[1:]) + else: + M = jlp_inputs[0].jac.shape[1] + jac_shape = (jlp_inputs[0].jac.shape[0], M) + tuple(shape_v) + jacs.append(jnp.zeros(jac_shape, dtype=jlp_inputs[0].jac.dtype)) + jac_out = lax.concatenate_p.bind(*jacs, dimension=jac_dim, **extra) + return JLP( + value=v_out, + jac=jac_out, + lap=lap_out, + level=level, + has_chunk_axis=has_chunk, + chunk_idx=chunk_idx, + ) + + +_RULE_REGISTRY[lax.concatenate_p] = _concatenate_rule + + +def _jit_p_rule(invals, params, W_levels): + + inner_jaxpr = params["jaxpr"] + j = inner_jaxpr.jaxpr + consts = inner_jaxpr.consts + env: dict = {} + for cv, c in zip(j.constvars, consts): + env[cv] = c + for iv, x in zip(j.invars, invals): + env[iv] = x + for eqn in j.eqns: + outvals = _eval_eqn(eqn, env, W_levels) + for ov, ov_val in zip(eqn.outvars, outvals): + env[ov] = ov_val + return [env[ov] for ov in j.outvars] + + +from jax._src import pjit as _pjit_module + +_RULE_REGISTRY[_pjit_module.jit_p] = _jit_p_rule + + +def _quadrilinear_merge_rule(invals, params, W_levels): + + T, u_a, u_b = invals + assert not _is_jlp(T), "quadrilinear_merge: T must be plain" + assert _is_jlp(u_a) and _is_jlp(u_b), "quadrilinear_merge: u_a, u_b must be JLPs" + assert u_a.level == u_b.level, ( + f"quadrilinear_merge: leg levels differ {u_a.level} vs {u_b.level}" + ) + assert u_a.has_chunk_axis and u_b.has_chunk_axis, ( + "quadrilinear_merge requires the compiled chunk axis" + ) + + k = u_a.level + new_k = 2 * k + G, d_r, _, _ = T.shape + d_m_eff = G * d_r + M = u_a.value.shape[0] + assert u_b.value.shape[0] == M, "quadrilinear_merge: chunked legs must agree on M" + u_a_2d = u_a.value.reshape(M, G, d_r) + u_b_2d = u_b.value.reshape(M, G, d_r) + raw_value_2d = jnp.einsum("ijkl,mik,mil->mij", T, u_a_2d, u_b_2d) + raw_value = raw_value_2d.reshape(M, d_m_eff) + ua_jac_2d = u_a.jac.reshape(u_a.jac.shape[0], M, G, d_r) + ub_jac_2d = u_b.jac.reshape(u_b.jac.shape[0], M, G, d_r) + jac_upper_2d = jnp.einsum( + "ijkl,Amik,mil->Amij", + T, + ua_jac_2d, + u_b_2d, + ) + jac_lower_2d = jnp.einsum( + "ijkl,mik,Bmil->Bmij", + T, + u_a_2d, + ub_jac_2d, + ) + jac_upper = jac_upper_2d.reshape(jac_upper_2d.shape[0], M, d_m_eff) + jac_lower = jac_lower_2d.reshape(jac_lower_2d.shape[0], M, d_m_eff) + jac_out = jnp.concatenate([jac_upper, jac_lower], axis=0) + ua_lap_2d = u_a.lap.reshape(M, G, d_r) + ub_lap_2d = u_b.lap.reshape(M, G, d_r) + term1_2d = jnp.einsum( + "ijkl,mik,mil->mij", + T, + ua_lap_2d, + u_b_2d, + ) + term2_2d = jnp.einsum( + "ijkl,mik,mil->mij", + T, + u_a_2d, + ub_lap_2d, + ) + W_2k = _W_at_level(W_levels, new_k) + W_off = W_2k[:, : 3 * k, 3 * k :] + cross_2d = jnp.einsum( + "ijkl,Amik,Bmil,mAB->mij", + T, + ua_jac_2d, + ub_jac_2d, + W_off, + ) + lap_out_2d = term1_2d + term2_2d + 2.0 * cross_2d + lap_out = lap_out_2d.reshape(M, d_m_eff) + return JLP( + value=raw_value, + jac=jac_out, + lap=lap_out, + level=new_k, + has_chunk_axis=True, + chunk_idx=-1, + ) + + +_RULE_REGISTRY[quadrilinear_merge_p] = _quadrilinear_merge_rule + + +def _eval_eqn(eqn, env, W_levels): + invals = [] + for v in eqn.invars: + if isinstance(v, Literal): + invals.append(v.val) + else: + invals.append(env[v]) + has_jlp = any(_is_jlp(x) for x in invals) + if not has_jlp: + bind_params = eqn.primitive.get_bind_params(eqn.params) + outvals = eqn.primitive.bind(*invals, **bind_params) + if not eqn.primitive.multiple_results: + outvals = [outvals] + return outvals + rule = _RULE_REGISTRY.get(eqn.primitive) + if rule is None: + raise NotImplementedError( + f"custom_lap: no rule for primitive {eqn.primitive.name!r}. " + f"eqn = {eqn}. Register a rule in energy/custom_lap.py." + ) + outvals = rule(invals, eqn.params, W_levels) + if not eqn.primitive.multiple_results: + outvals = [outvals] + return outvals + + +def _trace_z(fn, z, N, W_levels): + + z = jnp.asarray(z) + assert z.shape == (N, 3), f"z must be [N={N}, 3]; got {z.shape}" + eye3 = jnp.eye(3, dtype=z.dtype) + z_jac = jnp.broadcast_to(eye3[:, None, :], (3, N, 3)) + z_lap = jnp.zeros((N, 3), dtype=z.dtype) + z_jlp = JLP( + value=z, + jac=z_jac, + lap=z_lap, + level=1, + has_chunk_axis=True, + chunk_idx=-1, + ) + closed = jax.make_jaxpr(fn)(z) + jaxpr = closed.jaxpr + consts = closed.consts + env: dict = {} + for cv, c in zip(jaxpr.constvars, consts): + env[cv] = c + env[jaxpr.invars[0]] = z_jlp + for eqn in jaxpr.eqns: + outvals = _eval_eqn(eqn, env, W_levels) + for ov, ov_val in zip(eqn.outvars, outvals): + env[ov] = ov_val + return [env[ov] for ov in jaxpr.outvars] + + +def _canonical_jac(jlp: JLP): + + if jlp.jac.shape[1] != 1: + jac = jnp.swapaxes(jlp.jac, 0, 1) + return jac.reshape((jlp.jac.shape[0] * jlp.jac.shape[1],) + jlp.jac.shape[2:]) + return jlp.jac[:, 0] + + +def custom_forward_laplacian_with_jac(fn: Callable, W_levels: list, N: int) -> Callable: + + def lap_jac_fn(z): + outs = _trace_z(fn, z, N, W_levels) + if len(outs) == 1: + out = outs[0] + if _is_jlp(out): + return out.value, _canonical_jac(out), out.lap + value = out + return ( + value, + jnp.zeros((3 * N,) + value.shape, dtype=value.dtype), + jnp.zeros_like(value), + ) + values = tuple(o.value if _is_jlp(o) else o for o in outs) + jacs = tuple( + _canonical_jac(o) + if _is_jlp(o) + else jnp.zeros((3 * N,) + o.shape, dtype=o.dtype) + for o in outs + ) + laps = tuple(o.lap if _is_jlp(o) else jnp.zeros_like(o) for o in outs) + return values, jacs, laps + + return lap_jac_fn + + +__all__ = [ + "JLP", + "build_W_levels", + "custom_forward_laplacian_with_jac", + "custom_lap_active", + "use_custom_lap", +] diff --git a/src/hamiltonzero/energy/frame.py b/src/hamiltonzero/energy/frame.py new file mode 100644 index 0000000000000000000000000000000000000000..dacd72b7bfdd1f8cfea2e7e3a9dd26e5084250c1 --- /dev/null +++ b/src/hamiltonzero/energy/frame.py @@ -0,0 +1,152 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +from collections.abc import Sequence + +import jax +import jax.numpy as jnp +import numpy as np +from jaxtyping import Array + +from hamiltonzero.compiled.types import EnergyFrame, EnergyInputs, EnergyMasks +from hamiltonzero.energy.custom_lap import build_W_levels + + +def _real_dtype(dtype): + return jnp.real(jnp.zeros((), dtype)).dtype + + +def _host_eigvalsh(x): + out_dtype = _real_dtype(x.dtype) + out_shape = jax.ShapeDtypeStruct(x.shape[:-1], out_dtype) + + def callback(a): + values = np.linalg.eigvalsh(np.asarray(a)) + return values.astype(np.dtype(out_dtype)) + + return jax.pure_callback( + callback, + out_shape, + x, + vmap_method="sequential", + ) + + +def _compute_custom_lap_views_batched(J_full, mask, mu, eps): + n = J_full.shape[1] + n_systems = J_full.shape[0] + dtype = J_full.dtype + J_matrix = jnp.transpose(J_full, (0, 1, 3, 2, 4)).reshape(n_systems, 3 * n, 3 * n) + J_matrix_symmetric = 0.5 * (J_matrix + jnp.swapaxes(J_matrix, -1, -2)) + mask_3 = jnp.repeat(mask.astype(dtype), 3, axis=-1) + eye_3n = jnp.eye(3 * n, dtype=dtype) + masked_identity = eye_3n[None] * mask_3[:, None, :] + J_eff = (J_matrix_symmetric - (mu + eps)[:, None, None] * masked_identity) / 4.0 + J_eff = 0.5 * (J_eff + jnp.swapaxes(J_eff, -1, -2)) + + eigh_epsilon = jnp.asarray(1e-6, dtype=_real_dtype(dtype)) + eigenvalues = _host_eigvalsh(J_eff + eigh_epsilon * eye_3n[None]) - eigh_epsilon + lam_max = jnp.max(eigenvalues, axis=-1) + delta_mu = jax.nn.relu(4.0 * (lam_max + eps)) + shift = delta_mu / 4.0 + J_eff = J_eff - shift[:, None, None] * eye_3n[None] + mu_eff = mu + delta_mu + casimir_per_site = (mu_eff + eps) * jnp.asarray(0.75, dtype=dtype) + radial_const = -casimir_per_site * mask.astype(dtype).sum(axis=-1) + return J_eff, radial_const + + +def _compute_custom_lap_views(J_full, mask, mu, eps): + values = _compute_custom_lap_views_batched( + J_full[None], + mask[None], + jnp.asarray(mu)[None], + jnp.asarray(eps)[None], + ) + return tuple(value[0] for value in values) + + +def build_energy_inputs(J_full, h, mask, mu, eps) -> EnergyInputs: + dtype = J_full.dtype + mu_array = jnp.asarray(mu, dtype=dtype) + eps_array = jnp.asarray(eps, dtype=dtype) + J_eff, radial_const = _compute_custom_lap_views( + J_full, + mask, + mu_array, + eps_array, + ) + return EnergyInputs( + custom_lap_J_eff=J_eff, + custom_lap_radial_const=radial_const, + one_body_fields=(h,), + ) + + +def block_permutation3(perm: Array) -> Array: + + components = jnp.arange(3, dtype=perm.dtype) + return (perm[:, None] * 3 + components[None, :]).reshape(-1) + + +def route_one_body_field(field: Array, perm: Array) -> Array: + + return jnp.take(field, perm, axis=0) + + +def route_one_body_fields( + fields: Sequence[Array], + perm: Array, +) -> tuple[Array, ...]: + + return tuple(route_one_body_field(field, perm) for field in fields) + + +def route_J_eff(J_eff: Array, perm: Array) -> Array: + + block_perm = block_permutation3(perm) + return jnp.take(jnp.take(J_eff, block_perm, axis=0), block_perm, axis=1) + + +def route_energy_inputs(energy_inputs: EnergyInputs, perm: Array) -> EnergyInputs: + return EnergyInputs( + custom_lap_J_eff=route_J_eff(energy_inputs.custom_lap_J_eff, perm), + custom_lap_radial_const=energy_inputs.custom_lap_radial_const, + one_body_fields=route_one_body_fields(energy_inputs.one_body_fields, perm), + ) + + +def compile_energy_frame( + energy_inputs: EnergyInputs, + real_mask: Array, + balanced_mask: Array, + perm: Array, +) -> EnergyFrame: + + J_eff = route_J_eff(energy_inputs.custom_lap_J_eff, perm) + n_sites = int(perm.shape[0]) + frame = EnergyFrame( + custom_lap_J_eff=J_eff, + w_levels=tuple(build_W_levels(J_eff, n_sites)), + custom_lap_radial_const=energy_inputs.custom_lap_radial_const, + one_body_fields=route_one_body_fields(energy_inputs.one_body_fields, perm), + masks=EnergyMasks( + real=route_one_body_field(real_mask, perm), + balanced=balanced_mask, + ), + ) + return frame + + +__all__ = [ + "block_permutation3", + "build_energy_inputs", + "compile_energy_frame", + "route_energy_inputs", + "route_J_eff", + "route_one_body_field", + "route_one_body_fields", +] diff --git a/src/hamiltonzero/energy/kernel.py b/src/hamiltonzero/energy/kernel.py new file mode 100644 index 0000000000000000000000000000000000000000..940c3b7791896fbf00e0188aeb23c5d2e074823e --- /dev/null +++ b/src/hamiltonzero/energy/kernel.py @@ -0,0 +1,196 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +from typing import Any, Callable + +import jax +import jax.numpy as jnp +from jaxtyping import Array, Complex, Float + +from hamiltonzero.compiled.execute import execute_wavefunction +from hamiltonzero.energy.custom_lap import ( + custom_forward_laplacian_with_jac, + use_custom_lap, +) + + +def _right_su2_chart_jet( + q: Float[Array, "N 4"], + z: Float[Array, "N 3"], +) -> Float[Array, "N 4"]: + + q0, q1, q2, q3 = q[:, 0], q[:, 1], q[:, 2], q[:, 3] + zx, zy, zz = z[:, 0], z[:, 1], z[:, 2] + u0 = 1.0 - 0.5 * (zx * zx + zy * zy + zz * zz) + u1, u2, u3 = zz, zy, zx + return jnp.stack( + [ + q0 * u0 - q1 * u1 - q2 * u2 - q3 * u3, + q0 * u1 + q1 * u0 + q2 * u3 - q3 * u2, + q0 * u2 - q1 * u3 + q2 * u0 + q3 * u1, + q0 * u3 + q1 * u2 - q2 * u1 + q3 * u0, + ], + axis=-1, + ) + + +def _custom_lap_single_finetune( + model: Callable, + q: Float[Array, "N 4"], + t, + energy_frame: Any, +): + if len(energy_frame.one_body_fields) != 1: + raise ValueError("custom-Laplacian energy requires exactly one one-body field") + return _custom_lap_single_from_frame( + model, + None, + q, + t, + J_eff=energy_frame.custom_lap_J_eff, + W_levels=energy_frame.w_levels, + field_xyz=energy_frame.one_body_fields[0], + radial_const=energy_frame.custom_lap_radial_const, + ) + + +def _custom_lap_single_prebuilt( + model: Callable, + q: Float[Array, "N 4"], + t, + frame: Any, +): + if len(frame.one_body_fields) != 1: + raise ValueError("custom-Laplacian energy requires exactly one one-body field") + return _custom_lap_single_from_frame( + model, + None, + q, + t, + J_eff=frame.custom_lap_J_eff, + W_levels=frame.w_levels, + field_xyz=frame.one_body_fields[0], + radial_const=frame.custom_lap_radial_const, + ) + + +def _custom_lap_single_from_frame( + model: Callable, + ctx: Any, + q: Float[Array, "N 4"], + t, + *, + J_eff, + W_levels, + field_xyz, + radial_const, +): + N = q.shape[0] + + def f_entry(z): + q_pert = _right_su2_chart_jet(q, z) + re, im = model(q_pert, ctx, t) + return jnp.stack([re, im]) + + with use_custom_lap(): + _value, jac_pair, lap_pair = custom_forward_laplacian_with_jac( + f_entry, + W_levels, + N, + )(jnp.zeros((N, 3), dtype=q.dtype)) + + tr_total = lap_pair[0] + 1j * lap_pair[1] + + g_lie = jac_pair[:, 0] + 1j * jac_pair[:, 1] + quad_total = jnp.einsum("a,ab,b->", g_lie, J_eff.astype(g_lie.dtype), g_lie) + + la_xyz = 0.5 * g_lie.reshape(N, 3) + field = (1j * jnp.einsum("ic,ic->", field_xyz.astype(g_lie.dtype), la_xyz)).astype( + g_lie.dtype + ) + + total = tr_total + quad_total + radial_const.astype(g_lie.dtype) + field + + zero = jnp.zeros_like(total) + exchange = total - field + return total, exchange, zero, field + + +def _vmc_energy_custom_lap_finetune( + model: Callable, + energy_frame: Any, + q: Float[Array, "... N 4"], + t=0.0, + *, + chunk_size: int | None = None, +) -> tuple[ + Complex[Array, "..."], + Complex[Array, "..."], + Complex[Array, "..."], + Complex[Array, "..."], +]: + + def single(qq, tt): + return _custom_lap_single_finetune(model, qq, tt, energy_frame) + + return _run_custom_lap_batch(single, q, t, chunk_size) + + +def _vmc_energy_custom_lap_prebuilt( + kernel: Any, + tree: Any, + energy_frame: Any, + q: Float[Array, "... N 4"], + *, + chunk_size: int | None = None, +) -> tuple[ + Complex[Array, "..."], + Complex[Array, "..."], + Complex[Array, "..."], + Complex[Array, "..."], +]: + def model(q_pert, _ctx, _t): + return execute_wavefunction(kernel, tree, q_pert) + + def single(qq, tt): + return _custom_lap_single_prebuilt(model, qq, tt, energy_frame) + + return _run_custom_lap_batch(single, q, 0.0, chunk_size) + + +def _run_custom_lap_batch(single, q, t, chunk_size): + n_sites, n_dims = q.shape[-2], q.shape[-1] + assert n_dims == 4, f"expected quaternion last dim 4, got {n_dims}" + lead = q.shape[:-2] + n_items = 1 + for d in lead: + n_items *= d + + q_flat = q.reshape(n_items, n_sites, 4) + t_arr = jnp.asarray(t, dtype=q.dtype) + t_bcast = jnp.broadcast_to(t_arr, lead if lead else ()) + t_flat = t_bcast.reshape(n_items) if lead else jnp.broadcast_to(t_arr, (n_items,)) + + with jax.default_matmul_precision("highest"): + if chunk_size is None or chunk_size >= n_items: + total, exchange, casimir, field = jax.vmap(single)(q_flat, t_flat) + else: + total, exchange, casimir, field = jax.lax.map( + lambda x: single(x[0], x[1]), + (q_flat, t_flat), + batch_size=chunk_size, + ) + + out_shape = lead if lead else () + return ( + total.reshape(out_shape), + exchange.reshape(out_shape), + casimir.reshape(out_shape), + field.reshape(out_shape), + ) + + +__all__ = [] diff --git a/src/hamiltonzero/evaluation/__init__.py b/src/hamiltonzero/evaluation/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..b38d393f78ec7313dfc0fee0fdccfd7ab8cfad14 --- /dev/null +++ b/src/hamiltonzero/evaluation/__init__.py @@ -0,0 +1,42 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .backend import ( + BeamCandidates, + CanonicalContext, + EvalBackend, + LargeNCompilation, + MCMCPopulation, +) +from .runner import evaluate +from .runtime import DefaultEvalBackend, build_eval_backend +from .statistics import ( + ChannelMetrics, + EnergyWindow, + P01_TWO_SIDED, + select_winner_per_physical, + standard_error_from_tailstd, + welch_band_walkermean, +) +from .types import ContestCandidate, ContestResult, EvalMetric, EvalResult + +__all__ = [ + "BeamCandidates", + "CanonicalContext", + "ChannelMetrics", + "ContestCandidate", + "ContestResult", + "DefaultEvalBackend", + "EnergyWindow", + "EvalBackend", + "EvalMetric", + "EvalResult", + "LargeNCompilation", + "MCMCPopulation", + "P01_TWO_SIDED", + "evaluate", + "build_eval_backend", + "select_winner_per_physical", + "standard_error_from_tailstd", + "welch_band_walkermean", +] diff --git a/src/hamiltonzero/evaluation/backend.py b/src/hamiltonzero/evaluation/backend.py new file mode 100644 index 0000000000000000000000000000000000000000..9547283fcc8de5348c893356b512f18bf83036d6 --- /dev/null +++ b/src/hamiltonzero/evaluation/backend.py @@ -0,0 +1,148 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Protocol + +from hamiltonzero.config import EnergyConfig, EvalMCMCConfig, ModelConfig + + +@dataclass(frozen=True, slots=True) +class CanonicalContext: + context: Any + old_inverse: Any + + +@dataclass(frozen=True, slots=True) +class BeamCandidates: + permutations: Any + log_probabilities: Any + + +@dataclass(frozen=True, slots=True) +class LargeNCompilation: + wavefunction: Any + permutation: Any + log_probability: Any + + +@dataclass(frozen=True, slots=True) +class MCMCPopulation: + q: Any + sigma: Any + beta: Any + + +class EvalBackend(Protocol): + def load_system(self, path: Path, energy: EnergyConfig) -> Any: ... + + def load_model( + self, + checkpoint: Path, + config: ModelConfig, + key: Any, + context: Any, + *, + contextualizer_attention: str | None, + ) -> Any: ... + + def canonicalize_context(self, context: Any) -> CanonicalContext: ... + + def embedded_route(self, model: Any) -> Any | None: ... + + def route_context( + self, + context: Any, + permutation: Any, + *, + compact_custom_lap: bool, + ) -> Any: ... + + def release_context(self, context: Any) -> None: ... + + def virtual_context(self, context: Any, permutations: Any) -> Any: ... + + def beam_candidates( + self, + model: Any, + context: Any, + *, + beam_width: int, + top_k: int, + temperature: float, + ) -> BeamCandidates: ... + + def compile_single(self, model: Any, routed_context: Any) -> Any: ... + + def compile_embedded(self, model: Any) -> Any: ... + + def compile_candidates( + self, + model: Any, + canonical_context: Any, + permutations: Any, + ) -> Any: ... + + def select_candidate(self, wavefunctions: Any, winner: int) -> Any: ... + + def compile_large_n( + self, + model: Any, + canonical_context: Any, + *, + sequence_shards: int, + pair_tile_size: int, + temperature: float, + ) -> LargeNCompilation: ... + + def prepare_singular( + self, + model: Any, + context: Any, + state: Any, + ) -> tuple[Any, Any, Any]: ... + + def initialize_mcmc( + self, + key: Any, + model: Any, + context: Any, + config: EvalMCMCConfig, + ) -> Any: ... + + def step_mcmc( + self, + state: Any, + model: Any, + context: Any, + *, + replica_steps: int, + walker_chunk_size: int, + ) -> Any: ... + + def adapt_mcmc(self, state: Any, config: EvalMCMCConfig) -> Any: ... + + def route_mcmc(self, state: Any, permutation: Any) -> Any: ... + + def mcmc_population(self, state: Any) -> MCMCPopulation: ... + + def replace_mcmc_population( + self, + state: Any, + population: MCMCPopulation, + ) -> Any: ... + + def cold_walkers(self, state: Any) -> Any: ... + + def custom_lap_energy( + self, + model: Any, + context: Any, + q: Any, + config: EnergyConfig, + ) -> tuple[Any, Any, Any, Any]: ... + + def block_until_ready(self, value: Any) -> None: ... diff --git a/src/hamiltonzero/evaluation/greedy_router.py b/src/hamiltonzero/evaluation/greedy_router.py new file mode 100644 index 0000000000000000000000000000000000000000..3cea9f32819d84455479a68cc888abaf603eb81f --- /dev/null +++ b/src/hamiltonzero/evaluation/greedy_router.py @@ -0,0 +1,125 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +from collections.abc import Callable +from functools import partial + +import jax +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +from hamiltonzero.model.route_pointer import TreePrefixPointerMHSEA + + +def _ring_permute_rows_local(values, route_ids, *, axis_size: int): + + local_rows = values.shape[0] + lane = jax.lax.axis_index("seq").astype(jax.numpy.int32) + local_ids = jax.lax.dynamic_slice_in_dim( + route_ids, + lane * local_rows, + local_rows, + axis=0, + ) + owner = local_ids // local_rows + within_owner = local_ids % local_rows + output0 = jax.numpy.zeros_like(values) + + def take_from(panel, origin, output): + selected = panel[within_owner] + mask = owner == origin + while mask.ndim < selected.ndim: + mask = mask[..., None] + return jax.numpy.where(mask, selected, output) + + origin0 = lane + output0 = take_from(values, origin0, output0) + permutation = [(i, (i + 1) % axis_size) for i in range(axis_size)] + + def step(carry, _): + panel, origin, output = carry + panel = jax.lax.ppermute(panel, "seq", permutation) + origin = (origin - jax.numpy.asarray(1, jax.numpy.int32)) % axis_size + output = take_from(panel, origin, output) + return (panel, origin, output), None + + (_, _, output), _ = jax.lax.scan( + step, + (values, origin0, output0), + xs=None, + length=axis_size - 1, + ) + return output + + +def _make_row_permute(mesh: Mesh) -> Callable: + axis_size = int(mesh.shape["seq"]) + spec = P("seq", None, None) + if axis_size == 1: + return lambda values, route_ids: values[route_ids] + return jax.shard_map( + partial(_ring_permute_rows_local, axis_size=axis_size), + mesh=mesh, + in_specs=(spec, P()), + out_specs=spec, + check_vma=False, + ) + + +def build_compact_greedy_router( + *, + mesh: Mesh, + decoder_template: TreePrefixPointerMHSEA, + pair_tile_size: int = 128, +) -> Callable: + + if tuple(mesh.axis_names) != ("seq",): + raise ValueError("compact greedy router requires a one-dimensional 'seq' mesh") + if int(mesh.shape["seq"]) < 1: + raise ValueError("compact greedy router requires at least one seq lane") + if not isinstance(decoder_template, TreePrefixPointerMHSEA): + raise TypeError("decoder_template must be TreePrefixPointerMHSEA") + if isinstance(pair_tile_size, bool) or int(pair_tile_size) < 1: + raise ValueError("pair_tile_size must be a positive integer") + + rep = NamedSharding(mesh, P()) + seq_vec = NamedSharding(mesh, P("seq", None)) + seq_edge = NamedSharding(mesh, P("seq", None, None)) + decoder_rep = jax.tree_util.tree_map(lambda _leaf: rep, decoder_template) + row_permute = _make_row_permute(mesh) + + def decode(decoder, h, edge, mask, global_feat, tau, real_mask): + + h = jax.lax.with_sharding_constraint(h, seq_vec) + edge = jax.lax.with_sharding_constraint(edge, seq_edge) + perm, logp = decoder._decode_greedy_compact( + h, + edge, + mask, + global_feat=global_feat, + tau=tau, + real_mask=real_mask, + sequence_mesh=mesh, + pair_tile_size=int(pair_tile_size), + row_permute_fn=row_permute, + ) + return perm, logp + + return jax.jit( + decode, + in_shardings=( + decoder_rep, + seq_vec, + seq_edge, + rep, + rep, + rep, + rep, + ), + out_shardings=(rep, rep), + ) + + +__all__ = ["build_compact_greedy_router"] diff --git a/src/hamiltonzero/evaluation/large_n.py b/src/hamiltonzero/evaluation/large_n.py new file mode 100644 index 0000000000000000000000000000000000000000..ad98b84f37303e689e2cd2c96d6e62ca6dc4e5d5 --- /dev/null +++ b/src/hamiltonzero/evaluation/large_n.py @@ -0,0 +1,371 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +import gc +from typing import NamedTuple + +import jax +import jax.numpy as jnp +import numpy as np +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +from hamiltonzero.compiled.tree import bind_physical_compiler_kernel +from hamiltonzero.compiled.trunk import ( + bind_shared_kernel, + bind_trunk_compiler_kernel, +) +from hamiltonzero.compiled.types import CompiledWaveFunction +from hamiltonzero.model.route_pointer import TreePrefixPointerMHSEA +from .sequence_trunk import ( + build_sequence_pair_permute, + build_sequence_parallel_contextualizer, + build_sequence_parallel_edge_global_update, + build_sequence_parallel_physical_leaf, + build_sequence_parallel_physical_reducer, + build_sequence_parallel_shared_trunk, +) +from .greedy_router import build_compact_greedy_router + + +class LargeNCompiledEvalResult(NamedTuple): + wavefunction: CompiledWaveFunction + perm: jax.Array + logp: jax.Array + + +def _sequence_mesh(n: int, requested_shards: int) -> Mesh: + devices = tuple(jax.devices()) + shards = len(devices) if requested_shards == 0 else requested_shards + if shards < 1 or shards > len(devices): + raise ValueError( + f"compiled eval requested {shards} seq shards, but JAX exposes " + f"{len(devices)} devices" + ) + if n % shards: + raise ValueError(f"N={n} must be divisible by seq_shards={shards}") + local_rows = n // shards + if n & (n - 1) or local_rows & (local_rows - 1): + raise ValueError( + "large-N physical compilation requires a power-of-two padded " + f"width and power-of-two rows per lane; got N={n}, " + f"local_rows={local_rows}" + ) + return Mesh(np.asarray(devices[:shards], dtype=object), ("seq",)) + + +def _replicate(tree, sharding: NamedSharding): + return jax.device_put(tree, jax.tree_util.tree_map(lambda _leaf: sharding, tree)) + + +class _LargeNXlaTreePrefixPointer(TreePrefixPointerMHSEA): + def _resolve_heavy_attn_impl(self, n: int): + del n + return None + + def _resolve_tree_attn_impl(self): + return None + + def _route_attention( + self, + q, + k, + v, + edge_bias, + key_mask, + *, + impl, + key_mask_only=False, + attention_mask=None, + sequence_axis_name=None, + sequence_mesh=None, + ): + del impl, key_mask_only + dtype = q.dtype + valid = key_mask.astype(bool)[:, None, :] + if attention_mask is not None: + valid = valid & attention_mask.astype(bool) + has_key = jnp.any(valid, axis=-1) + q_c = q.astype(jnp.float32) + k_c = k.astype(jnp.float32) + v_c = v.astype(jnp.float32) + bias_c = edge_bias.astype(jnp.float32) + if sequence_axis_name is not None: + + def sharding(*axes): + spec = P(*axes) + return ( + NamedSharding(sequence_mesh, spec) + if sequence_mesh is not None + else spec + ) + + q_c = jax.lax.with_sharding_constraint( + q_c, sharding(None, sequence_axis_name, None, None) + ) + k_c = jax.lax.with_sharding_constraint( + k_c, sharding(None, None, None, None) + ) + v_c = jax.lax.with_sharding_constraint( + v_c, sharding(None, None, None, None) + ) + bias_c = jax.lax.with_sharding_constraint( + bias_c, sharding(None, sequence_axis_name, None, None) + ) + logits = jnp.einsum("bihd,bjhd->bhij", q_c, k_c) + logits = logits / jnp.sqrt(jnp.asarray(self.d_head, dtype=jnp.float32)) + logits = logits + jnp.transpose(bias_c, (0, 3, 1, 2)) + if sequence_axis_name is not None: + logits = jax.lax.with_sharding_constraint( + logits, sharding(None, None, sequence_axis_name, None) + ) + logits = jnp.where( + valid[:, None, :, :], + logits, + jnp.asarray(-1.0e30, dtype=jnp.float32), + ) + if sequence_axis_name is not None: + logits = jax.lax.with_sharding_constraint( + logits, sharding(None, None, sequence_axis_name, None) + ) + alpha = jax.nn.softmax(logits, axis=-1) + if sequence_axis_name is not None: + alpha = jax.lax.with_sharding_constraint( + alpha, sharding(None, None, sequence_axis_name, None) + ) + out = jnp.einsum("bhij,bjhd->bihd", alpha, v_c) + if sequence_axis_name is not None: + out = jax.lax.with_sharding_constraint( + out, sharding(None, sequence_axis_name, None, None) + ) + out = out.astype(dtype) + return jnp.where(has_key[..., None, None], out, jnp.zeros_like(out)) + + +def _large_n_xla_decoder_view(decoder: TreePrefixPointerMHSEA): + compiled = object.__new__(_LargeNXlaTreePrefixPointer) + compiled.__dict__.update(decoder.__dict__) + return compiled + + +def _validate_large_n_model(model) -> None: + decoder = getattr(model, "route_decoder", None) + failures: list[str] = [] + if not isinstance(decoder, TreePrefixPointerMHSEA): + failures.append("route decoder must be TreePrefixPointerMHSEA") + if getattr(model, "route_contextualizer", None) is None: + failures.append("route contextualizer must be enabled") + if getattr(model, "gladder_fork_route", None) is None: + failures.append("route global fork must be enabled") + if getattr(model, "readout_leaf_context", None) is None: + failures.append("physical contextualizer must be enabled") + if failures: + raise ValueError( + "unsupported large-N eval compiler model: " + "; ".join(failures) + ) + + +def _validate_context(ctx) -> int: + jdp = jnp.asarray(ctx.J_double_prime) + mask = jnp.asarray(ctx.mask) + bmask = jnp.asarray(ctx.bmask) + h_prime = jnp.asarray(ctx.h_prime) + if jdp.ndim != 3 or jdp.shape[-1] != 10: + raise ValueError("ctx.J_double_prime must have shape [N,N,10]") + n = int(jdp.shape[0]) + if jdp.shape[1] != n: + raise ValueError("ctx.J_double_prime pair axes must be square") + if mask.shape != (n,) or bmask.shape != (n,): + raise ValueError("ctx.mask and ctx.bmask must both have shape [N]") + if h_prime.shape != (n, 3): + raise ValueError("ctx.h_prime must have shape [N,3]") + if jdp.dtype != jnp.float32 or h_prime.dtype != jnp.float32: + raise TypeError( + "large-N eval keeps streamed pair/frontier arithmetic in fp32; " + f"got J={jdp.dtype}, h={h_prime.dtype}" + ) + return n + + +def compile_large_n_eval_wavefunction( + model, + ctx, + *, + seq_shards: int = 0, + pair_tile_size: int = 128, + tau: float = 1.0, +) -> LargeNCompiledEvalResult: + + _validate_large_n_model(model) + n = _validate_context(ctx) + if pair_tile_size < 1: + raise ValueError("pair_tile_size must be positive") + mesh = _sequence_mesh(n, int(seq_shards)) + rep = NamedSharding(mesh, P()) + seq_edge = NamedSharding(mesh, P("seq", None, None)) + + trunk_kernel = _replicate(bind_trunk_compiler_kernel(model), rep) + route_contextualizer = _replicate(model.route_contextualizer, rep) + route_global_fork = _replicate(model.gladder_fork_route, rep) + decoder = _replicate(_large_n_xla_decoder_view(model.route_decoder), rep) + physical_kernel = _replicate(bind_physical_compiler_kernel(model), rep) + + jdp = jax.device_put(jnp.asarray(ctx.J_double_prime), seq_edge) + h_prime, real_mask, structural_mask = jax.device_put( + ( + jnp.asarray(ctx.h_prime), + jnp.asarray(ctx.mask), + jnp.asarray(ctx.bmask), + ), + rep, + ) + + trunk_entry = build_sequence_parallel_shared_trunk( + mesh=mesh, + kernel_template=trunk_kernel, + featurizer_tile_size=int(pair_tile_size), + ) + trunk = trunk_entry(trunk_kernel, jdp, h_prime, real_mask, structural_mask) + + route_context_entry = build_sequence_parallel_contextualizer( + mesh=mesh, + contextualizer_template=route_contextualizer, + g_template=trunk.global_stream, + tile_size=int(pair_tile_size), + ) + route_node, route_edge, route_g = route_context_entry( + route_contextualizer, + trunk.node_raw, + trunk.edge_raw, + real_mask, + structural_mask, + trunk.global_stream, + ) + route_global_entry = build_sequence_parallel_edge_global_update( + mesh=mesh, + module_template=route_global_fork, + tile_size=int(pair_tile_size), + ) + route_global = route_global_entry( + route_global_fork, route_g, route_edge, structural_mask + ) + + route_entry = build_compact_greedy_router( + mesh=mesh, + decoder_template=decoder, + pair_tile_size=int(pair_tile_size), + ) + perm, logp = route_entry( + decoder, + route_node, + route_edge, + structural_mask, + route_global, + jnp.asarray(tau, dtype=jnp.float32), + real_mask, + ) + + jax.block_until_ready((perm, logp)) + del route_node, route_edge, route_g, route_global + + pair_permute = build_sequence_pair_permute(mesh=mesh) + routed_node, routed_edge = pair_permute(trunk.node_raw, trunk.edge_raw, perm) + leaf_real = real_mask[perm] + global_stream = trunk.global_stream + + jax.block_until_ready((routed_node, routed_edge, leaf_real, global_stream)) + del ( + trunk, + trunk_kernel, + route_contextualizer, + route_global_fork, + decoder, + jdp, + h_prime, + real_mask, + trunk_entry, + route_context_entry, + route_global_entry, + route_entry, + pair_permute, + ) + gc.collect() + + physical_context_entry = build_sequence_parallel_contextualizer( + mesh=mesh, + contextualizer_template=physical_kernel.contextualizer, + g_template=global_stream, + tile_size=int(pair_tile_size), + ) + physical_node, physical_edge, physical_context_g = physical_context_entry( + physical_kernel.contextualizer, + routed_node, + routed_edge, + leaf_real, + structural_mask, + global_stream, + ) + jax.block_until_ready((physical_node, physical_edge, physical_context_g)) + del routed_node, routed_edge, global_stream, physical_context_entry + gc.collect() + + physical_global_entry = build_sequence_parallel_edge_global_update( + mesh=mesh, + module_template=physical_kernel.global_fork, + tile_size=int(pair_tile_size), + ) + physical_global = physical_global_entry( + physical_kernel.global_fork, + physical_context_g, + physical_edge, + structural_mask, + ) + jax.block_until_ready(physical_global) + del physical_context_g, physical_global_entry + gc.collect() + + physical_leaf_entry = build_sequence_parallel_physical_leaf( + mesh=mesh, + kernel_template=physical_kernel, + ) + leaf_h, c_rows = physical_leaf_entry( + physical_kernel, + physical_node, + physical_global, + ) + jax.block_until_ready((leaf_h, c_rows)) + del physical_node, physical_leaf_entry + gc.collect() + + physical_reducer_entry = build_sequence_parallel_physical_reducer( + mesh=mesh, + kernel_template=physical_kernel, + edge_template=physical_edge, + replicate_threshold=min(512, n), + contextualizer_tile_size=int(pair_tile_size), + ) + tree = physical_reducer_entry( + physical_kernel, + physical_edge, + leaf_h, + c_rows, + leaf_real, + structural_mask, + physical_global, + perm, + ) + jax.block_until_ready(tree) + + return LargeNCompiledEvalResult( + wavefunction=CompiledWaveFunction(bind_shared_kernel(model), tree), + perm=perm, + logp=logp, + ) + + +__all__ = [ + "LargeNCompiledEvalResult", + "compile_large_n_eval_wavefunction", +] diff --git a/src/hamiltonzero/evaluation/pallas_mha.py b/src/hamiltonzero/evaluation/pallas_mha.py new file mode 100644 index 0000000000000000000000000000000000000000..30f0f691620610014a2ad36600793364c360ae1d --- /dev/null +++ b/src/hamiltonzero/evaluation/pallas_mha.py @@ -0,0 +1,164 @@ +# SPDX-License-Identifier: Apache-2.0 + +# Copyright 2023 The JAX Authors. +# Modifications copyright (c) 2026 Simulacra Research Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +from __future__ import annotations + +import dataclasses +import functools +import math +from typing import Any + +import jax +import jax.numpy as jnp +from jax import lax +from jax.experimental import pallas as pl +from jax.experimental.pallas import triton as plgpu + + +@dataclasses.dataclass(frozen=True, slots=True) +class BlockSizes: + block_q: int + block_k: int + + @classmethod + def get_default(cls): + return cls(block_q=128, block_k=128) + + +def noncausal_bias_mha_forward_kernel( + q_ref, + k_ref, + v_ref, + bias_ref, + o_ref: Any, + *, + block_q: int, + block_k: int, + head_dim: int, +): + seq_len = k_ref.shape[0] + start_q = pl.program_id(0) + head_dim_padded = q_ref.shape[-1] + m_i = jnp.zeros(block_q, dtype=jnp.float32) - float("inf") + l_i = jnp.zeros(block_q, dtype=jnp.float32) + o = jnp.zeros((block_q, head_dim_padded), dtype=jnp.float32) + curr_q_slice = pl.dslice(start_q * block_q, block_q) + head_mask = (jnp.arange(head_dim_padded) < head_dim)[None, :] + q = plgpu.load(q_ref, mask=head_mask, other=0.0) + + def body(start_k, carry): + o_prev, m_prev, l_prev = carry + curr_k_slice = pl.dslice(start_k * block_k, block_k) + k = plgpu.load(k_ref.at[curr_k_slice, :], mask=head_mask, other=0.0) + qk = pl.dot(q, k.T) + bias = plgpu.load(bias_ref.at[curr_q_slice, curr_k_slice]) + qk += bias + qk *= math.log2(math.e) + qk = qk.astype(q_ref.dtype) + m_curr = jnp.max(qk, axis=-1) + m_next = jnp.maximum(m_prev, m_curr) + correction = jnp.exp2(m_prev - m_next) + l_prev_corr = correction * l_prev + s_curr = jnp.exp2(qk - m_next[:, None]) + l_curr = s_curr.sum(axis=-1) + l_next = l_prev_corr + l_curr + o_prev_corr = correction[:, None] * o_prev + v = plgpu.load(v_ref.at[curr_k_slice, :], mask=head_mask) + o_curr = pl.dot(s_curr.astype(v.dtype), v) + o_next = o_prev_corr + o_curr + return o_next, m_next, l_next + + upper_bound = pl.cdiv(seq_len, block_k) + o, _m_i, l_i = lax.fori_loop(0, upper_bound, body, (o, m_i, l_i)) + o /= l_i[:, None] + plgpu.store(o_ref.at[:, : o.shape[-1]], o.astype(o_ref.dtype), mask=head_mask) + + +@functools.partial(jax.jit, static_argnames=["block_sizes"]) +def noncausal_bias_mha( + q, + k, + v, + bias, + *, + block_sizes: BlockSizes = BlockSizes.get_default(), +): + batch_size, q_seq_len, num_heads, head_dim = q.shape + kv_seq_len = k.shape[1] + block_q = min(block_sizes.block_q, q_seq_len) + block_k = min(block_sizes.block_k, kv_seq_len) + head_dim_padded = pl.next_power_of_2(head_dim) + if (q.shape[-1] != k.shape[-1]) or (q.shape[-1] != v.shape[-1]): + raise ValueError( + "This kernel expects q, k, and v to have the same head dimension, " + f"but found {q.shape=}, {k.shape=}, {v.shape=}." + ) + if bias.shape != (batch_size, q_seq_len, kv_seq_len, num_heads): + raise ValueError( + f"bias must have shape [batch, query, key, heads]; got {bias.shape}" + ) + if q_seq_len % block_q != 0: + raise ValueError(f"{q_seq_len=} must be a multiple of {block_q=}") + if kv_seq_len % block_k != 0: + raise ValueError(f"{kv_seq_len=} must be a multiple of {block_k=}") + grid = (pl.cdiv(q_seq_len, block_q), batch_size, num_heads) + num_warps = 4 if head_dim <= 64 else 8 + kernel = functools.partial( + noncausal_bias_mha_forward_kernel, + block_q=block_q, + block_k=block_k, + head_dim=head_dim, + ) + in_specs = [ + pl.BlockSpec( + (None, block_q, None, head_dim_padded), + lambda i, j, k: (j, i, k, 0), + ), + pl.BlockSpec( + (None, kv_seq_len, None, head_dim_padded), + lambda i, j, k: (j, 0, k, 0), + ), + pl.BlockSpec( + (None, kv_seq_len, None, head_dim_padded), + lambda i, j, k: (j, 0, k, 0), + ), + pl.BlockSpec( + (None, block_q, kv_seq_len, None), + lambda i, j, k: (j, i, 0, k), + ), + ] + out_shape = [q] + out_specs = [ + pl.BlockSpec( + (None, block_q, None, head_dim_padded), + lambda i, j, k: (j, i, k, 0), + ) + ] + out = pl.pallas_call( + kernel, + grid=grid, + in_specs=in_specs, + out_specs=out_specs, + compiler_params=plgpu.CompilerParams(num_warps=num_warps, num_stages=2), + out_shape=out_shape, + name="mha_forward", + )(q, k, v, bias) + return out[0] + + +__all__ = ["BlockSizes", "noncausal_bias_mha"] diff --git a/src/hamiltonzero/evaluation/runner.py b/src/hamiltonzero/evaluation/runner.py new file mode 100644 index 0000000000000000000000000000000000000000..3f059ee72add1177a788e5465add5b76cf075232 --- /dev/null +++ b/src/hamiltonzero/evaluation/runner.py @@ -0,0 +1,509 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import time +from dataclasses import replace +from typing import Any, Callable + +import jax +import jax.numpy as jnp +import numpy as np + +from hamiltonzero.config import EvalConfig + +from .backend import EvalBackend, MCMCPopulation +from .statistics import EnergyWindow, P01_TWO_SIDED, select_winner_per_physical +from .types import ContestCandidate, ContestResult, EvalMetric, EvalResult + + +def _permute_q_prefix(q, permutations): + index = jnp.broadcast_to(permutations[:, None, None, :, None], q.shape) + return jnp.take_along_axis(q, index, axis=3) + + +def _collapse_to_winner(q_virtual, permutations, winner_idx, *, P, K): + winner_idx = jnp.asarray(winner_idx, dtype=jnp.int32) + inverse = jnp.argsort(permutations, axis=-1) + q_canonical = _permute_q_prefix(q_virtual, inverse) + batch_per_candidate = q_canonical.shape[1] + tail = q_canonical.shape[2:] + q_physical = q_canonical.reshape((P, K, batch_per_candidate) + tail).reshape( + (P, K * batch_per_candidate) + tail + ) + permutations_pk = permutations.reshape((P, K, -1)) + route = jnp.take_along_axis( + permutations_pk, + winner_idx[:, None, None], + axis=1, + )[:, 0] + return _permute_q_prefix(q_physical, route), route + + +def _gather_winner_ladder(values, winner_idx, *, P, K): + winner_idx = jnp.asarray(winner_idx, dtype=jnp.int32) + values_pk = values.reshape((P, K) + values.shape[1:]) + index = winner_idx[(slice(None), None) + (None,) * (values_pk.ndim - 2)] + return jnp.take_along_axis(values_pk, index, axis=1)[:, 0] + + +def _as_batched_route(permutation): + value = jnp.asarray(permutation, dtype=jnp.int32) + if value.ndim == 1: + value = value[None, :] + if value.ndim != 2 or value.shape[0] != 1: + raise ValueError("single-system evaluation requires a route with shape [1, N]") + return value + + +def _compose_walker_route(old_inverse, route): + return jnp.take_along_axis( + jnp.asarray(old_inverse, dtype=jnp.int32), + route, + axis=-1, + ) + + +def _adapt(backend: EvalBackend, state: Any, config: EvalConfig): + return backend.adapt_mcmc(state, config.mcmc) + + +def _burn_in( + backend: EvalBackend, + state: Any, + model: Any, + context: Any, + config: EvalConfig, + *, + iterations: int, + replica_steps: int, +): + for _ in range(int(iterations)): + state = backend.step_mcmc( + state, + model, + context, + replica_steps=int(replica_steps), + walker_chunk_size=int(config.mcmc.walker_chunk_size), + ) + backend.block_until_ready(backend.cold_walkers(state)) + state = _adapt(backend, state, config) + return state + + +def _measure( + backend: EvalBackend, + state: Any, + model: Any, + context: Any, + config: EvalConfig, + *, + started: float, + metric_sink: Callable[[EvalMetric], None] | None, +): + window = EnergyWindow( + config.measurements, + systems=1, + batch_size=config.mcmc.batch_size, + ) + for step in range(config.measurements): + step_started = time.perf_counter() + state = backend.step_mcmc( + state, + model, + context, + replica_steps=int(config.mcmc.steps), + walker_chunk_size=int(config.mcmc.walker_chunk_size), + ) + q_cold = backend.cold_walkers(state) + total, exchange, _casimir, field = backend.custom_lap_energy( + model, + context, + q_cold, + config.energy, + ) + backend.block_until_ready(total) + window.push(total, exchange, field) + state = _adapt(backend, state, config) + if metric_sink is not None: + energy = np.asarray(total).real + metric_sink( + EvalMetric( + step=step, + energy=float(np.mean(energy)), + energy_std=float(np.std(energy)), + step_walltime=float(time.perf_counter() - step_started), + walltime=float(time.perf_counter() - started), + ) + ) + return state, window + + +def _ordinary( + backend: EvalBackend, + model: Any, + context: Any, + canonical, + mcmc_key, + config: EvalConfig, +): + state = backend.initialize_mcmc(mcmc_key, model, context, config.mcmc) + candidates = backend.beam_candidates( + model, + canonical.context, + beam_width=int(config.contest_beam_width), + top_k=1, + temperature=float(config.route_temperature), + ) + permutations = jnp.asarray(candidates.permutations, dtype=jnp.int32) + if permutations.ndim != 3 or permutations.shape[:2] != (1, 1): + raise ValueError("ordinary eval router must return shape [1, 1, N]") + route = permutations[:, 0] + walker_route = _compose_walker_route(canonical.old_inverse, route) + routed_context = backend.route_context( + canonical.context, + route, + compact_custom_lap=False, + ) + state = backend.route_mcmc(state, walker_route) + wavefunction = backend.compile_single(model, routed_context) + backend.block_until_ready(wavefunction) + logp = float(np.asarray(candidates.log_probabilities)[0, 0]) + return wavefunction, routed_context, state, route, logp, None + + +def _compiled_finetune_ordinary( + backend: EvalBackend, + model: Any, + context: Any, + canonical, + embedded_route, + mcmc_key, + config: EvalConfig, +): + state = backend.initialize_mcmc(mcmc_key, model, context, config.mcmc) + route = _as_batched_route(embedded_route) + if route.shape[-1] != canonical.context.mask.shape[-1]: + raise ValueError( + "compiled fine-tune route width does not match the evaluation system" + ) + walker_route = _compose_walker_route(canonical.old_inverse, route) + routed_context = backend.route_context( + canonical.context, + route, + compact_custom_lap=False, + ) + state = backend.route_mcmc(state, walker_route) + wavefunction = backend.compile_embedded(model) + backend.block_until_ready(wavefunction) + return wavefunction, routed_context, state, route, None, None + + +def _contest( + backend: EvalBackend, + model: Any, + canonical, + root_key, + config: EvalConfig, +): + K = int(config.contest_candidates) + batch_per_candidate = int(config.mcmc.batch_size) // K + candidates = backend.beam_candidates( + model, + canonical.context, + beam_width=int(config.contest_beam_width), + top_k=K, + temperature=float(config.route_temperature), + ) + beam_permutations = jnp.asarray(candidates.permutations, dtype=jnp.int32) + if beam_permutations.ndim != 3 or beam_permutations.shape[:2] != (1, K): + raise ValueError(f"contest router must return shape [1, {K}, N]") + n_sites = int(beam_permutations.shape[-1]) + permutations = beam_permutations.reshape((K, n_sites)) + virtual_context = backend.virtual_context(canonical.context, permutations) + race_mcmc = replace(config.mcmc, batch_size=batch_per_candidate) + state = backend.initialize_mcmc( + jax.random.fold_in(root_key, 7411), + model, + virtual_context, + race_mcmc, + ) + wavefunctions = backend.compile_candidates( + model, + canonical.context, + permutations, + ) + backend.block_until_ready(wavefunctions) + for _ in range(int(config.contest_preburn)): + state = backend.step_mcmc( + state, + wavefunctions, + virtual_context, + replica_steps=int(config.mcmc.burn_in_replica_steps), + walker_chunk_size=int(config.mcmc.walker_chunk_size), + ) + backend.block_until_ready(backend.cold_walkers(state)) + state = backend.adapt_mcmc(state, race_mcmc) + race_window = EnergyWindow( + config.contest_measurements, + systems=K, + batch_size=batch_per_candidate, + ) + for _ in range(int(config.contest_measurements)): + state = backend.step_mcmc( + state, + wavefunctions, + virtual_context, + replica_steps=int(config.mcmc.steps), + walker_chunk_size=int(config.mcmc.walker_chunk_size), + ) + state = backend.adapt_mcmc(state, race_mcmc) + q_cold = backend.cold_walkers(state) + total, exchange, _casimir, field = backend.custom_lap_energy( + wavefunctions, + virtual_context, + q_cold, + config.energy, + ) + backend.block_until_ready(total) + race_window.push(total, exchange, field) + energies = np.asarray( + [[race_window.tail_mean("total", candidate) for candidate in range(K)]] + ) + tailstd = np.asarray( + [[race_window.tail_std("total", candidate) for candidate in range(K)]] + ) + beam_logp = np.asarray(candidates.log_probabilities, dtype=float) + winners, ties, reasons, _bands, standard_errors = select_winner_per_physical( + energies, + tailstd, + beam_logp, + batch_per_candidate, + z=P01_TWO_SIDED, + ucb_z=float(config.contest_se_multiplier), + ) + winner = int(winners[0]) + wavefunction = backend.select_candidate(wavefunctions, winner) + population = backend.mcmc_population(state) + q_final, route = _collapse_to_winner( + population.q, + permutations, + winners, + P=1, + K=K, + ) + sigma = _gather_winner_ladder(population.sigma, winners, P=1, K=K) + beta = _gather_winner_ladder(population.beta, winners, P=1, K=K) + routed_context = backend.route_context( + canonical.context, + route, + compact_custom_lap=False, + ) + final_state = backend.initialize_mcmc( + jax.random.fold_in(root_key, 7919), + wavefunction, + routed_context, + config.mcmc, + ) + final_state = backend.replace_mcmc_population( + final_state, + MCMCPopulation(q=q_final, sigma=sigma, beta=beta), + ) + backend.block_until_ready(backend.cold_walkers(final_state)) + contest_candidates = tuple( + ContestCandidate( + index=index, + route_log_probability=float(beam_logp[0, index]), + energy=float(energies[0, index]), + standard_error=float(standard_errors[0, index]), + walker_tail_std=float(tailstd[0, index]), + in_tie_set=bool(ties[0, index]), + ) + for index in range(K) + ) + contest = ContestResult( + winner=winner, + reason=reasons[0], + candidates=contest_candidates, + ) + backend.release_context(virtual_context) + return ( + wavefunction, + routed_context, + final_state, + route, + float(beam_logp[0, winner]), + contest, + ) + + +def _large_n( + backend: EvalBackend, + model: Any, + context: Any, + canonical, + mcmc_key, + config: EvalConfig, +): + state = backend.initialize_mcmc(mcmc_key, model, context, config.mcmc) + compiled = backend.compile_large_n( + model, + canonical.context, + sequence_shards=int(config.large_n_sequence_shards), + pair_tile_size=int(config.large_n_pair_tile_size), + temperature=float(config.route_temperature), + ) + route = _as_batched_route(compiled.permutation) + walker_route = _compose_walker_route(canonical.old_inverse, route) + routed_context = backend.route_context( + canonical.context, + route, + compact_custom_lap=True, + ) + state = backend.route_mcmc(state, walker_route) + backend.block_until_ready(compiled.wavefunction) + logp = float(np.asarray(compiled.log_probability)) + return compiled.wavefunction, routed_context, state, route, logp, None + + +def _validate(config: EvalConfig) -> None: + if config.contest and config.large_n: + raise ValueError("contest and large_n are mutually exclusive") + if int(config.measurements) < 1: + raise ValueError("measurements must be positive") + if int(config.mcmc.batch_size) < 1: + raise ValueError("MCMC batch size must be positive") + if int(config.mcmc.replicas) < 2: + raise ValueError("MCMC requires at least two replicas") + if int(config.mcmc.steps) < 1: + raise ValueError("MCMC replica steps must be positive") + if int(config.mcmc.burn_in_replica_steps) < 1: + raise ValueError("burn-in replica steps must be positive") + if int(config.mcmc.walker_chunk_size) < 1: + raise ValueError("walker chunk size must be positive") + if config.contest: + K = int(config.contest_candidates) + W = int(config.contest_beam_width) + if K < 2 or W < K: + raise ValueError("contest requires beam_width >= candidates >= 2") + if int(config.mcmc.batch_size) % K: + raise ValueError("MCMC batch size must be divisible by candidates") + if int(config.mcmc.batch_size) // K < 32: + raise ValueError("contest requires at least 32 walkers per candidate") + if int(config.contest_preburn) < 0: + raise ValueError("contest preburn must be non-negative") + if int(config.contest_measurements) < 1: + raise ValueError("contest measurements must be positive") + if int(config.large_n_sequence_shards) < 0: + raise ValueError("large-N sequence shards must be non-negative") + if int(config.large_n_pair_tile_size) < 1: + raise ValueError("large-N pair tile size must be positive") + + +def evaluate( + config: EvalConfig, + backend: EvalBackend, + *, + metric_sink: Callable[[EvalMetric], None] | None = None, +) -> EvalResult: + _validate(config) + started = time.perf_counter() + root_key = jax.random.PRNGKey(int(config.seed)) + model_key, mcmc_key = jax.random.split(root_key) + context = backend.load_system(config.system, config.energy) + model = backend.load_model( + config.checkpoint, + config.model, + model_key, + context, + contextualizer_attention=config.contextualizer_attention, + ) + canonical = backend.canonicalize_context(context) + embedded_route = backend.embedded_route(model) + if embedded_route is not None and (config.contest or config.large_n): + raise ValueError( + "compiled fine-tune checkpoints support ordinary eval only; " + "contest and large_n require a router checkpoint" + ) + if embedded_route is not None: + prepared = _compiled_finetune_ordinary( + backend, + model, + context, + canonical, + embedded_route, + mcmc_key, + config, + ) + path = "ordinary" + elif config.contest: + prepared = _contest( + backend, + model, + canonical, + root_key, + config, + ) + path = "contest" + elif config.large_n: + prepared = _large_n( + backend, + model, + context, + canonical, + mcmc_key, + config, + ) + path = "large_n" + else: + prepared = _ordinary( + backend, + model, + context, + canonical, + mcmc_key, + config, + ) + path = "ordinary" + wavefunction, routed_context, state, route, route_logp, contest = prepared + del model, context, canonical, embedded_route, prepared + wavefunction, routed_context, state = backend.prepare_singular( + wavefunction, + routed_context, + state, + ) + state = _burn_in( + backend, + state, + wavefunction, + routed_context, + config, + iterations=int(config.mcmc.burn_in), + replica_steps=int(config.mcmc.burn_in_replica_steps), + ) + _state, window = _measure( + backend, + state, + wavefunction, + routed_context, + config, + started=started, + metric_sink=metric_sink, + ) + route_host = np.asarray(route, dtype=np.int32) + if route_host.shape[0] != 1: + raise ValueError("single-system eval produced more than one route") + return EvalResult( + path=path, + route=tuple(int(value) for value in route_host[0]), + route_log_probability=(None if route_logp is None else float(route_logp)), + measurements=int(window.count), + walltime_seconds=float(time.perf_counter() - started), + energy=window.metrics("total"), + channels={ + "exchange": window.metrics("exchange"), + "field": window.metrics("field"), + }, + contest=contest, + ) diff --git a/src/hamiltonzero/evaluation/runtime.py b/src/hamiltonzero/evaluation/runtime.py new file mode 100644 index 0000000000000000000000000000000000000000..6536b9d3bc67f9a3c0589afb086940a68ea40dbd --- /dev/null +++ b/src/hamiltonzero/evaluation/runtime.py @@ -0,0 +1,1037 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from functools import partial +from pathlib import Path +from typing import Any, NamedTuple + +import equinox as eqx +import jax +import jax.numpy as jnp +import numpy as np +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +from hamiltonzero.checkpoint import ( + load_model as load_checkpoint_model, + load_model_metadata, +) +from hamiltonzero.compiled.api import ( + compile_wavefunction, +) +from hamiltonzero.compiled.model import ( + CompiledFinetuneWaveFunction, + build_finetune_template_model, +) +from hamiltonzero.compiled.types import ( + CompiledWaveFunction, + CompiledWaveFunctions, + EnergyInputs, +) +from hamiltonzero.compiled.tree import ( + bind_physical_compiler_kernel, + compile_physical_tree_from_shared_trunk, +) +from hamiltonzero.compiled.trunk import bind_shared_kernel, compile_shared_trunk +from hamiltonzero.config import EnergyConfig, EvalMCMCConfig, ModelConfig +from hamiltonzero.data.systems import ( + build_context_and_energy, + load_system as load_spin_hamiltonian, +) +from hamiltonzero.energy import vmc_energy_custom_lap_compiled +from hamiltonzero.energy.frame import compile_energy_frame, route_energy_inputs +from hamiltonzero.mcmc.runtime import ( + adapt_batched, + cold_samples, + init_batched_state, + run_batched, +) +from hamiltonzero.model.api import build_model +from hamiltonzero.model.context import MultiSystemContext +from hamiltonzero.model.model import _shallow_replace +from hamiltonzero.router.api import ( + route_context as apply_route_context, + route_state, +) +from hamiltonzero.router.compiled import bind_router_kernel, compile_router_static +from hamiltonzero.router.permutation import ( + permute_ctx_prefix, + permute_multi_ctx_prefix, +) + +from .backend import ( + BeamCandidates, + CanonicalContext, + LargeNCompilation, + MCMCPopulation, +) +from .large_n import compile_large_n_eval_wavefunction + + +class CompiledEvalContext(eqx.Module): + mask: Any + bmask: Any + route_perm: Any + energy_frame: Any + + +class _BeamArrays(NamedTuple): + permutations: jax.Array + log_probabilities: jax.Array + + +def _attention_name(value: str) -> str: + if value == "tuned": + return "mhsea_tuned" + if value == "einsum": + return "einsum" + raise ValueError("contextualizer attention must be 'tuned' or 'einsum'") + + +def _with_contextualizer_attention(model, value: str | None): + if value is None: + return model + contextualizer = model.route_contextualizer + implementation = _attention_name(value) + layers = _shallow_replace( + contextualizer.layers, + attn_impl=implementation, + ) + return _shallow_replace( + model, + route_contextualizer=_shallow_replace(contextualizer, layers=layers), + ) + + +def _required_positive_int(metadata: dict[str, Any], name: str) -> int: + if name not in metadata or isinstance(metadata[name], bool): + raise ValueError(f"compiled fine-tune checkpoint metadata requires {name!r}") + value = int(metadata[name]) + if value < 1: + raise ValueError( + f"compiled fine-tune checkpoint metadata {name!r} must be positive" + ) + return value + + +def _compile_energy_frame(energy_inputs, mask, bmask, permutation): + return compile_energy_frame( + energy_inputs, + mask, + bmask, + jnp.asarray(permutation, dtype=jnp.int32), + ) + + +_compile_frame = jax.jit(_compile_energy_frame) + + +@jax.jit +def _compile_frame_rows(energy_inputs, mask, bmask, permutations): + return jax.vmap( + lambda permutation: _compile_energy_frame( + energy_inputs, + mask, + bmask, + permutation, + ) + )(permutations) + + +def _compiled_wavefunctions_vmap_axes(model: CompiledWaveFunctions): + return CompiledWaveFunctions( + kernel=jax.tree_util.tree_map(lambda _value: None, model.kernel), + trees=jax.tree_util.tree_map(lambda _value: 0, model.trees), + ) + + +@partial( + jax.jit, + static_argnames=("batch_size", "replicas", "initial_m"), +) +def _initialize_single( + key, + context, + initial_sigma, + *, + batch_size: int, + replicas: int, + initial_m: int, +): + return init_batched_state( + jax.random.fold_in(key, jnp.int32(0)), + context, + batch_size=batch_size, + n_replicas=replicas, + initial_m=initial_m, + initial_sigma=initial_sigma, + ) + + +@partial( + jax.jit, + static_argnames=("batch_size", "replicas", "initial_m"), +) +def _initialize_rows( + key, + context, + initial_sigma, + *, + batch_size: int, + replicas: int, + initial_m: int, +): + indices = jnp.arange(context.mask.shape[0], dtype=jnp.int32) + return jax.vmap( + lambda index, context_row: init_batched_state( + jax.random.fold_in(key, index), + context_row, + batch_size=batch_size, + n_replicas=replicas, + initial_m=initial_m, + initial_sigma=initial_sigma, + ) + )(indices, context) + + +def _step_single( + state, + model, + context, + *, + replica_steps: int, + walker_chunk_size: int, +): + return run_batched( + model, + context, + state, + replica_steps, + walker_chunk_size=walker_chunk_size, + ) + + +@partial(jax.jit, static_argnames=("replica_steps", "walker_chunk_size")) +def _step_compiled_rows( + state, + model, + context, + *, + replica_steps: int, + walker_chunk_size: int, +): + return jax.vmap( + lambda state_row, model_row, context_row: run_batched( + model_row, + context_row, + state_row, + replica_steps, + walker_chunk_size=walker_chunk_size, + ), + in_axes=(0, _compiled_wavefunctions_vmap_axes(model), 0), + )(state, model, context) + + +def _adapt_single( + state, + beta_history_weight, + sigma_target, + sigma_scale, + haar_target, +): + return adapt_batched( + state, + beta_history_weight=beta_history_weight, + sigma_target=sigma_target, + sigma_scale=sigma_scale, + haar_target=haar_target, + ) + + +@jax.jit +def _adapt_rows( + state, + beta_history_weight, + sigma_target, + sigma_scale, + haar_target, +): + return jax.vmap( + lambda row: adapt_batched( + row, + beta_history_weight=beta_history_weight, + sigma_target=sigma_target, + sigma_scale=sigma_scale, + haar_target=haar_target, + ) + )(state) + + +def _compiled_energy_single(kernel, tree, frame, q, *, chunk_size: int): + values = vmc_energy_custom_lap_compiled( + kernel, + tree, + frame, + q, + chunk_size=chunk_size, + ) + return tuple(jnp.expand_dims(value, axis=0) for value in values) + + +@partial(jax.jit, static_argnames=("chunk_size",)) +def _compiled_energy_rows(kernel, trees, frames, q, *, chunk_size: int): + return jax.vmap( + lambda tree, frame, q_row: vmc_energy_custom_lap_compiled( + kernel, + tree, + frame, + q_row, + chunk_size=chunk_size, + ) + )(trees, frames, q) + + +_compile_single = jax.jit(compile_wavefunction) + + +def _systems_mesh(count: int) -> Mesh: + devices = tuple(jax.devices()) + lanes = min(len(devices), int(count)) + if int(count) % lanes: + raise ValueError( + f"contest K={count} must be divisible by visible devices={lanes}" + ) + return Mesh(np.asarray(devices[:lanes], dtype=object), ("systems",)) + + +def _batch_mesh(batch_size: int) -> Mesh: + devices = tuple(jax.devices()) + lanes = min(len(devices), int(batch_size)) + while int(batch_size) % lanes: + lanes -= 1 + return Mesh(np.asarray(devices[:lanes], dtype=object), ("batch",)) + + +def _single_state_sharding(mesh: Mesh, state): + walkers = NamedSharding(mesh, P("batch")) + replicated = NamedSharding(mesh, P()) + return type(state)( + q=walkers, + log_p=walkers, + grad_log_p=walkers, + beta=replicated, + sigma=replicated, + step=replicated, + key=walkers, + n_local_accept=walkers, + n_local=walkers, + n_swap_accept=walkers, + n_swap=walkers, + mask=replicated, + m=replicated, + n_haar_accept=walkers, + n_haar=walkers, + ) + + +def _system_sharding(mesh: Mesh, value): + return jax.tree_util.tree_map( + lambda array: NamedSharding( + mesh, + P("systems", *([None] * (array.ndim - 1))), + ), + value, + ) + + +def _replicated_sharding(mesh: Mesh, value): + replicated = NamedSharding(mesh, P()) + return jax.tree_util.tree_map(lambda _array: replicated, value) + + +def _place_system_rows(mesh: Mesh, value): + return jax.device_put(value, _system_sharding(mesh, value)) + + +def _abstract(tree): + return jax.tree_util.tree_map( + lambda value: jax.ShapeDtypeStruct(value.shape, value.dtype), + tree, + ) + + +def _beam_mesh(width: int) -> Mesh: + devices = tuple(jax.devices()) + lanes = len(devices) if int(width) % len(devices) == 0 else 1 + return Mesh(np.asarray(devices[:lanes], dtype=object), ("systems",)) + + +def _distributed_beam_local(decoder, static, tau, *, width: int, lanes: int): + permutations, log_probabilities = decoder.beam_search( + static.node_input, + static.raw_edge, + static.routable_mask, + global_feat=static.global_input, + tau=tau, + beam_width=width, + real_mask=static.real_mask, + first_orbit_ids=( + static.quotient_node_key, + static.quotient_edge_key, + static.needs_fwl2, + ), + router_static=static, + distributed_axis_name="systems", + distributed_lanes=lanes, + ) + return _BeamArrays(permutations, log_probabilities) + + +def _build_distributed_beam(mesh: Mesh, decoder, static, width: int): + lanes = int(mesh.shape["systems"]) + if tuple(mesh.axis_names) != ("systems",) or int(width) % lanes: + raise ValueError("distributed eval beam requires a divisible systems mesh") + mapped = jax.shard_map( + partial( + _distributed_beam_local, + width=int(width), + lanes=lanes, + ), + mesh=mesh, + in_specs=( + jax.tree_util.tree_map(lambda _leaf: P(), decoder), + jax.tree_util.tree_map(lambda _leaf: P(), static), + P(), + ), + out_specs=_BeamArrays(P(), P()), + check_vma=False, + ) + replicated = NamedSharding(mesh, P()) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _leaf: replicated, decoder), + jax.tree_util.tree_map(lambda _leaf: replicated, static), + replicated, + ), + out_shardings=_BeamArrays( + replicated, + replicated, + ), + ) + + +def _compile_eval_router_static(model, context): + trunk = compile_shared_trunk(model, context) + kernel = bind_router_kernel(model) + return compile_router_static( + kernel, + trunk, + context.route_quotient_node_key, + context.route_quotient_edge_key, + context.needs_fwl2, + ) + + +class DefaultEvalBackend: + def __init__(self) -> None: + self._energy_inputs_by_context: dict[int, EnergyInputs] = {} + self._energy_frames: dict[int, Any] = {} + self._context_meshes: dict[int, Mesh] = {} + self._contest_mesh: Mesh | None = None + self._singular_mesh: Mesh | None = None + self._singular_state_sharding: Any | None = None + self._singular_model_sharding: Any | None = None + self._singular_context_sharding: Any | None = None + self._singular_step_entries: dict[tuple[int, int], Any] = {} + self._singular_adapt_entry: Any | None = None + self._singular_energy_entries: dict[int, Any] = {} + + def build_system(self, system, energy: EnergyConfig): + context, energy_inputs = build_context_and_energy( + system, + n_max=None, + mu=energy.mu, + eps=energy.eps, + ) + self._energy_inputs_by_context[id(context)] = energy_inputs + return context + + def load_system(self, path: Path, energy: EnergyConfig): + return self.build_system(load_spin_hamiltonian(path), energy) + + def load_model( + self, + checkpoint: Path, + config: ModelConfig, + key, + context, + *, + contextualizer_attention: str | None, + ): + metadata = load_model_metadata(checkpoint) or {} + kind = metadata.get("kind", "router") + n_sites = int(context.mask.shape[-1]) + eager_template = build_model(config, key, n_max=n_sites) + if kind == "router": + eager_template = _with_contextualizer_attention( + eager_template, + contextualizer_attention, + ) + return load_checkpoint_model(checkpoint, eager_template) + if kind != "compiled_finetune": + raise ValueError(f"unsupported checkpoint kind {kind!r}") + leaf_rank = _required_positive_int(metadata, "leaf_rank") + merge_rank = _required_positive_int(metadata, "merge_rank") + checkpoint_n = int(metadata.get("n_max", n_sites)) + if checkpoint_n != n_sites: + raise ValueError( + "compiled fine-tune checkpoint width does not match the " + f"evaluation system: checkpoint={checkpoint_n}, system={n_sites}" + ) + template = build_finetune_template_model( + eager_template, + n_sites, + leaf_rank=leaf_rank, + merge_rank=merge_rank, + ) + return load_checkpoint_model(checkpoint, template) + + def canonicalize_context(self, context) -> CanonicalContext: + route = jnp.asarray(context.route_perm, dtype=jnp.int32) + if route.ndim != 1: + raise ValueError("single-system context route must have shape [N]") + inverse = jnp.argsort(route).astype(jnp.int32) + canonical = permute_ctx_prefix(context, inverse) + identity = jnp.arange(route.shape[0], dtype=jnp.int32) + canonical = eqx.tree_at( + lambda value: value.route_perm, + canonical, + identity, + ) + energy_inputs = self._energy_inputs_by_context.get(id(context)) + if energy_inputs is None: + raise RuntimeError("energy inputs are unavailable for this context") + self._energy_inputs_by_context[id(canonical)] = route_energy_inputs( + energy_inputs, + inverse, + ) + return CanonicalContext( + context=canonical, + old_inverse=inverse[None, :], + ) + + def embedded_route(self, model): + if isinstance(model, CompiledFinetuneWaveFunction): + return jnp.asarray(model.perm, dtype=jnp.int32) + return None + + def route_context( + self, + context, + permutation, + *, + compact_custom_lap: bool, + ): + permutation = jnp.asarray(permutation, dtype=jnp.int32) + if permutation.ndim == 2: + if permutation.shape[0] != 1: + raise ValueError("single-system route must have shape [1, N]") + permutation = permutation[0] + if permutation.ndim != 1: + raise ValueError("single-system route must have shape [N]") + energy_inputs = self._energy_inputs_by_context.get(id(context)) + if energy_inputs is None: + raise RuntimeError("energy inputs are unavailable for this context") + frame = _compile_frame( + energy_inputs, + context.mask, + context.bmask, + permutation, + ) + if compact_custom_lap: + routed = CompiledEvalContext( + mask=frame.masks.real, + bmask=frame.masks.balanced, + route_perm=permutation, + energy_frame=frame, + ) + self._energy_frames[id(routed)] = frame + return routed + routed = apply_route_context(context, permutation) + self._energy_frames[id(routed)] = frame + return routed + + def virtual_context(self, context, permutations): + permutations = jnp.asarray(permutations, dtype=jnp.int32) + if permutations.ndim != 2: + raise ValueError("candidate permutations must have shape [K, N]") + count = permutations.shape[0] + batched = MultiSystemContext.from_single(context) + tiled = jax.tree_util.tree_map( + lambda value: ( + jnp.repeat(value, count, axis=0) + if eqx.is_array(value) and value.ndim >= 1 + else value + ), + batched, + ) + tiled = eqx.tree_at( + lambda value: value.route_perm, + tiled, + permutations, + ) + routed = permute_multi_ctx_prefix(tiled, permutations) + energy_inputs = self._energy_inputs_by_context.get(id(context)) + if energy_inputs is None: + raise RuntimeError("energy inputs are unavailable for this context") + frames = _compile_frame_rows( + energy_inputs, + context.mask, + context.bmask, + permutations, + ) + mesh = _systems_mesh(count) + routed = _place_system_rows(mesh, routed) + frames = _place_system_rows(mesh, frames) + self._energy_frames[id(routed)] = frames + self._context_meshes[id(routed)] = mesh + self._contest_mesh = mesh + return routed + + def release_context(self, context) -> None: + self._energy_inputs_by_context.pop(id(context), None) + self._energy_frames.pop(id(context), None) + self._context_meshes.pop(id(context), None) + self._contest_mesh = None + + def beam_candidates( + self, + model, + context, + *, + beam_width: int, + top_k: int, + temperature: float, + ) -> BeamCandidates: + decoder = getattr(model, "route_decoder", None) + from hamiltonzero.model.route_pointer import TreePrefixPointerMHSEA + + if not isinstance(decoder, TreePrefixPointerMHSEA): + raise ValueError("eval requires the learned-quotient TreePrefix decoder") + if int(top_k) > int(beam_width): + raise ValueError("top_k cannot exceed beam_width") + mesh = _beam_mesh(int(beam_width)) + static = eqx.filter_jit(_compile_eval_router_static)(model, context) + replicated = NamedSharding(mesh, P()) + decoder, static, tau = jax.device_put( + ( + decoder, + static, + jnp.asarray(temperature, dtype=jnp.float32), + ), + replicated, + ) + result = _build_distributed_beam( + mesh, + decoder, + static, + int(beam_width), + )(decoder, static, tau) + jax.block_until_ready(result.permutations) + permutations = result.permutations[: int(top_k)].astype(jnp.int32) + log_probabilities = result.log_probabilities[: int(top_k)].astype(jnp.float32) + return BeamCandidates( + permutations=permutations[None], + log_probabilities=log_probabilities[None], + ) + + def compile_single(self, model, routed_context): + n_sites = int(routed_context.mask.shape[-1]) + return _compile_single( + model, + routed_context, + jnp.arange(n_sites, dtype=jnp.int32), + ) + + def compile_embedded(self, model): + if not isinstance(model, CompiledFinetuneWaveFunction): + raise TypeError("embedded eval compilation requires a fine-tune checkpoint") + return CompiledWaveFunction( + kernel=model.kernel, + tree=model.as_compiled_tree(), + ) + + def compile_candidates(self, model, canonical_context, permutations): + if self._contest_mesh is None: + raise RuntimeError("contest context must be built before compilation") + mesh = self._contest_mesh + physical_kernel = bind_physical_compiler_kernel(model) + shared_trunk = jax.jit(compile_shared_trunk)(model, canonical_context) + jax.block_until_ready(shared_trunk) + physical_sharding = _replicated_sharding(mesh, physical_kernel) + trunk_sharding = _replicated_sharding(mesh, shared_trunk) + permutation_sharding = NamedSharding(mesh, P("systems", None)) + physical_kernel = jax.device_put(physical_kernel, physical_sharding) + shared_trunk = jax.device_put(shared_trunk, trunk_sharding) + permutations = jax.device_put( + jnp.asarray(permutations, dtype=jnp.int32), + permutation_sharding, + ) + + def compile_all(kernel, trunk, candidate_permutations): + return jax.vmap( + lambda permutation: compile_physical_tree_from_shared_trunk( + kernel, + trunk, + permutation, + ) + )(candidate_permutations) + + tree_template = jax.eval_shape( + compile_all, + _abstract(physical_kernel), + _abstract(shared_trunk), + _abstract(permutations), + ) + tree_sharding = _system_sharding(mesh, tree_template) + trees = jax.jit( + compile_all, + in_shardings=( + physical_sharding, + trunk_sharding, + permutation_sharding, + ), + out_shardings=tree_sharding, + )(physical_kernel, shared_trunk, permutations) + shared_kernel = bind_shared_kernel(model) + shared_kernel = jax.device_put( + shared_kernel, + _replicated_sharding(mesh, shared_kernel), + ) + compiled = CompiledWaveFunctions(shared_kernel, trees) + jax.block_until_ready(compiled) + return compiled + + def select_candidate(self, wavefunctions, winner: int): + if self._contest_mesh is None: + raise RuntimeError("contest mesh is unavailable for winner selection") + mesh = self._contest_mesh + count = int(wavefunctions.trees.perm.shape[0]) + index = int(winner) + if index < 0 or index >= count: + raise IndexError( + f"winner index {index} outside candidate range [0, {count})" + ) + tree_sharding = _system_sharding(mesh, wavefunctions.trees) + winner_sharding = NamedSharding(mesh, P()) + + def gather(trees, selected_index): + return jax.tree_util.tree_map( + lambda value: jax.lax.dynamic_index_in_dim( + value, + selected_index, + axis=0, + keepdims=False, + ), + trees, + ) + + output_template = jax.eval_shape( + gather, + _abstract(wavefunctions.trees), + jax.ShapeDtypeStruct((), jnp.int32), + ) + tree = jax.jit( + gather, + in_shardings=(tree_sharding, winner_sharding), + out_shardings=_replicated_sharding(mesh, output_template), + )( + wavefunctions.trees, + jax.device_put(jnp.asarray(index, jnp.int32), winner_sharding), + ) + selected = CompiledWaveFunction(wavefunctions.kernel, tree) + jax.block_until_ready(selected) + return selected + + def compile_large_n( + self, + model, + canonical_context, + *, + sequence_shards: int, + pair_tile_size: int, + temperature: float, + ) -> LargeNCompilation: + result = compile_large_n_eval_wavefunction( + model, + canonical_context, + seq_shards=int(sequence_shards), + pair_tile_size=int(pair_tile_size), + tau=float(temperature), + ) + device = jax.devices()[0] + return LargeNCompilation( + wavefunction=jax.device_put(result.wavefunction, device), + permutation=jax.device_put(result.perm, device), + log_probability=jax.device_put(result.logp, device), + ) + + def prepare_singular(self, model, context, state): + if state.q.ndim != 4: + raise ValueError( + "post-selection eval MCMC state must have shape [B, R, N, 4]" + ) + frame = self._frames(context) + mesh = _batch_mesh(int(state.q.shape[0])) + state_sharding = _single_state_sharding(mesh, state) + model_sharding = _replicated_sharding(mesh, model) + context_sharding = _replicated_sharding(mesh, context) + model = jax.device_put(model, model_sharding) + context = jax.device_put(context, context_sharding) + state = jax.device_put(state, state_sharding) + self._energy_frames[id(context)] = frame + jax.block_until_ready((model, context, state.q)) + self._singular_mesh = mesh + self._singular_state_sharding = state_sharding + self._singular_model_sharding = model_sharding + self._singular_context_sharding = context_sharding + self._singular_step_entries.clear() + self._singular_adapt_entry = None + self._singular_energy_entries.clear() + return model, context, state + + def initialize_mcmc( + self, + key, + model, + context, + config: EvalMCMCConfig, + ): + del model + if isinstance(context, MultiSystemContext): + state = _initialize_rows( + key, + context, + jnp.asarray(config.initial_sigma, dtype=jnp.float32), + batch_size=int(config.batch_size), + replicas=int(config.replicas), + initial_m=int(config.initial_haar_sites), + ) + mesh = self._context_meshes.get(id(context)) + return _place_system_rows(mesh, state) if mesh is not None else state + return _initialize_single( + key, + context, + batch_size=int(config.batch_size), + replicas=int(config.replicas), + initial_m=int(config.initial_haar_sites), + initial_sigma=jnp.asarray(config.initial_sigma, dtype=jnp.float32), + ) + + def step_mcmc( + self, + state, + model, + context, + *, + replica_steps: int, + walker_chunk_size: int, + ): + if state.q.ndim != 5: + if ( + self._singular_mesh is None + or self._singular_state_sharding is None + or self._singular_model_sharding is None + or self._singular_context_sharding is None + ): + raise RuntimeError("singular eval placement has not been prepared") + key = (int(replica_steps), int(walker_chunk_size)) + step = self._singular_step_entries.get(key) + if step is None: + step = jax.jit( + partial( + _step_single, + replica_steps=key[0], + walker_chunk_size=key[1], + ), + in_shardings=( + self._singular_state_sharding, + self._singular_model_sharding, + self._singular_context_sharding, + ), + out_shardings=self._singular_state_sharding, + donate_argnums=(0,), + ) + self._singular_step_entries[key] = step + return step(state, model, context) + if not isinstance(model, CompiledWaveFunctions): + raise TypeError("multirow eval MCMC requires compiled wavefunctions") + return _step_compiled_rows( + state, + model, + context, + replica_steps=int(replica_steps), + walker_chunk_size=int(walker_chunk_size), + ) + + def adapt_mcmc(self, state, config: EvalMCMCConfig): + arguments = ( + state, + jnp.asarray(config.beta_history_weight, dtype=jnp.float32), + jnp.asarray(config.langevin_target_acceptance, dtype=jnp.float32), + jnp.asarray(config.sigma_scale, dtype=jnp.float32), + jnp.asarray(config.haar_target_acceptance, dtype=jnp.float32), + ) + if state.q.ndim == 5: + return _adapt_rows(*arguments) + if self._singular_mesh is None or self._singular_state_sharding is None: + raise RuntimeError("singular eval placement has not been prepared") + if self._singular_adapt_entry is None: + replicated = NamedSharding(self._singular_mesh, P()) + self._singular_adapt_entry = jax.jit( + _adapt_single, + in_shardings=( + self._singular_state_sharding, + replicated, + replicated, + replicated, + replicated, + ), + out_shardings=self._singular_state_sharding, + ) + return self._singular_adapt_entry(*arguments) + + def route_mcmc(self, state, permutation): + return route_state(state, permutation) + + def mcmc_population(self, state) -> MCMCPopulation: + return MCMCPopulation(q=state.q, sigma=state.sigma, beta=state.beta) + + def replace_mcmc_population( + self, + state, + population: MCMCPopulation, + ): + q = population.q + sigma = population.sigma + beta = population.beta + if q.ndim == state.q.ndim + 1 and q.shape[0] == 1: + q = jax.device_put(q[0], jax.devices()[0]) + if sigma.ndim == state.sigma.ndim + 1 and sigma.shape[0] == 1: + sigma = jax.device_put(sigma[0], jax.devices()[0]) + if beta.ndim == state.beta.ndim + 1 and beta.shape[0] == 1: + beta = jax.device_put(beta[0], jax.devices()[0]) + return eqx.tree_at( + lambda value: (value.q, value.sigma, value.beta), + state, + ( + q.astype(state.q.dtype), + sigma.astype(state.sigma.dtype), + beta.astype(state.beta.dtype), + ), + ) + + def cold_walkers(self, state): + return ( + jax.vmap(cold_samples)(state) if state.q.ndim == 5 else cold_samples(state) + ) + + def _frames(self, context): + if isinstance(context, CompiledEvalContext): + return context.energy_frame + frames = self._energy_frames.get(id(context)) + if frames is not None: + return frames + energy_inputs = self._energy_inputs_by_context.get(id(context)) + if energy_inputs is None: + raise RuntimeError("energy frame is unavailable for this context") + if isinstance(context, MultiSystemContext): + raise RuntimeError("multi-system energy frames must be compiled explicitly") + n_sites = int(context.mask.shape[-1]) + frames = _compile_frame( + energy_inputs, + context.mask, + context.bmask, + jnp.arange(n_sites, dtype=jnp.int32), + ) + self._energy_frames[id(context)] = frames + return frames + + def custom_lap_energy( + self, + model, + context, + q, + config: EnergyConfig, + ): + chunk_size = int(config.chunk_size) + if isinstance(model, CompiledWaveFunctions): + frames = self._frames(context) + return _compiled_energy_rows( + model.kernel, + model.trees, + frames, + q, + chunk_size=chunk_size, + ) + if not isinstance(model, CompiledWaveFunction): + raise TypeError("singular eval energy requires a compiled wavefunction") + if ( + self._singular_mesh is None + or self._singular_model_sharding is None + or self._singular_context_sharding is None + ): + raise RuntimeError("singular eval placement has not been prepared") + q_sharding = NamedSharding( + self._singular_mesh, + P("batch", None, None), + ) + energy_sharding = NamedSharding( + self._singular_mesh, + P(None, "batch"), + ) + output_shardings = (energy_sharding,) * 4 + q = jax.device_put(q, q_sharding) + frames = self._frames(context) + frame_sharding = _replicated_sharding( + self._singular_mesh, + frames, + ) + frames = jax.device_put(frames, frame_sharding) + energy = self._singular_energy_entries.get(chunk_size) + if energy is None: + energy = jax.jit( + partial( + _compiled_energy_single, + chunk_size=chunk_size, + ), + in_shardings=( + self._singular_model_sharding.kernel, + self._singular_model_sharding.tree, + frame_sharding, + q_sharding, + ), + out_shardings=output_shardings, + ) + self._singular_energy_entries[chunk_size] = energy + return energy( + model.kernel, + model.tree, + frames, + q, + ) + + def block_until_ready(self, value) -> None: + jax.block_until_ready(value) + + +def build_eval_backend() -> DefaultEvalBackend: + return DefaultEvalBackend() + + +__all__ = [ + "DefaultEvalBackend", + "build_eval_backend", +] diff --git a/src/hamiltonzero/evaluation/sequence_parallel.py b/src/hamiltonzero/evaluation/sequence_parallel.py new file mode 100644 index 0000000000000000000000000000000000000000..7acc22fe560e021bdf97666c67db357b64a47712 --- /dev/null +++ b/src/hamiltonzero/evaluation/sequence_parallel.py @@ -0,0 +1,173 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +import jax +import jax.numpy as jnp + +from .pallas_mha import BlockSizes, noncausal_bias_mha + + +def _validate_rectangular_attention_inputs( + q_local: jax.Array, + k_global: jax.Array, + v_global: jax.Array, + edge_bias_local: jax.Array, + key_mask: jax.Array, +) -> None: + if q_local.ndim != 3 or k_global.ndim != 3 or v_global.ndim != 3: + raise ValueError("q, k, and v must have shapes [sequence, heads, head_dim]") + q_len, n_heads, head_dim = q_local.shape + kv_len = k_global.shape[0] + if q_len < 1 or kv_len < 1: + raise ValueError("query and key sequence lengths must both be positive") + if k_global.shape != v_global.shape: + raise ValueError( + f"k and v shapes must match; got {k_global.shape} and {v_global.shape}" + ) + if k_global.shape[1:] != (n_heads, head_dim): + raise ValueError( + "q, k, and v must have the same head count and head dimension; " + f"got {q_local.shape}, {k_global.shape}, and {v_global.shape}" + ) + if edge_bias_local.shape != (q_len, kv_len, n_heads): + raise ValueError( + "edge bias must have shape [local_queries, global_keys, heads]; " + f"got {edge_bias_local.shape}, expected {(q_len, kv_len, n_heads)}" + ) + if key_mask.shape != (kv_len,): + raise ValueError(f"key mask must have shape {(kv_len,)}, got {key_mask.shape}") + arrays = (q_local, k_global, v_global, edge_bias_local) + if any(value.dtype != jnp.float32 for value in arrays): + raise TypeError( + "large-N rectangular attention is fp32-only; got " + + ", ".join(str(value.dtype) for value in arrays) + ) + + +def _dividing_block_size(length: int, requested: int | None) -> int: + block = min(length, 128 if requested is None else requested) + if block < 1: + raise ValueError(f"block size must be positive, got {block}") + while length % block: + block //= 2 + return block + + +def pallas_rectangular_edge_attention( + q_local: jax.Array, + k_global: jax.Array, + v_global: jax.Array, + edge_bias_local: jax.Array, + key_mask: jax.Array, + *, + block_k: int | None = None, +) -> jax.Array: + + _validate_rectangular_attention_inputs( + q_local, k_global, v_global, edge_bias_local, key_mask + ) + sm_scale = float(q_local.shape[-1]) ** -0.5 + q_len = q_local.shape[0] + kv_len = k_global.shape[0] + bq = _dividing_block_size(q_len, None) + bk = _dividing_block_size(kv_len, block_k) + block_sizes = BlockSizes(block_q=bq, block_k=bk) + + masked_bias = jnp.where( + key_mask[None, :, None].astype(bool), + edge_bias_local, + jnp.asarray(-1.0e30, dtype=jnp.float32), + ) + return noncausal_bias_mha( + (q_local * jnp.asarray(sm_scale, dtype=jnp.float32))[None], + k_global[None], + v_global[None], + masked_bias[None], + block_sizes=block_sizes, + )[0] + + +def ring_learned_fwl2_local( + a_local: jax.Array, + b_local: jax.Array, + *, + axis_name: str, + axis_size: int, +) -> jax.Array: + + local_rows, global_columns, channels = a_local.shape + if b_local.shape != (local_rows, global_columns, channels): + raise ValueError( + f"local 2-FWL shapes must match; got {a_local.shape}, {b_local.shape}" + ) + return ring_learned_fwl2_columns_local( + a_local, + b_local, + axis_name=axis_name, + axis_size=axis_size, + ) + + +def ring_learned_fwl2_columns_local( + a_local: jax.Array, + b_local_columns: jax.Array, + *, + axis_name: str, + axis_size: int, +) -> jax.Array: + + if a_local.ndim != 3 or b_local_columns.ndim != 3: + raise ValueError( + "local 2-FWL operands must both be rank three; got " + f"{a_local.shape} and {b_local_columns.shape}" + ) + local_rows, global_columns, channels = a_local.shape + if b_local_columns.shape[0] != local_rows: + raise ValueError( + "local A/B row counts must match; got " + f"{local_rows} and {b_local_columns.shape[0]}" + ) + if b_local_columns.shape[2] != channels: + raise ValueError( + "local A/B channel counts must match; got " + f"{channels} and {b_local_columns.shape[2]}" + ) + if a_local.dtype != b_local_columns.dtype: + raise TypeError( + "local A/B dtypes must match; got " + f"{a_local.dtype} and {b_local_columns.dtype}" + ) + if global_columns != local_rows * axis_size: + raise ValueError( + "2-FWL row shards must evenly tile the contracted axis; " + f"got local_rows={local_rows}, columns={global_columns}, " + f"axis_size={axis_size}" + ) + + def contribution(b_panel, origin): + a_panel = jax.lax.dynamic_slice_in_dim( + a_local, origin * local_rows, local_rows, axis=1 + ) + return jnp.einsum("ikc,kjc->ijc", a_panel, b_panel) + + origin0 = jax.lax.axis_index(axis_name).astype(jnp.int32) + accumulator0 = contribution(b_local_columns, origin0) + ring_permutation = [(lane, (lane + 1) % axis_size) for lane in range(axis_size)] + + def ring_step(carry, _): + b_panel, origin, accumulator = carry + b_panel = jax.lax.ppermute(b_panel, axis_name=axis_name, perm=ring_permutation) + origin = (origin - jnp.asarray(1, jnp.int32)) % axis_size + accumulator = accumulator + contribution(b_panel, origin) + return (b_panel, origin, accumulator), None + + (_, _, result), _ = jax.lax.scan( + ring_step, + (b_local_columns, origin0, accumulator0), + xs=None, + length=axis_size - 1, + ) + return result diff --git a/src/hamiltonzero/evaluation/sequence_trunk.py b/src/hamiltonzero/evaluation/sequence_trunk.py new file mode 100644 index 0000000000000000000000000000000000000000..9333aa7da0f2ea6c7a0afc45d9e5c0a3789b3672 --- /dev/null +++ b/src/hamiltonzero/evaluation/sequence_trunk.py @@ -0,0 +1,1804 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +from collections.abc import Callable +from functools import partial +from math import gcd + +import equinox as eqx +import jax +import jax.numpy as jnp +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +from hamiltonzero.compiled.types import SharedTrunk, TrunkCompilerKernel +from hamiltonzero.model.fused_silu import fused_silu +from hamiltonzero.model.global_ladder import ( + BoundaryGlobalUpdate, + ResidualGlobalUpdate, + TreeGlobalUpdate, +) +from hamiltonzero.model.readout_leaf_context import ( + PhysicalReadoutContext, + PhysicalReadoutContextLayer, + RouterContext, + RouterContextLayer, + lca_alibi_bias, + lca_fixed_slopes, +) +from .sequence_parallel import ( + pallas_rectangular_edge_attention, + ring_learned_fwl2_columns_local, + ring_learned_fwl2_local, +) +from hamiltonzero.model.tree import _tree_ngpt_residual, _tree_sphere + + +def _local_rows(x: jax.Array, *, axis_name: str, local_size: int) -> jax.Array: + start = jax.lax.axis_index(axis_name) * local_size + return jax.lax.dynamic_slice_in_dim(x, start, local_size, axis=0) + + +def _global_row_indices(*, axis_name: str, local_size: int) -> jax.Array: + start = jax.lax.axis_index(axis_name) * local_size + return start + jnp.arange(local_size, dtype=jnp.int32) + + +def _gather_rows(x: jax.Array, *, axis_name: str) -> jax.Array: + + return jax.lax.all_gather(x, axis_name=axis_name, axis=0, tiled=True) + + +def ring_permute_rows_local( + values: jax.Array, + permutation: jax.Array, + *, + axis_name: str, + axis_size: int, +) -> jax.Array: + + local_size = values.shape[0] + lane = jax.lax.axis_index(axis_name).astype(jnp.int32) + output_ids = jax.lax.dynamic_slice_in_dim( + permutation, + lane * local_size, + local_size, + axis=0, + ).astype(jnp.int32) + owners = output_ids // local_size + offsets = output_ids % local_size + output = jnp.zeros_like(values) + + def select(panel, origin, current): + selected = panel[offsets] + take = owners == origin + while take.ndim < selected.ndim: + take = take[..., None] + return jnp.where(take, selected, current) + + origin = lane + output = select(values, origin, output) + ring = tuple((i, (i + 1) % axis_size) for i in range(axis_size)) + + def step(carry, _): + panel, panel_origin, current = carry + panel = jax.lax.ppermute(panel, axis_name, ring) + panel_origin = (panel_origin - jnp.asarray(1, dtype=jnp.int32)) % axis_size + return (panel, panel_origin, select(panel, panel_origin, current)), None + + (_, _, output), _ = jax.lax.scan( + step, + (values, origin, output), + xs=None, + length=axis_size - 1, + ) + return output + + +def permute_pair_rows_and_columns_local( + edge_rows: jax.Array, + permutation: jax.Array, + *, + axis_name: str, + axis_size: int, +) -> jax.Array: + + rows = ring_permute_rows_local( + edge_rows, + permutation, + axis_name=axis_name, + axis_size=axis_size, + ) + return rows[:, permutation] + + +def transpose_pair_rows_local( + edge_rows: jax.Array, + *, + axis_name: str, + axis_size: int, +) -> jax.Array: + + if axis_size == 1: + return jnp.swapaxes(edge_rows, 0, 1) + local_size, n = edge_rows.shape[:2] + if n != local_size * axis_size: + raise ValueError( + "pair transpose requires N == local_rows * axis_size, got " + f"N={n}, local_rows={local_size}, axis_size={axis_size}" + ) + trailing = edge_rows.shape[2:] + send = edge_rows.reshape((local_size, axis_size, local_size) + trailing) + send = jnp.swapaxes(send, 0, 1) + received = jax.lax.all_to_all( + send, + axis_name, + split_axis=0, + concat_axis=0, + ) + received = jnp.transpose( + received, + (2, 0, 1) + tuple(range(3, received.ndim)), + ) + return received.reshape((local_size, n) + trailing) + + +def _linear(module, x: jax.Array) -> jax.Array: + out = x @ module.weight + bias = getattr(module, "bias", None) + return out if bias is None else out + bias + + +def _norm(module, x: jax.Array) -> jax.Array: + + weight = getattr(module, "weight", None) + if weight is None: + return x + stats_dtype = jnp.promote_types(jnp.float32, x.dtype) + x_hi = x.astype(stats_dtype) if x.dtype != stats_dtype else x + bias = getattr(module, "bias", None) + if bias is None: + stat = jnp.mean(x_hi * x_hi, axis=-1, keepdims=True) + normalized = x_hi * jax.lax.rsqrt(stat + module.eps) + else: + mean = jnp.mean(x_hi, axis=-1, keepdims=True) + centered = x_hi - mean + stat = jnp.mean(centered * centered, axis=-1, keepdims=True) + normalized = centered * jax.lax.rsqrt(stat + module.eps) + normalized = normalized.astype(x.dtype) + out = normalized * weight.astype(x.dtype) + return out if bias is None else out + bias.astype(x.dtype) + + +def _raw_norm(x, scale, shift, *, eps: float): + stats_dtype = jnp.promote_types(jnp.float32, x.dtype) + x_hi = x.astype(stats_dtype) + if shift is None: + normalized = x_hi * jax.lax.rsqrt( + jnp.mean(x_hi * x_hi, axis=-1, keepdims=True) + eps + ) + else: + centered = x_hi - jnp.mean(x_hi, axis=-1, keepdims=True) + normalized = centered * jax.lax.rsqrt( + jnp.mean(centered * centered, axis=-1, keepdims=True) + eps + ) + out = normalized.astype(x.dtype) * scale.astype(x.dtype) + return out if shift is None else out + shift.astype(x.dtype) + + +def _mlp_after_input_projection(mlp, hidden: jax.Array) -> jax.Array: + for norm, l1, l2 in zip(mlp.block_norms, mlp.block_l1s, mlp.block_l2s, strict=True): + hidden = hidden + mlp.inner_gain * _linear( + l2, mlp._act(_linear(l1, _norm(norm, hidden))) + ) + return _linear(mlp.out_proj, _norm(mlp.out_norm, hidden)) + + +def _mlp(mlp, x: jax.Array) -> jax.Array: + return _mlp_after_input_projection(mlp, _linear(mlp.in_proj, x)) + + +def _unnormalized_mlp(mlp, x: jax.Array) -> jax.Array: + hidden = _linear(mlp.in_proj, x) + for l1, l2 in zip(mlp.block_l1s, mlp.block_l2s, strict=True): + hidden = hidden + mlp.inner_gain * _linear(l2, mlp._act(_linear(l1, hidden))) + return _linear(mlp.out_proj, hidden) + + +def _split_linear(module, parts: tuple[jax.Array, ...]) -> jax.Array: + + offset = 0 + out = module.bias + for part in parts: + width = part.shape[-1] + out = out + part @ module.weight[offset : offset + width] + offset += width + if offset != module.weight.shape[0]: + raise ValueError( + f"split input width {offset} does not match weight {module.weight.shape[0]}" + ) + return out + + +def _mlp_split_input(mlp, parts: tuple[jax.Array, ...]) -> jax.Array: + return _mlp_after_input_projection(mlp, _split_linear(mlp.in_proj, parts)) + + +def _g_descriptor_pool(pool, g, xs, mask): + + n = xs.shape[0] + xn = _norm(pool.ln_in, xs) + q = (g @ pool.W_q).reshape(pool.n_heads, pool.d_k) + k = _linear(pool.K, xn).reshape(n, pool.n_heads, pool.d_k) + v = _linear(pool.V, xn).reshape(n, pool.n_heads, pool.d_v) + scores = jnp.einsum("hd,nhd->hn", q, k) / jnp.sqrt( + jnp.asarray(pool.d_k, dtype=xs.dtype) + ) + scores = jnp.where( + mask[None, :] > 0, + scores, + jnp.asarray(-1.0e30, dtype=scores.dtype), + ) + weights = jax.nn.softmax(scores, axis=-1) + return jnp.einsum("hn,nhv->hv", weights, v).reshape(-1) + + +def _g_descriptor_pool_rows(pool, g, xs_rows, key_mask, *, tile_size: int = 128): + + r, n = xs_rows.shape[:2] + if tile_size < 1: + raise ValueError("descriptor-pool tile_size must be positive") + tile_width = gcd(n, min(n, int(tile_size))) + tile_count = n // tile_width + q = (g @ pool.W_q).reshape(pool.n_heads, pool.d_k) + scale = jnp.sqrt(jnp.asarray(pool.d_k, dtype=xs_rows.dtype)) + scores0 = jnp.zeros((r, pool.n_heads, n), dtype=xs_rows.dtype) + + def score_tile(tile_index, scores): + start = tile_index * tile_width + xs_tile = jax.lax.dynamic_slice_in_dim(xs_rows, start, tile_width, axis=1) + xn_tile = _norm(pool.ln_in, xs_tile) + k_tile = _linear(pool.K, xn_tile).reshape(r, tile_width, pool.n_heads, pool.d_k) + tile_scores = jnp.einsum("hd,rthd->rht", q, k_tile) / scale + mask_tile = jax.lax.dynamic_slice_in_dim(key_mask, start, tile_width, axis=0) + tile_scores = jnp.where( + mask_tile[None, None, :] > 0, + tile_scores, + jnp.asarray(-1.0e30, dtype=tile_scores.dtype), + ) + return jax.lax.dynamic_update_slice_in_dim(scores, tile_scores, start, axis=2) + + scores = jax.lax.fori_loop(0, tile_count, score_tile, scores0) + weights = jax.nn.softmax(scores, axis=-1) + pooled0 = jnp.zeros((r, pool.n_heads, pool.d_v), dtype=xs_rows.dtype) + + def value_tile(tile_index, pooled): + start = tile_index * tile_width + xs_tile = jax.lax.dynamic_slice_in_dim(xs_rows, start, tile_width, axis=1) + xn_tile = _norm(pool.ln_in, xs_tile) + v_tile = _linear(pool.V, xn_tile).reshape(r, tile_width, pool.n_heads, pool.d_v) + weight_tile = jax.lax.dynamic_slice_in_dim(weights, start, tile_width, axis=2) + return pooled + jnp.einsum("rht,rthv->rhv", weight_tile, v_tile) + + pooled = jax.lax.fori_loop(0, tile_count, value_tile, pooled0) + return pooled.reshape(r, -1) + + +def _g_update(update, g, pooled): + g_input = g @ update.g_tap_w + x = jnp.concatenate((g_input, pooled.astype(g.dtype))) + stats = jnp.mean(jnp.square(x), keepdims=True) + x = x * jax.lax.rsqrt(stats + 1.0e-5) * update.ln_s + hidden = fused_silu(x @ update.w1 + update.b1) + delta = hidden @ update.w2 + update.b2 + if isinstance(update, ResidualGlobalUpdate): + return g + update.residual_gain * delta + if isinstance(update, BoundaryGlobalUpdate): + return _tree_sphere(g + delta) + if not isinstance(update, TreeGlobalUpdate): + raise TypeError(f"unsupported global update {type(update)!r}") + skip = _tree_sphere(g) + proposal = _tree_sphere(delta) + gain = update.alpha_max * jax.nn.sigmoid(update.alpha) + return _tree_sphere(skip + gain * (proposal - skip)) + + +def _edge_row_col_global_update( + module, + g, + edge_rows, + mask, + row_mask, + *, + axis_name: str, + tile_size: int = 128, +): + row_desc_rows = _g_descriptor_pool_rows( + module.row_pool, g, edge_rows, mask, tile_size=tile_size + ) + row_desc = _gather_rows(row_desc_rows, axis_name=axis_name) + + edge_column_rows = transpose_pair_rows_local( + edge_rows, + axis_name=axis_name, + axis_size=edge_rows.shape[1] // edge_rows.shape[0], + ) + col_desc_rows = _g_descriptor_pool_rows( + module.col_pool, g, edge_column_rows, mask, tile_size=tile_size + ) + col_desc = _gather_rows(col_desc_rows, axis_name=axis_name) + descriptors = jnp.concatenate((row_desc, col_desc), axis=0) + descriptor_mask = jnp.concatenate((mask, mask), axis=0) + pooled = _g_descriptor_pool(module.set_pool, g, descriptors, descriptor_mask) + return _g_update(module.update, g, pooled) + + +def sequence_parallel_edge_global_update_local( + module, + g, + edge_rows, + mask, + *, + axis_name: str, + tile_size: int = 128, +): + + row_indices = _global_row_indices( + axis_name=axis_name, local_size=edge_rows.shape[0] + ) + return _edge_row_col_global_update( + module, + g, + edge_rows, + mask, + mask[row_indices], + axis_name=axis_name, + tile_size=tile_size, + ) + + +def _psi_project(edge_update, parts, *, left: bool): + if left: + linear_in = edge_update.psi_L_in + linear_out = edge_update.psi_L_out + else: + linear_in = edge_update.psi_R_in + linear_out = edge_update.psi_R_out + value = _split_linear(linear_in, parts) + hidden = fused_silu(value) + return _linear(linear_out, hidden) + + +def _context_edge_update( + edge_update, + edge_rows, + even_rows, + even_all, + mask, + row_mask, + *, + axis_name: str, + axis_size: int, +): + + edge_ln = _norm(edge_update.ln_edge, edge_rows) + even_rows_ln = _norm(edge_update.ln_even, even_rows) + even_all_ln = _norm(edge_update.ln_even, even_all) + if edge_update.node_ctx_proj is not None: + even_rows_ln = _linear(edge_update.node_ctx_proj, even_rows_ln) + even_all_ln = _linear(edge_update.node_ctx_proj, even_all_ln) + row_endpoint = even_rows_ln[:, None, :] + column_endpoint = even_all_ln[None, :, :] + + pair_parts = (edge_ln, row_endpoint, column_endpoint) + left = _psi_project(edge_update, pair_parts, left=True) + right = _psi_project(edge_update, pair_parts, left=False) + left = left * ( + row_mask[:, None, None].astype(left.dtype) + * mask[None, :, None].astype(left.dtype) + ) + right = right * mask[None, :, None].astype(right.dtype) + path = ring_learned_fwl2_local( + left, + right, + axis_name=axis_name, + axis_size=axis_size, + ) + n_eff = jnp.maximum(jnp.sum(mask), 1.0).astype(path.dtype) + path = _norm(edge_update.ln_path, path / jnp.sqrt(n_eff)) + return _mlp_split_input(edge_update.ffn, (*pair_parts, path)) + + +def _rectangular_block_attention( + attention, + even_rows, + edge_rows, + mask, + *, + axis_name: str, + block_k: int, +): + + r = even_rows.shape[0] + qkv_rows = _linear(attention.W_QKV, even_rows).reshape( + r, 3, attention.n_heads_kernel, attention.d_head + ) + qkv_all = _gather_rows(qkv_rows, axis_name=axis_name) + query = qkv_rows[:, 0] + key = qkv_all[:, 1] + value = qkv_all[:, 2] + + edge_pre = _norm(attention.ln_edge, edge_rows) + bias = _unnormalized_mlp(attention.bias_mlp, edge_pre) + bias = bias / jnp.sqrt(jnp.asarray(attention.d_head, bias.dtype)) + out = pallas_rectangular_edge_attention( + query, + key, + value, + bias, + mask, + block_k=block_k, + ) + gate = out[:, : attention.n_heads] + value_out = out[:, attention.n_heads :] + out = jax.nn.sigmoid(gate) * value_out + return _linear(attention.W_O, out.reshape(r, -1)) + + +def _sequence_transformer_block( + block, + even_rows, + edge_rows, + g, + mask, + row_mask, + *, + axis_name: str, + axis_size: int, + attention_block_k: int, +): + even_all = _gather_rows(even_rows, axis_name=axis_name) + edge_delta = _context_edge_update( + block.edge_update_ctx, + edge_rows, + even_rows, + even_all, + mask, + row_mask, + axis_name=axis_name, + axis_size=axis_size, + ) + edge_rows = edge_rows + block.residual_gain * edge_delta + + even_pre = _norm(block.ln_attn, even_rows) + attention_delta = _rectangular_block_attention( + block.attn, + even_pre, + edge_rows, + mask, + axis_name=axis_name, + block_k=attention_block_k, + ) + even_rows = even_rows + block.residual_gain * attention_delta + + even_pre = _norm(block.ln_ffn, even_rows) + even_pre = even_pre + (g @ block.g_ffn_proj_w)[None].astype(even_pre.dtype) + ffn_delta = _linear(block.ffn.l2, fused_silu(_linear(block.ffn.l1, even_pre))) + even_rows = even_rows + block.residual_gain * ffn_delta + + even_all = _gather_rows(even_rows, axis_name=axis_name) + pooled = _g_descriptor_pool(block.g_pool, g, even_all, mask) + g = _g_update(block.g_update, g, pooled) + return even_rows, edge_rows, g + + +def _sequence_trunk_local( + trunk, + local_rows, + edge_rows, + g, + mask, + row_mask, + *, + axis_name: str, + axis_size: int, + attention_block_k: int, +): + step = partial( + _sequence_transformer_block, + mask=mask, + row_mask=row_mask, + axis_name=axis_name, + axis_size=axis_size, + attention_block_k=attention_block_k, + ) + dynamic, static = eqx.partition(trunk.blocks, eqx.is_array) + + def scan_step(carry, layer_dynamic): + block = eqx.combine(layer_dynamic, static) + return step(block, *carry), None + + (local_rows, edge_rows, g), _ = jax.lax.scan( + scan_step, (local_rows, edge_rows, g), dynamic + ) + return local_rows, edge_rows, g + + +def _pair_mlp_tiled(mlp, values, *, tile_size: int): + + n = values.shape[1] + tile_width = min(n, int(tile_size)) + if tile_width < 1: + raise ValueError("pair MLP tile_size must be positive") + full_tiles = n // tile_width + tail_start = full_tiles * tile_width + output = jnp.zeros( + values.shape[:2] + (int(mlp.out_proj.weight.shape[1]),), + dtype=values.dtype, + ) + + def update_tile(start, width, current): + tile = jax.lax.dynamic_slice_in_dim(values, start, width, axis=1) + return jax.lax.dynamic_update_slice_in_dim( + current, _mlp(mlp, tile), start, axis=1 + ) + + output = jax.lax.fori_loop( + 0, + full_tiles, + lambda tile_index, current: update_tile( + tile_index * tile_width, tile_width, current + ), + output, + ) + if tail_start < n: + output = update_tile(tail_start, n - tail_start, output) + return output + + +def _pair_mlp_parts_tiled(mlp, parts, *, tile_size: int): + + parts = tuple(parts) + if not parts: + raise ValueError("pair MLP requires at least one input part") + pair_shape = parts[0].shape[:2] + if any(part.shape[:2] != pair_shape for part in parts): + raise ValueError("pair MLP input parts must share their [R,N] axes") + n = pair_shape[1] + tile_width = min(n, int(tile_size)) + if tile_width < 1: + raise ValueError("pair MLP tile_size must be positive") + full_tiles = n // tile_width + tail_start = full_tiles * tile_width + output = jnp.zeros( + pair_shape + (int(mlp.out_proj.weight.shape[1]),), + dtype=parts[0].dtype, + ) + + def update_tile(start, width, current): + tile = jnp.concatenate( + tuple( + jax.lax.dynamic_slice_in_dim(part, start, width, axis=1) + for part in parts + ), + axis=-1, + ) + return jax.lax.dynamic_update_slice_in_dim( + current, _mlp(mlp, tile), start, axis=1 + ) + + output = jax.lax.fori_loop( + 0, + full_tiles, + lambda tile_index, current: update_tile( + tile_index * tile_width, tile_width, current + ), + output, + ) + if tail_start < n: + output = update_tile(tail_start, n - tail_start, output) + return output + + +def _sequence_context_attention( + layer, + c_rows, + edge_rows, + bmask, + row_indices, + *, + axis_name: str, + block_k: int, + tile_size: int, +): + + r = c_rows.shape[0] + n = bmask.shape[0] + qkv_rows = _linear(layer.W_QKV, c_rows).reshape( + r, 3, layer.n_heads_kernel, layer.d_head + ) + qkv_all = _gather_rows(qkv_rows, axis_name=axis_name) + query = qkv_rows[:, 0] + key = qkv_all[:, 1] + value = qkv_all[:, 2] + col_indices = jnp.arange(n, dtype=jnp.int32) + rel = row_indices[:, None] - col_indices[None, :] + if isinstance(layer, PhysicalReadoutContextLayer): + direction = jnp.where(rel < 0, 1.0, jnp.where(rel > 0, -1.0, 0.0)).astype( + edge_rows.dtype + )[..., None] + elif isinstance(layer, RouterContextLayer): + direction = jnp.zeros((r, n, 1), dtype=edge_rows.dtype) + else: + raise TypeError("unsupported contextualizer layer") + bias = _pair_mlp_parts_tiled( + layer.bias_mlp, + (edge_rows, direction), + tile_size=tile_size, + ) + bias = bias / jnp.sqrt(jnp.asarray(layer.d_head, dtype=bias.dtype)) + if isinstance(layer, PhysicalReadoutContextLayer): + bias = bias + jnp.transpose( + lca_alibi_bias( + row_indices, + col_indices, + lca_fixed_slopes(layer.n_heads_kernel, dtype=bias.dtype), + ), + (1, 2, 0), + ) + out = pallas_rectangular_edge_attention( + query, + key, + value, + bias, + bmask, + block_k=block_k, + ) + out = jax.nn.sigmoid(out[:, : layer.n_heads]) * out[:, layer.n_heads :] + return _linear(layer.W_O, out.reshape(r, -1)) + + +def _sequence_context_layer( + layer, + c_rows, + edge_rows, + g, + bmask, + *, + axis_name: str, + axis_size: int, + tile_size: int, + attention_block_k: int, +): + + r, n = edge_rows.shape[:2] + dtype = edge_rows.dtype + row_indices = _global_row_indices(axis_name=axis_name, local_size=r) + row_mask = bmask[row_indices] + edge_n = _norm(layer.ln_edge, edge_rows) + if isinstance(layer, PhysicalReadoutContextLayer): + reverse_n = transpose_pair_rows_local( + edge_n, + axis_name=axis_name, + axis_size=axis_size, + ) + clock_rows = _local_rows( + layer._slot_clock(n, dtype, bmask), + axis_name=axis_name, + local_size=r, + ) + summary = layer.edge_summary_tiled( + edge_n, + bmask.astype(dtype), + edge_reverse_rows=reverse_n, + row_indices=row_indices, + tile_size=tile_size, + ) + parts = [ + _norm(layer.ln_c, c_rows + clock_rows), + _norm(layer.ln_summary, summary), + ] + elif isinstance(layer, RouterContextLayer): + clock_rows = None + parts = [_norm(layer.ln_c, c_rows)] + else: + raise TypeError("unsupported contextualizer layer") + delta_ctx = layer.residual_scale * _mlp( + layer.ctx_mlp, jnp.concatenate(parts, axis=-1) + ) + c1 = jnp.where( + row_mask[:, None].astype(bool), + c_rows + delta_ctx, + jnp.zeros_like(c_rows), + ) + + c1_all = _gather_rows(c1, axis_name=axis_name) + g = _g_update( + layer.g_update, + g, + _g_descriptor_pool(layer.g_pool, g, c1_all, bmask), + ) + + edge_ctx_rows = _linear( + layer.edge_node_ctx_proj, + _norm(layer.ln_edge_ctx, c1), + ) + edge_ctx_all = _gather_rows(edge_ctx_rows, axis_name=axis_name) + edge1 = layer.edge_update_tiled( + edge_rows, + edge_n, + edge_ctx_rows, + edge_ctx_all, + bmask.astype(dtype), + row_indices=row_indices, + g=g, + tile_size=tile_size, + ) + + attention_source = c1 + if clock_rows is not None: + attention_source = attention_source + clock_rows + delta_attn = layer.residual_scale * _sequence_context_attention( + layer, + _norm(layer.ln_attn, attention_source), + _norm(layer.ln_edge_attn, edge1), + bmask, + row_indices, + axis_name=axis_name, + block_k=attention_block_k, + tile_size=tile_size, + ) + c_out = jnp.where( + row_mask[:, None].astype(bool), + c1 + delta_attn, + jnp.zeros_like(c1), + ) + return c_out, edge1, g + + +def sequence_parallel_contextualizer_local( + contextualizer, + node_rows, + edge_rows, + real_mask, + structural_mask, + g=None, + *, + axis_name: str, + axis_size: int, + tile_size: int = 128, + attention_block_k: int = 128, +): + + if not isinstance(contextualizer, (PhysicalReadoutContext, RouterContext)): + raise TypeError("unsupported contextualizer") + node_rows = node_rows.astype(jnp.float32) + edge_rows = edge_rows.astype(jnp.float32) + r, n = edge_rows.shape[:2] + if node_rows.shape[0] != r or n != r * axis_size: + raise ValueError("contextualizer inputs do not match the seq row layout") + if real_mask.shape != (n,) or structural_mask.shape != (n,): + raise ValueError("contextualizer masks must have replicated shape [N]") + row_indices = _global_row_indices(axis_name=axis_name, local_size=r) + real_rows = real_mask[row_indices].astype(bool) + active_rows = structural_mask[row_indices].astype(bool) + virtual_rows = active_rows & ~real_rows + real = real_mask.astype(bool) + active = structural_mask.astype(bool) + virtual = active & ~real + + virtual_node = contextualizer.virtual_node[0].astype(node_rows.dtype) + node_rows = jnp.where( + real_rows[:, None], + node_rows, + jnp.where( + virtual_rows[:, None], + virtual_node[None, :], + jnp.zeros_like(node_rows), + ), + ) + row_real_pair = real_rows[:, None] & real[None, :] + mixed_pair = (real_rows[:, None] & virtual[None, :]) | ( + virtual_rows[:, None] & real[None, :] + ) + virtual_pair = virtual_rows[:, None] & virtual[None, :] + edge_rows = jnp.where( + row_real_pair[..., None], + edge_rows, + jnp.where( + mixed_pair[..., None], + contextualizer.edge_empty_nonempty[0].astype(edge_rows.dtype), + jnp.where( + virtual_pair[..., None], + contextualizer.edge_empty_empty[0].astype(edge_rows.dtype), + jnp.zeros_like(edge_rows), + ), + ), + ) + + step = partial( + _sequence_context_layer, + bmask=structural_mask, + axis_name=axis_name, + axis_size=axis_size, + tile_size=tile_size, + attention_block_k=attention_block_k, + ) + dynamic, static = eqx.partition(contextualizer.layers, eqx.is_array) + + def scan_step(carry, layer_dynamic): + layer = eqx.combine(layer_dynamic, static) + return step(layer, *carry), None + + (node_rows, edge_rows, g), _ = jax.lax.scan( + scan_step, (node_rows, edge_rows, g), dynamic + ) + return node_rows, edge_rows, g + + +def _sequence_tree_fwl( + module, + edge_rows, + c_rows, + c_all, + mask, + row_mask, + *, + axis_name: str, + axis_size: int, + tile_size: int, +): + + width = int(edge_rows.shape[1]) + if tile_size < 1: + raise ValueError("tree FWL tile_size must be positive") + + tile_width = gcd(width, min(width, int(tile_size))) + n_tiles = width // tile_width + c_rows_ctx = _norm(module.ln_c, c_rows) + c_all_ctx = _norm(module.ln_c, c_all) + if module.node_ctx_proj is not None: + c_rows_ctx = _linear(module.node_ctx_proj, c_rows_ctx) + c_all_ctx = _linear(module.node_ctx_proj, c_all_ctx) + + left0 = jnp.zeros( + ( + edge_rows.shape[0], + width, + int(module.psi_L_out.weight.shape[1]), + ), + dtype=edge_rows.dtype, + ) + + def project_left_tile(tile_index, left): + start = tile_index * tile_width + edge_tile = jax.lax.dynamic_slice_in_dim(edge_rows, start, tile_width, axis=1) + c_columns = jax.lax.dynamic_slice_in_dim(c_all_ctx, start, tile_width, axis=0) + mask_columns = jax.lax.dynamic_slice_in_dim(mask, start, tile_width, axis=0) + pair_parts = ( + _norm(module.ln_edge, edge_tile), + c_rows_ctx[:, None, :], + c_columns[None, :, :], + ) + projected = _psi_project(module, pair_parts, left=True) + projected = projected * ( + row_mask[:, None, None].astype(projected.dtype) + * mask_columns[None, :, None].astype(projected.dtype) + ) + return jax.lax.dynamic_update_slice_in_dim(left, projected, start, axis=1) + + left = jax.lax.fori_loop(0, n_tiles, project_left_tile, left0) + n_eff = jnp.maximum(jnp.sum(mask), 1.0).astype(edge_rows.dtype) + + def update_destination_tile(tile_index, updated_edges): + start = tile_index * tile_width + edge_tile = jax.lax.dynamic_slice_in_dim( + updated_edges, start, tile_width, axis=1 + ) + c_columns = jax.lax.dynamic_slice_in_dim(c_all_ctx, start, tile_width, axis=0) + mask_columns = jax.lax.dynamic_slice_in_dim(mask, start, tile_width, axis=0) + pair_parts = ( + _norm(module.ln_edge, edge_tile), + c_rows_ctx[:, None, :], + c_columns[None, :, :], + ) + right = _psi_project(module, pair_parts, left=False) + right = right * mask_columns[None, :, None].astype(right.dtype) + path = ring_learned_fwl2_columns_local( + left, + right, + axis_name=axis_name, + axis_size=axis_size, + ) + path = path / jnp.sqrt(n_eff) + path = _norm(module.ln_path, path) + hidden = fused_silu(_split_linear(module.ffn_in, (*pair_parts, path))) + delta = _linear(module.ffn_out, hidden) + update_mask = row_mask[:, None].astype(bool) & mask_columns[None, :].astype( + bool + ) + edge_tile = _tree_ngpt_residual( + edge_tile, + delta, + module.alpha, + max_gain=module.ngpt_alpha_max, + tag_id="", + update_mask=update_mask, + ) + return jax.lax.dynamic_update_slice_in_dim( + updated_edges, edge_tile, start, axis=1 + ) + + return jax.lax.fori_loop(0, n_tiles, update_destination_tile, edge_rows) + + +def _sequence_level_edge_attention( + module, + c_rows, + edge_rows, + mask, + row_mask, + *, + axis_name: str, + level: int, + tile_size: int, + attention_block_k: int, +): + + r, n = edge_rows.shape[:2] + x = _raw_norm( + c_rows, + module.ln_scale, + None, + eps=module.ln_eps, + ) + qkv_rows = (x @ module.w_qkv).reshape(r, 3, module.n_heads_kernel, module.d_head) + qkv_all = _gather_rows(qkv_rows, axis_name=axis_name) + query = qkv_rows[:, 0] + key = qkv_all[:, 1] + value = qkv_all[:, 2] + bias = _pair_mlp_tiled(module.bias_mlp, edge_rows, tile_size=tile_size) + bias = bias / jnp.sqrt(jnp.asarray(module.d_head, dtype=bias.dtype)) + row_indices = _global_row_indices(axis_name=axis_name, local_size=r) + bias = bias + jnp.transpose( + lca_alibi_bias( + row_indices, + jnp.arange(n, dtype=jnp.int32), + lca_fixed_slopes(module.n_heads_kernel, dtype=bias.dtype), + ), + (1, 2, 0), + ) + out = pallas_rectangular_edge_attention( + query, + key, + value, + bias, + mask, + block_k=attention_block_k, + ) + out = jax.nn.sigmoid(out[:, : module.n_heads]) * out[:, module.n_heads :] + proposal_attn = row_mask[:, None].astype(out.dtype) * ( + out.reshape(r, -1) @ module.w_o + ) + c_attn = _tree_ngpt_residual( + c_rows, + proposal_attn, + module.alpha_attn, + max_gain=module.ngpt_alpha_max, + tag_id="", + update_mask=row_mask, + ) + x_ffn = _raw_norm( + c_attn, + module.ffn_ln_scale, + None, + eps=module.ln_eps, + ) + proposal_ffn = row_mask[:, None].astype(x_ffn.dtype) * ( + fused_silu(x_ffn @ module.ffn_w1 + module.ffn_b1) @ module.ffn_w2 + + module.ffn_b2 + ) + return _tree_ngpt_residual( + c_attn, + proposal_ffn, + module.alpha_ffn, + max_gain=module.ngpt_alpha_max, + tag_id="", + update_mask=row_mask, + ) + + +def sequence_parallel_physical_leaf_local( + kernel, + contextualized_node_rows, + global_stream, + *, + axis_name: str, +): + + from hamiltonzero.compiled.tree import _project_global, compile_target_leaf_h + + leaf_g_emb = _project_global( + kernel.leaf_projection, + global_stream, + dense_tag="gladder.to_gemb", + norm_tag="gladder.gemb_ln", + ) + node_all = _gather_rows(contextualized_node_rows, axis_name=axis_name) + leaf_h = compile_target_leaf_h(kernel.leaf, node_all, leaf_g_emb) + c_rows = _linear(kernel.leaf.P_c, contextualized_node_rows) + c_rows = _tree_sphere(c_rows) + return leaf_h, c_rows + + +def sequence_parallel_reduce_physical_local( + kernel, + edge_rows, + leaf_h, + c_rows, + leaf_real, + structural_mask, + global_stream, + permutation, + *, + axis_name: str, + axis_size: int, + replicate_threshold: int = 512, + contextualizer_tile_size: int = 128, + attention_block_k: int = 128, +): + + from hamiltonzero.compiled.tree import ( + compile_merge_h, + compile_physical_tree_from_reduced_state, + ) + from hamiltonzero.compiled.types import CARRY_LEFT, CARRY_RIGHT, EMPTY, MERGE + from hamiltonzero.model.tree import ( + _tree_active_clock_depth, + _tree_depth_count_features, + edge_merge_masked, + ) + + if replicate_threshold < 1: + raise ValueError("replicate_threshold must be positive") + local_size, n = edge_rows.shape[:2] + if n != local_size * axis_size or n & (n - 1): + raise ValueError("physical sequence compiler requires power-of-two N") + if local_size & (local_size - 1): + raise ValueError("each seq lane must own a power-of-two row count") + + g = global_stream + + merge = kernel.merge + m = leaf_real.astype(c_rows.dtype) + k = structural_mask.astype(c_rows.dtype) + counts = m + n_total = jnp.sum(m) + feature_n_levels = _tree_active_clock_depth(m) + clock_depth = _tree_active_clock_depth(k) + edge_rows = _tree_sphere(edge_rows) + + early_merge_h = [] + early_opcodes = [] + width = n + rows_per_lane = local_size + level = 0 + while width > int(replicate_threshold): + if rows_per_lane < 2: + raise ValueError( + "replicate_threshold is too small for the available seq lanes" + ) + c_all = _gather_rows(c_rows, axis_name=axis_name) + c_a_rows, c_b_rows = c_rows[0::2], c_rows[1::2] + m_a, m_b = m[0::2], m[1::2] + k_a, k_b = k[0::2], k[1::2] + cnt_a, cnt_b = counts[0::2], counts[1::2] + m_rows = _local_rows(m, axis_name=axis_name, local_size=rows_per_lane) + k_rows = _local_rows(k, axis_name=axis_name, local_size=rows_per_lane) + m_a_rows, m_b_rows = m_rows[0::2], m_rows[1::2] + k_a_rows, k_b_rows = k_rows[0::2], k_rows[1::2] + both_struct = k_a * k_b + both_struct_rows = _local_rows( + both_struct, + axis_name=axis_name, + local_size=rows_per_lane // 2, + ) + pair_base = jnp.maximum( + jnp.sum((k_a + k_b - k_a * k_b).astype(jnp.int32)), + jnp.asarray(2, dtype=jnp.int32), + ) + depth = _tree_depth_count_features( + cnt_a, + cnt_b, + n_total, + level, + feature_n_levels, + c_rows.dtype, + ) + depth_rows = _local_rows( + depth, + axis_name=axis_name, + local_size=rows_per_lane // 2, + ) + + edge_blocks = edge_rows.reshape( + rows_per_lane // 2, + 2, + width // 2, + 2, + edge_rows.shape[-1], + ) + local_parent = jnp.arange(rows_per_lane // 2, dtype=jnp.int32) + global_parent = ( + jax.lax.axis_index(axis_name) * (rows_per_lane // 2) + local_parent + ) + sibling_lr = edge_blocks[local_parent, 0, global_parent, 1] + sibling_rl = edge_blocks[local_parent, 1, global_parent, 0] + level_active = jnp.any(both_struct.astype(bool)) + g_level = g @ kernel.tree_projection_weight + kernel.tree_projection_bias + + def candidate_one(ca, cb, elr, erl, dep, pidx, active): + return merge.context_candidate( + ca, + cb, + g_level, + sibling_edge_lr=elr, + sibling_edge_rl=erl, + level_idx=jnp.int32(level), + pair_idx=pidx, + pair_base=pair_base, + clock_depth=clock_depth, + depth_feats=dep, + kfac_structural_mask=active, + kfac_g_structural_mask=level_active, + kfac_scan_shared=False, + ) + + candidate = jax.vmap(candidate_one)( + c_a_rows, + c_b_rows, + sibling_lr, + sibling_rl, + depth_rows, + global_parent, + both_struct_rows, + ) + early_merge_h.append( + _gather_rows( + compile_merge_h( + merge, + candidate, + depth_rows, + ), + axis_name=axis_name, + ) + ) + early_opcodes.append( + jnp.where( + m_a.astype(bool), + jnp.where(m_b.astype(bool), MERGE, CARRY_LEFT), + jnp.where(m_b.astype(bool), CARRY_RIGHT, EMPTY), + ).astype(jnp.uint8) + ) + + gate_a_mask, gate_b_mask = k_a_rows, k_b_rows + gate_both = gate_a_mask * gate_b_mask + c_rows = ( + gate_both[:, None] * candidate + + (gate_a_mask * (1.0 - gate_b_mask))[:, None] * c_a_rows + + ((1.0 - gate_a_mask) * gate_b_mask)[:, None] * c_b_rows + ) + + c_a_all, c_b_all = c_all[0::2], c_all[1::2] + + parent_width = width // 2 + edge_tile_width = gcd( + parent_width, + min(parent_width, int(contextualizer_tile_size)), + ) + edge_tile_count = parent_width // edge_tile_width + edge_new0 = jnp.zeros( + (rows_per_lane // 2, parent_width, edge_rows.shape[-1]), + dtype=edge_rows.dtype, + ) + + def merge_destination_tile(tile_index, edge_output): + start = tile_index * edge_tile_width + + def take_columns(x): + return jax.lax.dynamic_slice_in_dim(x, start, edge_tile_width, axis=0) + + m_a_tile, m_b_tile = take_columns(m_a), take_columns(m_b) + k_a_tile, k_b_tile = take_columns(k_a), take_columns(k_b) + c_a_tile, c_b_tile = take_columns(c_a_all), take_columns(c_b_all) + + edge_block_tile = jax.lax.dynamic_slice_in_dim( + edge_blocks, start, edge_tile_width, axis=2 + ) + + def edge_row(e0, e1, e2, e3, ma, mb, ka, kb, ca, cb): + return jax.vmap( + lambda x0, x1, x2, x3, mqa, mqb, kqa, kqb, cqa, cqb: ( + edge_merge_masked( + x0, + x1, + x2, + x3, + ma, + mb, + mqa, + mqb, + ca, + cb, + cqa, + cqb, + merge.edge_merge, + k_2i=ka, + k_2i1=kb, + k_2j=kqa, + k_2j1=kqb, + kfac_scan_shared=False, + )[0] + ) + )( + e0, + e1, + e2, + e3, + m_a_tile, + m_b_tile, + k_a_tile, + k_b_tile, + c_a_tile, + c_b_tile, + ) + + edge_tile = jax.vmap(edge_row)( + edge_block_tile[:, 0, :, 0], + edge_block_tile[:, 0, :, 1], + edge_block_tile[:, 1, :, 0], + edge_block_tile[:, 1, :, 1], + m_a_rows, + m_b_rows, + k_a_rows, + k_b_rows, + c_a_rows, + c_b_rows, + ) + return jax.lax.dynamic_update_slice_in_dim( + edge_output, edge_tile, start, axis=1 + ) + + edge_new = jax.lax.fori_loop( + 0, edge_tile_count, merge_destination_tile, edge_new0 + ) + c_all_new = _gather_rows(c_rows, axis_name=axis_name) + edge_new = _sequence_tree_fwl( + merge.tree_edge_fwl, + edge_new, + c_rows, + c_all_new, + both_struct, + both_struct_rows, + axis_name=axis_name, + axis_size=axis_size, + tile_size=contextualizer_tile_size, + ) + edge_keep = both_struct_rows[:, None] * both_struct[None, :] + + def finalize_edge_tile(tile_index, current_edges): + start = tile_index * edge_tile_width + current_tile = jax.lax.dynamic_slice_in_dim( + current_edges, start, edge_tile_width, axis=1 + ) + fallback_tile = jax.lax.dynamic_slice_in_dim( + edge_blocks, start, edge_tile_width, axis=2 + )[:, 0, :, 0] + keep_tile = jax.lax.dynamic_slice_in_dim( + edge_keep, start, edge_tile_width, axis=1 + )[..., None].astype(bool) + current_tile = _tree_sphere(current_tile) + current_tile = jnp.where(keep_tile, current_tile, fallback_tile) + return jax.lax.dynamic_update_slice_in_dim( + current_edges, current_tile, start, axis=1 + ) + + edge_rows = jax.lax.fori_loop(0, edge_tile_count, finalize_edge_tile, edge_new) + c_skip = c_rows + c_rows = _sequence_level_edge_attention( + merge.level_edge_attn, + c_rows, + edge_rows, + both_struct, + both_struct_rows, + axis_name=axis_name, + level=level, + tile_size=contextualizer_tile_size, + attention_block_k=attention_block_k, + ) + c_rows = jnp.where( + both_struct_rows[:, None].astype(bool), + _tree_sphere(c_rows), + c_skip, + ) + + m = m_a + m_b - m_a * m_b + k = k_a + k_b - k_a * k_b + counts = cnt_a + cnt_b + level_mask = k + c_all_new = _gather_rows(c_rows, axis_name=axis_name) + updated_g = _g_update( + kernel.tree_update, + g, + _g_descriptor_pool(kernel.tree_pool, g, c_all_new, level_mask), + ) + g = jnp.where(level_active, updated_g, g) + width //= 2 + rows_per_lane //= 2 + level += 1 + + c_reduced = _gather_rows(c_rows, axis_name=axis_name) + edge_reduced = _gather_rows(edge_rows, axis_name=axis_name) + return compile_physical_tree_from_reduced_state( + kernel, + perm=permutation, + leaf_real=leaf_real, + leaf_h=leaf_h, + c_reduced=c_reduced, + edge_reduced=edge_reduced, + real_reduced=m, + structural_reduced=k, + counts_reduced=counts, + g_reduced=g, + early_merge_h=tuple(early_merge_h), + early_opcodes=tuple(early_opcodes), + full_structural_mask=structural_mask, + ) + + +def sequence_parallel_shared_trunk_local( + kernel: TrunkCompilerKernel, + j_double_prime_rows: jax.Array, + h_prime: jax.Array, + real_mask: jax.Array, + balanced_mask: jax.Array, + *, + axis_name: str, + axis_size: int, + featurizer_tile_size: int = 128, + attention_block_k: int = 128, +) -> SharedTrunk: + + if j_double_prime_rows.ndim != 3 or j_double_prime_rows.shape[-1] != 10: + raise ValueError("J rows must have shape [N/P,N,10]") + local_size, n = j_double_prime_rows.shape[:2] + if n != local_size * axis_size: + raise ValueError( + f"sequence rows must evenly tile N: local={local_size}, " + f"N={n}, shards={axis_size}" + ) + if h_prime.shape != (n, 3): + raise ValueError(f"h_prime must have shape {(n, 3)}, got {h_prime.shape}") + if real_mask.shape != (n,) or balanced_mask.shape != (n,): + raise ValueError("real and balanced masks must both have global shape [N]") + if j_double_prime_rows.dtype != jnp.float32: + raise TypeError("fresh sequence compiler currently requires fp32 graph inputs") + + row_indices = _global_row_indices(axis_name=axis_name, local_size=local_size) + row_mask = real_mask[row_indices] + featurizer = kernel.featurizer + bond_rows, descriptor_rows = featurizer.eval_embed_local_rows( + j_double_prime_rows, + row_mask, + real_mask, + tile_size=featurizer_tile_size, + ) + descriptor_all = _gather_rows(descriptor_rows, axis_name=axis_name) + + sum_j2, count_j = featurizer.eval_jh_stats_rows( + j_double_prime_rows, + row_mask, + real_mask, + row_indices=row_indices, + ) + jh_stats = ( + jax.lax.psum(sum_j2, axis_name), + jax.lax.psum(count_j, axis_name), + ) + + local_rows, global_raw = featurizer.eval_finalize_local_rows( + J_double_prime_rows=j_double_prime_rows, + local_desc_rows=descriptor_rows, + local_desc_all=descriptor_all, + row_indices=row_indices, + mask=real_mask, + h_prime=h_prime, + jh_stats=jh_stats, + ) + local_all = _gather_rows(local_rows, axis_name=axis_name) + edge_rows = featurizer.eval_edge_rows( + bond_emb_rows=bond_rows, + local_rows=local_rows, + local_final_all=local_all, + global_feat=global_raw, + row_indices=row_indices, + mask=real_mask, + tile_size=featurizer_tile_size, + ) + + g = _tree_sphere(global_raw.astype(local_rows.dtype)) + local_rows, edge_rows, g = _sequence_trunk_local( + kernel.trunk, + local_rows, + edge_rows, + g, + real_mask, + row_mask, + axis_name=axis_name, + axis_size=axis_size, + attention_block_k=attention_block_k, + ) + global_stream = _edge_row_col_global_update( + kernel.shared_global, + g, + edge_rows, + real_mask, + row_mask, + axis_name=axis_name, + ) + return SharedTrunk( + node_raw=local_rows, + edge_raw=edge_rows, + global_raw=global_raw, + global_stream=global_stream, + real_mask=real_mask, + balanced_mask=balanced_mask, + ) + + +def build_sequence_parallel_shared_trunk( + *, + mesh: Mesh, + kernel_template: TrunkCompilerKernel, + axis_name: str = "seq", + featurizer_tile_size: int = 128, + attention_block_k: int = 128, +) -> Callable[ + [TrunkCompilerKernel, jax.Array, jax.Array, jax.Array, jax.Array], + SharedTrunk, +]: + + if tuple(mesh.axis_names) != (axis_name,): + raise ValueError( + f"sequence trunk requires a one-dimensional {axis_name!r} mesh; " + f"got {mesh.axis_names}" + ) + axis_size = int(mesh.shape[axis_name]) + replicated_spec = P() + edge_spec = P(axis_name, None, None) + kernel_specs = jax.tree_util.tree_map(lambda _: replicated_spec, kernel_template) + output_specs = SharedTrunk( + node_raw=P(axis_name, None), + edge_raw=edge_spec, + global_raw=replicated_spec, + global_stream=replicated_spec, + real_mask=replicated_spec, + balanced_mask=replicated_spec, + ) + + local = partial( + sequence_parallel_shared_trunk_local, + axis_name=axis_name, + axis_size=axis_size, + featurizer_tile_size=featurizer_tile_size, + attention_block_k=attention_block_k, + ) + mapped = jax.shard_map( + local, + mesh=mesh, + in_specs=( + kernel_specs, + edge_spec, + replicated_spec, + replicated_spec, + replicated_spec, + ), + out_specs=output_specs, + check_vma=False, + ) + + replicated = NamedSharding(mesh, replicated_spec) + edge_sharding = NamedSharding(mesh, edge_spec) + output_shardings = SharedTrunk( + node_raw=NamedSharding(mesh, P(axis_name, None)), + edge_raw=edge_sharding, + global_raw=replicated, + global_stream=replicated, + real_mask=replicated, + balanced_mask=replicated, + ) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _: replicated, kernel_template), + edge_sharding, + replicated, + replicated, + replicated, + ), + out_shardings=output_shardings, + ) + + +def build_sequence_parallel_contextualizer( + *, + mesh: Mesh, + contextualizer_template, + g_template: jax.Array, + axis_name: str = "seq", + tile_size: int = 128, + attention_block_k: int = 128, +) -> Callable: + + if tuple(mesh.axis_names) != (axis_name,): + raise ValueError( + f"contextualizer requires a one-dimensional {axis_name!r} mesh" + ) + axis_size = int(mesh.shape[axis_name]) + rep_spec = P() + node_spec = P(axis_name, None) + edge_spec = P(axis_name, None, None) + context_specs = jax.tree_util.tree_map(lambda _: rep_spec, contextualizer_template) + local = partial( + sequence_parallel_contextualizer_local, + axis_name=axis_name, + axis_size=axis_size, + tile_size=tile_size, + attention_block_k=attention_block_k, + ) + mapped = jax.shard_map( + local, + mesh=mesh, + in_specs=( + context_specs, + node_spec, + edge_spec, + rep_spec, + rep_spec, + rep_spec, + ), + out_specs=(node_spec, edge_spec, rep_spec), + check_vma=False, + ) + rep = NamedSharding(mesh, rep_spec) + node_sharding = NamedSharding(mesh, node_spec) + edge_sharding = NamedSharding(mesh, edge_spec) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _: rep, contextualizer_template), + node_sharding, + edge_sharding, + rep, + rep, + rep, + ), + out_shardings=(node_sharding, edge_sharding, rep), + ) + + +def build_sequence_pair_permute(*, mesh: Mesh, axis_name: str = "seq") -> Callable: + + if tuple(mesh.axis_names) != (axis_name,): + raise ValueError( + f"pair permutation requires a one-dimensional {axis_name!r} mesh" + ) + axis_size = int(mesh.shape[axis_name]) + node_spec = P(axis_name, None) + edge_spec = P(axis_name, None, None) + + def local(node_rows, edge_rows, permutation): + return ( + ring_permute_rows_local( + node_rows, + permutation, + axis_name=axis_name, + axis_size=axis_size, + ), + permute_pair_rows_and_columns_local( + edge_rows, + permutation, + axis_name=axis_name, + axis_size=axis_size, + ), + ) + + return jax.jit( + jax.shard_map( + local, + mesh=mesh, + in_specs=(node_spec, edge_spec, P()), + out_specs=(node_spec, edge_spec), + check_vma=False, + ) + ) + + +def build_sequence_parallel_edge_global_update( + *, + mesh: Mesh, + module_template, + axis_name: str = "seq", + tile_size: int = 128, +) -> Callable: + + if tuple(mesh.axis_names) != (axis_name,): + raise ValueError( + f"global edge update requires a one-dimensional {axis_name!r} mesh" + ) + rep_spec = P() + edge_spec = P(axis_name, None, None) + mapped = jax.shard_map( + partial( + sequence_parallel_edge_global_update_local, + axis_name=axis_name, + tile_size=tile_size, + ), + mesh=mesh, + in_specs=( + jax.tree_util.tree_map(lambda _: rep_spec, module_template), + rep_spec, + edge_spec, + rep_spec, + ), + out_specs=rep_spec, + check_vma=False, + ) + rep = NamedSharding(mesh, rep_spec) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _: rep, module_template), + rep, + NamedSharding(mesh, edge_spec), + rep, + ), + out_shardings=rep, + ) + + +def build_sequence_parallel_physical_leaf( + *, + mesh: Mesh, + kernel_template, + axis_name: str = "seq", +) -> Callable: + + if tuple(mesh.axis_names) != (axis_name,): + raise ValueError( + f"physical leaf projection requires a one-dimensional {axis_name!r} mesh" + ) + rep_spec = P() + node_spec = P(axis_name, None) + kernel_specs = jax.tree_util.tree_map(lambda _: rep_spec, kernel_template) + mapped = jax.shard_map( + partial(sequence_parallel_physical_leaf_local, axis_name=axis_name), + mesh=mesh, + in_specs=(kernel_specs, node_spec, rep_spec), + out_specs=((rep_spec,), node_spec), + check_vma=False, + ) + rep = NamedSharding(mesh, rep_spec) + node_sharding = NamedSharding(mesh, node_spec) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _: rep, kernel_template), + node_sharding, + rep, + ), + out_shardings=((rep,), node_sharding), + ) + + +def build_sequence_parallel_physical_reducer( + *, + mesh: Mesh, + kernel_template, + edge_template, + axis_name: str = "seq", + replicate_threshold: int = 512, + contextualizer_tile_size: int = 128, + attention_block_k: int = 128, +) -> Callable: + + from hamiltonzero.compiled.types import CompiledTree + + if tuple(mesh.axis_names) != (axis_name,): + raise ValueError( + f"physical reducer requires a one-dimensional {axis_name!r} mesh" + ) + n = int(edge_template.shape[0]) + if tuple(edge_template.shape[:2]) != (n, n): + raise ValueError("edge_template must have square global pair axes") + axis_size = int(mesh.shape[axis_name]) + if n % axis_size or n & (n - 1): + raise ValueError("global N must be a power of two divisible by seq lanes") + rep_spec = P() + node_spec = P(axis_name, None) + edge_spec = P(axis_name, None, None) + kernel_specs = jax.tree_util.tree_map(lambda _: rep_spec, kernel_template) + n_levels = n.bit_length() - 1 + output_specs = CompiledTree( + perm=rep_spec, + inv_perm=rep_spec, + leaf_real=rep_spec, + leaf_h=(rep_spec,), + leaf_combiner_h=(), + merge_h=(rep_spec,) * n_levels, + opcodes=(rep_spec,) * n_levels, + readout_h=(rep_spec,), + readout_combiner_h=(), + ) + local = partial( + sequence_parallel_reduce_physical_local, + axis_name=axis_name, + axis_size=axis_size, + replicate_threshold=replicate_threshold, + contextualizer_tile_size=contextualizer_tile_size, + attention_block_k=attention_block_k, + ) + mapped = jax.shard_map( + local, + mesh=mesh, + in_specs=( + kernel_specs, + edge_spec, + (rep_spec,), + node_spec, + rep_spec, + rep_spec, + rep_spec, + rep_spec, + ), + out_specs=output_specs, + check_vma=False, + ) + rep = NamedSharding(mesh, rep_spec) + node_sharding = NamedSharding(mesh, node_spec) + edge_sharding = NamedSharding(mesh, edge_spec) + output_shardings = jax.tree_util.tree_map(lambda _: rep, output_specs) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _: rep, kernel_template), + edge_sharding, + (rep,), + node_sharding, + rep, + rep, + rep, + rep, + ), + out_shardings=output_shardings, + ) + + +__all__ = [ + "build_sequence_pair_permute", + "build_sequence_parallel_contextualizer", + "build_sequence_parallel_edge_global_update", + "build_sequence_parallel_physical_leaf", + "build_sequence_parallel_physical_reducer", + "build_sequence_parallel_shared_trunk", + "permute_pair_rows_and_columns_local", + "ring_permute_rows_local", + "sequence_parallel_contextualizer_local", + "sequence_parallel_edge_global_update_local", + "sequence_parallel_physical_leaf_local", + "sequence_parallel_reduce_physical_local", + "sequence_parallel_shared_trunk_local", + "transpose_pair_rows_local", +] diff --git a/src/hamiltonzero/evaluation/statistics.py b/src/hamiltonzero/evaluation/statistics.py new file mode 100644 index 0000000000000000000000000000000000000000..3bdcb47b03e8cb2cc0b0cd8212b1f5f44ec8188b --- /dev/null +++ b/src/hamiltonzero/evaluation/statistics.py @@ -0,0 +1,223 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import math +from dataclasses import dataclass + +import numpy as np + + +P01_TWO_SIDED = 2.576 + + +def standard_error_from_tailstd( + tailstd: float, + *, + n_systems: int, + batch_size: int, +) -> float: + if tailstd is None or not math.isfinite(float(tailstd)): + return float("nan") + denominator = max(1, int(n_systems) * int(batch_size)) + return float(tailstd) / math.sqrt(denominator) + + +def welch_band_walkermean( + tailstd_a: float, + tailstd_b: float, + *, + n_systems_a: int, + batch_size_a: int, + n_systems_b: int, + batch_size_b: int, + z: float = P01_TWO_SIDED, +) -> float: + se_a = standard_error_from_tailstd( + tailstd_a, + n_systems=n_systems_a, + batch_size=batch_size_a, + ) + se_b = standard_error_from_tailstd( + tailstd_b, + n_systems=n_systems_b, + batch_size=batch_size_b, + ) + if not (math.isfinite(se_a) and math.isfinite(se_b)): + return float("nan") + return float(z) * math.sqrt(se_a * se_a + se_b * se_b) + + +def select_winner_per_physical( + R, + tailstd, + beam_logp, + batch_per_candidate, + *, + z: float = P01_TWO_SIDED, + ucb_z: float = 2.0, +): + R = np.asarray(R, dtype=float) + tailstd = np.asarray(tailstd, dtype=float) + beam_logp = np.asarray(beam_logp, dtype=float) + P, K = R.shape + Bp = int(batch_per_candidate) + winner_idx = np.zeros(P, dtype=np.int64) + tie_mask = np.zeros((P, K), dtype=bool) + bands = np.full((P, K), np.nan, dtype=float) + se = np.array( + [ + [ + standard_error_from_tailstd( + tailstd[p, k], + n_systems=1, + batch_size=Bp, + ) + for k in range(K) + ] + for p in range(P) + ], + dtype=float, + ) + reason: list[str] = [] + for p in range(P): + Rp, sp, lp = R[p], tailstd[p], beam_logp[p] + finite = np.isfinite(Rp) + if not finite.any(): + w = int(np.argmax(lp)) + winner_idx[p] = w + tie_mask[p, w] = True + reason.append("energy_unavailable") + continue + istar = int(np.argmin(np.where(finite, Rp, np.inf))) + for k in range(K): + band = welch_band_walkermean( + sp[k], + sp[istar], + n_systems_a=1, + batch_size_a=Bp, + n_systems_b=1, + batch_size_b=Bp, + z=z, + ) + bands[p, k] = band + if k == istar: + tie_mask[p, k] = True + elif finite[k] and math.isfinite(band) and abs(Rp[k] - Rp[istar]) < band: + tie_mask[p, k] = True + tie_idx = np.where(tie_mask[p])[0] + se_tie = se[p, tie_idx] + ucb = Rp[tie_idx] + ucb_z * np.where(np.isfinite(se_tie), se_tie, np.inf) + w = int(tie_idx[int(np.argmin(ucb))]) if np.any(np.isfinite(ucb)) else istar + winner_idx[p] = w + reason.append("argmin_energy" if w == istar else "tie_broken_by_ucb") + return winner_idx, tie_mask, reason, bands, se + + +@dataclass(frozen=True, slots=True) +class ChannelMetrics: + mean: float + standard_error: float + walker_tail_std: float + local_energy_std: float + lag1_autocorrelation: float | None + + def as_dict(self) -> dict[str, float | None]: + return { + "mean": self.mean, + "standard_error": self.standard_error, + "walker_tail_std": self.walker_tail_std, + "local_energy_std": self.local_energy_std, + "lag1_autocorrelation": self.lag1_autocorrelation, + } + + +class EnergyWindow: + def __init__(self, steps: int, systems: int, batch_size: int): + self.steps = int(steps) + self.systems = int(systems) + self.batch_size = int(batch_size) + shape = (self.steps, self.systems, self.batch_size) + self.values = { + "total": np.empty(shape, dtype=np.float32), + "exchange": np.empty(shape, dtype=np.float32), + "field": np.empty(shape, dtype=np.float32), + } + self.sums = { + "total": np.zeros((self.systems, self.batch_size), dtype=np.float32), + "exchange": np.zeros((self.systems, self.batch_size), dtype=np.float32), + "field": np.zeros((self.systems, self.batch_size), dtype=np.float32), + } + self.count = 0 + + def push(self, total, exchange, field) -> None: + if self.count >= self.steps: + raise IndexError("energy window is full") + for name, value in ( + ("total", total), + ("exchange", exchange), + ("field", field), + ): + array = np.asarray(value).real.astype(np.float32, copy=False) + expected = (self.systems, self.batch_size) + if array.shape != expected: + raise ValueError( + f"{name} energy must have shape {expected}, got {array.shape}" + ) + self.values[name][self.count] = array + self.sums[name] += array + self.count += 1 + + def tail_means(self, channel: str) -> np.ndarray: + if self.count < 1: + raise RuntimeError("energy window is empty") + return (self.sums[channel] / float(self.count)).astype(np.float32, copy=False) + + def tail_mean(self, channel: str, system: int) -> float: + return float(np.mean(self.tail_means(channel)[system])) + + def tail_std(self, channel: str, system: int) -> float: + per_walker = np.asarray(self.tail_means(channel)[system]).ravel() + if per_walker.size < 2: + return float("nan") + return float(np.std(per_walker, ddof=1)) + + def metrics(self, channel: str) -> ChannelMetrics: + if self.systems != 1: + raise ValueError("bare eval metrics require one physical system") + samples = self.values[channel][: self.count, 0] + per_walker = self.tail_means(channel)[0] + mean = float(np.mean(per_walker)) + tailstd = ( + float(np.std(per_walker, ddof=1)) if per_walker.size >= 2 else float("nan") + ) + se = standard_error_from_tailstd( + tailstd, + n_systems=1, + batch_size=self.batch_size, + ) + step_std = np.std(samples, axis=1) + gap_sq = float(np.nanmean(step_std**2)) + local_std = ( + math.sqrt(gap_sq) + if math.isfinite(gap_sq) and gap_sq >= 0.0 + else float("nan") + ) + step_means = np.mean(samples, axis=1) + finite = step_means[np.isfinite(step_means)] + autocorrelation = None + if len(finite) >= 20: + centered = finite - np.mean(finite) + denominator = float(np.sum(centered * centered)) + if denominator > 0.0: + autocorrelation = float( + np.sum(centered[:-1] * centered[1:]) / denominator + ) + return ChannelMetrics( + mean=mean, + standard_error=se, + walker_tail_std=tailstd, + local_energy_std=local_std, + lag1_autocorrelation=autocorrelation, + ) diff --git a/src/hamiltonzero/evaluation/types.py b/src/hamiltonzero/evaluation/types.py new file mode 100644 index 0000000000000000000000000000000000000000..50fa7725cb4aa96eedc1c66f63542b65788a50af --- /dev/null +++ b/src/hamiltonzero/evaluation/types.py @@ -0,0 +1,79 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Literal + +from .statistics import ChannelMetrics + + +EvalPath = Literal["ordinary", "contest", "large_n"] + + +@dataclass(frozen=True, slots=True) +class ContestCandidate: + index: int + route_log_probability: float + energy: float + standard_error: float + walker_tail_std: float + in_tie_set: bool + + def as_dict(self) -> dict[str, int | float | bool]: + return { + "index": self.index, + "route_log_probability": self.route_log_probability, + "energy": self.energy, + "standard_error": self.standard_error, + "walker_tail_std": self.walker_tail_std, + "in_tie_set": self.in_tie_set, + } + + +@dataclass(frozen=True, slots=True) +class ContestResult: + winner: int + reason: str + candidates: tuple[ContestCandidate, ...] + + def as_dict(self) -> dict: + return { + "winner": self.winner, + "reason": self.reason, + "candidates": [candidate.as_dict() for candidate in self.candidates], + } + + +@dataclass(frozen=True, slots=True) +class EvalMetric: + step: int + energy: float + energy_std: float + step_walltime: float + walltime: float + + +@dataclass(frozen=True, slots=True) +class EvalResult: + path: EvalPath + route: tuple[int, ...] + route_log_probability: float | None + measurements: int + walltime_seconds: float + energy: ChannelMetrics + channels: dict[str, ChannelMetrics] + contest: ContestResult | None = None + + def as_dict(self) -> dict: + result = { + "measurements": self.measurements, + "walltime_seconds": self.walltime_seconds, + "energy": self.energy.mean, + "energy_std": self.energy.local_energy_std, + "channels": {name: metrics.mean for name, metrics in self.channels.items()}, + } + if self.energy.lag1_autocorrelation is not None: + result["energy_lag1_autocorrelation"] = self.energy.lag1_autocorrelation + return result diff --git a/src/hamiltonzero/hamiltonian.py b/src/hamiltonzero/hamiltonian.py new file mode 100644 index 0000000000000000000000000000000000000000..cbd129a1bfc9c87eb33dea4b2fad5204e8d7cdd9 --- /dev/null +++ b/src/hamiltonzero/hamiltonian.py @@ -0,0 +1,261 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Hashable, Iterable, Sequence + +import networkx as nx +import numpy as np + + +_PHI = 0.5 * (1.0 + np.sqrt(5.0)) + + +def _exchange_matrix(value: Any) -> np.ndarray: + array = np.asarray(value, dtype=np.float32) + if array.ndim == 0: + return np.eye(3, dtype=np.float32) * array + if array.shape == (3,): + return np.diag(array) + if array.shape == (3, 3): + return array + raise ValueError("exchange must be a scalar, length-3 vector, or 3x3 matrix") + + +def _field_vector(value: Any) -> np.ndarray: + array = np.asarray(value, dtype=np.float32) + if array.ndim == 0: + return np.asarray((0.0, 0.0, float(array)), dtype=np.float32) + if array.shape == (3,): + return array + raise ValueError("field must be a scalar z-field or length-3 vector") + + +def _pair_alpha(exchange: np.ndarray) -> float: + if not np.any(exchange): + return 0.0 + left, singular_values, right_t = np.linalg.svd(exchange) + operator_norm = float(singular_values[0]) + nuclear_norm = float(singular_values.sum()) + trace_component = float(np.trace(exchange) / 3.0) + trace_residual = exchange - trace_component * np.eye(3, dtype=exchange.dtype) + trace_singular_values = np.linalg.svd(trace_residual, compute_uv=False) + rotation = left @ right_t + polar_component = float(np.trace(rotation.T @ exchange) / 3.0) + polar_residual = exchange - polar_component * rotation + polar_singular_values = np.linalg.svd(polar_residual, compute_uv=False) + full_bound = min( + (_PHI / 2.0) * operator_norm, + 0.5 * nuclear_norm, + 0.75 * abs(trace_component) + (_PHI / 2.0) * float(trace_singular_values[0]), + 0.75 * abs(trace_component) + 0.5 * float(trace_singular_values.sum()), + 0.75 * abs(polar_component) + (_PHI / 2.0) * float(polar_singular_values[0]), + 0.75 * abs(polar_component) + 0.5 * float(polar_singular_values.sum()), + ) + symmetric = 0.5 * (exchange + exchange.T) + antisymmetric = exchange - symmetric + scale = max(float(np.linalg.norm(exchange)), 1e-9) + if float(np.linalg.norm(antisymmetric)) > 1e-9 * scale: + return full_bound + eigenvalues = np.linalg.eigvalsh(symmetric) + if not (np.all(eigenvalues >= -1e-12) or np.all(eigenvalues <= 1e-12)): + return full_bound + return min(full_bound, 0.5 * operator_norm) + + +def _mu_safe(exchange: np.ndarray, field: np.ndarray) -> float: + row_bound = np.zeros(exchange.shape[0], dtype=np.float64) + for left in range(exchange.shape[0]): + for right in range(left + 1, exchange.shape[0]): + pair_bound = _pair_alpha(exchange[left, right]) + row_bound[left] += pair_bound + row_bound[right] += pair_bound + field_bound = np.linalg.norm(field, axis=-1) + return float(np.max(row_bound + (2.0 / 3.0) * field_bound, initial=0.0)) + + +@dataclass(frozen=True, slots=True) +class SpinHamiltonian: + _coupling: np.ndarray + _J: np.ndarray + _h: np.ndarray + nodes: tuple[Hashable, ...] + _mu: float | None = None + _needs_fwl2: bool | None = None + _category: str | None = None + _tag: str | None = None + _topology_class: str | None = None + _j_class: str | None = None + + def __post_init__(self) -> None: + coupling = np.asarray(self._coupling, dtype=np.float32) + exchange = np.asarray(self._J, dtype=np.float32) + field = np.asarray(self._h, dtype=np.float32) + n_spins = len(self.nodes) + if coupling.shape != (n_spins, n_spins): + raise ValueError("coupling must have shape [N,N]") + if exchange.shape != (n_spins, n_spins, 3, 3): + raise ValueError("J must have shape [N,N,3,3]") + if field.shape != (n_spins, 3): + raise ValueError("h must have shape [N,3]") + if not np.array_equal(coupling, coupling.T): + raise ValueError("coupling must be symmetric") + if np.any(np.diag(coupling)) or np.any( + exchange[np.arange(n_spins), np.arange(n_spins)] + ): + raise ValueError("self-couplings are not supported") + if not np.allclose(exchange, exchange.transpose(1, 0, 3, 2), atol=0.0): + raise ValueError("J[j,i] must equal transpose(J[i,j])") + mu = None if self._mu is None else float(self._mu) + if mu is not None and (not np.isfinite(mu) or mu < 0.0): + raise ValueError("mu must be finite and non-negative") + object.__setattr__(self, "_coupling", coupling) + object.__setattr__(self, "_J", exchange) + object.__setattr__(self, "_h", field) + object.__setattr__(self, "_mu", mu) + object.__setattr__( + self, + "_needs_fwl2", + None if self._needs_fwl2 is None else bool(self._needs_fwl2), + ) + + @classmethod + def from_networkx( + cls, + graph: nx.Graph, + *, + J: Any = 1.0, + h: Any = 0.0, + edge_attribute: str = "J", + node_attribute: str = "h", + nodes: Iterable[Hashable] | None = None, + mu: float | None = None, + ) -> "SpinHamiltonian": + if graph.is_directed() or graph.is_multigraph(): + raise TypeError("expected a simple undirected NetworkX graph") + order = tuple(graph.nodes if nodes is None else nodes) + if len(order) != graph.number_of_nodes() or set(order) != set(graph.nodes): + raise ValueError("nodes must contain every graph node exactly once") + indices = {node: index for index, node in enumerate(order)} + n_spins = len(order) + coupling = np.zeros((n_spins, n_spins), dtype=np.float32) + exchange = np.zeros((n_spins, n_spins, 3, 3), dtype=np.float32) + field = np.zeros((n_spins, 3), dtype=np.float32) + for node, index in indices.items(): + public_field = graph.nodes[node].get(node_attribute, h) + field[index] = -_field_vector(public_field) + for left_node, right_node, attributes in graph.edges(data=True): + left = indices[left_node] + right = indices[right_node] + public_exchange = _exchange_matrix(attributes.get(edge_attribute, J)) + coupling[left, right] = coupling[right, left] = 1.0 + exchange[left, right] = -0.5 * public_exchange + exchange[right, left] = -0.5 * public_exchange.T + return cls( + coupling, + exchange, + field, + order, + mu, + graph.graph.get("needs_fwl2"), + graph.graph.get("category"), + graph.graph.get("tag"), + graph.graph.get("topology_class"), + graph.graph.get("j_class"), + ) + + @classmethod + def from_arrays( + cls, + J: Any, + h: Any | None = None, + *, + coupling: Any | None = None, + nodes: Sequence[Hashable] | None = None, + mu: float | None = None, + ) -> "SpinHamiltonian": + exchange = np.asarray(J, dtype=np.float32) + if exchange.ndim == 2 and exchange.shape[0] == exchange.shape[1]: + exchange = ( + exchange[:, :, None, None] * np.eye(3, dtype=np.float32)[None, None] + ) + elif exchange.ndim == 3 and exchange.shape[-1] == 3: + promoted = np.zeros((*exchange.shape[:2], 3, 3), dtype=np.float32) + diagonal = np.arange(3) + promoted[:, :, diagonal, diagonal] = exchange + exchange = promoted + if exchange.ndim != 4 or exchange.shape[-2:] != (3, 3): + raise ValueError("J must have shape [N,N,3] or [N,N,3,3]") + n_spins = exchange.shape[0] + if exchange.shape[1] != n_spins: + raise ValueError("J site axes must be square") + public_field = np.zeros((n_spins, 3), dtype=np.float32) + if h is not None: + h_array = np.asarray(h, dtype=np.float32) + if h_array.ndim == 0: + public_field[:, 2] = h_array + elif h_array.shape == (3,): + public_field[:] = h_array + elif h_array.shape == (n_spins,): + public_field[:, 2] = h_array + elif h_array.shape == (n_spins, 3): + public_field = h_array + else: + raise ValueError("h must be scalar, [3], [N], or [N,3]") + if coupling is None: + coupling_array = np.any(exchange != 0.0, axis=(-1, -2)).astype(np.float32) + else: + coupling_array = np.asarray(coupling, dtype=np.float32) + node_order = tuple(range(n_spins)) if nodes is None else tuple(nodes) + return cls( + coupling_array, + -0.5 * exchange, + -public_field, + node_order, + mu, + ) + + @property + def n_spins(self) -> int: + return len(self.nodes) + + @property + def coupling(self) -> np.ndarray: + return self._coupling.copy() + + @property + def J(self) -> np.ndarray: + return -2.0 * self._J.copy() + + @property + def h(self) -> np.ndarray: + return -self._h.copy() + + @property + def mu(self) -> float: + return _mu_safe(self._J, self._h) if self._mu is None else self._mu + + def model_arrays(self) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + return self._coupling.copy(), self._J.copy(), self._h.copy() + + def to_dict(self) -> dict[str, Any]: + result = { + "convention": "textbook", + "nodes": list(self.nodes), + "coupling": self._coupling.tolist(), + "J": self.J.tolist(), + "h": self.h.tolist(), + "mu": self.mu, + } + if self._needs_fwl2 is not None: + result["needs_fwl2"] = self._needs_fwl2 + for name in ("category", "tag", "topology_class", "j_class"): + value = getattr(self, f"_{name}") + if value is not None: + result[name] = value + return result + + +__all__ = ["SpinHamiltonian"] diff --git a/src/hamiltonzero/inference.py b/src/hamiltonzero/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..f13a571a3df7ed61194d0326914d8b03715cac2a --- /dev/null +++ b/src/hamiltonzero/inference.py @@ -0,0 +1,290 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, NamedTuple + +import jax +import jax.numpy as jnp +import numpy as np + +from hamiltonzero.config import EnergyConfig, EvalMCMCConfig, ModelConfig +from hamiltonzero.evaluation.runtime import DefaultEvalBackend +from hamiltonzero.evaluation.runner import _compose_walker_route +from hamiltonzero.hamiltonian import SpinHamiltonian +from hamiltonzero.observables import local_spin + + +class EnergySamples(NamedTuple): + total: Any + exchange: Any + casimir: Any + field: Any + + +class CompiledOrder(NamedTuple): + leaf_to_input: Any + input_to_leaf: Any + + +@dataclass(slots=True) +class _InferenceRuntime: + backend: DefaultEvalBackend + wavefunction: Any = None + context: Any = None + + +@dataclass(frozen=True, slots=True) +class PreparedInference: + wavefunction: Any = field(repr=False) + context: Any = field(repr=False) + energy_frame: Any = field(repr=False) + route: Any + route_log_probability: float + _initial_context: Any = field(repr=False) + _walker_route: Any = field(repr=False) + _runtime: _InferenceRuntime = field(repr=False) + + +def prepare( + system: SpinHamiltonian, + checkpoint: str | Path, + model_key, + *, + model: ModelConfig = ModelConfig(), + route_temperature: float = 4.0, + eps: float = 0.1, +) -> tuple[PreparedInference, CompiledOrder]: + if not isinstance(system, SpinHamiltonian): + raise TypeError("system must be a SpinHamiltonian") + temperature = float(route_temperature) + if not np.isfinite(temperature) or temperature <= 0.0: + raise ValueError("route_temperature must be finite and positive") + backend = DefaultEvalBackend() + initial_context = backend.build_system(system, EnergyConfig(eps=float(eps))) + foundation = backend.load_model( + Path(checkpoint), + model, + model_key, + initial_context, + contextualizer_attention=None, + ) + if backend.embedded_route(foundation) is not None: + raise ValueError("prepare requires a foundation router checkpoint") + canonical = backend.canonicalize_context(initial_context) + candidates = backend.beam_candidates( + foundation, + canonical.context, + beam_width=8, + top_k=1, + temperature=temperature, + ) + permutations = jnp.asarray(candidates.permutations, dtype=jnp.int32) + if permutations.ndim != 3 or permutations.shape[:2] != (1, 1): + raise ValueError("router must return one route with shape [1, 1, N]") + route = permutations[:, 0] + walker_route = _compose_walker_route(canonical.old_inverse, route) + routed_context = backend.route_context( + canonical.context, + route, + compact_custom_lap=False, + ) + wavefunction = backend.compile_single(foundation, routed_context) + compact_context = backend.route_context( + canonical.context, + route, + compact_custom_lap=True, + ) + backend.block_until_ready(wavefunction) + route = route[0] + order = CompiledOrder( + leaf_to_input=route, + input_to_leaf=jnp.argsort(route).astype(jnp.int32), + ) + prepared = PreparedInference( + wavefunction=wavefunction, + context=compact_context, + energy_frame=compact_context.energy_frame, + route=order.leaf_to_input, + route_log_probability=float(np.asarray(candidates.log_probabilities)[0, 0]), + _initial_context=initial_context, + _walker_route=walker_route, + _runtime=_InferenceRuntime(backend), + ) + return prepared, order + + +def _mcmc_config( + *, + batch_size: int, + replicas: int, + steps: int, + burn_in: int, + walker_chunk_size: int, + burn_in_replica_steps: int = 2, + initial_sigma: float = 0.3, + initial_haar_sites: int = 1, +) -> EvalMCMCConfig: + if int(batch_size) < 1: + raise ValueError("batch_size must be positive") + if int(replicas) < 2: + raise ValueError("replicas must be at least two") + if int(steps) < 1: + raise ValueError("steps must be positive") + if int(burn_in) < 0: + raise ValueError("burn_in must be non-negative") + if int(walker_chunk_size) < 1: + raise ValueError("walker_chunk_size must be positive") + if int(burn_in_replica_steps) < 1: + raise ValueError("replica_steps must be positive") + if not np.isfinite(float(initial_sigma)) or float(initial_sigma) <= 0.0: + raise ValueError("initial_sigma must be finite and positive") + if int(initial_haar_sites) < 1: + raise ValueError("initial_haar_sites must be positive") + return EvalMCMCConfig( + batch_size=int(batch_size), + replicas=int(replicas), + steps=int(steps), + burn_in=int(burn_in), + burn_in_replica_steps=int(burn_in_replica_steps), + walker_chunk_size=int(walker_chunk_size), + initial_sigma=float(initial_sigma), + initial_haar_sites=int(initial_haar_sites), + ) + + +def burn_in( + prepared: PreparedInference, + key, + *, + batch_size: int = 256, + replicas: int = 8, + burn_in: int = 1024, + replica_steps: int = 2, + walker_chunk_size: int = 16, + initial_sigma: float = 0.3, + initial_haar_sites: int = 1, +): + config = _mcmc_config( + batch_size=batch_size, + replicas=replicas, + steps=1, + burn_in=burn_in, + walker_chunk_size=walker_chunk_size, + burn_in_replica_steps=replica_steps, + initial_sigma=initial_sigma, + initial_haar_sites=initial_haar_sites, + ) + runtime = prepared._runtime + backend = runtime.backend + state = backend.initialize_mcmc( + key, + prepared.wavefunction, + prepared._initial_context, + config, + ) + state = backend.route_mcmc(state, prepared._walker_route) + runtime.wavefunction, runtime.context, state = backend.prepare_singular( + prepared.wavefunction, + prepared.context, + state, + ) + for _ in range(config.burn_in): + state = backend.step_mcmc( + state, + runtime.wavefunction, + runtime.context, + replica_steps=config.burn_in_replica_steps, + walker_chunk_size=config.walker_chunk_size, + ) + backend.block_until_ready(backend.cold_walkers(state)) + state = backend.adapt_mcmc(state, config) + q_cold = backend.cold_walkers(state) + backend.block_until_ready(q_cold) + return state, q_cold + + +def step( + prepared: PreparedInference, + state, + *, + steps: int = 24, + walker_chunk_size: int = 16, +): + config = _mcmc_config( + batch_size=int(state.q.shape[0]), + replicas=int(state.q.shape[1]), + steps=steps, + burn_in=0, + walker_chunk_size=walker_chunk_size, + ) + runtime = prepared._runtime + backend = runtime.backend + if runtime.wavefunction is None or runtime.context is None: + raise RuntimeError("burn_in must be called before step") + state = backend.step_mcmc( + state, + runtime.wavefunction, + runtime.context, + replica_steps=config.steps, + walker_chunk_size=config.walker_chunk_size, + ) + q_cold = backend.cold_walkers(state) + backend.block_until_ready(q_cold) + state = backend.adapt_mcmc(state, config) + return state, q_cold + + +def energy( + prepared: PreparedInference, + q, + *, + chunk_size: int = 512, +): + if int(chunk_size) < 1: + raise ValueError("chunk_size must be positive") + runtime = prepared._runtime + if runtime.wavefunction is None or runtime.context is None: + raise RuntimeError("burn_in must be called before energy") + values = runtime.backend.custom_lap_energy( + runtime.wavefunction, + runtime.context, + q, + EnergyConfig(chunk_size=int(chunk_size)), + ) + runtime.backend.block_until_ready(values) + return EnergySamples(*(value[0] for value in values)) + + +def spin( + prepared: PreparedInference, + q, + *, + chunk_size: int | None = 512, +): + runtime = prepared._runtime + if runtime.wavefunction is None or runtime.context is None: + raise RuntimeError("burn_in must be called before spin") + values = local_spin( + runtime.wavefunction, + runtime.context, + q, + chunk_size=chunk_size, + ) + runtime.backend.block_until_ready(values) + return values + + +__all__ = [ + "EnergySamples", + "CompiledOrder", + "PreparedInference", + "burn_in", + "energy", + "prepare", + "spin", + "step", +] diff --git a/src/hamiltonzero/mcmc/__init__.py b/src/hamiltonzero/mcmc/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..88600607428da32b3996f280789b298c0708292b --- /dev/null +++ b/src/hamiltonzero/mcmc/__init__.py @@ -0,0 +1,18 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .runtime import ( + REState, + adapt_batched, + cold_samples, + init_batched_state, + run_batched, +) + +__all__ = [ + "REState", + "adapt_batched", + "cold_samples", + "init_batched_state", + "run_batched", +] diff --git a/src/hamiltonzero/mcmc/_monotone_cubic.py b/src/hamiltonzero/mcmc/_monotone_cubic.py new file mode 100644 index 0000000000000000000000000000000000000000..61c8091d558162b660ee05ef7e3d00f8871fa6d8 --- /dev/null +++ b/src/hamiltonzero/mcmc/_monotone_cubic.py @@ -0,0 +1,85 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +import jax.numpy as jnp +from jaxtyping import Array, Float + + +def _endpoint_one_sided( + h0: Array, + h1: Array, + m0: Array, + m1: Array, +) -> Array: + d = ((2.0 * h0 + h1) * m0 - h0 * m1) / (h0 + h1) + d = jnp.where(jnp.sign(d) != jnp.sign(m0), jnp.zeros_like(d), d) + clamp = (jnp.sign(m0) != jnp.sign(m1)) & (jnp.abs(d) > 3.0 * jnp.abs(m0)) + return jnp.where(clamp, 3.0 * m0, d) + + +def _hyman_knot_derivs( + ts: Float[Array, "n"], + ys: Float[Array, "n"], +) -> Float[Array, "n"]: + h = jnp.diff(ts) + m = jnp.diff(ys) / h + h_prev = h[:-1] + h_next = h[1:] + m_prev = m[:-1] + m_next = m[1:] + d_natural = (h_next * m_prev + h_prev * m_next) / (h_prev + h_next) + same_sign = m_prev * m_next > 0 + abs_min = jnp.minimum(jnp.abs(m_prev), jnp.abs(m_next)) + sign_m = jnp.sign(m_next) + interior = jnp.where( + same_sign, + sign_m * jnp.minimum(jnp.abs(d_natural), 3.0 * abs_min), + jnp.zeros_like(d_natural), + ) + d_left = _endpoint_one_sided(h[0], h[1], m[0], m[1]) + d_right = _endpoint_one_sided(h[-1], h[-2], m[-1], m[-2]) + return jnp.concatenate([d_left[None], interior, d_right[None]]) + + +def _hermite_eval( + ts: Float[Array, "n"], + ys: Float[Array, "n"], + ds: Float[Array, "n"], + t: Float[Array, ""], +) -> Float[Array, ""]: + n = ts.shape[0] + idx = jnp.clip(jnp.searchsorted(ts, t, side="right") - 1, 0, n - 2) + x_lo = ts[idx] + x_hi = ts[idx + 1] + y_lo = ys[idx] + y_hi = ys[idx + 1] + d_lo = ds[idx] + d_hi = ds[idx + 1] + h = x_hi - x_lo + tau = (t - x_lo) / h + h00 = 2.0 * tau**3 - 3.0 * tau**2 + 1.0 + h10 = tau**3 - 2.0 * tau**2 + tau + h01 = -2.0 * tau**3 + 3.0 * tau**2 + h11 = tau**3 - tau**2 + return y_lo * h00 + h * d_lo * h10 + y_hi * h01 + h * d_hi * h11 + + +class MonotoneCubicInterpolation(eqx.Module): + ts: Float[Array, "n"] + ys: Float[Array, "n"] + ds: Float[Array, "n"] + + def __init__( + self, + ts: Float[Array, "n"], + ys: Float[Array, "n"], + ) -> None: + self.ts = ts + self.ys = ys + self.ds = _hyman_knot_derivs(ts, ys) + + def evaluate(self, t: Float[Array, ""]) -> Float[Array, ""]: + return _hermite_eval(self.ts, self.ys, self.ds, t) diff --git a/src/hamiltonzero/mcmc/langevin.py b/src/hamiltonzero/mcmc/langevin.py new file mode 100644 index 0000000000000000000000000000000000000000..49de33aa820fdbe7b40b527644a67607132cd47d --- /dev/null +++ b/src/hamiltonzero/mcmc/langevin.py @@ -0,0 +1,138 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Callable + +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray + +from .quaternion import ( + normalize_quaternion, + quaternion_conjugate, + quaternion_exp_tangent, + quaternion_log, + quaternion_multiply, +) + + +LogProbFn = Callable[[Float[Array, "N 4"]], Float[Array, ""]] +N_WRAP = 8 + + +def _su2_left_frame( + q: Float[Array, "N 4"], +) -> Float[Array, "N 3 4"]: + q0, q1, q2, q3 = q[:, 0], q[:, 1], q[:, 2], q[:, 3] + e1 = jnp.stack([-q1, q0, q3, -q2], axis=-1) + e2 = jnp.stack([-q2, -q3, q0, q1], axis=-1) + e3 = jnp.stack([-q3, q2, -q1, q0], axis=-1) + return jnp.stack([e1, e2, e3], axis=1) + + +def _drift_lie( + q: Float[Array, "N 4"], + grad_log_p: Float[Array, "N 4"], +) -> Float[Array, "N 3"]: + e_unit = _su2_left_frame(q) + return jnp.einsum("iaα,iα->ia", e_unit, grad_log_p) + + +def _log_jacobian_su2( + norm: Float[Array, "..."], +) -> Float[Array, "..."]: + eps = jnp.asarray(1e-9, dtype=norm.dtype) + norm_safe = jnp.maximum(norm, eps) + abs_sin = jnp.abs(jnp.sin(norm_safe)) + small = norm < jnp.asarray(1e-3, dtype=norm.dtype) + val_taylor = -norm * norm / 6.0 + val_full = jnp.log(jnp.maximum(abs_sin, eps)) - jnp.log(norm_safe) + log_sinc_abs = jnp.where(small, val_taylor, val_full) + return 2.0 * log_sinc_abs + + +def _wrapped_log_q( + xi_p: Float[Array, "N 3"], + mu: Float[Array, "N 3"], + sigma_sq: Float[Array, ""], + mask: Int[Array, "N"], +) -> Float[Array, ""]: + eps = jnp.asarray(1e-9, dtype=xi_p.dtype) + norm = jnp.linalg.norm(xi_p, axis=-1, keepdims=True) + small = norm < eps + safe_norm = jnp.where(small, jnp.ones_like(norm), norm) + default_u = jnp.broadcast_to( + jnp.asarray([1.0, 0.0, 0.0], dtype=xi_p.dtype), + xi_p.shape, + ) + u_hat = jnp.where(small, default_u, xi_p / safe_norm) + n_range = jnp.arange(-N_WRAP, N_WRAP + 1).astype(xi_p.dtype) + two_pi = jnp.asarray(2.0 * jnp.pi, dtype=xi_p.dtype) + branches = xi_p[None, :, :] + two_pi * n_range[:, None, None] * u_hat[None, :, :] + diffs = branches - mu[None, :, :] + log_gaussian = -jnp.sum(diffs * diffs, axis=-1) / (2.0 * sigma_sq) + norm_branches = jnp.linalg.norm(branches, axis=-1) + log_jacobian = _log_jacobian_su2(norm_branches) + log_kernel = log_gaussian - log_jacobian + log_q_per_site = jax.scipy.special.logsumexp(log_kernel, axis=0) + site_mask_f = mask.astype(xi_p.dtype) + return jnp.sum(log_q_per_site * site_mask_f) + + +def langevin_step_one_cached( + key: PRNGKeyArray, + q: Float[Array, "N 4"], + log_p: Float[Array, ""], + grad_log_p: Float[Array, "N 4"], + beta: Float[Array, ""], + sigma: Float[Array, ""], + mask: Int[Array, "N"], + log_p_fn: LogProbFn, +) -> tuple[ + Float[Array, "N 4"], + Float[Array, ""], + Float[Array, "N 4"], + Float[Array, ""], +]: + k_prop, k_acc = jax.random.split(key) + n_spins = q.shape[0] + site_mask_f = mask.astype(q.dtype)[:, None] + site_mask_b = mask.astype(jnp.bool_)[:, None] + drift_q = _drift_lie(q, grad_log_p) + drift_q = drift_q * site_mask_f + noise = jax.random.normal(k_prop, (n_spins, 3), dtype=q.dtype) + noise = noise * site_mask_f + sigma_sq = sigma * sigma + mu_fwd = (sigma_sq / 2.0) * beta * drift_q + xi_fwd = mu_fwd + sigma * noise + delta = quaternion_exp_tangent(xi_fwd) + q_prop = normalize_quaternion(quaternion_multiply(q, delta)) + delta_back = quaternion_multiply(quaternion_conjugate(q_prop), q) + xi_p_back = quaternion_log(delta_back) + xi_p_back = xi_p_back * site_mask_f + xi_p_fwd = -xi_p_back + log_p_at_qprop, grad_log_p_at_qprop = jax.value_and_grad(log_p_fn)(q_prop) + drift_qprop = _drift_lie(q_prop, grad_log_p_at_qprop) + drift_qprop = drift_qprop * site_mask_f + mu_back = (sigma_sq / 2.0) * beta * drift_qprop + log_q_fwd = _wrapped_log_q(xi_p_fwd, mu_fwd, sigma_sq, mask) + log_q_back = _wrapped_log_q(xi_p_fwd, -mu_back, sigma_sq, mask) + log_alpha = beta * (log_p_at_qprop - log_p) + log_q_back - log_q_fwd + u = jnp.log(jax.random.uniform(k_acc, dtype=q.dtype)) + mh_accept = u < log_alpha + nan_recovery = jnp.isnan(log_p) & jnp.logical_not(jnp.isnan(log_p_at_qprop)) + accept = mh_accept | nan_recovery + q_new = jnp.where(accept, q_prop, q) + q_new = jnp.where(site_mask_b, q_new, q) + log_p_new = jnp.where(accept, log_p_at_qprop, log_p) + grad_log_p_new = jnp.where( + accept, + grad_log_p_at_qprop, + grad_log_p, + ) + return q_new, log_p_new, grad_log_p_new, accept.astype(log_p.dtype) + + +__all__ = ["langevin_step_one_cached"] diff --git a/src/hamiltonzero/mcmc/quaternion.py b/src/hamiltonzero/mcmc/quaternion.py new file mode 100644 index 0000000000000000000000000000000000000000..e42fbb7add9d38b1fc1730064e271a8f8c3834db --- /dev/null +++ b/src/hamiltonzero/mcmc/quaternion.py @@ -0,0 +1,90 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import jax.numpy as jnp +from jaxtyping import Array, Float + + +def quaternion_multiply( + a: Float[Array, "... 4"], + b: Float[Array, "... 4"], +) -> Float[Array, "... 4"]: + w = ( + a[..., 0] * b[..., 0] + - a[..., 1] * b[..., 1] + - a[..., 2] * b[..., 2] + - a[..., 3] * b[..., 3] + ) + x = ( + a[..., 0] * b[..., 1] + + a[..., 1] * b[..., 0] + + a[..., 2] * b[..., 3] + - a[..., 3] * b[..., 2] + ) + y = ( + a[..., 0] * b[..., 2] + - a[..., 1] * b[..., 3] + + a[..., 2] * b[..., 0] + + a[..., 3] * b[..., 1] + ) + z = ( + a[..., 0] * b[..., 3] + + a[..., 1] * b[..., 2] + - a[..., 2] * b[..., 1] + + a[..., 3] * b[..., 0] + ) + return jnp.stack([w, x, y, z], axis=-1) + + +def quaternion_exp_tangent( + v: Float[Array, "... 3"], +) -> Float[Array, "... 4"]: + theta_sq = jnp.sum(v * v, axis=-1, keepdims=True) + theta = jnp.sqrt(theta_sq) + small = theta_sq < 1e-12 + sin_over_theta = jnp.where( + small, + 1.0 - theta_sq / 6.0 + theta_sq * theta_sq / 120.0, + jnp.sin(theta) / jnp.where(small, 1.0, theta), + ) + real = jnp.where( + small, + 1.0 - theta_sq / 2.0 + theta_sq * theta_sq / 24.0, + jnp.cos(theta), + ) + return jnp.concatenate([real, sin_over_theta * v], axis=-1) + + +def normalize_quaternion( + q: Float[Array, "... 4"], + eps: float = 1e-12, +) -> Float[Array, "... 4"]: + return q / jnp.sqrt(jnp.sum(q * q, axis=-1, keepdims=True) + eps) + + +def quaternion_conjugate( + q: Float[Array, "... 4"], +) -> Float[Array, "... 4"]: + return jnp.concatenate([q[..., 0:1], -q[..., 1:]], axis=-1) + + +def quaternion_log( + q: Float[Array, "... 4"], + eps: float = 1e-8, +) -> Float[Array, "... 3"]: + w = q[..., 0:1] + v = q[..., 1:] + sin_norm_sq = jnp.sum(v * v, axis=-1, keepdims=True) + sin_norm = jnp.sqrt(sin_norm_sq) + theta = jnp.arctan2(sin_norm, w) + small = sin_norm < eps + factor_small = 1.0 + sin_norm_sq / 6.0 + factor_normal = theta / jnp.where( + small, + jnp.ones_like(sin_norm), + sin_norm, + ) + factor = jnp.where(small, factor_small, factor_normal) + return v * factor diff --git a/src/hamiltonzero/mcmc/replica_exchange.py b/src/hamiltonzero/mcmc/replica_exchange.py new file mode 100644 index 0000000000000000000000000000000000000000..317bb240b19fabde9e6cc351848e113bbeee88dc --- /dev/null +++ b/src/hamiltonzero/mcmc/replica_exchange.py @@ -0,0 +1,487 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Callable + +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray + +from ._monotone_cubic import MonotoneCubicInterpolation +from .langevin import langevin_step_one_cached +from .quaternion import normalize_quaternion + + +LogProbFn = Callable[[Float[Array, "N 4"]], Float[Array, ""]] + + +class REState(eqx.Module): + q: Float[Array, "R N 4"] + log_p: Float[Array, "R"] + grad_log_p: Float[Array, "R N 4"] + beta: Float[Array, "R"] + sigma: Float[Array, "R"] + step: Int[Array, ""] + key: PRNGKeyArray + n_local_accept: Float[Array, "R"] + n_local: Int[Array, ""] + n_swap_accept: Float[Array, "R_minus_1"] + n_swap: Int[Array, "R_minus_1"] + mask: Int[Array, "N"] + m: Int[Array, "R"] + n_haar_accept: Float[Array, "R"] + n_haar: Int[Array, ""] + + +def annealing_geometric( + n_replicas: int, + base: float = 2.0**0.5, +) -> Float[Array, "R"]: + if n_replicas < 2: + raise ValueError("n_replicas must be at least two") + if n_replicas == 2: + return jnp.array([0.0, 1.0]) + powers = jnp.power(jnp.asarray(base), -jnp.arange(n_replicas - 1)) + tail = jnp.append(powers, 0.0) + return tail[::-1] + + +def init_state( + key: PRNGKeyArray, + *, + n_replicas: int, + n_spins: int, + sigma: float, + mask: Int[Array, "N"], + initial_m: int, +) -> REState: + k_q, k_state = jax.random.split(key) + q_raw = jax.random.normal(k_q, (n_replicas, n_spins, 4)) + q_init = normalize_quaternion(q_raw) + beta = annealing_geometric(n_replicas) + dtype = q_init.dtype + sigma_arr = jnp.full((n_replicas,), sigma, dtype=dtype) + log_p = jnp.full((n_replicas,), jnp.asarray(-1e30, dtype=dtype)) + grad_log_p = jnp.zeros_like(q_init) + n_active = jnp.sum(mask).astype(jnp.int32) + initial_m_clipped = jnp.minimum( + jnp.maximum(jnp.asarray(initial_m, dtype=jnp.int32), jnp.int32(1)), + jnp.maximum(n_active, jnp.int32(1)), + ) + m_arr = jnp.broadcast_to(initial_m_clipped, (n_replicas,)) + is_hot = beta.astype(dtype) == jnp.asarray(0.0, dtype=dtype) + m_arr = jnp.where(is_hot, n_active, m_arr) + return REState( + q=q_init, + log_p=log_p, + grad_log_p=grad_log_p, + beta=beta.astype(dtype), + sigma=sigma_arr, + step=jnp.int32(0), + key=k_state, + n_local_accept=jnp.zeros(n_replicas, dtype=dtype), + n_local=jnp.int32(0), + n_swap_accept=jnp.zeros(n_replicas - 1, dtype=dtype), + n_swap=jnp.zeros(n_replicas - 1, dtype=jnp.int32), + mask=mask.astype(jnp.int32), + m=m_arr, + n_haar_accept=jnp.zeros(n_replicas, dtype=dtype), + n_haar=jnp.int32(0), + ) + + +def _refresh_state_log_p_grad( + state: REState, + log_p_fn: LogProbFn, +) -> REState: + log_p_and_grad = jax.vmap(jax.value_and_grad(log_p_fn)) + new_log_p, new_grad = log_p_and_grad(state.q) + return eqx.tree_at( + lambda value: (value.log_p, value.grad_log_p), + state, + ( + new_log_p.astype(state.log_p.dtype), + new_grad.astype(state.grad_log_p.dtype), + ), + ) + + +def _deo_permutation( + edge_accept: Int[Array, "R_minus_1"], +) -> Int[Array, "R"]: + replicas = edge_accept.shape[0] + 1 + zero = jnp.zeros((1,), dtype=edge_accept.dtype) + left = jnp.concatenate([zero, edge_accept]) + right = jnp.concatenate([edge_accept, zero]) + index = jnp.arange(replicas, dtype=jnp.int32) + index = jnp.where(left.astype(bool), index - 1, index) + return jnp.where(right.astype(bool), index + 1, index) + + +def _swap_move( + key: PRNGKeyArray, + q: Float[Array, "R N 4"], + log_p: Float[Array, "R"], + grad_log_p: Float[Array, "R N 4"], + beta: Float[Array, "R"], + step: Int[Array, ""], +): + replicas = q.shape[0] + index = jnp.arange(replicas - 1) + parity = step % 2 + edge_eligible = (index % 2) == parity + lp_i, lp_j = log_p[:-1], log_p[1:] + beta_i, beta_j = beta[:-1], beta[1:] + log_alpha = (beta_j - beta_i) * (lp_i - lp_j) + uniform = jnp.log(jax.random.uniform(key, (replicas - 1,), dtype=log_p.dtype)) + mh_edge = uniform < log_alpha + nan_i = jnp.isnan(lp_i) + nan_j = jnp.isnan(lp_j) + nan_higher_only = nan_j & jnp.logical_not(nan_i) + edge_accept = edge_eligible & (mh_edge | nan_higher_only) + permutation = _deo_permutation(edge_accept.astype(jnp.int32)) + return ( + q[permutation], + log_p[permutation], + grad_log_p[permutation], + edge_accept.astype(log_p.dtype), + edge_eligible.astype(jnp.int32), + ) + + +def _global_haar_step( + key: PRNGKeyArray, + state: REState, + log_p_fn: LogProbFn, +): + replicas, n_sites, _ = state.q.shape + dtype = state.q.dtype + k_subset, k_haar, k_accept = jax.random.split(key, 3) + scores = jax.random.uniform(k_subset, (replicas, n_sites), dtype=dtype) + active = state.mask.astype(jnp.bool_) + scores = jnp.where( + active[None, :], + scores, + jnp.asarray(jnp.inf, dtype=dtype), + ) + ranks = jnp.argsort(jnp.argsort(scores, axis=-1), axis=-1) + move_mask = active[None, :] & (ranks < state.m[:, None]) + raw = jax.random.normal(k_haar, state.q.shape, dtype=dtype) + q_haar = normalize_quaternion(raw) + q_prop = jnp.where(move_mask[..., None], q_haar, state.q) + log_p_prop, grad_log_p_prop = jax.vmap(jax.value_and_grad(log_p_fn))(q_prop) + log_alpha = jnp.minimum( + jnp.asarray(0.0, dtype=dtype), + state.beta.astype(dtype) * (log_p_prop.astype(dtype) - state.log_p), + ) + uniform = jax.random.uniform(k_accept, (replicas,), dtype=dtype) + accept = jnp.log(uniform) <= log_alpha + q_new = jnp.where(accept[:, None, None], q_prop, state.q) + log_p_new = jnp.where( + accept, + log_p_prop.astype(dtype), + state.log_p, + ) + grad_log_p_new = jnp.where( + accept[:, None, None], + grad_log_p_prop, + state.grad_log_p, + ) + return ( + q_new, + log_p_new, + grad_log_p_new, + accept.astype(state.n_haar_accept.dtype), + ) + + +def global_haar_step( + state: REState, + log_p_fn: LogProbFn, +) -> REState: + k_haar, k_next = jax.random.split(state.key, 2) + q_new, log_p_new, grad_log_p_new, accept = _global_haar_step( + k_haar, + state, + log_p_fn, + ) + return eqx.tree_at( + lambda value: ( + value.q, + value.log_p, + value.grad_log_p, + value.key, + value.n_haar_accept, + value.n_haar, + ), + state, + ( + q_new, + log_p_new, + grad_log_p_new, + k_next, + state.n_haar_accept + accept, + state.n_haar + jnp.int32(1), + ), + ) + + +def step_re_langevin_cached( + state: REState, + log_p_fn: LogProbFn, +) -> REState: + k_local, k_swap, k_next = jax.random.split(state.key, 3) + local_keys = jax.random.split(k_local, state.q.shape[0]) + + def local_move(key, q, log_p, grad_log_p, beta, sigma): + return langevin_step_one_cached( + key, + q, + log_p, + grad_log_p, + beta, + sigma, + state.mask, + log_p_fn, + ) + + q_local, log_p_local, grad_local, accepts = jax.vmap(local_move)( + local_keys, + state.q, + state.log_p, + state.grad_log_p, + state.beta, + state.sigma, + ) + q_swap, log_p_swap, grad_swap, edge_accept, edge_eligible = _swap_move( + k_swap, + q_local, + log_p_local, + grad_local, + state.beta, + state.step, + ) + return REState( + q=q_swap, + log_p=log_p_swap, + grad_log_p=grad_swap, + beta=state.beta, + sigma=state.sigma, + step=state.step + 1, + key=k_next, + n_local_accept=state.n_local_accept + accepts, + n_local=state.n_local + 1, + n_swap_accept=state.n_swap_accept + edge_accept, + n_swap=state.n_swap + edge_eligible, + mask=state.mask, + m=state.m, + n_haar_accept=state.n_haar_accept, + n_haar=state.n_haar, + ) + + +def run_re_langevin_cached( + state: REState, + log_p_fn: LogProbFn, + n_steps: int, +) -> REState: + state = _refresh_state_log_p_grad(state, log_p_fn) + state = global_haar_step(state, log_p_fn) + + def body(carry, _unused): + return step_re_langevin_cached(carry, log_p_fn), None + + state, _unused = jax.lax.scan( + body, + state, + xs=None, + length=n_steps, + ) + return state + + +def _local_accept_rate(state: REState) -> Float[Array, "R"]: + attempts = state.n_local.astype(state.sigma.dtype)[..., None] + return state.n_local_accept / jnp.maximum(attempts, 1.0) + + +def _swap_accept_rate(state: REState) -> Float[Array, "R_minus_1"]: + attempts = state.n_swap.astype(state.sigma.dtype) + return state.n_swap_accept / jnp.maximum(attempts, 1.0) + + +def adapt_sigma( + state: REState, + *, + target: float = 0.234, + factor: float = 1.1, + lo: float = 1e-4, + hi: float = 3.14, +) -> REState: + rate = _local_accept_rate(state) + sigma_new = jnp.where( + rate > target, + state.sigma * factor, + state.sigma / factor, + ) + sigma_new = jnp.clip(sigma_new, lo, hi) + return eqx.tree_at( + lambda value: ( + value.sigma, + value.n_local_accept, + value.n_local, + ), + state, + ( + sigma_new, + jnp.zeros_like(state.n_local_accept), + jnp.zeros_like(state.n_local), + ), + ) + + +def _haar_accept_rate(state: REState) -> Float[Array, "R"]: + denominator = jnp.maximum( + state.n_haar.astype(state.n_haar_accept.dtype), + jnp.asarray(1.0, dtype=state.n_haar_accept.dtype), + ) + return state.n_haar_accept / denominator[..., None] + + +def adapt_m( + state: REState, + *, + target: float = 0.234, +) -> REState: + rate = _haar_accept_rate(state) + n_active = jnp.sum(state.mask).astype(jnp.int32) + delta = jnp.where(rate > target, jnp.int32(1), jnp.int32(-1)) + m_new = jnp.clip(state.m + delta, jnp.int32(1), n_active) + is_hot = state.beta == jnp.asarray(0.0, dtype=state.beta.dtype) + m_new = jnp.where(is_hot, n_active, m_new) + return eqx.tree_at( + lambda value: ( + value.m, + value.n_haar_accept, + value.n_haar, + ), + state, + ( + m_new, + jnp.zeros_like(state.n_haar_accept), + jnp.zeros_like(state.n_haar), + ), + ) + + +def _estimate_lambda_values( + rejection_rates: Float[Array, "R_minus_1"], + offset: float = 1e-4, +) -> Float[Array, "R"]: + rejection_rates = jnp.maximum(rejection_rates, offset) + extended = jnp.concatenate( + [jnp.zeros_like(rejection_rates[..., :1]), rejection_rates], + axis=-1, + ) + return jnp.cumsum(extended, axis=-1) + + +def _annealing_optimal_hyman( + n_replicas: int, + previous_schedule: Float[Array, "R"], + rejection_rates: Float[Array, "R_minus_1"], + *, + offset: float = 1e-4, + ema: float = 0.99, +) -> Float[Array, "R"]: + lambda_values = _estimate_lambda_values( + rejection_rates, + offset=offset, + ) + lambda_fn = MonotoneCubicInterpolation( + ts=previous_schedule, + ys=lambda_values, + ) + lambda_max = lambda_values[-1] + dtype = previous_schedule.dtype + indices = jnp.arange(1, n_replicas - 1, dtype=dtype) + targets = indices * (lambda_max / (n_replicas - 1)) + lower = jnp.asarray(offset, dtype=dtype) + upper = jnp.asarray(1.0 - offset, dtype=dtype) + tolerance = jnp.asarray(offset * 1e-3, dtype=dtype) + max_iterations = jnp.int32(50) + + def bisect(target): + def condition(state): + lo, hi, iteration = state + return (iteration < max_iterations) & ((hi - lo) > tolerance) + + def body(state): + lo, hi, iteration = state + midpoint = (lo + hi) * jnp.asarray(0.5, dtype=dtype) + value = lambda_fn.evaluate(midpoint) - target + move_lower = value < 0.0 + new_lower = jnp.where(move_lower, midpoint, lo) + new_upper = jnp.where(move_lower, hi, midpoint) + return new_lower, new_upper, iteration + 1 + + lo_final, hi_final, _iteration = jax.lax.while_loop( + condition, + body, + (lower, upper, jnp.int32(0)), + ) + return (lo_final + hi_final) * jnp.asarray(0.5, dtype=dtype) + + interior = jax.vmap(bisect)(targets) + new_schedule = jnp.concatenate( + [ + jnp.zeros((1,), dtype=dtype), + interior, + jnp.ones((1,), dtype=dtype), + ] + ) + return jnp.clip( + (1.0 - ema) * new_schedule + ema * previous_schedule, + 0.0, + 1.0, + ).astype(dtype) + + +def adapt_beta_equi_rej( + state: REState, + *, + ema: float = 0.99, +) -> REState: + rejection = 1.0 - _swap_accept_rate(state) + beta_new = _annealing_optimal_hyman( + n_replicas=int(state.beta.shape[0]), + previous_schedule=state.beta, + rejection_rates=rejection, + ema=ema, + ) + return eqx.tree_at( + lambda value: ( + value.beta, + value.n_swap_accept, + value.n_swap, + ), + state, + ( + beta_new, + jnp.zeros_like(state.n_swap_accept), + jnp.zeros_like(state.n_swap), + ), + ) + + +__all__ = [ + "REState", + "adapt_beta_equi_rej", + "adapt_m", + "adapt_sigma", + "init_state", + "run_re_langevin_cached", +] diff --git a/src/hamiltonzero/mcmc/runtime.py b/src/hamiltonzero/mcmc/runtime.py new file mode 100644 index 0000000000000000000000000000000000000000..f1ccf375d80429effcb7ae8b84fd84ff1b22a22b --- /dev/null +++ b/src/hamiltonzero/mcmc/runtime.py @@ -0,0 +1,208 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import fields +from typing import Any, Callable + +import equinox as eqx +import jax +import jax.numpy as jnp + +from .replica_exchange import ( + REState, + adapt_beta_equi_rej, + adapt_m, + adapt_sigma, + init_state, + run_re_langevin_cached, +) + + +def _state_axes(mask_axis: int | None) -> REState: + return REState( + q=0, + log_p=0, + grad_log_p=0, + beta=None, + sigma=None, + step=None, + key=0, + n_local_accept=0, + n_local=0, + n_swap_accept=0, + n_swap=0, + mask=mask_axis, + m=None, + n_haar_accept=0, + n_haar=0, + ) + + +def init_batched_state( + key: jax.Array, + context: Any, + batch_size: int, + n_replicas: int, + initial_m: int = 1, + initial_sigma: float = 0.3, +) -> REState: + mask = jnp.asarray(context.mask, dtype=jnp.int32) + keys = jax.random.split(key, batch_size) + return jax.vmap( + lambda walker_key: init_state( + walker_key, + n_replicas=n_replicas, + n_spins=int(mask.shape[-1]), + sigma=initial_sigma, + mask=mask, + initial_m=initial_m, + ), + out_axes=_state_axes(None), + )(keys) + + +def _log_probability( + model: Callable, + context: Any, + q: jax.Array, +) -> jax.Array: + real, _phase = model(q, context, 0.0) + return 2.0 * real + + +def _run_one( + state: REState, + model: Callable, + context: Any, + n_steps: int, +) -> REState: + return run_re_langevin_cached( + state, + lambda q: _log_probability(model, context, q), + n_steps, + ) + + +def run_batched( + model: Callable, + context: Any, + state: REState, + n_steps: int, + walker_chunk_size: int | None = None, +) -> REState: + if walker_chunk_size is None or walker_chunk_size >= state.q.shape[0]: + return jax.vmap( + lambda walker: _run_one(walker, model, context, n_steps), + in_axes=(_state_axes(None),), + out_axes=_state_axes(None), + )(state) + + shared_names = {"mask", "beta", "sigma", "step", "m"} + shared = {name: getattr(state, name) for name in shared_names} + walker_fields = { + item.name: getattr(state, item.name) + for item in fields(state) + if item.name not in shared_names + } + + def run_walker(walker: dict[str, jax.Array]) -> dict[str, jax.Array]: + new_state = _run_one( + REState(**shared, **walker), + model, + context, + n_steps, + ) + return { + item.name: getattr(new_state, item.name) + for item in fields(new_state) + if item.name not in {"mask", "beta", "sigma", "m"} + } + + mapped = jax.lax.map( + run_walker, + walker_fields, + batch_size=int(walker_chunk_size), + ) + step = mapped.pop("step")[0] + return REState(**{**shared, "step": step}, **mapped) + + +def adapt_batched( + state: REState, + *, + beta_history_weight: float = 0.9, + sigma_target: float = 0.574, + sigma_scale: float = 1.1, + haar_target: float = 0.234, +) -> REState: + pooled = REState( + q=state.q[0], + log_p=state.log_p[0], + grad_log_p=state.grad_log_p[0], + beta=state.beta, + sigma=state.sigma, + step=state.step, + key=state.key[0], + n_local_accept=state.n_local_accept.sum(axis=0), + n_local=state.n_local.sum(axis=0).astype(state.n_local.dtype), + n_swap_accept=state.n_swap_accept.sum(axis=0), + n_swap=state.n_swap.sum(axis=0).astype(state.n_swap.dtype), + mask=state.mask, + m=state.m, + n_haar_accept=state.n_haar_accept.sum(axis=0), + n_haar=state.n_haar.sum(axis=0).astype(state.n_haar.dtype), + ) + adapted = adapt_sigma( + pooled, + target=sigma_target, + factor=sigma_scale, + ) + adapted = adapt_m(adapted, target=haar_target) + adapted = adapt_beta_equi_rej( + adapted, + ema=beta_history_weight, + ) + + def target(value: REState): + return ( + value.sigma, + value.beta, + value.m, + value.n_local_accept, + value.n_local, + value.n_swap_accept, + value.n_swap, + value.n_haar_accept, + value.n_haar, + ) + + return eqx.tree_at( + target, + state, + ( + adapted.sigma, + adapted.beta, + adapted.m, + jnp.zeros_like(state.n_local_accept), + jnp.zeros_like(state.n_local), + jnp.zeros_like(state.n_swap_accept), + jnp.zeros_like(state.n_swap), + jnp.zeros_like(state.n_haar_accept), + jnp.zeros_like(state.n_haar), + ), + ) + + +def cold_samples(state: REState) -> jax.Array: + return state.q[:, -1] + + +__all__ = [ + "REState", + "adapt_batched", + "cold_samples", + "init_batched_state", + "run_batched", +] diff --git a/src/hamiltonzero/model/__init__.py b/src/hamiltonzero/model/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..dda08d2b263f514af49e54d5c9377def51423d2f --- /dev/null +++ b/src/hamiltonzero/model/__init__.py @@ -0,0 +1,37 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .api import ( + SpinAnsatz, + balanced_subtree_mask, + build_model, + edge_merge_masked, + normalize_leaf_carriers, + quadrilinear_merge, + tagged_dense, + tagged_dense_no_bias, + tagged_rms_eqx_style, + tree_active_clock_depth, + tree_depth_count_features, + tree_sphere, +) +from .context import MultiSystemContext, SpinContext +from .route_quotient import route_quotient_keys + +__all__ = [ + "MultiSystemContext", + "SpinAnsatz", + "SpinContext", + "balanced_subtree_mask", + "build_model", + "edge_merge_masked", + "normalize_leaf_carriers", + "quadrilinear_merge", + "route_quotient_keys", + "tagged_dense", + "tagged_dense_no_bias", + "tagged_rms_eqx_style", + "tree_active_clock_depth", + "tree_depth_count_features", + "tree_sphere", +] diff --git a/src/hamiltonzero/model/_custom_lap_primitives.py b/src/hamiltonzero/model/_custom_lap_primitives.py new file mode 100644 index 0000000000000000000000000000000000000000..21171917aaa144a7a0fa3ca541bcb864a34d2661 --- /dev/null +++ b/src/hamiltonzero/model/_custom_lap_primitives.py @@ -0,0 +1,91 @@ +# 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", +] diff --git a/src/hamiltonzero/model/_pallas_attn/__init__.py b/src/hamiltonzero/model/_pallas_attn/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..ba6b74af888e75a928f9db4ea20c77515f740359 --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .mhsea_tuned import mhsea_with_tuned_ad + +__all__ = ["mhsea_with_tuned_ad"] diff --git a/src/hamiltonzero/model/_pallas_attn/custom_gradients.py b/src/hamiltonzero/model/_pallas_attn/custom_gradients.py new file mode 100644 index 0000000000000000000000000000000000000000..b5892e277d71b2aea7cb34b05e9c81d578b7720c --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/custom_gradients.py @@ -0,0 +1,347 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# Modifications copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: MIT + +from functools import partial +from typing import Tuple +import jax +import jax.numpy as jnp +from jax.experimental import pallas as pl +from .mhsea import mhsea_kernel +from .utils import ( + big_number, + compiler_params, + create_grid, + get_key_value_block_spec, + get_lse_block_spec, + get_mask_block_spec, + get_query_block_spec, + sum_columns, +) + + +def mhsea_forward( + q: jax.Array, + k: jax.Array, + e: jax.Array, + v: jax.Array, + mask: jax.Array, + *, + q_block_len: int, + num_warps: int, + num_stages: int, + precision, +) -> Tuple[ + jax.Array, + Tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array, jax.Array, jax.Array], +]: + batch_len, seq_len, num_heads, head_len = q.shape + v_dim = v.shape[-1] + qdtype = q.dtype + dtype = jnp.float32 + kernel_fn = pl.pallas_call( + partial(mhsea_kernel, q_block_len=q_block_len, precision=precision), + grid=create_grid(batch_len, seq_len, num_heads, q_block_len), + in_specs=[ + get_query_block_spec(q_block_len, head_len), + get_key_value_block_spec(seq_len, head_len), + get_query_block_spec(q_block_len, seq_len), + get_key_value_block_spec(seq_len, v_dim), + get_mask_block_spec(seq_len), + ], + out_specs=[ + get_query_block_spec(q_block_len, v_dim), + get_lse_block_spec(q_block_len), + ], + out_shape=[ + jax.ShapeDtypeStruct( + shape=(batch_len, seq_len, num_heads, v_dim), dtype=dtype + ), + jax.ShapeDtypeStruct(shape=(batch_len, seq_len, num_heads), dtype=dtype), + ], + compiler_params=compiler_params(num_warps=num_warps, num_stages=num_stages), + debug=False, + interpret=False, + name="mhea_forward", + ) + o, lse = kernel_fn( + q.astype(dtype), + k.astype(dtype), + e.astype(dtype), + v.astype(dtype), + mask.astype(jnp.bool_), + ) + return (o.astype(qdtype), (q, k, e, v, mask, lse, o)) + + +def mhsea_backward( + q_block_len: int, + num_warps: int, + num_stages: int, + precision, + fwd_cache: Tuple[ + jax.Array, jax.Array, jax.Array, jax.Array, jax.Array, jax.Array, jax.Array + ], + o_vjp: jax.Array, +) -> Tuple[jax.Array, jax.Array, jax.Array, jax.Array]: + q, k, e, v, mask, lse, o = fwd_cache + mask = mask.astype(jnp.bool_) + batch_len, seq_len, num_heads, head_len = q.shape + block_len = q_block_len + dtype = jnp.float32 + dq, de = pl.pallas_call( + partial(mhsea_q_vjp_kernel, block_len=block_len, precision=precision), + grid=(batch_len, seq_len // block_len, num_heads), + in_specs=[ + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, 0, k, 0), + block_shape=(None, seq_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, seq_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, 0, k, 0), + block_shape=(None, seq_len, None, head_len), + ), + pl.BlockSpec(index_map=lambda i, j, k: (i, 0), block_shape=(None, seq_len)), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k), block_shape=(None, block_len, None) + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + ], + out_specs=[ + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, seq_len), + ), + ], + out_shape=[ + jax.ShapeDtypeStruct( + shape=(batch_len, seq_len, num_heads, head_len), dtype=dtype + ), + jax.ShapeDtypeStruct( + shape=(batch_len, seq_len, num_heads, seq_len), dtype=dtype + ), + ], + compiler_params=compiler_params(num_warps=num_warps, num_stages=num_stages), + debug=False, + interpret=False, + name="mhsea_backward_q_vjp", + )( + q.astype(dtype), + k.astype(dtype), + e.astype(dtype), + v.astype(dtype), + mask, + lse.astype(dtype), + o.astype(dtype), + o_vjp.astype(dtype), + ) + dk, dv = pl.pallas_call( + partial(mhsea_kv_vjp_kernel, block_len=block_len, precision=precision), + grid=(batch_len, seq_len // block_len, num_heads), + in_specs=[ + pl.BlockSpec( + index_map=lambda i, j, k: (i, 0, k, 0), + block_shape=(None, seq_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, 0, k, j), + block_shape=(None, seq_len, None, block_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + pl.BlockSpec(index_map=lambda i, j, k: (i, 0), block_shape=(None, seq_len)), + pl.BlockSpec( + index_map=lambda i, j, k: (i, 0, k), block_shape=(None, seq_len, None) + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, 0, k, 0), + block_shape=(None, seq_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, 0, k, 0), + block_shape=(None, seq_len, None, head_len), + ), + ], + out_specs=[ + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, head_len), + ), + ], + out_shape=[ + jax.ShapeDtypeStruct( + shape=(batch_len, seq_len, num_heads, head_len), dtype=dtype + ), + jax.ShapeDtypeStruct( + shape=(batch_len, seq_len, num_heads, head_len), dtype=dtype + ), + ], + compiler_params=compiler_params(num_warps=num_warps, num_stages=num_stages), + debug=False, + interpret=False, + name="mhsea_backward_kv_vjp", + )( + q.astype(dtype), + k.astype(dtype), + e.astype(dtype), + v.astype(dtype), + mask.astype(jnp.bool_), + lse.astype(dtype), + o.astype(dtype), + o_vjp.astype(dtype), + ) + return ( + dq.astype(o_vjp.dtype), + dk.astype(o_vjp.dtype), + de.astype(o_vjp.dtype), + dv.astype(o_vjp.dtype), + ) + + +def mhsea_q_vjp_kernel( + q_ref, + k_ref, + e_ref, + v_ref, + mask_ref, + lse_ref, + o_ref, + o_vjp_ref, + q_vjp_ref, + e_vjp_ref, + block_len, + precision, +): + qix = pl.program_id(1) + q_vjp = jnp.zeros(q_vjp_ref.shape, dtype=q_vjp_ref.dtype) + q_slice = pl.dslice(qix * block_len, block_len) + + def _kaxis_loop(kix, q_vjp): + mask_q = mask_ref[q_slice] + k_slice = pl.dslice(kix * block_len, block_len) + mask_k = mask_ref[k_slice] + square_mask = mask_q[:, None] & mask_k[None, :] + q = jnp.where(mask_q[:, None], q_ref[:, :], jnp.asarray(0, dtype=q_ref.dtype)) + k = jnp.where( + mask_k[:, None], k_ref[k_slice, :], jnp.asarray(0, dtype=k_ref.dtype) + ) + v = jnp.where( + mask_k[:, None], v_ref[k_slice, :], jnp.asarray(0, dtype=v_ref.dtype) + ) + e = e_ref[:, k_slice] + lse = lse_ref[:] + s = jnp.where( + square_mask, + pl.dot(q, k, trans_b=True, precision=precision) + e, + big_number(), + ) + p = jnp.exp(s - lse[:, None]) + o = o_ref[:, :] + o_vjp = jnp.where( + mask_q[:, None], o_vjp_ref[:, :], jnp.asarray(0, dtype=o_vjp_ref.dtype) + ) + s_vjp = ( + pl.dot(o_vjp, v, trans_b=True, precision=precision) - sum_columns(o * o_vjp) + ) * p + s_vjp *= square_mask.astype(s_vjp.dtype) + q_vjp_kblock = pl.dot(s_vjp, k.astype(s_vjp.dtype), precision=precision) + q_vjp += q_vjp_kblock + e_vjp_ref[:, k_slice] = s_vjp.astype(e_vjp_ref.dtype) + return q_vjp.astype(jnp.float32) + + q_vjp = jax.lax.fori_loop( + 0, k_ref.shape[0] // block_len, _kaxis_loop, q_vjp.astype(jnp.float32) + ) + q_vjp_ref[:, :] = q_vjp.astype(q_vjp_ref.dtype) + + +def mhsea_kv_vjp_kernel( + q_ref, + k_ref, + e_ref, + v_ref, + mask_ref, + lse_ref, + o_ref, + o_vjp_ref, + k_vjp_ref, + v_vjp_ref, + block_len, + precision, +): + kix = pl.program_id(1) + k_vjp = jnp.zeros((block_len, k_vjp_ref.shape[-1]), dtype=k_vjp_ref.dtype) + v_vjp = jnp.zeros((block_len, v_vjp_ref.shape[-1]), dtype=v_vjp_ref.dtype) + k_slice = pl.dslice(kix * block_len, block_len) + + def _qaxis_loop(qix, store): + k_vjp, v_vjp = store + q_slice = pl.dslice(qix * block_len, block_len) + mask_q = mask_ref[q_slice] + mask_k = mask_ref[k_slice] + square_mask = mask_q[:, None] & mask_k[None, :] + q = jnp.where( + mask_q[:, None], q_ref[q_slice, :], jnp.asarray(0, dtype=q_ref.dtype) + ) + k = jnp.where(mask_k[:, None], k_ref[:, :], jnp.asarray(0, dtype=k_ref.dtype)) + v = jnp.where(mask_k[:, None], v_ref[:, :], jnp.asarray(0, dtype=v_ref.dtype)) + e = e_ref[q_slice, :] + lse = lse_ref[q_slice] + s = jnp.where( + square_mask, + pl.dot(q, k, trans_b=True, precision=precision) + e, + big_number(), + ) + p = jnp.exp(s - lse[:, None]) + o = o_ref[q_slice, :] + o_vjp = jnp.where( + mask_q[:, None], + o_vjp_ref[q_slice, :], + jnp.asarray(0, dtype=o_vjp_ref.dtype), + ) + s_vjp = ( + pl.dot(o_vjp, v, trans_b=True, precision=precision) - sum_columns(o * o_vjp) + ) * p + s_vjp *= square_mask.astype(s_vjp.dtype) + k_vjp += pl.dot(s_vjp, q.astype(s_vjp.dtype), trans_a=True, precision=precision) + v_vjp += pl.dot(p, o_vjp.astype(p.dtype), trans_a=True, precision=precision) + return (k_vjp.astype(jnp.float32), v_vjp.astype(jnp.float32)) + + k_vjp, v_vjp = jax.lax.fori_loop( + 0, + q_ref.shape[0] // block_len, + _qaxis_loop, + (k_vjp.astype(jnp.float32), v_vjp.astype(jnp.float32)), + ) + k_vjp_ref[:, :] = k_vjp.astype(k_vjp_ref.dtype) + v_vjp_ref[:, :] = v_vjp.astype(v_vjp_ref.dtype) diff --git a/src/hamiltonzero/model/_pallas_attn/mhsea.py b/src/hamiltonzero/model/_pallas_attn/mhsea.py new file mode 100644 index 0000000000000000000000000000000000000000..addaec4ccdf30332ce728ee1d53ace5f40569fa7 --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/mhsea.py @@ -0,0 +1,31 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# Modifications copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: MIT + +import jax.numpy as jnp +from jax.experimental import pallas as pl +from .utils import big_number + + +def mhsea_kernel( + q_ref, k_ref, e_ref, v_ref, mask_ref, o_ref, lse_ref, q_block_len: int, precision +): + q_idx = pl.program_id(1) + kv_mask = mask_ref[:] + q_slice = pl.dslice(q_idx * q_block_len, q_block_len) + q_mask = mask_ref[q_slice] + square_mask = q_mask[:, None] & kv_mask[None, :] + q = jnp.where(q_mask[:, None], q_ref[:, :], jnp.asarray(0, dtype=q_ref.dtype)) + k = jnp.where(kv_mask[:, None], k_ref[:, :], jnp.asarray(0, dtype=k_ref.dtype)) + e = e_ref[:, :] + v = jnp.where(kv_mask[:, None], v_ref[:, :], jnp.asarray(0, dtype=v_ref.dtype)) + s = jnp.where( + square_mask, pl.dot(q, k, trans_b=True, precision=precision) + e, big_number() + ) + max_val = jnp.max(s, axis=1, keepdims=False) + lse = max_val + jnp.log(jnp.sum(jnp.exp(s - max_val[:, None]), axis=1)) + p = jnp.exp(s - lse[:, None]) + lse_ref[:] = lse.astype(lse_ref.dtype) + o = pl.dot(p, v.astype(p.dtype), precision=precision) + o_ref[:, :] = o.astype(o_ref.dtype) diff --git a/src/hamiltonzero/model/_pallas_attn/mhsea_full_ad.py b/src/hamiltonzero/model/_pallas_attn/mhsea_full_ad.py new file mode 100644 index 0000000000000000000000000000000000000000..2effab0f3d68640fe80119bfa429cf7d63b866e3 --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/mhsea_full_ad.py @@ -0,0 +1,194 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import jax +import jax.core as _shaped_core +import jax.numpy as jnp +from jax.extend import core as _ext_core +from jax.interpreters import ad, batching, mlir + +from .custom_gradients import mhsea_backward, mhsea_forward +from .mhsea_full_ad_jvp import mhsea_jvp_pallas + + +_TF32_PRECISION = jax.lax.Precision.DEFAULT + + +def _validate_primals(q, k, e, v, mask) -> None: + assert q.ndim == k.ndim == v.ndim == 4 + assert e.ndim == 4 + assert mask.ndim == 2 + assert q.dtype == k.dtype == e.dtype == v.dtype + assert jnp.dtype(q.dtype) == jnp.dtype(jnp.float32) + batch_size, seq_len, num_heads, _head_dim = q.shape + assert k.shape == q.shape + assert v.shape == q.shape + assert e.shape == (batch_size, seq_len, num_heads, seq_len) + assert mask.shape == (batch_size, seq_len) + + +mhsea_full_ad_p = _ext_core.Primitive("mhsea_full_ad") +mhsea_full_ad_p.multiple_results = True + + +def _outer_impl(q, k, e, v, mask): + _validate_primals(q, k, e, v, mask) + seq_len = q.shape[1] + out, residuals = mhsea_forward( + q, + k, + e, + v, + mask, + q_block_len=seq_len, + num_warps=8, + num_stages=3, + precision=_TF32_PRECISION, + ) + return out, residuals[5] + + +def _outer_abstract_eval(q, k, e, v, mask): + _validate_primals(q, k, e, v, mask) + batch_size, seq_len, num_heads, _head_dim = q.shape + return ( + _shaped_core.ShapedArray(v.shape, v.dtype), + _shaped_core.ShapedArray((batch_size, seq_len, num_heads), q.dtype), + ) + + +def _outer_jvp(primals, tangents): + q, k, e, v, mask = primals + qt, kt, et, vt, _mask_t = tangents + qt = ad.instantiate_zeros(qt) + kt = ad.instantiate_zeros(kt) + et = ad.instantiate_zeros(et) + vt = ad.instantiate_zeros(vt) + primal_out, lse = mhsea_full_ad_p.bind(q, k, e, v, mask) + tangent_out = mhsea_full_ad_lin_p.bind( + q, k, e, v, mask, primal_out, lse, qt, kt, et, vt + ) + return (primal_out, lse), ( + tangent_out, + ad.Zero(_shaped_core.ShapedArray(lse.shape, lse.dtype)), + ) + + +def _merge_first_two(value): + return value.reshape(value.shape[0] * value.shape[1], *value.shape[2:]) + + +def _split_first_two(value, mapped_size: int): + return value.reshape(mapped_size, value.shape[0] // mapped_size, *value.shape[1:]) + + +def _prepare_batched_args(args, axes): + mapped_size = None + moved = [] + for value, axis in zip(args, axes): + if axis is None: + moved.append(value) + continue + value = jnp.moveaxis(value, axis, 0) + if mapped_size is None: + mapped_size = value.shape[0] + else: + assert value.shape[0] == mapped_size + moved.append(value) + if mapped_size is None: + return None, None + broadcasted = [ + jnp.broadcast_to(value[None], (mapped_size, *value.shape)) + if axis is None + else value + for value, axis in zip(moved, axes) + ] + return mapped_size, [_merge_first_two(value) for value in broadcasted] + + +def _outer_batch(args, axes): + mapped_size, flat = _prepare_batched_args(args, axes) + if mapped_size is None: + out, lse = mhsea_full_ad_p.bind(*args) + return (out, lse), (None, None) + out_flat, lse_flat = mhsea_full_ad_p.bind(*flat) + return ( + _split_first_two(out_flat, mapped_size), + _split_first_two(lse_flat, mapped_size), + ), (0, 0) + + +def _outer_mlir(ctx, *args, **kwargs): + return mlir.lower_fun(_outer_impl, multiple_results=True)(ctx, *args, **kwargs) + + +mhsea_full_ad_p.def_impl(_outer_impl) +mhsea_full_ad_p.def_abstract_eval(_outer_abstract_eval) +ad.primitive_jvps[mhsea_full_ad_p] = _outer_jvp +batching.primitive_batchers[mhsea_full_ad_p] = _outer_batch +mlir.register_lowering(mhsea_full_ad_p, _outer_mlir) + + +mhsea_full_ad_lin_p = _ext_core.Primitive("mhsea_full_ad_lin") +mhsea_full_ad_lin_p.multiple_results = False + + +def _inner_impl(q, k, e, v, mask, primal_out, lse, qt, kt, et, vt): + del primal_out, lse + _out, tangent = mhsea_jvp_pallas( + q, k, e, v, mask, qt, kt, et, vt, precision=_TF32_PRECISION + ) + return tangent + + +def _inner_abstract_eval(q, k, e, v, mask, primal_out, lse, qt, kt, et, vt): + del q, k, e, mask, primal_out, lse, qt, kt, et, vt + return _shaped_core.ShapedArray(v.shape, v.dtype) + + +def _inner_transpose(cot, q, k, e, v, mask, primal_out, lse, qt, kt, et, vt): + if isinstance(cot, ad.Zero): + zeros = tuple( + ad.Zero(_shaped_core.ShapedArray(value.shape, value.dtype)) + for value in (q, k, e, v) + ) + return (None,) * 7 + zeros + dq, dk, de, dv = mhsea_backward( + q.shape[1], + 2, + 1, + _TF32_PRECISION, + (q, k, e, v, mask, lse, primal_out), + cot, + ) + return (None,) * 7 + (dq, dk, de, dv) + + +def _inner_batch(args, axes): + mapped_size, flat = _prepare_batched_args(args, axes) + if mapped_size is None: + return mhsea_full_ad_lin_p.bind(*args), None + out_flat = mhsea_full_ad_lin_p.bind(*flat) + return _split_first_two(out_flat, mapped_size), 0 + + +def _inner_mlir(ctx, *args, **kwargs): + return mlir.lower_fun(_inner_impl, multiple_results=False)(ctx, *args, **kwargs) + + +mhsea_full_ad_lin_p.def_impl(_inner_impl) +mhsea_full_ad_lin_p.def_abstract_eval(_inner_abstract_eval) +ad.primitive_transposes[mhsea_full_ad_lin_p] = _inner_transpose +batching.primitive_batchers[mhsea_full_ad_lin_p] = _inner_batch +mlir.register_lowering(mhsea_full_ad_lin_p, _inner_mlir) + + +def mhsea_with_full_ad(q, k, e, v, mask) -> jax.Array: + _validate_primals(q, k, e, v, mask) + out, _lse = mhsea_full_ad_p.bind(q, k, e, v, mask) + return out + + +__all__ = ["mhsea_with_full_ad"] diff --git a/src/hamiltonzero/model/_pallas_attn/mhsea_full_ad_jvp.py b/src/hamiltonzero/model/_pallas_attn/mhsea_full_ad_jvp.py new file mode 100644 index 0000000000000000000000000000000000000000..8798b959dfefb404cfd85576c595d95a70e60c5c --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/mhsea_full_ad_jvp.py @@ -0,0 +1,112 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from functools import partial + +import jax +import jax.numpy as jnp +from jax.experimental import pallas as pl + +from .utils import compiler_params + + +def _mask_value() -> jax.Array: + return jnp.float32(-10000.0) + + +def _mhsea_jvp_kernel( + q_ref, + k_ref, + e_ref, + v_ref, + mask_ref, + qd_ref, + kd_ref, + ed_ref, + vd_ref, + o_ref, + od_ref, + *, + precision, +): + mask = mask_ref[:] + square_mask = mask[:, None] & mask[None, :] + + q = jnp.where(mask[:, None], q_ref[:, :], jnp.float32(0.0)) + k = jnp.where(mask[:, None], k_ref[:, :], jnp.float32(0.0)) + v = jnp.where(mask[:, None], v_ref[:, :], jnp.float32(0.0)) + e = e_ref[:, :] + + qd = jnp.where(mask[:, None], qd_ref[:, :], jnp.float32(0.0)) + kd = jnp.where(mask[:, None], kd_ref[:, :], jnp.float32(0.0)) + vd = jnp.where(mask[:, None], vd_ref[:, :], jnp.float32(0.0)) + ed = ed_ref[:, :] + + scores_raw = pl.dot(q, k, trans_b=True, precision=precision) + e + scores = jnp.where(square_mask, scores_raw, _mask_value()) + scores_max = jnp.max(scores, axis=-1, keepdims=True) + probabilities_unscaled = jnp.exp(scores - scores_max) + probabilities = probabilities_unscaled / jnp.sum( + probabilities_unscaled, axis=-1, keepdims=True + ) + + scores_d_raw = ( + pl.dot(qd, k, trans_b=True, precision=precision) + + pl.dot(q, kd, trans_b=True, precision=precision) + + ed + ) + scores_d = jnp.where(square_mask, scores_d_raw, jnp.float32(0.0)) + scores_d_mean = jnp.sum(probabilities * scores_d, axis=-1, keepdims=True) + probabilities_d = probabilities * (scores_d - scores_d_mean) + + output = pl.dot(probabilities, v, precision=precision) + output_d = pl.dot(probabilities_d, v, precision=precision) + pl.dot( + probabilities, vd, precision=precision + ) + o_ref[:, :] = output.astype(o_ref.dtype) + od_ref[:, :] = output_d.astype(od_ref.dtype) + + +def mhsea_jvp_pallas(q, k, e, v, mask, qd, kd, ed, vd, *, precision): + batch_size, seq_len, num_heads, head_dim = q.shape + qkv_shape = (batch_size, seq_len, num_heads, head_dim) + edge_shape = (batch_size, seq_len, num_heads, seq_len) + qkv_spec = pl.BlockSpec( + (None, seq_len, None, head_dim), + lambda batch, head: (batch, 0, head, 0), + ) + edge_spec = pl.BlockSpec( + (None, seq_len, None, seq_len), + lambda batch, head: (batch, 0, head, 0), + ) + mask_spec = pl.BlockSpec((None, seq_len), lambda batch, _head: (batch, 0)) + kernel = pl.pallas_call( + partial(_mhsea_jvp_kernel, precision=precision), + grid=(batch_size, num_heads), + in_specs=[ + qkv_spec, + qkv_spec, + edge_spec, + qkv_spec, + mask_spec, + qkv_spec, + qkv_spec, + edge_spec, + qkv_spec, + ], + out_specs=[qkv_spec, qkv_spec], + out_shape=[ + jax.ShapeDtypeStruct(qkv_shape, jnp.float32), + jax.ShapeDtypeStruct(qkv_shape, jnp.float32), + ], + compiler_params=compiler_params(num_warps=4, num_stages=1), + debug=False, + interpret=False, + name="mhsea_full_ad_jvp", + ) + return kernel(q, k, e, v, mask, qd, kd, ed, vd) + + +__all__ = ["mhsea_jvp_pallas"] diff --git a/src/hamiltonzero/model/_pallas_attn/mhsea_tuned.py b/src/hamiltonzero/model/_pallas_attn/mhsea_tuned.py new file mode 100644 index 0000000000000000000000000000000000000000..2bca55ab654aa1724f3c05e36181e5481ed2157f --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/mhsea_tuned.py @@ -0,0 +1,258 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +import jax +import jax.core as _shaped_core +import jax.numpy as jnp +from jax.extend import core as _ext_core +from jax.interpreters import ad, batching, mlir +from .custom_gradients import mhsea_backward, mhsea_forward +from .mhsea_tuned_jvp import mhsea_tangent_pallas + +_SUPPORTED_TUNED_HEAD_DIMS = frozenset((8, 16, 32, 64)) +_TF32_PRECISION = jax.lax.Precision.DEFAULT + + +def _forward_config(seq_len: int) -> tuple[int, int, int]: + if seq_len <= 128: + default_q_block, default_warps, default_stages = (32, 4, 2) + else: + default_q_block, default_warps, default_stages = (16, 8, 3) + return (default_q_block, default_warps, default_stages) + + +_BACKWARD_CONFIG = (32, 4, 1) +_JVP_CONFIG = (32, 32, 4, 2) + + +def _validate_primals(q, k, e, v, mask) -> None: + assert q.ndim == k.ndim == v.ndim == 4 + assert e.ndim == 4 + assert mask.ndim == 2 + assert q.dtype == k.dtype == e.dtype == v.dtype, ( + f"mhsea_tuned_ad requires q/k/e/v to share one storage dtype; got {q.dtype}, {k.dtype}, {e.dtype}, {v.dtype}" + ) + assert jnp.dtype(q.dtype) == jnp.dtype(jnp.float32), ( + f"mhsea_tuned_ad currently requires float32 public tensors (TF32 dot execution); got {q.dtype}" + ) + batch_len, seq_len, num_heads, head_len = q.shape + assert k.shape == q.shape, f"k shape {k.shape} != q shape {q.shape}" + assert v.shape == q.shape, ( + f"mhsea_tuned_ad currently requires Vdim == q/k head dim and matching leading dimensions; got v={v.shape}, q={q.shape}" + ) + assert e.shape == (batch_len, seq_len, num_heads, seq_len), ( + f"e shape {e.shape} != ({batch_len}, {seq_len}, {num_heads}, {seq_len})" + ) + assert mask.shape == (batch_len, seq_len), ( + f"mask shape {mask.shape} != ({batch_len}, {seq_len})" + ) + assert head_len == v.shape[-1] + + +def _supports_tuned_shape(q) -> bool: + return ( + jnp.dtype(q.dtype) == jnp.dtype(jnp.float32) + and q.shape[1] >= 64 + and (q.shape[1] % 32 == 0) + and (q.shape[-1] in _SUPPORTED_TUNED_HEAD_DIMS) + ) + + +def _supports_full_ad_shape(q) -> bool: + seq_len = q.shape[1] + return ( + jnp.dtype(q.dtype) == jnp.dtype(jnp.float32) + and 0 < seq_len < 64 + and not (seq_len & (seq_len - 1)) + and (q.shape[-1] in _SUPPORTED_TUNED_HEAD_DIMS) + ) + + +mhsea_tuned_ad_p = _ext_core.Primitive("mhsea_tuned_ad") +mhsea_tuned_ad_p.multiple_results = True + + +def _outer_impl(q, k, e, v, mask): + _validate_primals(q, k, e, v, mask) + q_block_len, num_warps, num_stages = _forward_config(q.shape[1]) + out, residuals = mhsea_forward( + q, + k, + e, + v, + mask, + q_block_len=q_block_len, + num_warps=num_warps, + num_stages=num_stages, + precision=_TF32_PRECISION, + ) + lse = residuals[5] + out = jnp.where(mask[:, :, None, None].astype(jnp.bool_), out, 0.0) + return (out, lse) + + +def _outer_abstract_eval(q, k, e, v, mask): + _validate_primals(q, k, e, v, mask) + batch_len, seq_len, num_heads, _head_len = q.shape + return ( + _shaped_core.ShapedArray(v.shape, v.dtype), + _shaped_core.ShapedArray((batch_len, seq_len, num_heads), q.dtype), + ) + + +def _outer_jvp(primals, tangents): + q, k, e, v, mask = primals + qt, kt, et, vt, _mask_t = tangents + qt = ad.instantiate_zeros(qt) + kt = ad.instantiate_zeros(kt) + et = ad.instantiate_zeros(et) + vt = ad.instantiate_zeros(vt) + primal_out, lse = mhsea_tuned_ad_p.bind(q, k, e, v, mask) + tangent_out = mhsea_tuned_ad_lin_p.bind( + q, k, e, v, mask, primal_out, lse, qt, kt, et, vt + ) + return ( + (primal_out, lse), + (tangent_out, ad.Zero(_shaped_core.ShapedArray(lse.shape, lse.dtype))), + ) + + +def _merge_first_two(x): + return x.reshape(x.shape[0] * x.shape[1], *x.shape[2:]) + + +def _split_first_two(x, mapped_size: int): + return x.reshape(mapped_size, x.shape[0] // mapped_size, *x.shape[1:]) + + +def _prepare_batched_args(args, axes): + mapped_size = None + moved = [] + for x, axis in zip(args, axes): + if axis is None: + moved.append(x) + continue + x = jnp.moveaxis(x, axis, 0) + if mapped_size is None: + mapped_size = x.shape[0] + else: + assert x.shape[0] == mapped_size + moved.append(x) + if mapped_size is None: + return (None, None) + broadcasted = [ + jnp.broadcast_to(x[None], (mapped_size, *x.shape)) if axis is None else x + for x, axis in zip(moved, axes) + ] + return (mapped_size, [_merge_first_two(x) for x in broadcasted]) + + +def _outer_batch(args, axes): + mapped_size, flat = _prepare_batched_args(args, axes) + if mapped_size is None: + out, lse = mhsea_tuned_ad_p.bind(*args) + return ((out, lse), (None, None)) + out_flat, lse_flat = mhsea_tuned_ad_p.bind(*flat) + return ( + ( + _split_first_two(out_flat, mapped_size), + _split_first_two(lse_flat, mapped_size), + ), + (0, 0), + ) + + +def _outer_mlir(ctx, *args, **kwargs): + return mlir.lower_fun(_outer_impl, multiple_results=True)(ctx, *args, **kwargs) + + +mhsea_tuned_ad_p.def_impl(_outer_impl) +mhsea_tuned_ad_p.def_abstract_eval(_outer_abstract_eval) +ad.primitive_jvps[mhsea_tuned_ad_p] = _outer_jvp +batching.primitive_batchers[mhsea_tuned_ad_p] = _outer_batch +mlir.register_lowering(mhsea_tuned_ad_p, _outer_mlir) +mhsea_tuned_ad_lin_p = _ext_core.Primitive("mhsea_tuned_ad_lin") +mhsea_tuned_ad_lin_p.multiple_results = False + + +def _inner_impl(q, k, e, v, mask, primal_out, lse, qt, kt, et, vt): + q_block_len, kv_block_len, num_warps, num_stages = _JVP_CONFIG + return mhsea_tangent_pallas( + q, + k, + e, + v, + mask, + primal_out, + lse, + qt, + kt, + et, + vt, + q_block_len=q_block_len, + kv_block_len=kv_block_len, + num_warps=num_warps, + num_stages=num_stages, + precision=_TF32_PRECISION, + ) + + +def _inner_abstract_eval(q, k, e, v, mask, primal_out, lse, qt, kt, et, vt): + del q, k, e, mask, primal_out, lse, qt, kt, et, vt + return _shaped_core.ShapedArray(v.shape, v.dtype) + + +def _inner_transpose(cot, q, k, e, v, mask, primal_out, lse, qt, kt, et, vt): + if isinstance(cot, ad.Zero): + zeros = tuple( + (ad.Zero(_shaped_core.ShapedArray(x.shape, x.dtype)) for x in (q, k, e, v)) + ) + return (None,) * 7 + zeros + q_block_len, num_warps, num_stages = _BACKWARD_CONFIG + dq, dk, de, dv = mhsea_backward( + q_block_len, + num_warps, + num_stages, + _TF32_PRECISION, + (q, k, e, v, mask, lse, primal_out), + cot, + ) + return (None,) * 7 + (dq, dk, de, dv) + + +def _inner_batch(args, axes): + mapped_size, flat = _prepare_batched_args(args, axes) + if mapped_size is None: + return (mhsea_tuned_ad_lin_p.bind(*args), None) + out_flat = mhsea_tuned_ad_lin_p.bind(*flat) + return (_split_first_two(out_flat, mapped_size), 0) + + +def _inner_mlir(ctx, *args, **kwargs): + return mlir.lower_fun(_inner_impl, multiple_results=False)(ctx, *args, **kwargs) + + +mhsea_tuned_ad_lin_p.def_impl(_inner_impl) +mhsea_tuned_ad_lin_p.def_abstract_eval(_inner_abstract_eval) +ad.primitive_transposes[mhsea_tuned_ad_lin_p] = _inner_transpose +batching.primitive_batchers[mhsea_tuned_ad_lin_p] = _inner_batch +mlir.register_lowering(mhsea_tuned_ad_lin_p, _inner_mlir) + + +def mhsea_with_tuned_ad(q, k, e, v, mask) -> jax.Array: + if not _supports_tuned_shape(q): + if not _supports_full_ad_shape(q): + raise ValueError( + f"mhsea requires a power-of-two N and head dimension in {sorted(_SUPPORTED_TUNED_HEAD_DIMS)}" + ) + from .mhsea_full_ad import mhsea_with_full_ad + + out = mhsea_with_full_ad(q, k, e, v, mask) + return jnp.where(mask[:, :, None, None].astype(jnp.bool_), out, 0.0) + _validate_primals(q, k, e, v, mask) + out, _lse = mhsea_tuned_ad_p.bind(q, k, e, v, mask) + return out + + +__all__ = ["mhsea_with_tuned_ad"] diff --git a/src/hamiltonzero/model/_pallas_attn/mhsea_tuned_jvp.py b/src/hamiltonzero/model/_pallas_attn/mhsea_tuned_jvp.py new file mode 100644 index 0000000000000000000000000000000000000000..dadf1525015bdaa56dadb2f90fb1acab1ac01399 --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/mhsea_tuned_jvp.py @@ -0,0 +1,222 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +from functools import partial +import jax +import jax.numpy as jnp +from jax.experimental import pallas as pl +from .utils import compiler_params + +_MIN_SEQUENCE_LENGTH = 64 +_SUPPORTED_HEAD_DIMS = frozenset((8, 16, 32, 64)) +_INTERNAL_DTYPE = jnp.dtype(jnp.float32) + + +def _mask_value() -> jax.Array: + return jnp.float32(-10000.0) + + +def _mhsea_tangent_kernel( + q_ref, + k_ref, + e_ref, + v_ref, + mask_ref, + primal_out_ref, + lse_ref, + qt_ref, + kt_ref, + et_ref, + vt_ref, + tangent_ref, + *, + q_block_len: int, + kv_block_len: int, + precision, +): + q_tile = pl.program_id(1) + head_dim = q_ref.shape[-1] + q_slice = pl.dslice(q_tile * q_block_len, q_block_len) + q_mask = mask_ref[q_slice] + q_zero = jnp.asarray(0, dtype=q_ref.dtype) + qt_zero = jnp.asarray(0, dtype=qt_ref.dtype) + q = jnp.where(q_mask[:, None], q_ref[:, :], q_zero) + qt = jnp.where(q_mask[:, None], qt_ref[:, :], qt_zero) + mean = jnp.zeros((q_block_len,), dtype=jnp.float32) + score_d_v = jnp.zeros((q_block_len, head_dim), dtype=jnp.float32) + p_vt = jnp.zeros((q_block_len, head_dim), dtype=jnp.float32) + lse = lse_ref[:].astype(jnp.float32) + + def _visit_kv_tile(kv_tile, carry): + mean, score_d_v, p_vt = carry + kv_slice = pl.dslice(kv_tile * kv_block_len, kv_block_len) + kv_mask = mask_ref[kv_slice] + square_mask = q_mask[:, None] & kv_mask[None, :] + k_zero = jnp.asarray(0, dtype=k_ref.dtype) + v_zero = jnp.asarray(0, dtype=v_ref.dtype) + kt_zero = jnp.asarray(0, dtype=kt_ref.dtype) + vt_zero = jnp.asarray(0, dtype=vt_ref.dtype) + k = jnp.where(kv_mask[:, None], k_ref[kv_slice, :], k_zero) + v = jnp.where(kv_mask[:, None], v_ref[kv_slice, :], v_zero) + kt = jnp.where(kv_mask[:, None], kt_ref[kv_slice, :], kt_zero) + vt = jnp.where(kv_mask[:, None], vt_ref[kv_slice, :], vt_zero) + e = e_ref[:, kv_slice] + et = et_ref[:, kv_slice] + score_raw = pl.dot(q, k, trans_b=True, precision=precision) + e + score = jnp.where(square_mask, score_raw, _mask_value()) + p = jnp.exp(score.astype(jnp.float32) - lse[:, None]) + score_d_raw = ( + pl.dot(qt, k, trans_b=True, precision=precision) + + pl.dot(q, kt, trans_b=True, precision=precision) + + et + ) + score_d = jnp.where(square_mask, score_d_raw, jnp.float32(0)) + p_score_d = p * score_d.astype(jnp.float32) + mean += jnp.sum(p_score_d, axis=1, dtype=jnp.float32) + score_d_v += pl.dot( + p_score_d.astype(v_ref.dtype), v, precision=precision + ).astype(jnp.float32) + p_vt += pl.dot(p.astype(vt_ref.dtype), vt, precision=precision).astype( + jnp.float32 + ) + return (mean, score_d_v, p_vt) + + mean, score_d_v, p_vt = jax.lax.fori_loop( + 0, k_ref.shape[0] // kv_block_len, _visit_kv_tile, (mean, score_d_v, p_vt) + ) + primal_out = primal_out_ref[:, :].astype(jnp.float32) + tangent = score_d_v + p_vt - mean[:, None] * primal_out + tangent = jnp.where(q_mask[:, None], tangent, jnp.float32(0)) + tangent_ref[:, :] = tangent.astype(tangent_ref.dtype) + + +def _require_shape(name: str, value: jax.Array, expected: tuple[int, ...]) -> None: + if value.shape != expected: + raise ValueError(f"{name}.shape must be {expected}, got {value.shape}") + + +def _require_external_f32(name: str, value: jax.Array) -> None: + if jnp.dtype(value.dtype) != jnp.dtype(jnp.float32): + raise TypeError(f"{name}.dtype must be float32, got {value.dtype}") + + +def mhsea_tangent_pallas( + q, + k, + e, + v, + mask, + primal_out, + lse, + qt, + kt, + et, + vt, + *, + q_block_len, + kv_block_len, + num_warps, + num_stages, + precision, +): + if q.ndim != 4: + raise ValueError(f"q must have rank 4 [B, N, H, D], got shape {q.shape}") + batch_size, seq_len, num_heads, head_dim = q.shape + qkv_shape = (batch_size, seq_len, num_heads, head_dim) + edge_shape = (batch_size, seq_len, num_heads, seq_len) + if seq_len < _MIN_SEQUENCE_LENGTH: + raise ValueError( + f"sequence length must be >= {_MIN_SEQUENCE_LENGTH}, got {seq_len}; q_block_len and kv_block_len must divide it exactly" + ) + if head_dim not in _SUPPORTED_HEAD_DIMS: + raise ValueError( + f"head dimension must be one of {sorted(_SUPPORTED_HEAD_DIMS)}, got {head_dim}" + ) + if q_block_len <= 0 or seq_len % q_block_len: + raise ValueError( + f"q_block_len must be a positive divisor of N={seq_len}, got {q_block_len}" + ) + if kv_block_len <= 0 or seq_len % kv_block_len: + raise ValueError( + f"kv_block_len must be a positive divisor of N={seq_len}, got {kv_block_len}" + ) + if num_warps <= 0: + raise ValueError(f"num_warps must be positive, got {num_warps}") + if num_stages <= 0: + raise ValueError(f"num_stages must be positive, got {num_stages}") + internal_dtype = _INTERNAL_DTYPE + for name, value in ( + ("q", q), + ("k", k), + ("v", v), + ("primal_out", primal_out), + ("qt", qt), + ("kt", kt), + ("vt", vt), + ): + _require_shape(name, value, qkv_shape) + _require_external_f32(name, value) + for name, value in (("e", e), ("et", et)): + _require_shape(name, value, edge_shape) + _require_external_f32(name, value) + _require_shape("mask", mask, (batch_size, seq_len)) + _require_shape("lse", lse, (batch_size, seq_len, num_heads)) + _require_external_f32("lse", lse) + q_like_spec = pl.BlockSpec( + (None, q_block_len, None, head_dim), + lambda batch, q_tile, head: (batch, q_tile, head, 0), + ) + kv_like_spec = pl.BlockSpec( + (None, seq_len, None, head_dim), + lambda batch, _q_tile, head: (batch, 0, head, 0), + ) + edge_spec = pl.BlockSpec( + (None, q_block_len, None, seq_len), + lambda batch, q_tile, head: (batch, q_tile, head, 0), + ) + mask_spec = pl.BlockSpec((None, seq_len), lambda batch, _q_tile, _head: (batch, 0)) + lse_spec = pl.BlockSpec( + (None, q_block_len, None), lambda batch, q_tile, head: (batch, q_tile, head) + ) + kernel = pl.pallas_call( + partial( + _mhsea_tangent_kernel, + q_block_len=q_block_len, + kv_block_len=kv_block_len, + precision=precision, + ), + grid=(batch_size, seq_len // q_block_len, num_heads), + in_specs=[ + q_like_spec, + kv_like_spec, + edge_spec, + kv_like_spec, + mask_spec, + q_like_spec, + lse_spec, + q_like_spec, + kv_like_spec, + edge_spec, + kv_like_spec, + ], + out_specs=q_like_spec, + out_shape=jax.ShapeDtypeStruct(qkv_shape, jnp.float32), + compiler_params=compiler_params(num_warps=num_warps, num_stages=num_stages), + debug=False, + interpret=False, + name="mhsea_tuned_tangent", + ) + return kernel( + q.astype(internal_dtype), + k.astype(internal_dtype), + e.astype(internal_dtype), + v.astype(internal_dtype), + mask.astype(jnp.bool_), + primal_out, + lse, + qt.astype(internal_dtype), + kt.astype(internal_dtype), + et.astype(internal_dtype), + vt.astype(internal_dtype), + ) diff --git a/src/hamiltonzero/model/_pallas_attn/utils.py b/src/hamiltonzero/model/_pallas_attn/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..16ee28c18d4044e8896daba5770d79cbc12ab432 --- /dev/null +++ b/src/hamiltonzero/model/_pallas_attn/utils.py @@ -0,0 +1,64 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# Modifications copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: MIT + +import jax +import jax.numpy as jnp +from jax.experimental import pallas as pl + +try: + from jax.experimental.pallas import triton as plgpu +except ImportError: + from jax.experimental.pallas import gpu as plgpu +from packaging.version import Version + + +def sum_columns(x: jax.Array) -> jax.Array: + return x.astype(jnp.float32).sum(axis=1, keepdims=True, dtype=jnp.float32) + + +def get_query_block_spec(block_len: int, width: int): + return pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k, 0), + block_shape=(None, block_len, None, width), + ) + + +def get_key_value_block_spec(seq_len: int, width: int): + return pl.BlockSpec( + index_map=lambda i, _j, k: (i, 0, k, 0), + block_shape=(None, seq_len, None, width), + ) + + +def get_mask_block_spec(seq_len: int): + return pl.BlockSpec(index_map=lambda i, _j, _k: (i, 0), block_shape=(None, seq_len)) + + +def get_lse_block_spec(block_len: int) -> pl.BlockSpec: + return pl.BlockSpec( + index_map=lambda i, j, k: (i, j, k), block_shape=(None, block_len, None) + ) + + +def create_grid( + batch_len: int, seq_len: int, num_heads: int, q_block_len: int +) -> tuple[int, int, int]: + return (batch_len, seq_len // q_block_len, num_heads) + + +def big_number() -> float: + return jnp.float32(-10000.0) + + +def compiler_params(num_warps, num_stages): + if Version(jax.__version__) >= Version("0.4.34"): + if hasattr(plgpu, "CompilerParams"): + return plgpu.CompilerParams(num_warps=num_warps, num_stages=num_stages) + elif hasattr(plgpu, "TritonCompilerParams"): + return plgpu.TritonCompilerParams( + num_warps=num_warps, num_stages=num_stages + ) + else: + return dict(triton=dict(num_warps=num_warps, num_stages=num_stages)) diff --git a/src/hamiltonzero/model/api.py b/src/hamiltonzero/model/api.py new file mode 100644 index 0000000000000000000000000000000000000000..701cbdd45969523b8ebe9df51494bedf5bbfcef6 --- /dev/null +++ b/src/hamiltonzero/model/api.py @@ -0,0 +1,148 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +from .model import ( + SpinAnsatz, + _normalize_leaf_carriers, +) +from .tree import ( + _balanced_subtree_mask, + _quadrilinear_merge, + _tagged_dense, + _tagged_dense_no_bias, + _tagged_rms_eqx_style, + _tree_active_clock_depth, + _tree_depth_count_features, + _tree_sphere, + edge_merge_masked, +) + + +normalize_leaf_carriers = _normalize_leaf_carriers +quadrilinear_merge = _quadrilinear_merge +tagged_dense_no_bias = _tagged_dense_no_bias +tree_sphere = _tree_sphere +balanced_subtree_mask = _balanced_subtree_mask +tree_active_clock_depth = _tree_active_clock_depth +tree_depth_count_features = _tree_depth_count_features +tagged_dense = _tagged_dense +tagged_rms_eqx_style = _tagged_rms_eqx_style + + +def _attention_name(value: str) -> str: + if value == "tuned": + return "mhsea_tuned" + if value == "einsum": + return "einsum" + raise ValueError("attention must be 'tuned' or 'einsum'") + + +def build_model(config: Any, key, *, n_max: int) -> SpinAnsatz: + attention = _attention_name(str(config.attention)) + model = SpinAnsatz( + d_e=int(config.d_e), + d_o=int(config.d_o), + d_c=int(config.d_c), + d_r=int(config.d_r), + n_heads=int(config.n_heads), + n_layers=int(config.n_layers), + rank=int(config.rank), + n_edge=int(config.edge_channels), + d_e_attn=int(config.attention_qk_dim), + d_c_attn=int(config.attention_v_dim), + trunk_edge_node_ctx_dim=int(config.trunk_edge_node_context_dim), + trunk_edge_hidden_dim=int(config.trunk_edge_hidden_dim), + trunk_attn_bias_hidden_dim=int(config.trunk_attention_bias_hidden_dim), + trunk_ffn_hidden_dim=int(config.trunk_ffn_hidden_dim), + trunk_two_hop_hidden_dim=int(config.trunk_two_hop_hidden_dim), + tree_edge_node_ctx_dim=int(config.tree_edge_node_context_dim), + attn_impl=attention, + global_d_g=int(config.global_dim), + d_m_merge=int(config.merge_dim), + merge_chain_hypernet_rank=int(config.merge_hypernet_rank), + feat_d_bond=int(config.featurizer_bond_dim), + feat_n_heads=int(config.featurizer_heads), + feat_head_dim=int(config.featurizer_head_dim), + feat_n_global_q=int(config.featurizer_global_queries), + feat_edge_hidden_dim=int(config.featurizer_edge_hidden_dim), + feat_zeeman_hidden_dim=int(config.featurizer_zeeman_hidden_dim), + feat_global_hidden_dim=int(config.featurizer_global_hidden_dim), + feat_combine_hidden_dim=int(config.featurizer_combine_hidden_dim), + feat_token_initial_scale=float(config.featurizer_token_initial_scale), + feat_d_edge=int(config.edge_channels), + polar_group_norm_tau=float(config.polar_group_norm_tau), + polar_group_norm_bond_hidden=int(config.polar_bond_hidden_dim), + polar_group_norm_n_bond_groups=int(config.polar_bond_groups), + polar_group_norm_d_bond_group=int(config.polar_bond_group_dim), + polar_group_norm_n_zeeman_groups=int(config.polar_zeeman_groups), + polar_group_norm_d_zeeman_group=int(config.polar_zeeman_group_dim), + route_pointer_max_n=max(int(config.router_max_n), int(n_max)), + route_pointer_d_model=int(config.router_model_dim), + route_pointer_n_heads=int(config.router_heads), + route_pointer_attn_dim=int(config.router_attention_dim), + route_pointer_score_dim=int(config.router_score_dim), + route_pointer_candidate_hidden=int(config.router_candidate_dim), + route_pointer_summary_hidden=int(config.router_summary_dim), + route_pointer_ffn_hidden=int(config.router_ffn_dim), + route_pointer_score_init_scale=float(config.router_score_initial_scale), + route_pointer_rope_base=float(config.router_rope_base), + route_pointer_rope_scaling=float(config.router_rope_scaling), + route_tree_prefix_layers=int(config.router_tree_prefix_layers), + route_tree_prefix_candidate_layers=int(config.router_tree_candidate_layers), + route_tree_prefix_merge_hidden=int(config.router_tree_merge_dim), + route_tree_prefix_post_prefix_suffix_layers=int(config.router_tree_post_layers), + route_contextualizer_layers=int(config.router_context_layers), + route_contextualizer_n_heads=int(config.router_context_heads), + route_contextualizer_attn_dim=int(config.router_context_attention_dim), + route_contextualizer_edge_node_ctx_dim=int(config.router_context_edge_node_dim), + level_edge_attn_n_heads=int(config.level_edge_heads), + level_edge_attn_edge_mlp_hidden=int(config.level_edge_mlp_dim), + level_edge_attn_edge_mlp_n_blocks=int(config.level_edge_mlp_blocks), + level_edge_attn_ffn_d_hidden=int(config.level_edge_ffn_dim), + level_edge_attn_rope_base=float(config.level_edge_rope_base), + level_edge_attn_rope_scaling=float(config.level_edge_rope_scaling), + root_readout_edge_rank=int(config.root_readout_edge_rank), + ngpt_alpha_initial=float(config.ngpt_alpha_initial), + ngpt_alpha_initial_fraction=float(config.ngpt_alpha_initial_fraction), + ngpt_alpha_maximum=float(config.ngpt_alpha_maximum), + global_ladder_tap_dim=int(config.global_ladder_tap_dim), + level_edge_attn_bias_mlp_hidden=int(config.level_edge_bias_mlp_dim), + level_edge_attn_bias_mlp_n_blocks=int(config.level_edge_bias_mlp_blocks), + merge_c_mlp_hidden=int(config.merge_context_mlp_dim), + readout_leaf_context_layers=int(config.readout_context_layers), + readout_leaf_context_n_heads=int(config.readout_context_heads), + readout_leaf_context_attn_dim=int(config.readout_context_attention_dim), + readout_leaf_context_edge_node_ctx_dim=int( + config.readout_context_edge_node_dim + ), + readout_leaf_context_summary_hidden=int(config.readout_context_summary_dim), + readout_leaf_context_mlp_hidden=int(config.readout_context_mlp_dim), + readout_leaf_context_bias_hidden=int(config.readout_context_bias_dim), + readout_leaf_context_edge_ffn_hidden=int(config.readout_context_edge_ffn_dim), + readout_leaf_context_rope_base=float(config.readout_context_rope_base), + readout_leaf_context_rope_scaling=float(config.readout_context_rope_scaling), + two_hop_channels=int(config.two_hop_channels), + tree_edge_fwl_channels=int(config.tree_fwl_channels), + key=key, + ) + return model + + +__all__ = [ + "SpinAnsatz", + "balanced_subtree_mask", + "build_model", + "edge_merge_masked", + "normalize_leaf_carriers", + "quadrilinear_merge", + "tagged_dense", + "tagged_dense_no_bias", + "tagged_rms_eqx_style", + "tree_active_clock_depth", + "tree_depth_count_features", + "tree_sphere", +] diff --git a/src/hamiltonzero/model/context.py b/src/hamiltonzero/model/context.py new file mode 100644 index 0000000000000000000000000000000000000000..c2ecbeb6175e6fde9063cc3003bde30502bdc105 --- /dev/null +++ b/src/hamiltonzero/model/context.py @@ -0,0 +1,421 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +import jax +import jax.numpy as jnp +import numpy as np +from jaxtyping import Array, Float, Int + + +def _balanced_mask(mask: Int[Array, "n"]) -> Int[Array, "n"]: + width = int(mask.shape[-1]) + n_real = jnp.sum(mask.astype(jnp.int32)) + max_power = max(1, (width - 1).bit_length()) + powers = 2 ** jnp.arange(max_power + 1, dtype=jnp.int32) + sentinel = jnp.asarray(1 << 30, dtype=jnp.int32) + next_power = jnp.min(jnp.where(powers >= jnp.maximum(n_real, 1), powers, sentinel)) + return (jnp.arange(width, dtype=jnp.int32) < next_power).astype(jnp.int32) + + +_EPS_ABC = jnp.asarray( + [ + [[0.0, 0.0, 0.0], [0.0, 0.0, 1.0], [0.0, -1.0, 0.0]], + [[0.0, 0.0, -1.0], [0.0, 0.0, 0.0], [1.0, 0.0, 0.0]], + [[0.0, 1.0, 0.0], [-1.0, 0.0, 0.0], [0.0, 0.0, 0.0]], + ], + dtype=jnp.float32, +) + + +def _real_dtype(dtype): + return jnp.real(jnp.zeros((), dtype)).dtype + + +def _host_eigvalsh(x): + out_dtype = _real_dtype(x.dtype) + out_shape = jax.ShapeDtypeStruct(x.shape[:-1], out_dtype) + + def callback(a): + values = np.linalg.eigvalsh(np.asarray(a)) + return values.astype(np.dtype(out_dtype)) + + return jax.pure_callback( + callback, + out_shape, + x, + vmap_method="sequential", + ) + + +def _compute_J_double_prime_batched( + J_full: Float[Array, "s n n 3 3"], + h: Float[Array, "s n 3"], + mask: Int[Array, "s n"], +) -> tuple[Float[Array, "s n n 10"], Float[Array, "s"]]: + n = J_full.shape[1] + n_systems = J_full.shape[0] + dtype = J_full.dtype + + J_filled = (J_full + jnp.conj(jnp.transpose(J_full, (0, 2, 1, 4, 3)))) / 2.0 + + eps_abc = _EPS_ABC.astype(_real_dtype(dtype)) + M_h = jnp.einsum("abc,sic->siab", eps_abc, h.astype(eps_abc.dtype)) * 0.5 + M_h = M_h * mask.astype(M_h.dtype)[:, :, None, None] + diagonal = jnp.arange(n) + + complex_dtype = jnp.result_type(J_filled.dtype, jnp.complex64) + M_h_diagonal = ( + jnp.zeros(J_filled.shape, dtype=complex_dtype) + .at[:, diagonal, diagonal] + .set((2j * M_h).astype(complex_dtype)) + ) + J_for_norm = J_filled.astype(complex_dtype) + M_h_diagonal + J_matrix = jnp.transpose(J_for_norm, (0, 1, 3, 2, 4)).reshape( + n_systems, 3 * n, 3 * n + ) + eigh_epsilon = jnp.asarray(1e-6, dtype=_real_dtype(J_matrix.dtype)) + eye_3n = jnp.eye(3 * n, dtype=J_matrix.dtype) + eigenvalues = _host_eigvalsh(J_matrix + eigh_epsilon * eye_3n[None]) - eigh_epsilon + s_norm = jnp.maximum( + jnp.max(jnp.abs(eigenvalues), axis=-1), + jnp.asarray(1e-12, eigenvalues.dtype), + ).astype(dtype) + + J_normalized = J_filled / s_norm[:, None, None, None, None] + J_normalized = J_normalized.at[:, diagonal, diagonal, :, :].set(0.0) + J_flat = jnp.real(J_normalized).reshape(n_systems, n, n, 9).astype(dtype) + identity_column = jnp.broadcast_to( + jnp.eye(n, dtype=dtype)[None, ..., None], + (n_systems, n, n, 1), + ) + return jnp.concatenate([J_flat, identity_column], axis=-1), s_norm + + +def _compute_J_double_prime( + J_full: Float[Array, "n n 3 3"], + h: Float[Array, "n 3"], + mask: Int[Array, "n"], +) -> tuple[Float[Array, "n n 10"], Float[Array, ""]]: + J_double_prime, s_norm = _compute_J_double_prime_batched( + J_full[None], h[None], mask[None] + ) + return J_double_prime[0], s_norm[0] + + +class SpinContext(eqx.Module): + mask: Int[Array, "n"] + bmask: Int[Array, "n"] + J_double_prime: Float[Array, "n n 10"] + s_norm: Float[Array, ""] + h_prime: Float[Array, "n 3"] + route_quotient_node_key: Int[Array, "n"] + route_quotient_edge_key: Int[Array, "n n"] + needs_fwl2: Array + route_perm: Int[Array, "n"] + + def __init__( + self, + J_full: Float[Array, "n n 3 3"], + h: Float[Array, "n 3"], + mask: Int[Array, "n"], + *, + needs_fwl2: Array | bool, + ) -> None: + from .route_quotient import route_quotient_keys + + self.mask = mask.astype(jnp.int32) + self.bmask = _balanced_mask(self.mask) + self.J_double_prime, self.s_norm = _compute_J_double_prime(J_full, h, self.mask) + self.h_prime = h / jnp.real(self.s_norm).astype(h.dtype) + ( + self.route_quotient_node_key, + self.route_quotient_edge_key, + ) = route_quotient_keys(J_full, h, self.mask, self.bmask) + self.needs_fwl2 = jnp.asarray(needs_fwl2, dtype=jnp.bool_) + self.route_perm = jnp.arange(self.mask.shape[0], dtype=jnp.int32) + + @classmethod + def from_precomputed( + cls, + *, + mask, + bmask, + J_double_prime, + s_norm, + h_prime, + route_quotient_node_key, + route_quotient_edge_key, + needs_fwl2, + route_perm, + ) -> "SpinContext": + self = object.__new__(cls) + fields = { + "mask": mask, + "bmask": bmask, + "J_double_prime": J_double_prime, + "s_norm": s_norm, + "h_prime": h_prime, + "route_quotient_node_key": route_quotient_node_key, + "route_quotient_edge_key": route_quotient_edge_key, + "needs_fwl2": needs_fwl2, + "route_perm": route_perm, + } + for name, value in fields.items(): + dtype = ( + jnp.bool_ + if name == "needs_fwl2" + else jnp.int32 + if name + in { + "mask", + "bmask", + "route_quotient_node_key", + "route_quotient_edge_key", + "route_perm", + } + else None + ) + object.__setattr__(self, name, jnp.asarray(value, dtype=dtype)) + return self + + @property + def n_sites(self) -> int: + return int(self.mask.shape[0]) + + +class MultiSystemContext(eqx.Module): + mask: Int[Array, "s n"] + bmask: Int[Array, "s n"] + J_double_prime: Float[Array, "s n n 10"] + s_norm: Float[Array, "s"] + h_prime: Float[Array, "s n 3"] + route_quotient_node_key: Int[Array, "s n"] + route_quotient_edge_key: Int[Array, "s n n"] + needs_fwl2: Array + route_perm: Int[Array, "s n"] + + def __init__( + self, + J_full: Float[Array, "s n n 3 3"], + h: Float[Array, "s n 3"], + mask: Int[Array, "s n"], + *, + needs_fwl2: Array | bool, + ) -> None: + from .route_quotient import route_quotient_keys + + self.mask = mask.astype(jnp.int32) + self.bmask = jax.vmap(_balanced_mask)(self.mask) + self.J_double_prime, self.s_norm = _compute_J_double_prime_batched( + J_full, h, self.mask + ) + self.h_prime = h / jnp.real(self.s_norm).astype(h.dtype)[:, None, None] + n_systems = self.mask.shape[0] + ( + self.route_quotient_node_key, + self.route_quotient_edge_key, + ) = jax.jit(jax.vmap(route_quotient_keys))(J_full, h, self.mask, self.bmask) + self.needs_fwl2 = jnp.broadcast_to( + jnp.asarray(needs_fwl2, dtype=jnp.bool_), + (n_systems,), + ) + self.route_perm = jnp.broadcast_to( + jnp.arange(self.mask.shape[1], dtype=jnp.int32)[None, :], + self.mask.shape, + ) + + @classmethod + def from_precomputed( + cls, + *, + mask, + bmask, + J_double_prime, + s_norm, + h_prime, + route_quotient_node_key, + route_quotient_edge_key, + needs_fwl2, + route_perm, + ) -> "MultiSystemContext": + self = object.__new__(cls) + fields = { + "mask": mask, + "bmask": bmask, + "J_double_prime": J_double_prime, + "s_norm": s_norm, + "h_prime": h_prime, + "route_quotient_node_key": route_quotient_node_key, + "route_quotient_edge_key": route_quotient_edge_key, + "needs_fwl2": needs_fwl2, + "route_perm": route_perm, + } + for name, value in fields.items(): + dtype = ( + jnp.bool_ + if name == "needs_fwl2" + else jnp.int32 + if name + in { + "mask", + "bmask", + "route_quotient_node_key", + "route_quotient_edge_key", + "route_perm", + } + else None + ) + object.__setattr__(self, name, jnp.asarray(value, dtype=dtype)) + return self + + @classmethod + def from_single(cls, context: SpinContext) -> "MultiSystemContext": + return cls.from_precomputed( + mask=context.mask[None], + bmask=context.bmask[None], + J_double_prime=context.J_double_prime[None], + s_norm=context.s_norm[None], + h_prime=context.h_prime[None], + route_quotient_node_key=context.route_quotient_node_key[None], + route_quotient_edge_key=context.route_quotient_edge_key[None], + needs_fwl2=context.needs_fwl2[None], + route_perm=context.route_perm[None], + ) + + @classmethod + def stack(cls, contexts: list[SpinContext]) -> "MultiSystemContext": + if not contexts: + raise ValueError("MultiSystemContext.stack requires a context") + widths = [int(context.mask.shape[0]) for context in contexts] + n_max = max(widths) + + def pad_sites(value, n): + return ( + value + if n == n_max + else jnp.pad(value, ((0, n_max - n),) + ((0, 0),) * (value.ndim - 1)) + ) + + def pad_pairs(value, n): + padding = n_max - n + return ( + value + if padding == 0 + else jnp.pad( + value, + ((0, padding), (0, padding)) + ((0, 0),) * (value.ndim - 2), + ) + ) + + mask = jnp.stack( + [ + pad_sites(context.mask, width) + for context, width in zip(contexts, widths, strict=True) + ] + ) + bmask = jax.vmap(_balanced_mask)(mask) + J_double_prime = jnp.stack( + [ + pad_pairs(context.J_double_prime, width) + for context, width in zip(contexts, widths, strict=True) + ] + ) + diagonal = jnp.arange(n_max, dtype=jnp.int32) + J_double_prime = J_double_prime.at[:, diagonal, diagonal, 9].set(1.0) + s_norm = jnp.stack([context.s_norm for context in contexts]) + h_prime = jnp.stack( + [ + pad_sites(context.h_prime, width) + for context, width in zip(contexts, widths, strict=True) + ] + ) + + edge_shapes = { + tuple(context.route_quotient_edge_key.shape) for context in contexts + } + compact_edges = all(shape == (0, 0) for shape in edge_shapes) + full_edges = all( + shape == (width, width) + for shape, width in zip( + [context.route_quotient_edge_key.shape for context in contexts], + widths, + strict=True, + ) + ) + if not (compact_edges or full_edges): + raise ValueError("cannot stack mixed quotient edge carriers") + route_edge = ( + jnp.zeros((len(contexts), 0, 0), dtype=jnp.int32) + if compact_edges + else jnp.stack( + [ + pad_pairs(context.route_quotient_edge_key, width) + for context, width in zip(contexts, widths, strict=True) + ] + ) + ) + + def pad_node_key(value, n): + padding = n_max - n + return ( + value + if padding == 0 + else jnp.pad(value, ((0, padding),), constant_values=-1) + ) + + def pad_perm(value, n): + if n == n_max: + return value + return jnp.concatenate([value, jnp.arange(n, n_max, dtype=value.dtype)]) + + return cls.from_precomputed( + mask=mask, + bmask=bmask, + J_double_prime=J_double_prime, + s_norm=s_norm, + h_prime=h_prime, + route_quotient_node_key=jnp.stack( + [ + pad_node_key(context.route_quotient_node_key, width) + for context, width in zip(contexts, widths, strict=True) + ] + ), + route_quotient_edge_key=route_edge, + needs_fwl2=jnp.stack([context.needs_fwl2 for context in contexts]), + route_perm=jnp.stack( + [ + pad_perm(context.route_perm, width) + for context, width in zip(contexts, widths, strict=True) + ] + ), + ) + + def select(self, system_id: int) -> SpinContext: + return SpinContext.from_precomputed( + mask=self.mask[system_id], + bmask=self.bmask[system_id], + J_double_prime=self.J_double_prime[system_id], + s_norm=self.s_norm[system_id], + h_prime=self.h_prime[system_id], + route_quotient_node_key=self.route_quotient_node_key[system_id], + route_quotient_edge_key=self.route_quotient_edge_key[system_id], + needs_fwl2=self.needs_fwl2[system_id], + route_perm=self.route_perm[system_id], + ) + + @property + def n_systems(self) -> int: + return int(self.mask.shape[0]) + + @property + def n_sites(self) -> int: + return int(self.mask.shape[1]) + + +__all__ = [ + "MultiSystemContext", + "SpinContext", +] diff --git a/src/hamiltonzero/model/featurizer.py b/src/hamiltonzero/model/featurizer.py new file mode 100644 index 0000000000000000000000000000000000000000..455202abf6f53ea0801b3618d01219ab2715a31f --- /dev/null +++ b/src/hamiltonzero/model/featurizer.py @@ -0,0 +1,1119 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray +from hamiltonzero.model.odd_ops import BiasFreeLinear, Linear, _RMS +from hamiltonzero.model.fused_silu import fused_silu + + +class GroupLayerNorm(eqx.Module): + weight: Float[Array, "n_groups d_group"] + eps: float = eqx.field(static=True, default=1e-05) + _use_id: str = eqx.field(static=True, default="") + n_groups: int = eqx.field(static=True) + d_group: int = eqx.field(static=True) + + def __init__(self, n_groups: int, d_group: int, eps: float = 1e-05): + self.weight = jnp.ones((n_groups, d_group)) + self.eps = float(eps) + self._use_id = "" + self.n_groups = int(n_groups) + self.d_group = int(d_group) + + def __call__( + self, + x: Float[Array, "... n_groups d_group"], + *, + pathway: str | None = None, + kfac_structural_mask=None, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ) -> Float[Array, "... n_groups d_group"]: + if x.shape[-2] != self.n_groups or x.shape[-1] != self.d_group: + raise ValueError( + f"GroupLayerNorm expected trailing shape ({self.n_groups}, {self.d_group}), got {x.shape[-2:]}." + ) + if pathway is None: + pathway = "even" + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + from hamiltonzero.model.tree import _kfac_name_kw + + out_cdtype = ( + _compute_dtype() if pathway in ("even", "hypernet_eside") else jnp.float32 + ) + stats_dtype = jnp.promote_types(jnp.float32, x.dtype) + x_hi = x.astype(stats_dtype) if x.dtype != stats_dtype else x + ms = jnp.mean(x_hi * x_hi, axis=-1, keepdims=True) + normalized_hi = x_hi * jax.lax.rsqrt(ms + self.eps) + normalized = ( + normalized_hi.astype(out_cdtype) + if normalized_hi.dtype != out_cdtype + else normalized_hi + ) + w = ( + self.weight.astype(out_cdtype) + if self.weight.dtype != out_cdtype + else self.weight + ) + y = normalized * w + from hamiltonzero.optim.blocks import ( + register_structural_trailing_stacked_scale_and_shift, + ) + + return register_structural_trailing_stacked_scale_and_shift( + y, + normalized, + kfac_structural_mask, + self.weight, + repeat_ndim=kfac_repeat_ndim, + context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + **_kfac_name_kw(self._use_id), + ) + + +class SystemFeaturizer(eqx.Module): + bond_group_w1: Linear + bond_group_w2: Linear + bond_group_post: Linear + bond_group_out: Linear + bond_group_ln: GroupLayerNorm + Q_col: Float[Array, "n_heads head_dim"] + K_col: BiasFreeLinear + V_col: BiasFreeLinear + ln_bond: _RMS + Q_row: Float[Array, "n_global_q n_heads head_dim"] + K_row: BiasFreeLinear + V_row: BiasFreeLinear + ln_local_pre_row: _RMS + ln_edge_cond: _RMS + edge_cond_w1: Linear + edge_cond_w2: Linear + ln_edge_global: _RMS + edge_global_film: Linear + edge_residual_proj: BiasFreeLinear + ln_edge_out: _RMS + ln_local_out: _RMS + ln_global_out: _RMS + zeeman_w1: Linear + zeeman_group_proj: Linear + zeeman_group_ln: GroupLayerNorm + zeeman_group_post: Linear + zeeman_w2: Linear + Q_row2: Float[Array, "n_global_q n_heads head_dim"] + K_row2: BiasFreeLinear + V_row2: BiasFreeLinear + ln_local_pre_row2: _RMS + ln_g_edge: _RMS + ln_g_zee: _RMS + global_w1: Linear + global_w2: Linear + ln_c1: _RMS + ln_c2: _RMS + ln_c3: _RMS + combine_w1: Linear + combine_w2: Linear + tok_bond_gln: Float[Array, "gb db"] + tok_zeeman_gln: Float[Array, "gz dz"] + tok_bond_key: Float[Array, "d_bond"] + tok_field_row: Float[Array, "d_local"] + tok_field_global: Float[Array, "d_global"] + tok_field_combine: Float[Array, "d_local"] + _use_id_Q_col: str = eqx.field(static=True, default="") + _use_id_Q_row: str = eqx.field(static=True, default="") + _use_id_Q_row2: str = eqx.field(static=True, default="") + _use_id_tok_bond_gln: str = eqx.field(static=True, default="") + _use_id_tok_zeeman_gln: str = eqx.field(static=True, default="") + _use_id_tok_bond_key: str = eqx.field(static=True, default="") + _use_id_tok_field_row: str = eqx.field(static=True, default="") + _use_id_tok_field_global: str = eqx.field(static=True, default="") + _use_id_tok_field_combine: str = eqx.field(static=True, default="") + d_bond: int = eqx.field(static=True) + d_local: int = eqx.field(static=True) + d_global: int = eqx.field(static=True) + d_edge: int = eqx.field(static=True) + n_heads: int = eqx.field(static=True) + head_dim: int = eqx.field(static=True) + n_global_q: int = eqx.field(static=True) + d_hidden_edge: int = eqx.field(static=True) + polar_group_norm_tau: float = eqx.field(static=True, default=0.001) + polar_group_norm_bond_hidden: int = eqx.field(static=True, default=128) + polar_group_norm_n_bond_groups: int = eqx.field(static=True, default=16) + polar_group_norm_d_bond_group: int = eqx.field(static=True, default=16) + polar_group_norm_n_zeeman_groups: int = eqx.field(static=True, default=16) + polar_group_norm_d_zeeman_group: int = eqx.field(static=True, default=16) + + def __init__( + self, + *, + key: PRNGKeyArray, + d_bond: int, + n_heads: int, + head_dim: int, + n_global_q: int, + d_edge: int, + d_hidden_edge: int, + polar_group_norm_tau: float, + polar_group_norm_bond_hidden: int, + polar_group_norm_n_bond_groups: int, + polar_group_norm_d_bond_group: int, + polar_group_norm_n_zeeman_groups: int, + polar_group_norm_d_zeeman_group: int, + zeeman_hidden_dim: int, + global_hidden_dim: int, + combine_hidden_dim: int, + token_initial_scale: float, + ): + d_local = n_heads * head_dim + d_global = n_global_q * n_heads * head_dim + n_h_kernel = 2 * n_heads + d_local_kernel = n_h_kernel * head_dim + edge_cond_in = d_bond + 2 * d_local + self.polar_group_norm_tau = float(polar_group_norm_tau) + self.polar_group_norm_bond_hidden = int(polar_group_norm_bond_hidden) + self.polar_group_norm_n_bond_groups = int(polar_group_norm_n_bond_groups) + self.polar_group_norm_d_bond_group = int(polar_group_norm_d_bond_group) + self.polar_group_norm_n_zeeman_groups = int(polar_group_norm_n_zeeman_groups) + self.polar_group_norm_d_zeeman_group = int(polar_group_norm_d_zeeman_group) + keys = jax.random.split(key, 25) + bond_group_dim = polar_group_norm_n_bond_groups * polar_group_norm_d_bond_group + self.bond_group_w1 = Linear(20, polar_group_norm_bond_hidden, key=keys[0]) + self.bond_group_w2 = Linear( + polar_group_norm_bond_hidden, bond_group_dim, key=keys[1] + ) + self.bond_group_post = Linear( + bond_group_dim, polar_group_norm_bond_hidden, key=keys[2] + ) + self.bond_group_ln = GroupLayerNorm( + polar_group_norm_n_bond_groups, polar_group_norm_d_bond_group + ) + self.bond_group_out = Linear(polar_group_norm_bond_hidden, d_bond, key=keys[3]) + self.Q_col = jax.random.normal(keys[4], (n_h_kernel, head_dim)) * head_dim ** ( + -0.5 + ) + self.K_col = BiasFreeLinear(d_bond, d_local_kernel, key=keys[5]) + self.V_col = BiasFreeLinear(d_bond, d_local_kernel, key=keys[6]) + self.ln_bond = _RMS(d_bond) + self.Q_row = jax.random.normal( + keys[7], (n_global_q, n_h_kernel, head_dim) + ) * head_dim ** (-0.5) + self.K_row = BiasFreeLinear(d_local, d_local_kernel, key=keys[8]) + self.V_row = BiasFreeLinear(d_local, d_local_kernel, key=keys[9]) + self.ln_local_pre_row = _RMS(d_local) + self.ln_edge_cond = _RMS(edge_cond_in) + self.edge_cond_w1 = Linear(edge_cond_in, d_hidden_edge, key=keys[10]) + self.edge_cond_w2 = Linear(d_hidden_edge, d_edge, key=keys[11]) + self.ln_edge_global = _RMS(d_global) + self.edge_global_film = Linear(d_global, 2 * d_edge, key=keys[12]) + self.edge_residual_proj = BiasFreeLinear(d_bond, d_edge, key=keys[13]) + zeeman_group_dim = ( + polar_group_norm_n_zeeman_groups * polar_group_norm_d_zeeman_group + ) + self.zeeman_w1 = Linear(7, zeeman_hidden_dim, key=keys[14]) + self.zeeman_group_proj = Linear( + zeeman_hidden_dim, zeeman_group_dim, key=keys[15] + ) + self.zeeman_group_ln = GroupLayerNorm( + polar_group_norm_n_zeeman_groups, polar_group_norm_d_zeeman_group + ) + self.zeeman_group_post = Linear( + zeeman_group_dim, zeeman_hidden_dim, key=keys[16] + ) + self.zeeman_w2 = Linear(zeeman_hidden_dim, d_local, key=keys[17]) + self.Q_row2 = jax.random.normal( + keys[18], (n_global_q, n_h_kernel, head_dim) + ) * head_dim ** (-0.5) + self.K_row2 = BiasFreeLinear(d_local, d_local_kernel, key=keys[19]) + self.V_row2 = BiasFreeLinear(d_local, d_local_kernel, key=keys[20]) + self.ln_local_pre_row2 = _RMS(d_local) + self.ln_g_edge = _RMS(d_global) + self.ln_g_zee = _RMS(d_global) + self.global_w1 = Linear(2 * d_global + 8, global_hidden_dim, key=keys[21]) + self.global_w2 = Linear(global_hidden_dim, d_global, key=keys[22]) + combine_in = 2 * d_local + d_global + self.ln_c1 = _RMS(d_local) + self.ln_c2 = _RMS(d_local) + self.ln_c3 = _RMS(d_global) + self.combine_w1 = Linear(combine_in, combine_hidden_dim, key=keys[23]) + self.combine_w2 = Linear(combine_hidden_dim, d_local, key=keys[24]) + self.ln_edge_out = _RMS(d_edge) + self.ln_local_out = _RMS(d_local) + self.ln_global_out = _RMS(d_global) + tok_keys = jax.random.split(jax.random.fold_in(key, 7389448), 6) + self.tok_bond_gln = token_initial_scale * jax.random.normal( + tok_keys[0], (polar_group_norm_n_bond_groups, polar_group_norm_d_bond_group) + ) + self.tok_zeeman_gln = token_initial_scale * jax.random.normal( + tok_keys[1], + (polar_group_norm_n_zeeman_groups, polar_group_norm_d_zeeman_group), + ) + self.tok_bond_key = token_initial_scale * jax.random.normal( + tok_keys[2], (d_bond,) + ) + self.tok_field_row = token_initial_scale * jax.random.normal( + tok_keys[3], (d_local,) + ) + self.tok_field_global = token_initial_scale * jax.random.normal( + tok_keys[4], (d_global,) + ) + self.tok_field_combine = token_initial_scale * jax.random.normal( + tok_keys[5], (d_local,) + ) + self.d_bond = d_bond + self.d_local = d_local + self.d_global = d_global + self.d_edge = d_edge + self.n_heads = n_heads + self.head_dim = head_dim + self.n_global_q = n_global_q + self.d_hidden_edge = d_hidden_edge + + def _polar_split( + self, x: Float[Array, "... d"] + ) -> tuple[Float[Array, "... 1"], Float[Array, "... d"]]: + x_f32 = x.astype(jnp.float32) + sq = jnp.sum(x_f32 * x_f32, axis=-1, keepdims=True) + tau = jnp.asarray(self.polar_group_norm_tau, dtype=jnp.float32) + r = jnp.sqrt(sq).astype(x.dtype) + direction = (x_f32 * jax.lax.rsqrt(sq + tau * tau)).astype(x.dtype) + return (r, direction) + + def _bond_input( + self, J_double_prime: Float[Array, "n n 10"] + ) -> Float[Array, "n n 20"]: + J9 = J_double_prime[..., :9] + eye = J_double_prime[..., 9:] + rJ, uJ = self._polar_split(J9) + return jnp.concatenate([J9, rJ, uJ, eye], axis=-1) + + def _zeeman_input(self, h_prime: Float[Array, "n 3"]) -> Float[Array, "n d_in"]: + rh, uh = self._polar_split(h_prime) + return jnp.concatenate([h_prime, rh, uh], axis=-1) + + def _embed_bonds( + self, + J_double_prime: Float[Array, "n n 10"], + *, + pathway: str, + structural_mask: Float[Array, "n n"] | None = None, + ) -> Float[Array, "n n d_bond"]: + if structural_mask is None: + structural_mask = jnp.ones(J_double_prime.shape[:-1], dtype=bool) + z = fused_silu( + self._dense_structural( + self.bond_group_w1, + self._bond_input(J_double_prime), + structural_mask, + repeat_ndim=2, + pathway=pathway, + ) + ) + z = self._dense_structural( + self.bond_group_w2, z, structural_mask, repeat_ndim=2, pathway=pathway + ) + pair_shape = J_double_prime.shape[:-1] + z = z.reshape( + *pair_shape, + self.polar_group_norm_n_bond_groups, + self.polar_group_norm_d_bond_group, + ) + z = self.bond_group_ln( + z, + pathway=pathway, + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + from hamiltonzero.optim.spin_blocks import register_small_full + + present = jnp.any(J_double_prime[..., :9] != 0, axis=-1) | ( + J_double_prime[..., 9] > 0.5 + ) + tok = register_small_full( + self.tok_bond_gln, tag_id=self._use_id_tok_bond_gln + ).astype(z.dtype) + z = jnp.where(present[..., None, None], z, tok) + z = z.reshape(*pair_shape, -1) + z = fused_silu( + self._dense_structural( + self.bond_group_post, z, structural_mask, repeat_ndim=2, pathway=pathway + ) + ) + return fused_silu( + self._dense_structural( + self.bond_group_out, z, structural_mask, repeat_ndim=2, pathway=pathway + ) + ) + + def _embed_zeeman( + self, + h_prime: Float[Array, "n 3"], + *, + pathway: str, + structural_mask: Float[Array, "n"] | None = None, + ) -> Float[Array, "n d_local"]: + if structural_mask is None: + structural_mask = jnp.ones(h_prime.shape[:-1], dtype=bool) + z = fused_silu( + self._dense_structural( + self.zeeman_w1, + self._zeeman_input(h_prime), + structural_mask, + repeat_ndim=1, + pathway=pathway, + ) + ) + z = self._dense_structural( + self.zeeman_group_proj, z, structural_mask, repeat_ndim=1, pathway=pathway + ) + n = h_prime.shape[0] + z = z.reshape( + n, + self.polar_group_norm_n_zeeman_groups, + self.polar_group_norm_d_zeeman_group, + ) + z = self.zeeman_group_ln( + z, + pathway=pathway, + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + from hamiltonzero.optim.spin_blocks import register_small_full + + present = jnp.any(h_prime != 0, axis=-1) + tok = register_small_full( + self.tok_zeeman_gln, tag_id=self._use_id_tok_zeeman_gln + ).astype(z.dtype) + z = jnp.where(present[:, None, None], z, tok) + z = z.reshape(n, -1) + z = fused_silu( + self._dense_structural( + self.zeeman_group_post, + z, + structural_mask, + repeat_ndim=1, + pathway=pathway, + ) + ) + return self._dense_structural( + self.zeeman_w2, z, structural_mask, repeat_ndim=1, pathway=pathway + ) + + def _bond_norm_gated( + self, + bond_emb: Float[Array, "n n d_bond"], + J_double_prime: Float[Array, "n n 10"], + *, + pathway: str, + structural_mask: Float[Array, "n n"] | None = None, + ) -> Float[Array, "n n d_bond"]: + from hamiltonzero.optim.spin_blocks import register_small_full + + present = jnp.any(J_double_prime[..., :9] != 0, axis=-1) | ( + J_double_prime[..., 9] > 0.5 + ) + if structural_mask is None: + structural_mask = jnp.ones(present.shape, dtype=bool) + tok = register_small_full( + self.tok_bond_key, tag_id=self._use_id_tok_bond_key + ).astype(bond_emb.dtype) + return jnp.where( + present[..., None], + self._norm_structural( + self.ln_bond, bond_emb, structural_mask, repeat_ndim=2, pathway=pathway + ), + tok, + ) + + @staticmethod + def _dense_structural(lin, x, structural_mask, *, repeat_ndim: int, pathway: str): + return lin( + x, + pathway=pathway, + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=repeat_ndim, + kfac_context_primal_reused_over_walkers=True, + ) + + @staticmethod + def _norm_structural(norm, x, structural_mask, *, repeat_ndim: int, pathway: str): + return norm( + x, + pathway=pathway, + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=repeat_ndim, + kfac_context_primal_reused_over_walkers=True, + ) + + def _dense_bare( + self, lin: Linear, x: Float[Array, "..."], *, pathway: str + ) -> Float[Array, "..."]: + from hamiltonzero.model.tree import _tagged_dense + + return _tagged_dense( + lin.weight, + lin.bias, + x, + tag_id=getattr(lin, "_use_id", ""), + pathway=pathway, + kfac_structural_mask=jnp.asarray(True), + kfac_scan_shared=False, + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=True, + ) + + @staticmethod + def _eval_tiles(n: int, tile_size: int): + tile_size = int(tile_size) + if tile_size < 1: + raise ValueError(f"tile_size must be positive, got {tile_size}") + return tuple((slice(j, min(j + tile_size, n)) for j in range(0, n, tile_size))) + + def eval_embed_local_rows( + self, + J_double_prime_rows: Float[Array, "r n 10"], + row_mask: Float[Array, "r"], + mask: Float[Array, "n"], + *, + tile_size: int = 128, + ) -> tuple[Float[Array, "r n d_bond"], Float[Array, "r d_local"]]: + if J_double_prime_rows.ndim != 3 or J_double_prime_rows.shape[-1] != 10: + raise ValueError( + f"J_double_prime_rows must have shape [R,N,10], got {J_double_prime_rows.shape}" + ) + r, n = J_double_prime_rows.shape[:2] + if row_mask.shape != (r,) or mask.shape != (n,): + raise ValueError( + f"row_mask/mask must match J rows/columns, got {row_mask.shape}, {mask.shape}, {J_double_prime_rows.shape}" + ) + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + from hamiltonzero.optim.spin_blocks import register_small_full + + PW = "even" + dtype = _compute_dtype() + mr = row_mask.astype(dtype) + mc = mask.astype(dtype) + pair_mask = mr[:, None] * mc[None, :] + bond_tiles = [] + tiles = self._eval_tiles(n, tile_size) + for sl in tiles: + pm = pair_mask[:, sl] + bond_tile = self._embed_bonds( + J_double_prime_rows[:, sl], pathway=PW, structural_mask=pm + ) + bond_tiles.append(bond_tile * pm[..., None]) + bond_emb = jnp.concatenate(bond_tiles, axis=1) + n_h_kernel = 2 * self.n_heads + scale = self.head_dim ** (-0.5) + Q_col = register_small_full(self.Q_col, tag_id=self._use_id_Q_col).astype(dtype) + score_max = jnp.full((r, n_h_kernel), -jnp.inf, dtype=dtype) + for sl in tiles: + pm = pair_mask[:, sl] + bn = self._bond_norm_gated( + bond_emb[:, sl], + J_double_prime_rows[:, sl], + pathway=PW, + structural_mask=pm, + ) + k_tile = self._dense_structural( + self.K_col, bn, pm, repeat_ndim=2, pathway=PW + ).reshape(r, sl.stop - sl.start, n_h_kernel, self.head_dim) + scores = jnp.einsum("hd,ijhd->ijh", Q_col, k_tile) * scale + scores = jnp.where( + pm[..., None] > 0, + scores, + jnp.asarray(-1000000000.0, dtype=scores.dtype), + ) + score_max = jnp.maximum(score_max, jnp.max(scores, axis=1)) + denom = jnp.zeros((r, n_h_kernel), dtype=dtype) + numer = jnp.zeros((r, n_h_kernel, self.head_dim), dtype=dtype) + for sl in tiles: + width = sl.stop - sl.start + pm = pair_mask[:, sl] + bn = self._bond_norm_gated( + bond_emb[:, sl], + J_double_prime_rows[:, sl], + pathway=PW, + structural_mask=pm, + ) + k_tile = self._dense_structural( + self.K_col, bn, pm, repeat_ndim=2, pathway=PW + ).reshape(r, width, n_h_kernel, self.head_dim) + scores = jnp.einsum("hd,ijhd->ijh", Q_col, k_tile) * scale + scores = jnp.where( + pm[..., None] > 0, + scores, + jnp.asarray(-1000000000.0, dtype=scores.dtype), + ) + weight = jnp.exp(scores - score_max[:, None, :]) + v_tile = self._dense_structural( + self.V_col, bn, pm, repeat_ndim=2, pathway=PW + ).reshape(r, width, n_h_kernel, self.head_dim) + denom = denom + jnp.sum(weight, axis=1) + numer = numer + jnp.einsum("ijh,ijhd->ihd", weight, v_tile) + col_out = numer / jnp.maximum( + denom[..., None], jnp.asarray(1e-30, dtype=numer.dtype) + ) + gate = col_out[:, : self.n_heads, :] + val = col_out[:, self.n_heads :, :] + local_desc_rows = (jax.nn.sigmoid(gate) * val).reshape(r, self.d_local) * mr[ + :, None + ] + return (bond_emb, local_desc_rows) + + def eval_jh_stats_rows( + self, + J_double_prime_rows: Float[Array, "r n 10"], + row_mask: Float[Array, "r"], + mask: Float[Array, "n"], + *, + row_indices: Int[Array, "r"] | None = None, + ) -> tuple[Float[Array, ""], Float[Array, ""]]: + r, n = J_double_prime_rows.shape[:2] + if row_indices is None: + if r != n: + raise ValueError("row_indices is required when R != N") + row_indices = jnp.arange(n, dtype=jnp.int32) + row_indices = jnp.asarray(row_indices, dtype=jnp.int32) + col_indices = jnp.arange(n, dtype=jnp.int32) + active = ( + row_mask[:, None].astype(J_double_prime_rows.dtype) + * mask[None, :].astype(J_double_prime_rows.dtype) + * (row_indices[:, None] != col_indices[None, :]).astype( + J_double_prime_rows.dtype + ) + ) + norm2 = jnp.sum(jnp.square(J_double_prime_rows[..., :9]), axis=-1) + return (jnp.sum(norm2 * active), jnp.sum(active)) + + def eval_edge_rows( + self, + *, + bond_emb_rows: Float[Array, "r n d_bond"], + local_rows: Float[Array, "r d_local"], + local_final_all: Float[Array, "n d_local"], + global_feat: Float[Array, "d_global"], + row_indices: Int[Array, "r"], + mask: Float[Array, "n"], + tile_size: int = 128, + ) -> Float[Array, "r n d_edge"]: + PW = "even" + r, n = bond_emb_rows.shape[:2] + row_indices = jnp.asarray(row_indices, dtype=jnp.int32) + if row_indices.shape != (r,): + raise ValueError("row_indices must have shape [R]") + if local_rows.shape != (r, self.d_local): + raise ValueError("local_rows must have shape [R,d_local]") + if local_final_all.shape != (n, self.d_local): + raise ValueError("local_final_all must have shape [N,d_local]") + if mask.shape != (n,): + raise ValueError("mask must have shape [N]") + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + dtype = _compute_dtype() + m = mask.astype(dtype) + mr = m[row_indices] + pair_mask = mr[:, None] * m[None, :] + global_mask = jnp.asarray(True) + film_in = self._norm_structural( + self.ln_edge_global, global_feat, global_mask, repeat_ndim=0, pathway=PW + ) + film = self._dense_bare(self.edge_global_film, film_in, pathway=PW) + gamma, beta = jnp.split(film, 2, axis=-1) + gamma = 0.1 * jnp.tanh(gamma) + beta = 0.1 * beta + edge_tiles = [] + for sl in self._eval_tiles(n, tile_size): + pm = pair_mask[:, sl] + bond = bond_emb_rows[:, sl] + width = sl.stop - sl.start + li = jnp.broadcast_to(local_rows[:, None, :], (r, width, self.d_local)) + lj = jnp.broadcast_to( + local_final_all[sl][None, :, :], (r, width, self.d_local) + ) + edge_in = jnp.concatenate([bond, li, lj], axis=-1) + edge_in = self._norm_structural( + self.ln_edge_cond, edge_in, pm, repeat_ndim=2, pathway=PW + ) + core = self._dense_structural( + self.edge_cond_w2, + fused_silu( + self._dense_structural( + self.edge_cond_w1, edge_in, pm, repeat_ndim=2, pathway=PW + ) + ), + pm, + repeat_ndim=2, + pathway=PW, + ) + update_edge = core * (1.0 + gamma[None, None, :]) + beta[None, None, :] + residual = self._dense_structural( + self.edge_residual_proj, bond, pm, repeat_ndim=2, pathway=PW + ) + edge_tile = ( + self._norm_structural( + self.ln_edge_out, + residual + update_edge, + pm, + repeat_ndim=2, + pathway=PW, + ) + * pm[..., None] + ) + edge_tiles.append(edge_tile) + return jnp.concatenate(edge_tiles, axis=1) + + def eval_finalize_local_rows( + self, + *, + J_double_prime_rows: Float[Array, "r n 10"], + local_desc_rows: Float[Array, "r d_local"], + local_desc_all: Float[Array, "n d_local"], + row_indices: Int[Array, "r"], + mask: Float[Array, "n"], + h_prime: Float[Array, "n 3"], + jh_stats: tuple[Float[Array, ""], Float[Array, ""]] | None = None, + ) -> tuple[Float[Array, "r d_local"], Float[Array, "d_global"]]: + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + from hamiltonzero.optim.spin_blocks import register_small_full + + PW = "even" + dtype = _compute_dtype() + r, n = J_double_prime_rows.shape[:2] + row_indices = jnp.asarray(row_indices, dtype=jnp.int32) + if row_indices.shape != (r,): + raise ValueError("row_indices must have shape [R]") + if local_desc_rows.shape != (r, self.d_local): + raise ValueError("local_desc_rows has the wrong shape") + if local_desc_all.shape != (n, self.d_local) or mask.shape != (n,): + raise ValueError("local_desc_all/mask must have global width N") + m = mask.astype(dtype) + mr = m[row_indices] + n_h_kernel = 2 * self.n_heads + scale = self.head_dim ** (-0.5) + global_mask = jnp.asarray(True) + + def row_descriptor(x, ln_pre, Q_arr, K_lin, V_lin, q_use_id, present=None): + x_norm = self._norm_structural(ln_pre, x, m, repeat_ndim=1, pathway=PW) + if present is not None: + tok_row = register_small_full( + self.tok_field_row, tag_id=self._use_id_tok_field_row + ).astype(x_norm.dtype) + x_norm = jnp.where(present[:, None], x_norm, tok_row) + qr = register_small_full(Q_arr, tag_id=q_use_id).astype(dtype) + kr = self._dense_structural( + K_lin, x_norm, m, repeat_ndim=1, pathway=PW + ).reshape(n, n_h_kernel, self.head_dim) + vr = self._dense_structural( + V_lin, x_norm, m, repeat_ndim=1, pathway=PW + ).reshape(n, n_h_kernel, self.head_dim) + sc = jnp.einsum("qhd,ihd->qhi", qr, kr) * scale + sc = jnp.where( + m[None, None, :] > 0, sc, jnp.asarray(-1000000000.0, dtype=sc.dtype) + ) + a = jax.nn.softmax(sc, axis=-1) + out = jnp.einsum("qhi,ihd->qhd", a, vr) + out = jax.nn.sigmoid(out[:, : self.n_heads, :]) * out[:, self.n_heads :, :] + return out.reshape(self.d_global) + + h_prime = h_prime.astype(dtype) + local_prime_all = ( + self._embed_zeeman(h_prime, pathway=PW, structural_mask=m) * m[:, None] + ) + local_prime_rows = local_prime_all[row_indices] + h_present = jnp.any(h_prime != 0, axis=-1) + g_edge = row_descriptor( + local_desc_all, + self.ln_local_pre_row, + self.Q_row, + self.K_row, + self.V_row, + self._use_id_Q_row, + ) + g_zee = row_descriptor( + local_prime_all, + self.ln_local_pre_row2, + self.Q_row2, + self.K_row2, + self.V_row2, + self._use_id_Q_row2, + present=h_present, + ) + if jh_stats is None: + if r != n: + raise ValueError( + "row-sharded featurization requires the psum result from eval_jh_stats_rows" + ) + jh_stats = self.eval_jh_stats_rows( + J_double_prime_rows, mr, m, row_indices=row_indices + ) + sum_j2, count_j = jh_stats + s_j2 = sum_j2 / jnp.maximum(count_j, 1.0) + s_h2 = jnp.sum(jnp.sum(jnp.square(h_prime), axis=-1) * m) / jnp.maximum( + jnp.sum(m), 1.0 + ) + + def safe_sqrt(x): + ok = x > 0 + return jnp.where(ok, jnp.sqrt(jnp.where(ok, x, 1.0)), 0.0) + + rj = safe_sqrt(s_j2) + rh = safe_sqrt(s_h2) + ok = rj + rh > 0 + theta = jnp.arctan2(jnp.where(ok, rh, 0.0), jnp.where(ok, rj, 1.0)) + jh_features = jnp.stack( + [ + jnp.log1p(rj), + jnp.log1p(rh), + jnp.sin(2.0 * theta), + jnp.cos(2.0 * theta), + jnp.sin(4.0 * theta), + jnp.cos(4.0 * theta), + jnp.sin(8.0 * theta), + jnp.cos(8.0 * theta), + ] + ).astype(dtype) + global_cat = jnp.concatenate( + [ + self._norm_structural( + self.ln_g_edge, g_edge, global_mask, repeat_ndim=0, pathway=PW + ), + jnp.where( + jnp.any(h_present), + self._norm_structural( + self.ln_g_zee, g_zee, global_mask, repeat_ndim=0, pathway=PW + ), + register_small_full( + self.tok_field_global, tag_id=self._use_id_tok_field_global + ).astype(dtype), + ), + jh_features, + ], + axis=-1, + ) + global_raw = self._dense_bare( + self.global_w2, + fused_silu(self._dense_bare(self.global_w1, global_cat, pathway=PW)), + pathway=PW, + ) + global_feat = self._norm_structural( + self.ln_global_out, global_raw, global_mask, repeat_ndim=0, pathway=PW + ) + combine = jnp.concatenate( + [ + jnp.where( + h_present[row_indices, None], + self._norm_structural( + self.ln_c1, local_prime_rows, mr, repeat_ndim=1, pathway=PW + ), + register_small_full( + self.tok_field_combine, tag_id=self._use_id_tok_field_combine + ).astype(dtype), + ), + self._norm_structural( + self.ln_c2, local_desc_rows, mr, repeat_ndim=1, pathway=PW + ), + jnp.broadcast_to( + self._norm_structural( + self.ln_c3, global_feat, global_mask, repeat_ndim=0, pathway=PW + ), + (r, self.d_global), + ), + ], + axis=-1, + ) + update = self._dense_structural( + self.combine_w2, + fused_silu( + self._dense_structural( + self.combine_w1, combine, mr, repeat_ndim=1, pathway=PW + ) + ), + mr, + repeat_ndim=1, + pathway=PW, + ) + local_rows = ( + self._norm_structural( + self.ln_local_out, + local_prime_rows + update, + mr, + repeat_ndim=1, + pathway=PW, + ) + * mr[:, None] + ) + return (local_rows, global_feat) + + def eval_streamed( + self, + J_double_prime: Float[Array, "n n 10"], + mask: Float[Array, "n"], + h_prime: Float[Array, "n 3"], + *, + tile_size: int = 128, + ): + n = J_double_prime.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + bond, local_desc = self.eval_embed_local_rows( + J_double_prime, mask, mask, tile_size=tile_size + ) + stats = self.eval_jh_stats_rows(J_double_prime, mask, mask, row_indices=idx) + local, global_feat = self.eval_finalize_local_rows( + J_double_prime_rows=J_double_prime, + local_desc_rows=local_desc, + local_desc_all=local_desc, + row_indices=idx, + mask=mask, + h_prime=h_prime, + jh_stats=stats, + ) + edge = self.eval_edge_rows( + bond_emb_rows=bond, + local_rows=local, + local_final_all=local, + global_feat=global_feat, + row_indices=idx, + mask=mask, + tile_size=tile_size, + ) + return (edge, local, global_feat) + + def __call__( + self, + J_double_prime: Float[Array, "n n 10"], + mask: Float[Array, "n"], + h_prime: Float[Array, "n 3"], + ) -> tuple[ + Float[Array, "n n d_edge"], Float[Array, "n d_local"], Float[Array, "d_global"] + ]: + PW = "even" + n = J_double_prime.shape[0] + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + dtype = _compute_dtype() + m = mask.astype(dtype) + pair_mask = m[:, None] * m[None, :] + global_mask = jnp.asarray(True) + head_dim = self.head_dim + n_heads = self.n_heads + n_global_q = self.n_global_q + scale = head_dim ** (-0.5) + bond_emb = self._embed_bonds( + J_double_prime, pathway=PW, structural_mask=pair_mask + ) + bond_emb = bond_emb * pair_mask[..., None] + n_h_kernel = 2 * n_heads + from hamiltonzero.optim.spin_blocks import register_small_full + + bond_norm = self._bond_norm_gated( + bond_emb, J_double_prime, pathway=PW, structural_mask=pair_mask + ) + Q_col = register_small_full(self.Q_col, tag_id=self._use_id_Q_col).astype(dtype) + K = self._dense_structural( + self.K_col, bond_norm, pair_mask, repeat_ndim=2, pathway=PW + ).reshape(n, n, n_h_kernel, head_dim) + V = self._dense_structural( + self.V_col, bond_norm, pair_mask, repeat_ndim=2, pathway=PW + ).reshape(n, n, n_h_kernel, head_dim) + scores = jnp.einsum("hd,ijhd->ijh", Q_col, K) * scale + scores = jnp.where( + pair_mask[..., None] > 0, + scores, + jnp.asarray(-1000000000.0, dtype=scores.dtype), + ) + attn = jax.nn.softmax(scores, axis=1) + col_out = jnp.einsum("ijh,ijhd->ihd", attn, V) + gate = col_out[:, :n_heads, :] + val = col_out[:, n_heads:, :] + col_out = jax.nn.sigmoid(gate) * val + local_i_raw = col_out.reshape(n, self.d_local) * m[:, None] + h_prime = h_prime.astype(dtype) + + def _row_descriptor(x, ln_pre, Q_arr, K_lin, V_lin, q_use_id, present=None): + x_norm = self._norm_structural(ln_pre, x, m, repeat_ndim=1, pathway=PW) + if present is not None: + tok_row = register_small_full( + self.tok_field_row, tag_id=self._use_id_tok_field_row + ).astype(x_norm.dtype) + x_norm = jnp.where(present[:, None], x_norm, tok_row) + Qr = register_small_full(Q_arr, tag_id=q_use_id).astype(dtype) + Kr = self._dense_structural( + K_lin, x_norm, m, repeat_ndim=1, pathway=PW + ).reshape(n, n_h_kernel, head_dim) + Vr = self._dense_structural( + V_lin, x_norm, m, repeat_ndim=1, pathway=PW + ).reshape(n, n_h_kernel, head_dim) + sc = jnp.einsum("qhd,ihd->qhi", Qr, Kr) * scale + sc = jnp.where( + m[None, None, :] > 0, sc, jnp.asarray(-1000000000.0, dtype=sc.dtype) + ) + a = jax.nn.softmax(sc, axis=-1) + ro = jnp.einsum("qhi,ihd->qhd", a, Vr) + ro = jax.nn.sigmoid(ro[:, :n_heads, :]) * ro[:, n_heads:, :] + return ro.reshape(self.d_global) + + local_desc_i = local_i_raw + local_i_prime = self._embed_zeeman(h_prime, pathway=PW, structural_mask=m) + local_i_prime = local_i_prime * m[:, None] + h_present = jnp.any(h_prime != 0, axis=-1) + g_edge = _row_descriptor( + local_desc_i, + self.ln_local_pre_row, + self.Q_row, + self.K_row, + self.V_row, + self._use_id_Q_row, + ) + g_zee = _row_descriptor( + local_i_prime, + self.ln_local_pre_row2, + self.Q_row2, + self.K_row2, + self.V_row2, + self._use_id_Q_row2, + present=h_present, + ) + off_diag = pair_mask * (1.0 - jnp.eye(n, dtype=dtype)) + bond_magnitude2 = jnp.sum( + jnp.square(J_double_prime[..., :9].astype(dtype)), axis=-1 + ) + mean_j2 = jnp.sum(bond_magnitude2 * off_diag) / jnp.maximum( + jnp.sum(off_diag), 1.0 + ) + mean_h2 = jnp.sum(jnp.sum(jnp.square(h_prime), axis=-1) * m) / jnp.maximum( + jnp.sum(m), 1.0 + ) + + def _safe_sqrt(x): + ok = x > 0 + return jnp.where(ok, jnp.sqrt(jnp.where(ok, x, 1.0)), 0.0) + + rj = _safe_sqrt(mean_j2) + rh = _safe_sqrt(mean_h2) + nonzero = rj + rh > 0 + theta = jnp.arctan2(jnp.where(nonzero, rh, 0.0), jnp.where(nonzero, rj, 1.0)) + jh_features = jnp.stack( + [ + jnp.log1p(rj), + jnp.log1p(rh), + jnp.sin(2.0 * theta), + jnp.cos(2.0 * theta), + jnp.sin(4.0 * theta), + jnp.cos(4.0 * theta), + jnp.sin(8.0 * theta), + jnp.cos(8.0 * theta), + ] + ).astype(dtype) + global_cat = jnp.concatenate( + [ + self._norm_structural( + self.ln_g_edge, g_edge, global_mask, repeat_ndim=0, pathway=PW + ), + jnp.where( + jnp.any(h_present), + self._norm_structural( + self.ln_g_zee, g_zee, global_mask, repeat_ndim=0, pathway=PW + ), + register_small_full( + self.tok_field_global, tag_id=self._use_id_tok_field_global + ).astype(dtype), + ), + jh_features, + ], + axis=-1, + ) + global_raw = self._dense_bare( + self.global_w2, + fused_silu(self._dense_bare(self.global_w1, global_cat, pathway=PW)), + pathway=PW, + ) + global_feat = self._norm_structural( + self.ln_global_out, global_raw, global_mask, repeat_ndim=0, pathway=PW + ) + combine_in = jnp.concatenate( + [ + jnp.where( + h_present[:, None], + self._norm_structural( + self.ln_c1, local_i_prime, m, repeat_ndim=1, pathway=PW + ), + register_small_full( + self.tok_field_combine, tag_id=self._use_id_tok_field_combine + ).astype(dtype), + ), + self._norm_structural( + self.ln_c2, local_desc_i, m, repeat_ndim=1, pathway=PW + ), + jnp.broadcast_to( + self._norm_structural( + self.ln_c3, global_feat, global_mask, repeat_ndim=0, pathway=PW + )[None, :], + (n, self.d_global), + ), + ], + axis=-1, + ) + local_update = self._dense_structural( + self.combine_w2, + fused_silu( + self._dense_structural( + self.combine_w1, combine_in, m, repeat_ndim=1, pathway=PW + ) + ), + m, + repeat_ndim=1, + pathway=PW, + ) + local_i = self._norm_structural( + self.ln_local_out, + local_i_prime + local_update, + m, + repeat_ndim=1, + pathway=PW, + ) + local_i = local_i * m[:, None] + local_i_for_i = jnp.broadcast_to(local_i[:, None, :], (n, n, self.d_local)) + local_i_for_j = jnp.broadcast_to(local_i[None, :, :], (n, n, self.d_local)) + edge_cond_in = jnp.concatenate( + [bond_emb, local_i_for_i, local_i_for_j], axis=-1 + ) + edge_cond_norm = self._norm_structural( + self.ln_edge_cond, edge_cond_in, pair_mask, repeat_ndim=2, pathway=PW + ) + edge_core = self._dense_structural( + self.edge_cond_w2, + fused_silu( + self._dense_structural( + self.edge_cond_w1, + edge_cond_norm, + pair_mask, + repeat_ndim=2, + pathway=PW, + ) + ), + pair_mask, + repeat_ndim=2, + pathway=PW, + ) + ln_global_for_film = self._norm_structural( + self.ln_edge_global, global_feat, global_mask, repeat_ndim=0, pathway=PW + ) + film = self._dense_bare(self.edge_global_film, ln_global_for_film, pathway=PW) + gamma, beta = jnp.split(film, 2, axis=-1) + gamma = 0.1 * jnp.tanh(gamma) + beta = 0.1 * beta + edge_update = edge_core * (1.0 + gamma[None, None, :]) + beta[None, None, :] + bond_emb_proj = self._dense_structural( + self.edge_residual_proj, bond_emb, pair_mask, repeat_ndim=2, pathway=PW + ) + edge_ij = bond_emb_proj + edge_update + edge_ij = self._norm_structural( + self.ln_edge_out, edge_ij, pair_mask, repeat_ndim=2, pathway=PW + ) + edge_ij = edge_ij * pair_mask[..., None] + return (edge_ij, local_i, global_feat) diff --git a/src/hamiltonzero/model/fp32.py b/src/hamiltonzero/model/fp32.py new file mode 100644 index 0000000000000000000000000000000000000000..4f4501ee0de1c16ced00700bebd2268496caddfd --- /dev/null +++ b/src/hamiltonzero/model/fp32.py @@ -0,0 +1,11 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +import jax.numpy as jnp + + +def compute_dtype(*_args, **_kwargs): + return jnp.float32 + + +__all__ = ["compute_dtype"] diff --git a/src/hamiltonzero/model/fused_silu.py b/src/hamiltonzero/model/fused_silu.py new file mode 100644 index 0000000000000000000000000000000000000000..e8cc150f56a0bc4a5a5974000abcc5183be69403 --- /dev/null +++ b/src/hamiltonzero/model/fused_silu.py @@ -0,0 +1,11 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +import jax + + +def fused_silu(x): + return jax.nn.silu(x) + + +__all__ = ["fused_silu"] diff --git a/src/hamiltonzero/model/global_ladder.py b/src/hamiltonzero/model/global_ladder.py new file mode 100644 index 0000000000000000000000000000000000000000..089ca0b38222337e400bffec3a8b34b215ae41ea --- /dev/null +++ b/src/hamiltonzero/model/global_ladder.py @@ -0,0 +1,503 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray + +from .odd_ops import Linear, _RMS +from .tree import ( + _inline_norm_forward, + _tagged_bounded_ngpt_gain, + _tagged_dense, + _tagged_dense_no_bias, + _tree_sphere, +) + + +def _pool_masked_softmax(scores, mask, axis): + neg = jnp.asarray(-1.0e30, dtype=scores.dtype) + scores = jnp.where(mask > 0, scores, neg) + return jax.nn.softmax(scores, axis=axis) + + +class GDescriptorPool(eqx.Module): + W_q: Float[Array, "d_g hq"] + K: Linear + V: Linear + ln_in: _RMS + n_heads: int = eqx.field(static=True) + d_k: int = eqx.field(static=True) + d_v: int = eqx.field(static=True) + tag: str = eqx.field(static=True, default="") + + def __init__( + self, + d_g: int, + d_x: int, + *, + key: PRNGKeyArray, + n_heads: int = 4, + d_k: int = 64, + d_v: int = 64, + tag: str = "", + ): + kq, kk, kv = jax.random.split(key, 3) + self.W_q = jax.random.normal(kq, (d_g, n_heads * d_k)) * (d_g**-0.5) + self.K = Linear(d_x, n_heads * d_k, key=kk) + self.V = Linear(d_x, n_heads * d_v, key=kv) + self.ln_in = _RMS(d_x) + self.n_heads = n_heads + self.d_k = d_k + self.d_v = d_v + self.tag = tag + + @property + def d_out(self) -> int: + return self.n_heads * self.d_v + + def __call__( + self, + g, + xs, + mask, + *, + kfac_structural_mask=None, + kfac_update_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 1, + kfac_context_primal_reused_over_walkers: bool = False, + ): + + n = xs.shape[0] + xn = _inline_norm_forward( + self.ln_in, + xs, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + q_structural_mask = ( + jnp.any(jnp.asarray(kfac_structural_mask).astype(bool)) + if kfac_update_mask is None and kfac_structural_mask is not None + else kfac_update_mask + ) + q = _tagged_dense_no_bias( + self.W_q, + g, + tag_id=f"{self.tag}.W_q", + pathway="even", + kfac_structural_mask=q_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ).reshape(self.n_heads, self.d_k) + k = _tagged_dense( + self.K.weight, + self.K.bias, + xn, + tag_id=self.K._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ).reshape(n, self.n_heads, self.d_k) + v = _tagged_dense( + self.V.weight, + self.V.bias, + xn, + tag_id=self.V._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ).reshape(n, self.n_heads, self.d_v) + scores = jnp.einsum("hd,nhd->hn", q, k) / jnp.sqrt( + jnp.asarray(self.d_k, dtype=xs.dtype) + ) + attn = _pool_masked_softmax(scores, mask[None, :], axis=-1) + out = jnp.einsum("hn,nhv->hv", attn, v) + return out.reshape(-1) + + +def _global_update_parameters(d_g: int, d_pool: int, tap_dim: int, key): + tap = int(tap_dim) + if tap < 1 or tap >= int(d_g): + raise ValueError("global tap dimension must be positive and smaller than d_g") + d_in = tap + int(d_pool) + d_hidden = 2 * int(d_g) + k1, k2, k3 = jax.random.split(key, 3) + return ( + jax.random.normal(k1, (d_in, d_hidden)) * (d_in**-0.5), + jnp.zeros((d_hidden,)), + jax.random.normal(k2, (d_hidden, d_g)) * (d_hidden**-0.5), + jnp.zeros((d_g,)), + jnp.ones((d_in,)), + jax.random.normal(k3, (d_g, tap)) * (d_g**-0.5), + ) + + +def _global_delta( + update, + g, + pool, + *, + kfac_structural_mask, + kfac_g_structural_mask, + kfac_scan_shared, + kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers, +): + from hamiltonzero.model.tree import _tagged_rms_eqx_style + from hamiltonzero.model.fused_silu import fused_silu + + g_mask = ( + kfac_structural_mask + if kfac_g_structural_mask is None + else kfac_g_structural_mask + ) + g_in = _tagged_dense_no_bias( + update.g_tap_w, + g, + tag_id=f"{update.tag}.gtap", + pathway="even", + kfac_structural_mask=g_mask, + kfac_scan_shared=( + kfac_scan_shared if kfac_g_structural_mask is None else False + ), + kfac_repeat_ndim=(kfac_repeat_ndim if kfac_g_structural_mask is None else 0), + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + x = jnp.concatenate([g_in, pool.astype(g.dtype)]) + x = _tagged_rms_eqx_style( + update.ln_s, + x, + tag_id=f"{update.tag}.ln", + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + hidden = fused_silu( + _tagged_dense( + update.w1, + update.b1, + x, + tag_id=f"{update.tag}.ffn1", + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + ) + return _tagged_dense( + update.w2, + update.b2, + hidden, + tag_id=f"{update.tag}.ffn2", + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + +class ResidualGlobalUpdate(eqx.Module): + w1: Float[Array, "d_in d_hidden"] + b1: Float[Array, "d_hidden"] + w2: Float[Array, "d_hidden d_g"] + b2: Float[Array, "d_g"] + ln_s: Float[Array, "d_in"] + g_tap_w: Float[Array, "d_g tap_dim"] + residual_gain: float = eqx.field(static=True) + tag: str = eqx.field(static=True) + + def __init__( + self, + d_g: int, + d_pool: int, + *, + key: PRNGKeyArray, + tap_dim: int, + residual_gain: float, + tag: str, + ): + ( + self.w1, + self.b1, + self.w2, + self.b2, + self.ln_s, + self.g_tap_w, + ) = _global_update_parameters(d_g, d_pool, tap_dim, key) + self.residual_gain = float(residual_gain) + self.tag = tag + + def __call__( + self, + g, + pool, + *, + kfac_structural_mask=None, + kfac_g_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ): + delta = _global_delta( + self, + g, + pool, + kfac_structural_mask=kfac_structural_mask, + kfac_g_structural_mask=kfac_g_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + return g + self.residual_gain * delta + + +class BoundaryGlobalUpdate(eqx.Module): + w1: Float[Array, "d_in d_hidden"] + b1: Float[Array, "d_hidden"] + w2: Float[Array, "d_hidden d_g"] + b2: Float[Array, "d_g"] + ln_s: Float[Array, "d_in"] + g_tap_w: Float[Array, "d_g tap_dim"] + tag: str = eqx.field(static=True) + + def __init__( + self, d_g: int, d_pool: int, *, key: PRNGKeyArray, tap_dim: int, tag: str + ): + ( + self.w1, + self.b1, + self.w2, + self.b2, + self.ln_s, + self.g_tap_w, + ) = _global_update_parameters(d_g, d_pool, tap_dim, key) + self.tag = tag + + def __call__( + self, + g, + pool, + *, + kfac_structural_mask=None, + kfac_g_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ): + delta = _global_delta( + self, + g, + pool, + kfac_structural_mask=kfac_structural_mask, + kfac_g_structural_mask=kfac_g_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + return _tree_sphere(g + delta) + + +class TreeGlobalUpdate(eqx.Module): + w1: Float[Array, "d_in d_hidden"] + b1: Float[Array, "d_hidden"] + w2: Float[Array, "d_hidden d_g"] + b2: Float[Array, "d_g"] + ln_s: Float[Array, "d_in"] + alpha: Float[Array, "d_g"] + g_tap_w: Float[Array, "d_g tap_dim"] + alpha_max: float = eqx.field(static=True) + tag: str = eqx.field(static=True) + + def __init__( + self, + d_g: int, + d_pool: int, + *, + key: PRNGKeyArray, + tap_dim: int, + alpha_init: float, + alpha_max: float, + tag: str, + ): + ( + self.w1, + self.b1, + self.w2, + self.b2, + self.ln_s, + self.g_tap_w, + ) = _global_update_parameters(d_g, d_pool, tap_dim, key) + self.alpha = float(alpha_init) * jnp.ones((d_g,)) + self.alpha_max = float(alpha_max) + self.tag = tag + + def __call__( + self, + g, + pool, + update_mask=None, + *, + kfac_structural_mask=None, + kfac_g_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ): + delta = _global_delta( + self, + g, + pool, + kfac_structural_mask=kfac_structural_mask, + kfac_g_structural_mask=kfac_g_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + skip = _tree_sphere(g) + proposal = _tree_sphere(delta) + direction = proposal - skip + gain = _tagged_bounded_ngpt_gain( + self.alpha, + direction, + max_gain=self.alpha_max, + tag_id=f"{self.tag}.alpha", + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + updated = _tree_sphere(skip + gain * direction) + if update_mask is None: + return updated + active = jnp.asarray(update_mask).astype(bool) + while active.ndim < updated.ndim: + active = active[..., None] + return jnp.where(active, updated, g) + + +class EdgeRowColGlobalUpdate(eqx.Module): + row_pool: GDescriptorPool + col_pool: GDescriptorPool + set_pool: GDescriptorPool + update: BoundaryGlobalUpdate + + def __init__( + self, + d_g: int, + d_edge: int, + *, + key: PRNGKeyArray, + n_heads: int = 4, + d_k: int = 64, + d_v: int = 64, + tag: str = "", + tap_dim: int, + ): + kr, kc, ks, ku = jax.random.split(key, 4) + self.row_pool = GDescriptorPool( + d_g, d_edge, key=kr, n_heads=n_heads, d_k=d_k, d_v=d_v, tag=f"{tag}.row" + ) + self.col_pool = GDescriptorPool( + d_g, d_edge, key=kc, n_heads=n_heads, d_k=d_k, d_v=d_v, tag=f"{tag}.col" + ) + self.set_pool = GDescriptorPool( + d_g, + n_heads * d_v, + key=ks, + n_heads=n_heads, + d_k=d_k, + d_v=d_v, + tag=f"{tag}.set", + ) + self.update = BoundaryGlobalUpdate( + d_g, + n_heads * d_v, + key=ku, + tag=f"{tag}.upd", + tap_dim=tap_dim, + ) + + def __call__(self, g, edge, mask): + + system_active = jnp.any(mask.astype(bool)) + rows = jax.vmap( + lambda ei, mi: self.row_pool( + g, + ei, + mask, + kfac_structural_mask=mi * mask, + kfac_update_mask=system_active, + kfac_scan_shared=False, + kfac_repeat_ndim=2, + ) + )(edge, mask) + cols = jax.vmap( + lambda ej, mj: self.col_pool( + g, + ej, + mask, + kfac_structural_mask=mj * mask, + kfac_update_mask=system_active, + kfac_scan_shared=False, + kfac_repeat_ndim=2, + ) + )(jnp.swapaxes(edge, 0, 1), mask) + descs = jnp.concatenate([rows, cols], axis=0) + dmask = jnp.concatenate([mask, mask], axis=0) + pooled = self.set_pool( + g, + descs, + dmask, + kfac_structural_mask=dmask, + kfac_update_mask=system_active, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + ) + return self.update( + g, + pooled, + kfac_structural_mask=system_active, + kfac_scan_shared=False, + ) diff --git a/src/hamiltonzero/model/model.py b/src/hamiltonzero/model/model.py new file mode 100644 index 0000000000000000000000000000000000000000..4176d1fba23c9621f19e0d19aae5645268598813 --- /dev/null +++ b/src/hamiltonzero/model/model.py @@ -0,0 +1,641 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +from typing import NamedTuple +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray +from .context import SpinContext +from .featurizer import SystemFeaturizer +from .odd_ops import BiasFreeLinear +from .readout_leaf_context import PhysicalReadoutContext, RouterContext +from .route_pointer import TreePrefixPointerMHSEA +from .tree import ( + LeafBuilder, + MergeOp, + RootReadout, + balanced_tree_reduce_masked_scan as balanced_tree_reduce_masked, +) +from .trunk import Trunk + + +class PerSystemInvariants(NamedTuple): + g_emb: Float[Array, "d_g"] + e_leaf: Float[Array, "n d_e"] + edge_leaf: Float[Array, "n n d_edge"] + g_stream: Float[Array, "d_global_feat"] + + +def _normalize_leaf_carriers(u_all): + eps = 1e-30 + u32 = u_all.astype(jnp.float32) + rms = jnp.sqrt(jnp.mean(u32 * u32, axis=-1, keepdims=True) + eps) + u_all = u_all / rms.astype(u_all.dtype) + log_rms = jnp.log(rms)[..., 0] + return (u_all, log_rms) + + +def _shallow_replace(obj, **changes): + cls = type(obj) + new_obj = object.__new__(cls) + new_obj.__dict__.update(obj.__dict__) + for k, v in changes.items(): + object.__setattr__(new_obj, k, v) + return new_obj + + +class SpinAnsatz(eqx.Module): + featurizer: SystemFeaturizer + trunk: Trunk + leaf: LeafBuilder + merge: MergeOp + readout: RootReadout + readout_leaf_context: PhysicalReadoutContext + route_contextualizer: RouterContext + gladder_post: "EdgeRowColGlobalUpdate" + gladder_to_gemb_w: Float[Array, "d_global_feat d_g"] + gladder_to_gemb_b: Float[Array, "d_g"] + gladder_gemb_ln_s: Float[Array, "d_g"] + gladder_tree_pool: "GDescriptorPool" + gladder_tree_update: "TreeGlobalUpdate" + gladder_tree_proj_w: Float[Array, "d_global_feat d_g"] + gladder_tree_proj_b: Float[Array, "d_g"] + gladder_root_proj_w: Float[Array, "d_global_feat d_g"] + gladder_root_proj_b: Float[Array, "d_g"] + gladder_root_ln_s: Float[Array, "d_g"] + gladder_fork_phys: "EdgeRowColGlobalUpdate" + gladder_fork_route: "EdgeRowColGlobalUpdate" + route_decoder: TreePrefixPointerMHSEA + q_to_odd: BiasFreeLinear + + def __init__( + self, + *, + d_e: int, + d_o: int, + d_c: int, + d_r: int, + n_heads: int, + n_layers: int, + rank: int, + n_edge: int, + d_e_attn: int, + d_c_attn: int, + trunk_edge_node_ctx_dim: int, + trunk_edge_hidden_dim: int, + trunk_attn_bias_hidden_dim: int, + trunk_ffn_hidden_dim: int, + trunk_two_hop_hidden_dim: int, + tree_edge_node_ctx_dim: int, + global_d_g: int, + d_m_merge: int, + merge_chain_hypernet_rank: int, + feat_d_bond: int, + feat_n_heads: int, + feat_head_dim: int, + feat_n_global_q: int, + feat_edge_hidden_dim: int, + feat_zeeman_hidden_dim: int, + feat_global_hidden_dim: int, + feat_combine_hidden_dim: int, + feat_token_initial_scale: float, + feat_d_edge: int, + polar_group_norm_tau: float, + polar_group_norm_bond_hidden: int, + polar_group_norm_n_bond_groups: int, + polar_group_norm_d_bond_group: int, + polar_group_norm_n_zeeman_groups: int, + polar_group_norm_d_zeeman_group: int, + route_pointer_max_n: int, + route_pointer_d_model: int, + route_pointer_n_heads: int, + route_pointer_attn_dim: int, + route_pointer_score_dim: int, + route_pointer_candidate_hidden: int, + route_pointer_summary_hidden: int, + route_pointer_ffn_hidden: int, + route_pointer_score_init_scale: float, + route_pointer_rope_base: float, + route_pointer_rope_scaling: float, + route_tree_prefix_layers: int, + route_tree_prefix_candidate_layers: int, + route_tree_prefix_merge_hidden: int, + route_tree_prefix_post_prefix_suffix_layers: int, + route_contextualizer_layers: int, + route_contextualizer_n_heads: int, + route_contextualizer_attn_dim: int, + route_contextualizer_edge_node_ctx_dim: int, + level_edge_attn_n_heads: int, + level_edge_attn_edge_mlp_hidden: int, + level_edge_attn_edge_mlp_n_blocks: int, + level_edge_attn_ffn_d_hidden: int, + level_edge_attn_rope_base: float, + level_edge_attn_rope_scaling: float, + root_readout_edge_rank: int, + ngpt_alpha_initial: float, + ngpt_alpha_initial_fraction: float, + ngpt_alpha_maximum: float, + global_ladder_tap_dim: int, + level_edge_attn_bias_mlp_hidden: int, + level_edge_attn_bias_mlp_n_blocks: int, + merge_c_mlp_hidden: int, + readout_leaf_context_layers: int, + readout_leaf_context_n_heads: int, + readout_leaf_context_attn_dim: int, + readout_leaf_context_edge_node_ctx_dim: int, + readout_leaf_context_summary_hidden: int, + readout_leaf_context_mlp_hidden: int, + readout_leaf_context_bias_hidden: int, + readout_leaf_context_edge_ffn_hidden: int, + readout_leaf_context_rope_base: float, + readout_leaf_context_rope_scaling: float, + two_hop_channels: int, + tree_edge_fwl_channels: int, + attn_impl: str, + key: PRNGKeyArray, + ): + from .odd_ops import bounded_gain_logit + + alpha_init = bounded_gain_logit( + ngpt_alpha_initial, + max_gain=ngpt_alpha_maximum, + init_fraction=ngpt_alpha_initial_fraction, + ) + k_feat, k_tr, k_lf, k_mg, k_ro, k_ge, k_route, k_extras = jax.random.split( + key, 8 + ) + k_leaf_ctx = jax.random.fold_in(k_route, 85897159) + self.featurizer = SystemFeaturizer( + key=k_feat, + d_bond=feat_d_bond, + n_heads=feat_n_heads, + head_dim=feat_head_dim, + n_global_q=feat_n_global_q, + d_edge=feat_d_edge, + d_hidden_edge=feat_edge_hidden_dim, + polar_group_norm_tau=polar_group_norm_tau, + polar_group_norm_bond_hidden=polar_group_norm_bond_hidden, + polar_group_norm_n_bond_groups=polar_group_norm_n_bond_groups, + polar_group_norm_d_bond_group=polar_group_norm_d_bond_group, + polar_group_norm_n_zeeman_groups=polar_group_norm_n_zeeman_groups, + polar_group_norm_d_zeeman_group=polar_group_norm_d_zeeman_group, + zeeman_hidden_dim=feat_zeeman_hidden_dim, + global_hidden_dim=feat_global_hidden_dim, + combine_hidden_dim=feat_combine_hidden_dim, + token_initial_scale=feat_token_initial_scale, + ) + d_local = feat_n_heads * feat_head_dim + d_global_feat = feat_n_global_q * feat_n_heads * feat_head_dim + self.trunk = Trunk( + d_e=d_e, + n_heads=n_heads, + n_layers=n_layers, + n_edge=n_edge, + d_local_in=d_local, + d_edge_in=feat_d_edge, + key=k_tr, + gladder_d_g=d_global_feat, + global_tap_dim=global_ladder_tap_dim, + attn_impl=attn_impl, + attn_dim=d_e_attn, + attn_bias_hidden_dim=trunk_attn_bias_hidden_dim, + ffn_hidden_dim=trunk_ffn_hidden_dim, + edge_hidden_dim=trunk_edge_hidden_dim, + edge_node_ctx_dim=trunk_edge_node_ctx_dim, + two_hop_channels=two_hop_channels, + two_hop_hidden_dim=trunk_two_hop_hidden_dim, + ) + self.q_to_odd = BiasFreeLinear(4, d_o, key=jax.random.fold_in(k_tr, 2430463726)) + tree_d_c = d_c + tree_d_r = d_r + self.leaf = LeafBuilder( + d_e=d_e, + d_o=d_o, + d_c=tree_d_c, + d_r=tree_d_r, + rank=rank, + key=k_lf, + d_g=global_d_g, + leaf_hypernet_rank=merge_chain_hypernet_rank, + d_m_merge=d_m_merge, + ) + self.merge = MergeOp( + d_r=tree_d_r, + d_c=tree_d_c, + key=k_mg, + d_g=global_d_g, + alpha_init=alpha_init, + alpha_max=ngpt_alpha_maximum, + d_m_merge=d_m_merge, + merge_output_hypernet_rank=merge_chain_hypernet_rank, + level_edge_attn_d_edge=feat_d_edge, + level_edge_attn_n_heads=level_edge_attn_n_heads, + level_edge_attn_attn_dim=d_c_attn, + tree_edge_node_ctx_dim=tree_edge_node_ctx_dim, + level_edge_attn_attn_impl=attn_impl, + level_edge_attn_edge_mlp_hidden=level_edge_attn_edge_mlp_hidden, + level_edge_attn_edge_mlp_n_blocks=level_edge_attn_edge_mlp_n_blocks, + level_edge_attn_ffn_d_hidden=level_edge_attn_ffn_d_hidden, + level_edge_attn_max_n=int(route_pointer_max_n), + level_edge_attn_rope_base=float(level_edge_attn_rope_base), + level_edge_attn_rope_scaling=float(level_edge_attn_rope_scaling), + tree_edge_fwl_channels=tree_edge_fwl_channels, + level_edge_attn_bias_mlp_hidden=level_edge_attn_bias_mlp_hidden, + level_edge_attn_bias_mlp_n_blocks=level_edge_attn_bias_mlp_n_blocks, + merge_c_mlp_hidden=merge_c_mlp_hidden, + ) + self.readout = RootReadout( + d_r=tree_d_r, + key=k_ro, + d_m_merge=d_m_merge, + d_edge=feat_d_edge, + edge_rank=root_readout_edge_rank, + d_g=global_d_g, + d_c=tree_d_c, + ) + self.readout_leaf_context = PhysicalReadoutContext( + d_e=d_e, + d_edge=n_edge, + n_layers=int(readout_leaf_context_layers), + n_heads=int(readout_leaf_context_n_heads), + summary_hidden=readout_leaf_context_summary_hidden, + mlp_hidden=readout_leaf_context_mlp_hidden, + bias_hidden=readout_leaf_context_bias_hidden, + edge_ffn_hidden=readout_leaf_context_edge_ffn_hidden, + attn_dim=readout_leaf_context_attn_dim, + edge_node_ctx_dim=readout_leaf_context_edge_node_ctx_dim, + attn_impl=attn_impl, + rope_base=float(readout_leaf_context_rope_base), + rope_scaling=float(readout_leaf_context_rope_scaling), + gladder_d_g=d_global_feat, + global_tap_dim=global_ladder_tap_dim, + key=k_leaf_ctx, + ) + self.route_contextualizer = RouterContext( + d_e=d_e, + d_edge=n_edge, + n_layers=int(route_contextualizer_layers), + n_heads=int(route_contextualizer_n_heads), + mlp_hidden=readout_leaf_context_mlp_hidden, + bias_hidden=readout_leaf_context_bias_hidden, + edge_ffn_hidden=readout_leaf_context_edge_ffn_hidden, + attn_dim=route_contextualizer_attn_dim, + edge_node_ctx_dim=route_contextualizer_edge_node_ctx_dim, + attn_impl=attn_impl, + gladder_d_g=d_global_feat, + global_tap_dim=global_ladder_tap_dim, + key=jax.random.fold_in(k_route, 2802764542), + ) + from .global_ladder import ( + EdgeRowColGlobalUpdate, + GDescriptorPool, + TreeGlobalUpdate, + ) + + _k_gl = jax.random.split(jax.random.fold_in(key, 25005), 2) + self.gladder_post = EdgeRowColGlobalUpdate( + d_global_feat, + feat_d_edge, + key=_k_gl[0], + tag="gladder.post_trunk", + tap_dim=global_ladder_tap_dim, + ) + self.gladder_to_gemb_w = jax.random.normal( + _k_gl[1], (d_global_feat, global_d_g) + ) * d_global_feat ** (-0.5) + self.gladder_to_gemb_b = jnp.zeros((global_d_g,)) + self.gladder_gemb_ln_s = jnp.ones((global_d_g,)) + _k_gt = jax.random.split(jax.random.fold_in(key, 25006), 3) + self.gladder_tree_pool = GDescriptorPool( + d_global_feat, tree_d_c, key=_k_gt[0], tag="gladder.tree.pool" + ) + self.gladder_tree_update = TreeGlobalUpdate( + d_global_feat, + self.gladder_tree_pool.d_out, + key=_k_gt[1], + tag="gladder.tree.upd", + tap_dim=global_ladder_tap_dim, + alpha_init=alpha_init, + alpha_max=ngpt_alpha_maximum, + ) + self.gladder_tree_proj_w = jax.random.normal( + _k_gt[2], (d_global_feat, global_d_g) + ) * d_global_feat ** (-0.5) + self.gladder_tree_proj_b = jnp.zeros((global_d_g,)) + _k_rp = jax.random.fold_in(key, 25010) + self.gladder_root_proj_w = jax.random.normal( + _k_rp, (d_global_feat, global_d_g) + ) * d_global_feat ** (-0.5) + self.gladder_root_proj_b = jnp.zeros((global_d_g,)) + self.gladder_root_ln_s = jnp.ones((global_d_g,)) + _k_gf = jax.random.split(jax.random.fold_in(key, 25008), 2) + self.gladder_fork_phys = EdgeRowColGlobalUpdate( + d_global_feat, + feat_d_edge, + key=_k_gf[0], + tag="gladder.fork_phys", + tap_dim=global_ladder_tap_dim, + ) + self.gladder_fork_route = EdgeRowColGlobalUpdate( + d_global_feat, + feat_d_edge, + key=_k_gf[1], + tag="gladder.fork_route", + tap_dim=global_ladder_tap_dim, + ) + d_global_route = d_global_feat + route_d_model = int(route_pointer_d_model) + route_n_heads = int(route_pointer_n_heads) + self.route_decoder = TreePrefixPointerMHSEA( + d_in=d_e, + d_edge=n_edge, + d_global=d_global_route, + d_model=route_d_model, + n_heads=route_n_heads, + attention_dim=route_pointer_attn_dim, + pointer_score_dim=route_pointer_score_dim, + candidate_hidden=route_pointer_candidate_hidden, + summary_hidden=route_pointer_summary_hidden, + ffn_hidden=route_pointer_ffn_hidden, + global_tap_dim=global_ladder_tap_dim, + alpha_init=alpha_init, + alpha_max=ngpt_alpha_maximum, + max_n=int(route_pointer_max_n), + score_init_scale=float(route_pointer_score_init_scale), + route_tree_prefix_layers=int(route_tree_prefix_layers), + route_tree_prefix_candidate_layers=int(route_tree_prefix_candidate_layers), + route_tree_prefix_merge_hidden=route_tree_prefix_merge_hidden, + route_tree_prefix_post_prefix_suffix_layers=int( + route_tree_prefix_post_prefix_suffix_layers + ), + route_decoder_attn_impl=attn_impl, + rope_base=float(route_pointer_rope_base), + rope_scaling=float(route_pointer_rope_scaling), + key=k_route, + ) + + def _gladder_g_stream(self, edge, mask, g_in): + return self.gladder_post(g_in.astype(edge.dtype), edge, mask) + + def _gladder_tree_refs(self): + return ( + self.gladder_tree_pool, + self.gladder_tree_update, + self.gladder_tree_proj_w, + self.gladder_tree_proj_b, + ) + + def _gladder_project(self, g_stream): + from .tree import _tagged_dense, _tagged_rms_eqx_style + + structural_active = jnp.asarray(True) + out = _tagged_dense( + self.gladder_to_gemb_w, + self.gladder_to_gemb_b, + g_stream, + tag_id="gladder.to_gemb", + pathway="even", + kfac_structural_mask=structural_active, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + return _tagged_rms_eqx_style( + self.gladder_gemb_ln_s, + out, + tag_id="gladder.gemb_ln", + pathway="even", + kfac_structural_mask=structural_active, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + + def _gladder_root_project(self, g_final): + from .tree import _tagged_dense, _tagged_rms_eqx_style + + structural_active = jnp.asarray(True) + out = _tagged_dense( + self.gladder_root_proj_w, + self.gladder_root_proj_b, + g_final, + tag_id="gladder.root_proj", + pathway="even", + kfac_structural_mask=structural_active, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + return _tagged_rms_eqx_style( + self.gladder_root_ln_s, + out, + tag_id="gladder.root_ln", + pathway="even", + kfac_structural_mask=structural_active, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + + def _contextualize_leaf_even_with_edge_g(self, e, edge, mask, bmask, g): + return self.readout_leaf_context.with_edge(e, edge, mask, bmask, g=g) + + def route_features( + self, ctx: SpinContext + ) -> tuple[ + Float[Array, "n d_e"], Float[Array, "n n d_edge"], Float[Array, "d_global_feat"] + ]: + edge_feat, local_feat, global_feat = self.featurizer( + ctx.J_double_prime, ctx.mask, ctx.h_prime + ) + from .tree import _tree_sphere + + g_trunk = _tree_sphere(global_feat.astype(local_feat.dtype)) + e, edge, g_trunk = self.trunk(ctx, edge_feat, local_feat, g_trunk) + g_route = self.gladder_post(g_trunk, edge, ctx.mask) + e, edge, g_route = self.route_contextualizer.with_edge( + e, edge, ctx.mask, ctx.bmask, g=g_route + ) + g_route = self.gladder_fork_route(g_route, edge, ctx.bmask) + return (e, edge, g_route.astype(e.dtype)) + + def call_with_route_logprob( + self, + q: Float[Array, "n 4"], + ctx: SpinContext, + t: Float[Array, ""] | float = 0.0, + *, + tau: float = 1.0, + ) -> tuple[Float[Array, ""], Float[Array, ""], Float[Array, ""]]: + t_val = jnp.asarray(t, dtype=q.dtype) + first_orbit_ids = ( + ctx.route_quotient_node_key, + ctx.route_quotient_edge_key, + ctx.needs_fwl2, + ) + edge_feat, local_feat, global_feat = self.featurizer( + ctx.J_double_prime, ctx.mask, ctx.h_prime + ) + from .tree import _tree_sphere + + g_trunk = _tree_sphere(global_feat.astype(local_feat.dtype)) + e, edge_route, g_trunk = self.trunk(ctx, edge_feat, local_feat, g_trunk) + z = self.q_to_odd( + q, + pathway="odd", + kfac_structural_mask=ctx.mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=False, + ) + g_post = self.gladder_post(g_trunk, edge_route, ctx.mask) + e_route, edge_route_for_policy, g_route = self.route_contextualizer.with_edge( + e, edge_route, ctx.mask, ctx.bmask, g=g_post + ) + g_route = self.gladder_fork_route(g_route, edge_route_for_policy, ctx.bmask) + route_logp = self.route_decoder.logprob_identity( + e_route, + edge_route_for_policy, + ctx.bmask, + global_feat=g_route.astype(e.dtype), + tau=tau, + real_mask=ctx.mask, + first_orbit_ids=first_orbit_ids, + ) + e_leaf, edge_leaf, g_stream = self.readout_leaf_context.with_edge( + e, edge_route, ctx.mask, ctx.bmask, g=g_post + ) + g_stream = self.gladder_fork_phys(g_stream, edge_leaf, ctx.bmask) + g_emb = self._gladder_project(g_stream) + re, im = self._forward_leaf_to_readout( + q=q, + ctx=ctx, + t=t_val, + z=z, + e_leaf=e_leaf, + g_emb=g_emb, + edge_leaf=edge_leaf, + g_stream=g_stream, + ) + return (re, im, route_logp) + + def compute_per_system_invariants( + self, ctx: SpinContext, t: Float[Array, ""] | float = 0.0 + ) -> PerSystemInvariants: + del t + edge_feat, local_feat, global_feat = self.featurizer( + ctx.J_double_prime, ctx.mask, ctx.h_prime + ) + from .tree import _tree_sphere + + g_trunk = _tree_sphere(global_feat.astype(local_feat.dtype)) + e, edge_trunk, g_trunk = self.trunk(ctx, edge_feat, local_feat, g_trunk) + g_stream = self.gladder_post(g_trunk, edge_trunk, ctx.mask) + e_leaf, edge_leaf, g_stream = self.readout_leaf_context.with_edge( + e, edge_trunk, ctx.mask, ctx.bmask, g=g_stream + ) + g_stream = self.gladder_fork_phys(g_stream, edge_leaf, ctx.bmask) + g_emb = self._gladder_project(g_stream) + return PerSystemInvariants( + g_emb=g_emb, e_leaf=e_leaf, edge_leaf=edge_leaf, g_stream=g_stream + ) + + def __call__( + self, + q: Float[Array, "n 4"], + ctx: SpinContext, + t: Float[Array, ""] | float = 0.0, + ) -> tuple[Float[Array, ""], Float[Array, ""]]: + t_val = jnp.asarray(t, dtype=q.dtype) + edge_feat, local_feat, global_feat = self.featurizer( + ctx.J_double_prime, ctx.mask, ctx.h_prime + ) + from .tree import _tree_sphere + + g_trunk = _tree_sphere(global_feat.astype(local_feat.dtype)) + e, edge_trunk, g_trunk = self.trunk(ctx, edge_feat, local_feat, g_trunk) + z = self.q_to_odd( + q, + pathway="odd", + kfac_structural_mask=ctx.mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=False, + ) + g_stream = self.gladder_post(g_trunk, edge_trunk, ctx.mask) + e_leaf, edge_leaf, g_stream = self.readout_leaf_context.with_edge( + e, edge_trunk, ctx.mask, ctx.bmask, g=g_stream + ) + g_stream = self.gladder_fork_phys(g_stream, edge_leaf, ctx.bmask) + g_emb = self._gladder_project(g_stream) + return self._forward_leaf_to_readout( + q=q, + ctx=ctx, + t=t_val, + z=z, + e_leaf=e_leaf, + g_emb=g_emb, + edge_leaf=edge_leaf, + g_stream=g_stream, + ) + + def forward_with_precomputed( + self, + q: Float[Array, "n 4"], + ctx: SpinContext, + t: Float[Array, ""] | float = 0.0, + *, + precomputed: PerSystemInvariants, + ) -> tuple[Float[Array, ""], Float[Array, ""]]: + t_val = jnp.asarray(t, dtype=q.dtype) + z = self.q_to_odd( + q, + pathway="odd", + kfac_structural_mask=ctx.mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=False, + ) + return self._forward_leaf_to_readout( + q=q, + ctx=ctx, + t=t_val, + z=z, + e_leaf=precomputed.e_leaf, + g_emb=precomputed.g_emb, + edge_leaf=precomputed.edge_leaf, + g_stream=precomputed.g_stream, + ) + + def _forward_leaf_to_readout( + self, *, q, ctx, t, z, e_leaf, g_emb, edge_leaf, g_stream + ): + del q, t + c_all, u_all, s_all = self.leaf( + e_leaf, + z, + g_emb=g_emb, + kfac_structural_mask=ctx.bmask, + kfac_odd_structural_mask=ctx.mask, + ) + u_all, log_rms = _normalize_leaf_carriers(u_all) + s_all = s_all + log_rms.astype(s_all.dtype) + mask = ctx.mask.astype(c_all.dtype) + edges = edge_leaf.astype(c_all.dtype) + reduced = balanced_tree_reduce_masked( + c_all, + u_all, + s_all, + mask, + self.merge, + g_emb, + edges_init=edges, + gladder=self._gladder_tree_refs(), + g_stream0=g_stream, + ) + reduced, final_stream = (reduced[:-1], reduced[-1]) + c_root, u_root, s_root, _, edge_root = reduced + g_emb = self._gladder_root_project(final_stream) + return self.readout( + u_root, + s_root, + e_root=edge_root, + g_emb=g_emb, + c_root=c_root, + kfac_structural_mask=jnp.any(ctx.bmask.astype(bool)), + ) diff --git a/src/hamiltonzero/model/odd_ops.py b/src/hamiltonzero/model/odd_ops.py new file mode 100644 index 0000000000000000000000000000000000000000..eafb22f354d83603899615190f4349ac85d44311 --- /dev/null +++ b/src/hamiltonzero/model/odd_ops.py @@ -0,0 +1,331 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +import math +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, PRNGKeyArray + + +def bounded_gain_logit( + value: float, + *, + max_gain: float, + init_fraction: float | None, +) -> float: + maximum = float(max_gain) + initial = float(value) + if not (math.isfinite(maximum) and maximum > 0.0): + raise ValueError(f"bounded update maximum must be positive, got {maximum}") + if init_fraction is None: + if not (math.isfinite(initial) and 0.0 < initial < maximum): + raise ValueError( + f"bounded update initial value must be in (0, {maximum}), got {initial}" + ) + fraction = initial / maximum + else: + fraction = float(init_fraction) + if not (math.isfinite(fraction) and 0.0 < fraction < 1.0): + raise ValueError( + f"bounded update initial fraction must be in (0, 1), got {fraction}" + ) + return math.log(fraction) - math.log1p(-fraction) + + +class BiasFreeLinear(eqx.Module): + weight: Float[Array, "in out"] + _use_id: str = eqx.field(static=True, default="") + + def __init__( + self, + in_features: int, + out_features: int, + *, + key: PRNGKeyArray, + scale: float | None = None, + ): + std = scale if scale is not None else in_features ** (-0.5) + self.weight = jax.random.normal(key, (in_features, out_features)) * std + self._use_id = "" + + def __call__( + self, + x: Float[Array, "... in"], + *, + pathway: str | None = None, + kfac_structural_mask=None, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ) -> Float[Array, "... out"]: + from hamiltonzero.model.tree import _tagged_dense_no_bias + + if pathway is None: + pathway = "even" + return _tagged_dense_no_bias( + self.weight, + x, + tag_id=self._use_id, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=False, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + + +class Linear(eqx.Module): + weight: Float[Array, "in out"] + bias: Float[Array, "out"] + _use_id: str = eqx.field(static=True, default="") + + def __init__(self, in_features: int, out_features: int, *, key: PRNGKeyArray): + std = in_features ** (-0.5) + self.weight = jax.random.normal(key, (in_features, out_features)) * std + self.bias = jnp.zeros((out_features,)) + self._use_id = "" + + def __call__( + self, + x: Float[Array, "... in"], + *, + pathway: str | None = None, + kfac_structural_mask=None, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ) -> Float[Array, "... out"]: + from hamiltonzero.model.tree import _tagged_dense + + if pathway is None: + pathway = "even" + return _tagged_dense( + self.weight, + self.bias, + x, + tag_id=self._use_id, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=False, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + + +class HypernetMatrix(eqx.Module): + U: Float[Array, "R d_out"] + V: Float[Array, "d_in R"] + W_h: Float[Array, "d_e R"] + _use_id_U: str = eqx.field(static=True, default="") + _use_id_V: str = eqx.field(static=True, default="") + _use_id_W_h: str = eqx.field(static=True, default="") + + def __init__( + self, d_in: int, d_out: int, d_e: int, rank: int, *, key: PRNGKeyArray + ): + k_u, k_v, k_h = jax.random.split(key, 3) + self.U = jax.random.normal(k_u, (rank, d_out)) * rank ** (-0.5) + self.V = jax.random.normal(k_v, (d_in, rank)) * d_in ** (-0.5) + self.W_h = jax.random.normal(k_h, (d_e, rank)) * d_e ** (-0.5) + self._use_id_U = "" + self._use_id_V = "" + self._use_id_W_h = "" + + def apply( + self, + e: Float[Array, "d_e"], + z: Float[Array, "d_in"], + *, + e_pathway: str | None = None, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + kfac_all_primals_reused_over_walkers: bool = False, + ) -> Float[Array, "d_out"]: + from hamiltonzero.model.tree import _tagged_dense_no_bias + + eff_e_pathway = e_pathway if e_pathway is not None else "even" + h = _tagged_dense_no_bias( + self.W_h, + e, + tag_id=self._use_id_W_h, + pathway=eff_e_pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers + or kfac_all_primals_reused_over_walkers, + ) + Vz = _tagged_dense_no_bias( + self.V, + z, + tag_id=self._use_id_V, + pathway="odd", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_all_primals_reused_over_walkers, + ) + m = h * Vz + return _tagged_dense_no_bias( + self.U, + m, + tag_id=self._use_id_U, + pathway="odd", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_all_primals_reused_over_walkers, + ) + + +class _RMS(eqx.Module): + weight: Float[Array, "d"] + eps: float = eqx.field(static=True, default=1e-05) + _use_id: str = eqx.field(static=True, default="") + + def __init__(self, d: int): + self.weight = jnp.ones((d,)) + self._use_id = "" + + def __call__( + self, + x: Float[Array, "... d"], + *, + pathway: str | None = None, + kfac_structural_mask=None, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ) -> Float[Array, "... d"]: + from hamiltonzero.model.tree import _tagged_rms_eqx_style + + if pathway is None: + pathway = "even" + return _tagged_rms_eqx_style( + self.weight, + x, + self.eps, + tag_id=self._use_id, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=False, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + + +class MLP(eqx.Module): + in_proj: Linear + block_norms: list + block_l1s: list + block_l2s: list + out_norm: _RMS + out_proj: Linear + inner_gain: float = eqx.field(static=True) + + def __init__( + self, + d_in: int, + d_hidden: int, + d_out: int, + *, + key: PRNGKeyArray, + n_blocks: int = 2, + inner_gain: float = 1.0, + ): + keys = jax.random.split(key, 2 + 2 * n_blocks) + self.in_proj = Linear(d_in, d_hidden, key=keys[0]) + self.block_norms = [_RMS(d_hidden) for _ in range(n_blocks)] + self.block_l1s = [ + Linear(d_hidden, d_hidden, key=keys[1 + 2 * i]) for i in range(n_blocks) + ] + self.block_l2s = [ + Linear(d_hidden, d_hidden, key=keys[2 + 2 * i]) for i in range(n_blocks) + ] + self.out_norm = _RMS(d_hidden) + self.out_proj = Linear(d_hidden, d_out, key=keys[-1]) + self.inner_gain = float(inner_gain) + + def _act(self, x): + return x * jax.nn.sigmoid(x) + + def __call__( + self, + x: Float[Array, "... d_in"], + *, + pathway: str | None = None, + kfac_structural_mask=None, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ) -> Float[Array, "... d_out"]: + kfac_kwargs = dict( + kfac_structural_mask=kfac_structural_mask, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + x = self.in_proj(x, pathway=pathway, **kfac_kwargs) + for nrm, l1, l2 in zip(self.block_norms, self.block_l1s, self.block_l2s): + normed = nrm(x, pathway=pathway, **kfac_kwargs) + x = x + self.inner_gain * l2( + self._act(l1(normed, pathway=pathway, **kfac_kwargs)), + pathway=pathway, + **kfac_kwargs, + ) + out_normed = self.out_norm(x, pathway=pathway, **kfac_kwargs) + return self.out_proj(out_normed, pathway=pathway, **kfac_kwargs) + + +class UnnormalizedMLP(eqx.Module): + in_proj: Linear + block_l1s: list + block_l2s: list + out_proj: Linear + inner_gain: float = eqx.field(static=True) + + def __init__( + self, + d_in: int, + d_hidden: int, + d_out: int, + *, + key: PRNGKeyArray, + n_blocks: int = 1, + inner_gain: float = 1.0, + ): + keys = jax.random.split(key, 2 + 2 * n_blocks) + self.in_proj = Linear(d_in, d_hidden, key=keys[0]) + self.block_l1s = [ + Linear(d_hidden, d_hidden, key=keys[1 + 2 * i]) for i in range(n_blocks) + ] + self.block_l2s = [ + Linear(d_hidden, d_hidden, key=keys[2 + 2 * i]) for i in range(n_blocks) + ] + self.out_proj = Linear(d_hidden, d_out, key=keys[-1]) + self.inner_gain = float(inner_gain) + + def _act(self, x): + return x * jax.nn.sigmoid(x) + + def __call__( + self, + x: Float[Array, "... d_in"], + *, + pathway: str | None = None, + kfac_structural_mask=None, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ) -> Float[Array, "... d_out"]: + kfac_kwargs = dict( + kfac_structural_mask=kfac_structural_mask, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + x = self.in_proj(x, pathway=pathway, **kfac_kwargs) + for l1, l2 in zip(self.block_l1s, self.block_l2s): + x = x + self.inner_gain * l2( + self._act(l1(x, pathway=pathway, **kfac_kwargs)), + pathway=pathway, + **kfac_kwargs, + ) + return self.out_proj(x, pathway=pathway, **kfac_kwargs) diff --git a/src/hamiltonzero/model/pallas_attention.py b/src/hamiltonzero/model/pallas_attention.py new file mode 100644 index 0000000000000000000000000000000000000000..27fc56011bf9f9fdb0593e9cc3c2d9a6cfe8707d --- /dev/null +++ b/src/hamiltonzero/model/pallas_attention.py @@ -0,0 +1,70 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int + +from ._pallas_attn import mhsea_with_tuned_ad + + +_SUPPORTED_MHSEA_HEAD_DIMS = frozenset((8, 16, 32, 64)) + + +def _require_fp32(*values: Array) -> None: + for value in values: + if jnp.dtype(value.dtype) != jnp.dtype(jnp.float32): + raise TypeError(f"attention tensors must be float32, got {value.dtype}") + + +def reference_edge_attention( + Q: Float[Array, "n h d"], + K: Float[Array, "n h d"], + V: Float[Array, "n h d"], + edge_bias: Float[Array, "n n h"], + mask: Int[Array, "n"], +) -> Float[Array, "n h d"]: + _require_fp32(Q, K, V, edge_bias) + _n, _n_heads, d_head = Q.shape + sm_scale = float(d_head) ** -0.5 + logits = jnp.einsum("ihd,jhd->hij", Q, K) * sm_scale + logits = logits + jnp.transpose(edge_bias, (2, 0, 1)) + key_mask = (1.0 - mask.astype(Q.dtype))[None, None, :] * -1e9 + logits = logits + key_mask + alpha = jax.nn.softmax(logits, axis=-1) + return jnp.einsum("hij,jhd->ihd", alpha, V) + + +def mhsea_tuned_edge_attention( + Q: Float[Array, "n h d"], + K: Float[Array, "n h d"], + V: Float[Array, "n h d"], + edge_bias: Float[Array, "n n h"], + mask: Int[Array, "n"], +) -> Float[Array, "n h d"]: + _require_fp32(Q, K, V, edge_bias) + _n, _n_heads, d_head = Q.shape + sm_scale = float(d_head) ** -0.5 + square = ( + K.shape == Q.shape + and V.shape == Q.shape + and edge_bias.shape == (_n, _n, _n_heads) + and mask.shape == (_n,) + ) + supports_tuned = square and _n >= 64 and _n % 32 == 0 + supports_full_ad = square and 0 < _n < 64 and not (_n & (_n - 1)) + if d_head not in _SUPPORTED_MHSEA_HEAD_DIMS or not ( + supports_tuned or supports_full_ad + ): + return reference_edge_attention(Q, K, V, edge_bias, mask) + Q4 = (Q * jnp.asarray(sm_scale, dtype=Q.dtype))[None] + K4 = K[None] + V4 = V[None] + e4 = jnp.transpose(edge_bias, (0, 2, 1))[None] + mask4 = mask.astype(jnp.bool_)[None] + return mhsea_with_tuned_ad(Q4, K4, e4, V4, mask4)[0] + + +__all__ = ["mhsea_tuned_edge_attention", "reference_edge_attention"] diff --git a/src/hamiltonzero/model/readout_leaf_context.py b/src/hamiltonzero/model/readout_leaf_context.py new file mode 100644 index 0000000000000000000000000000000000000000..32339c33f9d83625089b4f9e11f6708f643558da --- /dev/null +++ b/src/hamiltonzero/model/readout_leaf_context.py @@ -0,0 +1,1300 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import math + +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray + +from .odd_ops import BiasFreeLinear, MLP, _RMS + + +def _inline_norm( + norm, + x, + *, + pathway="even", + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from .tree import _tagged_rms_eqx_style + + return _tagged_rms_eqx_style( + norm.weight, + x, + eps=norm.eps, + tag_id=norm._use_id, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + +def _inline_mlp_forward( + mlp, + x, + *, + pathway="even", + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from .tree import _tagged_dense + + arguments = dict( + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + x = _tagged_dense( + mlp.in_proj.weight, + mlp.in_proj.bias, + x, + tag_id=mlp.in_proj._use_id, + **arguments, + ) + for norm, linear_1, linear_2 in zip( + mlp.block_norms, + mlp.block_l1s, + mlp.block_l2s, + strict=True, + ): + normalized = _inline_norm(norm, x, **arguments) + inner = _tagged_dense( + linear_1.weight, + linear_1.bias, + normalized, + tag_id=linear_1._use_id, + **arguments, + ) + inner = mlp._act(inner) + inner = _tagged_dense( + linear_2.weight, + linear_2.bias, + inner, + tag_id=linear_2._use_id, + **arguments, + ) + x = x + mlp.inner_gain * inner + normalized = _inline_norm(mlp.out_norm, x, **arguments) + return _tagged_dense( + mlp.out_proj.weight, + mlp.out_proj.bias, + normalized, + tag_id=mlp.out_proj._use_id, + **arguments, + ) + + +def _inline_bias_free_linear( + layer: BiasFreeLinear, + x, + *, + pathway="even", + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from .tree import _tagged_dense_no_bias + + return _tagged_dense_no_bias( + layer.weight, + x, + tag_id=layer._use_id, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + +_ALLOWED_ATTN_IMPLS = ("einsum", "mhsea_tuned") +_LCA_MAX_LEVELS = 13 + + +def default_tree_depth(max_n: int) -> int: + max_n = max(1, int(max_n)) + return max(1, int(math.ceil(math.log2(max_n)))) + + +def _relative_positions(n: int) -> tuple[Int[Array, "n n"], Float[Array, "n n 1"]]: + index = jnp.arange(n, dtype=jnp.int32) + relative = index[:, None] - index[None, :] + sign = jnp.where( + relative < 0, + 1.0, + jnp.where(relative > 0, -1.0, 0.0), + ) + return relative, sign[..., None] + + +def lca_level(pos_q, pos_k): + xor = jnp.bitwise_xor( + jnp.asarray(pos_q, jnp.int32)[..., :, None], + jnp.asarray(pos_k, jnp.int32)[..., None, :], + ) + thresholds = jnp.exp2(jnp.arange(_LCA_MAX_LEVELS, dtype=jnp.float32)).astype( + jnp.int32 + ) + return jnp.sum((xor[..., None] >= thresholds).astype(jnp.int32), axis=-1) + + +def lca_alibi_bias(pos_q, pos_k, slopes): + level = lca_level(pos_q, pos_k).astype(slopes.dtype) + return -(slopes[:, None, None] * level[None, :, :]) + + +def lca_fixed_slopes(n_heads, dtype=jnp.float32): + head = jnp.arange(n_heads, dtype=jnp.float32) + denominator = jnp.maximum(jnp.asarray(n_heads - 1, jnp.float32), 1.0) + return (1.5 * jnp.exp2(-(head * 4.0 / denominator))).astype(dtype) + + +def lca_gaussian_decay(pos_q, pos_k, w_raw, b): + level = lca_level(pos_q, pos_k).astype(b.dtype) + weight = jax.nn.softplus(w_raw) + distance = level[:, :, None] - b[None, None, :] + return jnp.exp(-(weight[None, None, :] * distance * distance)) + + +def lca_gaussian_decay_row(pos_q_scalar, pos_k, w_raw, b): + xor = jnp.bitwise_xor( + jnp.asarray(pos_q_scalar, jnp.int32), + jnp.asarray(pos_k, jnp.int32), + ) + thresholds = jnp.exp2(jnp.arange(_LCA_MAX_LEVELS, dtype=jnp.float32)).astype( + jnp.int32 + ) + level = jnp.sum( + (xor[:, None] >= thresholds[None, :]).astype(jnp.int32), + axis=-1, + ).astype(b.dtype) + weight = jax.nn.softplus(w_raw) + distance = level[:, None] - b[None, :] + return jnp.exp(-(weight[None, :] * distance * distance)) + + +def register_vector_as_dense(w2d, *, tag_id=""): + from hamiltonzero.optim.spin_blocks import register_small_full + + return register_small_full(w2d, tag_id=tag_id) + + +def lca_order_init_w_b(d_model, dtype=jnp.float32): + width = int(d_model) + centers = ( + (jnp.arange(width, dtype=jnp.float32) % 9.0).astype(dtype).reshape(1, width) + ) + raw_weight = jnp.full( + (1, width), + float(jnp.log(jnp.expm1(jnp.asarray(0.7)))), + dtype=dtype, + ) + return raw_weight, centers + + +def _attention_dimensions( + *, + d_e: int, + n_heads: int, + attn_dim: int, + attn_impl: str, + require_even_model: bool, +) -> tuple[int, int, int]: + if n_heads < 1: + raise ValueError("contextualizer n_heads must be positive") + if attn_dim < 1 or attn_dim % n_heads: + raise ValueError("contextualizer attention width must divide by n_heads") + if require_even_model and d_e % 2: + raise ValueError("physical contextualizer width must be even") + d_head = attn_dim // n_heads + if d_head % 2: + raise ValueError("contextualizer attention head width must be even") + if attn_impl not in _ALLOWED_ATTN_IMPLS: + raise ValueError("attention must be 'einsum' or 'mhsea_tuned'") + return 2 * n_heads, d_head, attn_dim + + +def _global_modules( + key, d_g: int, d_e: int, residual_scale: float, global_tap_dim: int +): + from .global_ladder import GDescriptorPool, ResidualGlobalUpdate + + keys = jax.random.split(jax.random.fold_in(key, 25007), 3) + pool = GDescriptorPool(d_g, d_e, key=keys[0], tag="gladder.ctx.pool") + update = ResidualGlobalUpdate( + d_g, + pool.d_out, + key=keys[1], + tag="gladder.ctx.upd", + tap_dim=global_tap_dim, + residual_gain=residual_scale, + ) + projection = jax.random.normal(keys[2], (d_g, 64)) * d_g ** (-0.5) + return pool, update, projection + + +def _run_attention(query, key, value, bias, mask, *, implementation: str, d_head: int): + from .pallas_attention import mhsea_tuned_edge_attention, reference_edge_attention + + if implementation == "einsum": + return reference_edge_attention(query, key, value, bias, mask) + padded_width = max(16, d_head) + padding = padded_width - d_head + scale = jnp.sqrt(jnp.asarray(padded_width / d_head, dtype=jnp.float32)).astype( + query.dtype + ) + zeros = jnp.zeros( + (query.shape[0], query.shape[1], padding), + dtype=query.dtype, + ) + query = jnp.concatenate([query * scale, zeros], axis=-1) + key = jnp.concatenate([key, zeros], axis=-1) + value = jnp.concatenate([value, zeros], axis=-1) + return mhsea_tuned_edge_attention(query, key, value, bias, mask)[..., :d_head] + + +def _edge_update_dense(layer, edge, edge_n, edge_context, bmask_f, g): + n = edge.shape[0] + structural = bmask_f.astype(bool) + pair_structural = structural[:, None] & structural[None, :] + context_i = jnp.broadcast_to( + edge_context[:, None, :], + (n, n, edge_context.shape[-1]), + ) + context_j = jnp.broadcast_to( + edge_context[None, :, :], + (n, n, edge_context.shape[-1]), + ) + pair_mask = (bmask_f[:, None] * bmask_f[None, :])[..., None] + from .tree import _tagged_dense_no_bias + + global_edge = _tagged_dense_no_bias( + layer.g_edge_proj_w, + g, + tag_id="gladder.ctx.eproj", + pathway="even", + kfac_structural_mask=jnp.any(structural), + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=True, + ) + edge_input = jnp.concatenate( + [ + pair_mask * edge_n, + context_i, + context_j, + jnp.broadcast_to( + global_edge[None, None, :], + (n, n, global_edge.shape[-1]), + ).astype(edge_n.dtype), + ], + axis=-1, + ) + delta = _inline_mlp_forward( + layer.edge_ffn, + edge_input, + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + return jnp.where( + pair_structural[..., None], + edge + layer.residual_scale * delta, + jnp.zeros_like(edge), + ) + + +def _edge_update_tiled( + layer, + edge_rows, + edge_n_rows, + edge_context_rows, + edge_context_all, + bmask, + *, + row_indices=None, + g, + tile_size: int, +): + rows, n, _ = edge_rows.shape + if edge_n_rows.shape != edge_rows.shape: + raise ValueError("normalized edge rows must match edge rows") + if edge_context_all.shape[0] != n or bmask.shape != (n,): + raise ValueError("context and mask must have global width N") + if row_indices is None: + if rows != n: + raise ValueError("row_indices is required for row-sharded edges") + row_indices = jnp.arange(n, dtype=jnp.int32) + row_indices = jnp.asarray(row_indices, dtype=jnp.int32) + mask = bmask.astype(edge_rows.dtype) + mask_rows = mask[row_indices] + from .tree import _tagged_dense_no_bias + + global_edge = _tagged_dense_no_bias( + layer.g_edge_proj_w, + g, + tag_id="gladder.ctx.eproj", + pathway="even", + kfac_structural_mask=jnp.any(bmask.astype(bool)), + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=True, + ) + tile_width = min(n, int(tile_size)) + if tile_width < 1: + raise ValueError("tile_size must be positive") + full_tiles = n // tile_width + tail_start = full_tiles * tile_width + + def update_tile(start, width, output): + mask_tile = jax.lax.dynamic_slice_in_dim(mask, start, width, axis=0) + edge_n_tile = jax.lax.dynamic_slice_in_dim(edge_n_rows, start, width, axis=1) + context_tile = jax.lax.dynamic_slice_in_dim( + edge_context_all, start, width, axis=0 + ) + edge_tile = jax.lax.dynamic_slice_in_dim(output, start, width, axis=1) + pair_structural = mask_rows[:, None].astype(bool) & mask_tile[None, :].astype( + bool + ) + context_i = jnp.broadcast_to( + edge_context_rows[:, None, :], + (rows, width, edge_context_rows.shape[-1]), + ) + context_j = jnp.broadcast_to( + context_tile[None, :, :], + (rows, width, edge_context_all.shape[-1]), + ) + edge_input = jnp.concatenate( + [ + (mask_rows[:, None] * mask_tile[None, :])[..., None] * edge_n_tile, + context_i, + context_j, + jnp.broadcast_to( + global_edge, + (rows, width, global_edge.shape[-1]), + ).astype(edge_rows.dtype), + ], + axis=-1, + ) + delta = _inline_mlp_forward( + layer.edge_ffn, + edge_input, + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + updated = jnp.where( + pair_structural[..., None], + edge_tile + layer.residual_scale * delta, + jnp.zeros_like(edge_tile), + ) + return jax.lax.dynamic_update_slice_in_dim(output, updated, start, axis=1) + + output = jax.lax.fori_loop( + 0, + full_tiles, + lambda tile_index, current: update_tile( + tile_index * tile_width, tile_width, current + ), + edge_rows, + ) + if tail_start < n: + output = update_tile(tail_start, n - tail_start, output) + return output + + +class PhysicalReadoutContextLayer(eqx.Module): + ln_edge: _RMS + ln_edge_attn: _RMS + ln_c: _RMS + ln_summary: _RMS + ln_attn: _RMS + ln_edge_ctx: _RMS + summary_mlp: MLP + ctx_mlp: MLP + edge_node_ctx_proj: BiasFreeLinear + edge_ffn: MLP + g_pool: "GDescriptorPool" + g_update: "ResidualGlobalUpdate" + g_edge_proj_w: Float[Array, "d_g 64"] + W_QKV: BiasFreeLinear + W_O: BiasFreeLinear + bias_mlp: MLP + d_e: int = eqx.field(static=True) + n_heads: int = eqx.field(static=True) + n_heads_kernel: int = eqx.field(static=True) + d_head: int = eqx.field(static=True) + d_attn: int = eqx.field(static=True) + attn_impl: str = eqx.field(static=True) + rope_base: float = eqx.field(static=True) + rope_scaling: float = eqx.field(static=True) + residual_scale: float = eqx.field(static=True) + + def __init__( + self, + *, + d_e: int, + d_edge: int, + n_heads: int, + summary_hidden: int, + mlp_hidden: int, + bias_hidden: int, + edge_ffn_hidden: int, + attn_dim: int, + edge_node_ctx_dim: int, + attn_impl: str, + rope_base: float, + rope_scaling: float, + gladder_d_g: int, + global_tap_dim: int, + residual_scale: float, + key: PRNGKeyArray, + ): + n_heads_kernel, d_head, d_attn = _attention_dimensions( + d_e=d_e, + n_heads=n_heads, + attn_dim=attn_dim, + attn_impl=attn_impl, + require_even_model=True, + ) + if rope_base <= 0.0 or rope_scaling <= 0.0: + raise ValueError("physical contextualizer rope values must be positive") + key_sum, key_ctx, key_edge, key_qkv, key_out, key_bias = jax.random.split( + key, 6 + ) + self.ln_edge = _RMS(d_edge) + self.ln_edge_attn = _RMS(d_edge) + self.ln_c = _RMS(d_e) + self.ln_summary = _RMS(d_e) + self.ln_attn = _RMS(d_e) + self.ln_edge_ctx = _RMS(d_e) + self.summary_mlp = MLP( + 2 * d_edge + 1, + int(summary_hidden), + d_e, + key=key_sum, + n_blocks=1, + ) + self.ctx_mlp = MLP( + 2 * d_e, + int(mlp_hidden), + d_e, + key=key_ctx, + n_blocks=2, + ) + self.edge_node_ctx_proj = BiasFreeLinear( + d_e, + int(edge_node_ctx_dim), + key=jax.random.fold_in(key, 3694), + ) + self.edge_ffn = MLP( + d_edge + 2 * int(edge_node_ctx_dim) + 64, + int(edge_ffn_hidden), + d_edge, + key=key_edge, + n_blocks=1, + ) + self.g_pool, self.g_update, self.g_edge_proj_w = _global_modules( + key, + int(gladder_d_g), + d_e, + residual_scale, + global_tap_dim, + ) + self.W_QKV = BiasFreeLinear(d_e, 3 * n_heads_kernel * d_head, key=key_qkv) + self.W_O = BiasFreeLinear(d_attn, d_e, key=key_out) + self.bias_mlp = MLP( + d_edge + 1, + int(bias_hidden), + n_heads_kernel, + key=key_bias, + n_blocks=1, + ) + self.d_e = int(d_e) + self.n_heads = int(n_heads) + self.n_heads_kernel = int(n_heads_kernel) + self.d_head = int(d_head) + self.d_attn = int(d_attn) + self.attn_impl = str(attn_impl) + self.rope_base = float(rope_base) + self.rope_scaling = float(rope_scaling) + self.residual_scale = float(residual_scale) + + def _slot_clock(self, n: int, dtype, bmask) -> Float[Array, "n d_e"]: + position = jnp.arange(n, dtype=jnp.float32) + n_active = jnp.maximum(jnp.sum(jnp.asarray(bmask, dtype=jnp.int32)), 1) + depth = jnp.ceil(jnp.log2(n_active.astype(jnp.float32))).astype(jnp.int32) + span = jnp.power( + jnp.asarray(2.0, dtype=jnp.float32), + depth.astype(jnp.float32), + ) + position = position - (span - jnp.asarray(1.0, jnp.float32)) * 0.5 + position = position / jnp.asarray(self.rope_scaling, jnp.float32) + half = (self.d_e + 1) // 2 + band = jnp.arange(half, dtype=jnp.float32) + inverse_frequency = jnp.exp( + -jnp.log(jnp.asarray(self.rope_base, jnp.float32)) + * band + / jnp.asarray(max(half, 1), jnp.float32) + ) + angle = position[:, None] * inverse_frequency[None, :] + embedding = jnp.concatenate([jnp.sin(angle), jnp.cos(angle)], axis=-1) + return embedding[:, : self.d_e].astype(dtype) + + def _edge_summary(self, edge_n, direction, bmask_f): + n = edge_n.shape[0] + dtype = edge_n.dtype + structural = bmask_f.astype(bool) + pair_structural = ( + structural[:, None] & structural[None, :] & ~jnp.eye(n, dtype=bool) + ) + pair = jnp.concatenate( + [edge_n, jnp.swapaxes(edge_n, 0, 1), direction.astype(dtype)], + axis=-1, + ) + message = _inline_mlp_forward( + self.summary_mlp, + pair, + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + off_diagonal = 1.0 - jnp.eye(n, dtype=dtype) + weight = bmask_f[:, None] * bmask_f[None, :] * off_diagonal + index = jnp.arange(n, dtype=jnp.int32) + weight = weight * jnp.exp(-0.5 * lca_level(index, index).astype(dtype)) + denominator = jnp.maximum(jnp.sum(weight, axis=1, keepdims=True), 1.0) + return jnp.sum(message * weight[..., None], axis=1) / denominator + + def edge_summary_tiled( + self, + edge_n_rows, + summary_mask, + *, + edge_reverse_rows=None, + row_indices=None, + tile_size: int = 128, + ): + rows, n, _ = edge_n_rows.shape + dtype = edge_n_rows.dtype + if summary_mask.shape != (n,): + raise ValueError("summary_mask must have global shape [N]") + if row_indices is None: + if rows != n: + raise ValueError("row_indices is required for sharded summaries") + row_indices = jnp.arange(n, dtype=jnp.int32) + row_indices = jnp.asarray(row_indices, dtype=jnp.int32) + if edge_reverse_rows is None: + if rows != n: + raise ValueError("reverse rows are required for sharded summaries") + edge_reverse_rows = jnp.swapaxes(edge_n_rows, 0, 1) + mask = summary_mask.astype(dtype) + mask_rows = mask[row_indices] + numerator = jnp.zeros((rows, self.d_e), dtype=dtype) + denominator = jnp.zeros((rows, 1), dtype=dtype) + tile_width = min(n, int(tile_size)) + if tile_width < 1: + raise ValueError("tile_size must be positive") + full_tiles = n // tile_width + tail_start = full_tiles * tile_width + + def accumulate(start, width, carry): + numerator, denominator = carry + column = start + jnp.arange(width, dtype=jnp.int32) + relative = row_indices[:, None] - column[None, :] + direction = jnp.where( + relative < 0, + 1.0, + jnp.where(relative > 0, -1.0, 0.0), + ).astype(dtype)[..., None] + off_diagonal = row_indices[:, None] != column[None, :] + mask_tile = jax.lax.dynamic_slice_in_dim(mask, start, width, axis=0) + edge_tile = jax.lax.dynamic_slice_in_dim(edge_n_rows, start, width, axis=1) + reverse_tile = jax.lax.dynamic_slice_in_dim( + edge_reverse_rows, start, width, axis=1 + ) + structural = ( + mask_rows[:, None].astype(bool) + & mask_tile[None, :].astype(bool) + & off_diagonal + ) + pair = jnp.concatenate([edge_tile, reverse_tile, direction], axis=-1) + message = _inline_mlp_forward( + self.summary_mlp, + pair, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + weight = ( + mask_rows[:, None] * mask_tile[None, :] * off_diagonal.astype(dtype) + ) + weight = weight * jnp.exp( + -0.5 * lca_level(row_indices, column).astype(dtype) + ) + numerator = numerator + jnp.sum(message * weight[..., None], axis=1) + denominator = denominator + jnp.sum(weight, axis=1, keepdims=True) + return numerator, denominator + + numerator, denominator = jax.lax.fori_loop( + 0, + full_tiles, + lambda tile_index, carry: accumulate( + tile_index * tile_width, tile_width, carry + ), + (numerator, denominator), + ) + if tail_start < n: + numerator, denominator = accumulate( + tail_start, + n - tail_start, + (numerator, denominator), + ) + return numerator / jnp.maximum(denominator, 1.0) + + def edge_update_tiled(self, *args, **kwargs): + return _edge_update_tiled(self, *args, **kwargs) + + def _attend(self, c, edge_n, direction, bmask): + n = c.shape[0] + dtype = c.dtype + structural = bmask.astype(bool) + pair_structural = structural[:, None] & structural[None, :] + qkv = _inline_bias_free_linear( + self.W_QKV, + c, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ).reshape(n, 3, self.n_heads_kernel, self.d_head) + query, key, value = qkv[:, 0], qkv[:, 1], qkv[:, 2] + bias_input = jnp.concatenate([edge_n, direction.astype(dtype)], axis=-1) + bias = _inline_mlp_forward( + self.bias_mlp, + bias_input, + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + bias = bias / jnp.sqrt(jnp.asarray(self.d_head, dtype=dtype)) + index = jnp.arange(n, dtype=jnp.int32) + bias = bias + jnp.transpose( + lca_alibi_bias( + index, + index, + lca_fixed_slopes(self.n_heads_kernel, dtype=dtype), + ), + (1, 2, 0), + ) + output = _run_attention( + query, + key, + value, + bias, + bmask.astype(dtype), + implementation=self.attn_impl, + d_head=self.d_head, + ) + output = jax.nn.sigmoid(output[:, : self.n_heads]) * output[:, self.n_heads :] + return _inline_bias_free_linear( + self.W_O, + output.reshape(n, -1), + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + + def __call__(self, c, edge, mask, bmask, g): + del mask + n = c.shape[0] + dtype = c.dtype + bmask_f = bmask.astype(dtype) + structural = bmask.astype(bool) + pair_structural = structural[:, None] & structural[None, :] + edge_n = _inline_norm( + self.ln_edge, + edge.astype(dtype), + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + _relative, direction = _relative_positions(n) + direction = direction.astype(dtype) + clock = self._slot_clock(n, dtype, bmask) + summary = self._edge_summary(edge_n, direction, bmask_f) + context_input = jnp.concatenate( + [ + _inline_norm( + self.ln_c, + c + clock, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ), + _inline_norm( + self.ln_summary, + summary, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ), + ], + axis=-1, + ) + delta = self.residual_scale * _inline_mlp_forward( + self.ctx_mlp, + context_input, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + c1 = jnp.where(structural[:, None], c + delta, jnp.zeros_like(c)) + global_active = jnp.any(structural) + g = self.g_update( + g, + self.g_pool( + g, + c1, + bmask_f, + kfac_structural_mask=structural, + kfac_update_mask=global_active, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ), + kfac_structural_mask=global_active, + kfac_context_primal_reused_over_walkers=True, + ) + edge_context = _inline_bias_free_linear( + self.edge_node_ctx_proj, + _inline_norm( + self.ln_edge_ctx, + c1, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ), + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + edge1 = _edge_update_dense(self, edge, edge_n, edge_context, bmask_f, g) + attention_input = _inline_norm( + self.ln_attn, + c1 + clock, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + edge_attention = _inline_norm( + self.ln_edge_attn, + edge1.astype(dtype), + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + delta_attention = self.residual_scale * self._attend( + attention_input, edge_attention, direction, bmask + ) + output = jnp.where( + structural[:, None], + c1 + delta_attention, + jnp.zeros_like(c1), + ) + return output, edge1, g + + +class RouterContextLayer(eqx.Module): + ln_edge: _RMS + ln_edge_attn: _RMS + ln_c: _RMS + ln_attn: _RMS + ln_edge_ctx: _RMS + ctx_mlp: MLP + edge_node_ctx_proj: BiasFreeLinear + edge_ffn: MLP + g_pool: "GDescriptorPool" + g_update: "ResidualGlobalUpdate" + g_edge_proj_w: Float[Array, "d_g 64"] + W_QKV: BiasFreeLinear + W_O: BiasFreeLinear + bias_mlp: MLP + n_heads: int = eqx.field(static=True) + n_heads_kernel: int = eqx.field(static=True) + d_head: int = eqx.field(static=True) + d_attn: int = eqx.field(static=True) + attn_impl: str = eqx.field(static=True) + residual_scale: float = eqx.field(static=True) + + def __init__( + self, + *, + d_e: int, + d_edge: int, + n_heads: int, + mlp_hidden: int, + bias_hidden: int, + edge_ffn_hidden: int, + attn_dim: int, + edge_node_ctx_dim: int, + attn_impl: str, + gladder_d_g: int, + global_tap_dim: int, + residual_scale: float, + key: PRNGKeyArray, + ): + n_heads_kernel, d_head, d_attn = _attention_dimensions( + d_e=d_e, + n_heads=n_heads, + attn_dim=attn_dim, + attn_impl=attn_impl, + require_even_model=False, + ) + _key_sum, key_ctx, key_edge, key_qkv, key_out, key_bias = jax.random.split( + key, 6 + ) + self.ln_edge = _RMS(d_edge) + self.ln_edge_attn = _RMS(d_edge) + self.ln_c = _RMS(d_e) + self.ln_attn = _RMS(d_e) + self.ln_edge_ctx = _RMS(d_e) + self.ctx_mlp = MLP( + d_e, + int(mlp_hidden), + d_e, + key=key_ctx, + n_blocks=2, + ) + self.edge_node_ctx_proj = BiasFreeLinear( + d_e, + int(edge_node_ctx_dim), + key=jax.random.fold_in(key, 3694), + ) + self.edge_ffn = MLP( + d_edge + 2 * int(edge_node_ctx_dim) + 64, + int(edge_ffn_hidden), + d_edge, + key=key_edge, + n_blocks=1, + ) + self.g_pool, self.g_update, self.g_edge_proj_w = _global_modules( + key, + int(gladder_d_g), + d_e, + residual_scale, + global_tap_dim, + ) + self.W_QKV = BiasFreeLinear(d_e, 3 * n_heads_kernel * d_head, key=key_qkv) + self.W_O = BiasFreeLinear(d_attn, d_e, key=key_out) + self.bias_mlp = MLP( + d_edge + 1, + int(bias_hidden), + n_heads_kernel, + key=key_bias, + n_blocks=1, + ) + self.n_heads = int(n_heads) + self.n_heads_kernel = int(n_heads_kernel) + self.d_head = int(d_head) + self.d_attn = int(d_attn) + self.attn_impl = str(attn_impl) + self.residual_scale = float(residual_scale) + + def edge_update_tiled(self, *args, **kwargs): + return _edge_update_tiled(self, *args, **kwargs) + + def _attend(self, c, edge_n, bmask): + n = c.shape[0] + dtype = c.dtype + structural = bmask.astype(bool) + pair_structural = structural[:, None] & structural[None, :] + qkv = _inline_bias_free_linear( + self.W_QKV, + c, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ).reshape(n, 3, self.n_heads_kernel, self.d_head) + query, key, value = qkv[:, 0], qkv[:, 1], qkv[:, 2] + direction = jnp.zeros((n, n, 1), dtype=dtype) + bias = _inline_mlp_forward( + self.bias_mlp, + jnp.concatenate([edge_n, direction], axis=-1), + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + bias = bias / jnp.sqrt(jnp.asarray(self.d_head, dtype=dtype)) + output = _run_attention( + query, + key, + value, + bias, + bmask.astype(dtype), + implementation=self.attn_impl, + d_head=self.d_head, + ) + output = jax.nn.sigmoid(output[:, : self.n_heads]) * output[:, self.n_heads :] + return _inline_bias_free_linear( + self.W_O, + output.reshape(n, -1), + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + + def __call__(self, c, edge, mask, bmask, g): + del mask + dtype = c.dtype + bmask_f = bmask.astype(dtype) + structural = bmask.astype(bool) + pair_structural = structural[:, None] & structural[None, :] + edge_n = _inline_norm( + self.ln_edge, + edge.astype(dtype), + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + context_input = _inline_norm( + self.ln_c, + c, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + delta = self.residual_scale * _inline_mlp_forward( + self.ctx_mlp, + context_input, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + c1 = jnp.where(structural[:, None], c + delta, jnp.zeros_like(c)) + global_active = jnp.any(structural) + g = self.g_update( + g, + self.g_pool( + g, + c1, + bmask_f, + kfac_structural_mask=structural, + kfac_update_mask=global_active, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ), + kfac_structural_mask=global_active, + kfac_context_primal_reused_over_walkers=True, + ) + edge_context = _inline_bias_free_linear( + self.edge_node_ctx_proj, + _inline_norm( + self.ln_edge_ctx, + c1, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ), + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + edge1 = _edge_update_dense(self, edge, edge_n, edge_context, bmask_f, g) + attention_input = _inline_norm( + self.ln_attn, + c1, + pathway="even", + kfac_structural_mask=structural, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + edge_attention = _inline_norm( + self.ln_edge_attn, + edge1.astype(dtype), + pathway="even", + kfac_structural_mask=pair_structural, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + delta_attention = self.residual_scale * self._attend( + attention_input, edge_attention, bmask + ) + output = jnp.where( + structural[:, None], + c1 + delta_attention, + jnp.zeros_like(c1), + ) + return output, edge1, g + + +def _stack_layers(layers): + dynamic_static = [eqx.partition(layer, eqx.is_array) for layer in layers] + dynamic = [item[0] for item in dynamic_static] + static = dynamic_static[0][1] + stacked = jax.tree.map(lambda *values: jnp.stack(values, axis=0), *dynamic) + return eqx.combine(stacked, static) + + +def _initialize_structural_inputs(contextualizer, c, edge, mask, bmask): + real = mask.astype(bool) + active = bmask.astype(bool) + virtual = active & ~real + virtual_node = register_vector_as_dense( + contextualizer.virtual_node, + tag_id=contextualizer._use_id_virtual_node, + )[0].astype(c.dtype) + empty_nonempty = register_vector_as_dense( + contextualizer.edge_empty_nonempty, + tag_id=contextualizer._use_id_edge_empty_nonempty, + )[0].astype(edge.dtype) + empty_empty = register_vector_as_dense( + contextualizer.edge_empty_empty, + tag_id=contextualizer._use_id_edge_empty_empty, + )[0].astype(edge.dtype) + c = jnp.where( + real[:, None], + c, + jnp.where( + virtual[:, None], + virtual_node[None, :], + jnp.zeros_like(c), + ), + ) + real_pair = real[:, None] & real[None, :] + mixed_pair = (real[:, None] & virtual[None, :]) | (virtual[:, None] & real[None, :]) + virtual_pair = virtual[:, None] & virtual[None, :] + edge = jnp.where( + real_pair[..., None], + edge, + jnp.where( + mixed_pair[..., None], + empty_nonempty[None, None, :], + jnp.where( + virtual_pair[..., None], + empty_empty[None, None, :], + jnp.zeros_like(edge), + ), + ), + ) + return c, edge + + +def _run_context_layers(contextualizer, c, edge, mask, bmask, g): + dynamic, static = eqx.partition(contextualizer.layers, eqx.is_array) + + def scan_step(carry, layer_dynamic): + layer = eqx.combine(layer_dynamic, static) + c_value, edge_value, g_value = carry + return ( + layer(c_value, edge_value, mask, bmask, g_value), + None, + ) + + (c, edge, g), _ = jax.lax.scan( + scan_step, + (c, edge, g), + dynamic, + ) + return c, edge, g + + +class PhysicalReadoutContext(eqx.Module): + layers: PhysicalReadoutContextLayer + virtual_node: Float[Array, "one d_e"] + edge_empty_nonempty: Float[Array, "one d_edge"] + edge_empty_empty: Float[Array, "one d_edge"] + _use_id_virtual_node: str = eqx.field(static=True, default="") + _use_id_edge_empty_nonempty: str = eqx.field(static=True, default="") + _use_id_edge_empty_empty: str = eqx.field(static=True, default="") + + def __init__( + self, + *, + d_e: int, + d_edge: int, + n_layers: int, + n_heads: int, + summary_hidden: int, + mlp_hidden: int, + bias_hidden: int, + edge_ffn_hidden: int, + attn_dim: int, + edge_node_ctx_dim: int, + attn_impl: str, + rope_base: float, + rope_scaling: float, + gladder_d_g: int, + global_tap_dim: int, + key: PRNGKeyArray, + ): + if n_layers < 1: + raise ValueError("physical contextualizer layers must be positive") + residual_scale = float(n_layers) ** (-0.5) + keys = jax.random.split(key, n_layers) + self.layers = _stack_layers( + [ + PhysicalReadoutContextLayer( + d_e=d_e, + d_edge=d_edge, + n_heads=n_heads, + summary_hidden=summary_hidden, + mlp_hidden=mlp_hidden, + bias_hidden=bias_hidden, + edge_ffn_hidden=edge_ffn_hidden, + attn_dim=attn_dim, + edge_node_ctx_dim=edge_node_ctx_dim, + attn_impl=attn_impl, + rope_base=rope_base, + rope_scaling=rope_scaling, + gladder_d_g=gladder_d_g, + global_tap_dim=global_tap_dim, + residual_scale=residual_scale, + key=layer_key, + ) + for layer_key in keys + ] + ) + self.virtual_node = jax.random.normal( + jax.random.fold_in(key, 201793223), (1, d_e) + ) * d_e ** (-0.5) + self.edge_empty_nonempty = jax.random.normal( + jax.random.fold_in(key, 235798529), (1, d_edge) + ) * d_edge ** (-0.5) + self.edge_empty_empty = jax.random.normal( + jax.random.fold_in(key, 235798530), (1, d_edge) + ) * d_edge ** (-0.5) + self._use_id_virtual_node = "" + self._use_id_edge_empty_nonempty = "" + self._use_id_edge_empty_empty = "" + + def with_edge(self, c, edge, mask, bmask, *, g): + c = c.astype(jnp.float32) + edge = edge.astype(jnp.float32) + c, edge = _initialize_structural_inputs(self, c, edge, mask, bmask) + return _run_context_layers(self, c, edge, mask, bmask, g) + + +class RouterContext(eqx.Module): + layers: RouterContextLayer + virtual_node: Float[Array, "one d_e"] + edge_empty_nonempty: Float[Array, "one d_edge"] + edge_empty_empty: Float[Array, "one d_edge"] + _use_id_virtual_node: str = eqx.field(static=True, default="") + _use_id_edge_empty_nonempty: str = eqx.field(static=True, default="") + _use_id_edge_empty_empty: str = eqx.field(static=True, default="") + + def __init__( + self, + *, + d_e: int, + d_edge: int, + n_layers: int, + n_heads: int, + mlp_hidden: int, + bias_hidden: int, + edge_ffn_hidden: int, + attn_dim: int, + edge_node_ctx_dim: int, + attn_impl: str, + gladder_d_g: int, + global_tap_dim: int, + key: PRNGKeyArray, + ): + if n_layers < 1: + raise ValueError("router contextualizer layers must be positive") + residual_scale = float(n_layers) ** (-0.5) + keys = jax.random.split(key, n_layers) + self.layers = _stack_layers( + [ + RouterContextLayer( + d_e=d_e, + d_edge=d_edge, + n_heads=n_heads, + mlp_hidden=mlp_hidden, + bias_hidden=bias_hidden, + edge_ffn_hidden=edge_ffn_hidden, + attn_dim=attn_dim, + edge_node_ctx_dim=edge_node_ctx_dim, + attn_impl=attn_impl, + gladder_d_g=gladder_d_g, + global_tap_dim=global_tap_dim, + residual_scale=residual_scale, + key=layer_key, + ) + for layer_key in keys + ] + ) + self.virtual_node = jax.random.normal( + jax.random.fold_in(key, 201793223), (1, d_e) + ) * d_e ** (-0.5) + self.edge_empty_nonempty = jax.random.normal( + jax.random.fold_in(key, 235798529), (1, d_edge) + ) * d_edge ** (-0.5) + self.edge_empty_empty = jax.random.normal( + jax.random.fold_in(key, 235798530), (1, d_edge) + ) * d_edge ** (-0.5) + self._use_id_virtual_node = "" + self._use_id_edge_empty_nonempty = "" + self._use_id_edge_empty_empty = "" + + def with_edge(self, c, edge, mask, bmask, *, g): + c = c.astype(jnp.float32) + edge = edge.astype(jnp.float32) + c, edge = _initialize_structural_inputs(self, c, edge, mask, bmask) + return _run_context_layers(self, c, edge, mask, bmask, g) + + +__all__ = [ + "PhysicalReadoutContext", + "PhysicalReadoutContextLayer", + "RouterContext", + "RouterContextLayer", + "default_tree_depth", + "lca_alibi_bias", + "lca_fixed_slopes", + "lca_gaussian_decay", + "lca_gaussian_decay_row", + "lca_level", + "lca_order_init_w_b", + "register_vector_as_dense", +] diff --git a/src/hamiltonzero/model/route_pointer.py b/src/hamiltonzero/model/route_pointer.py new file mode 100644 index 0000000000000000000000000000000000000000..903a398abc1d96ea12502e95c67f63c8a27b6732 --- /dev/null +++ b/src/hamiltonzero/model/route_pointer.py @@ -0,0 +1,6233 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray + +from .fused_silu import fused_silu +from .readout_leaf_context import ( + default_tree_depth, + lca_alibi_bias, + lca_fixed_slopes, + lca_gaussian_decay, + lca_gaussian_decay_row, +) +from .route_quotient import ( + conditional_orbit_ids_from_keys, + conditional_orbit_pair_ids_from_keys, +) +from .tree import ( + CausalRouterEdgeFWLUpdate, + EdgeMergeOp, + _TREE_NGPT_DEPTH_FEAT_DIM, + _tree_ngpt_residual, + _tree_clock_root_center_from_depth, + _tree_dyadic_segment_clock, + _tree_ngpt_level_counts, + _tree_sphere, +) + +QuotientCarrier = tuple[Array, Array] | tuple[Array, Array, Array] + + +def _route_next_pow2(n: int) -> int: + return 1 << (int(n) - 1).bit_length() + + +def _dyadic_frontier_add(frontier, value, position): + + carry = jnp.asarray(value, dtype=frontier.dtype) + active = jnp.asarray(True) + pos = jnp.asarray(position, dtype=jnp.int32) + for level in range(frontier.shape[0]): + old = frontier[level] + occupied = jnp.bitwise_and(jnp.right_shift(pos, level), 1) == 1 + merge = active & occupied + place = active & ~occupied + frontier = frontier.at[level].set( + jnp.where(place, carry, jnp.where(merge, jnp.zeros_like(old), old)) + ) + carry = jnp.where(merge, old + carry, carry) + active = merge + return frontier + + +def _dyadic_lca_frontier_sum(frontier, position, w_raw, b): + + depth = frontier.shape[0] + pos = jnp.asarray(position, dtype=jnp.int32) + levels = jnp.arange(depth, dtype=jnp.int32) + widths = jnp.left_shift(jnp.ones((depth,), dtype=jnp.int32), levels + 1) + starts = jnp.bitwise_and(pos, jnp.bitwise_not(widths - 1)) + present = jnp.bitwise_and(jnp.right_shift(pos, levels), 1) == 1 + decay = lca_gaussian_decay_row(pos, starts, w_raw, b) + scale = jnp.where(present[:, None], decay, jnp.zeros_like(decay)) + while scale.ndim < frontier.ndim: + scale = scale[:, None, :] + return jnp.sum(frontier * scale, axis=0) + + +def _replace_square_row_column(matrix, index, row_value, column_value): + + idx = jnp.arange(matrix.shape[0], dtype=jnp.int32) + select = idx == jnp.asarray(index, dtype=jnp.int32) + diagonal = row_value[index] + column_value = jnp.where(select[:, None], diagonal, column_value) + matrix = jnp.where(select[:, None, None], row_value[None, :, :], matrix) + return jnp.where( + select[None, :, None], + column_value[:, None, :], + matrix, + ) + + +def _square_row_by_reduction(matrix, index): + + idx = jnp.arange(matrix.shape[0], dtype=jnp.int32) + select = idx == jnp.asarray(index, dtype=jnp.int32) + return jnp.sum( + jnp.where(select[:, None, None], matrix, jnp.zeros_like(matrix)), + axis=0, + ) + + +def _square_column_local(matrix, index): + + idx = jnp.arange(matrix.shape[1], dtype=jnp.int32) + select = idx == jnp.asarray(index, dtype=jnp.int32) + return jnp.sum( + jnp.where(select[None, :, None], matrix, jnp.zeros_like(matrix)), + axis=1, + ) + + +def _route_clock(pos, width: int, dtype, *, base: float | Array = 10000.0, scale=None): + if int(width) <= 0: + pos_arr = jnp.asarray(pos) + return jnp.zeros(pos_arr.shape + (0,), dtype=dtype) + pos_f = jnp.asarray(pos, dtype=jnp.float32) + if scale is not None: + denom = jnp.maximum( + jnp.asarray(scale, dtype=jnp.float32) - jnp.asarray(1.0, dtype=jnp.float32), + jnp.asarray(1.0, dtype=jnp.float32), + ) + pos_f = pos_f / denom + half = (int(width) + 1) // 2 + band = jnp.arange(half, dtype=jnp.float32) + base_f = jnp.maximum( + jnp.asarray(base, dtype=jnp.float32), + jnp.asarray(2.0, dtype=jnp.float32), + ) + inv_freq = jnp.exp( + -jnp.log(base_f) * band / jnp.asarray(max(half, 1), dtype=jnp.float32) + ) + angle = pos_f[..., None] * inv_freq + emb = jnp.concatenate([jnp.sin(angle), jnp.cos(angle)], axis=-1) + return emb[..., : int(width)].astype(dtype) + + +def _route_merge_clock( + level_idx, + pair_idx, + pair_base, + width: int, + max_depth, + dtype, + *, + root_centered: bool = False, +): + del pair_base + root_center = ( + _tree_clock_root_center_from_depth(max_depth, dtype) if root_centered else None + ) + return _tree_dyadic_segment_clock( + level_idx, + pair_idx, + width, + dtype, + root_center=root_center, + ) + + +class _RoutePointerBase(eqx.Module): + w_global: Float[Array, "d_global d_model"] + b_global: Float[Array, "d_model"] + + pref_msg_ln_scale: Float[Array, "two_d_edge"] + pref_msg_w1: Float[Array, "two_d_edge d_msg_hidden"] + pref_msg_b1: Float[Array, "d_msg_hidden"] + pref_msg_w2: Float[Array, "d_msg_hidden d_model"] + pref_msg_b2: Float[Array, "d_model"] + suff_msg_ln_scale: Float[Array, "two_d_edge"] + suff_msg_w1: Float[Array, "two_d_edge d_msg_hidden"] + suff_msg_b1: Float[Array, "d_msg_hidden"] + suff_msg_w2: Float[Array, "d_msg_hidden d_model"] + suff_msg_b2: Float[Array, "d_model"] + + virt_emb: Float[Array, "one d_model"] + + order_decay_w: Float[Array, "one d_model"] + order_decay_b: Float[Array, "one d_model"] + + virt_decay_w: Float[Array, "one d_model"] + virt_decay_b: Float[Array, "one d_model"] + cand_node_ln_scale: Float[Array, "d_in"] + cand_global_ln_scale: Float[Array, "d_model"] + + cand_graw_ln_scale: Float[Array, "d_graw"] + + cand_g_tap_w: Float[Array, "d_global d_graw"] + cand_pref_ln_scale: Float[Array, "d_model"] + cand_pref_order_ln_scale: Float[Array, "d_model"] + cand_suff_ln_scale: Float[Array, "d_model"] + cand_virt_pref_ln_scale: Float[Array, "d_model"] + cand_node_w: Float[Array, "d_in d_cand_hidden"] + cand_global_w: Float[Array, "d_model d_cand_hidden"] + cand_graw_w: Float[Array, "d_graw d_cand_hidden"] + cand_pref_w: Float[Array, "d_model d_cand_hidden"] + cand_pref_order_w: Float[Array, "d_model d_cand_hidden"] + cand_suff_w: Float[Array, "d_model d_cand_hidden"] + cand_virt_pref_w: Float[Array, "d_model d_cand_hidden"] + cand_virt_ratios_w: Float[Array, "three d_cand_hidden"] + cand_b_in: Float[Array, "d_cand_hidden"] + cand_block_ln_scale: Float[Array, "b d_cand_hidden"] + cand_block_w1: Float[Array, "b d_cand_hidden d_cand_hidden"] + cand_block_b1: Float[Array, "b d_cand_hidden"] + cand_block_w2: Float[Array, "b d_cand_hidden d_cand_hidden"] + cand_block_b2: Float[Array, "b d_cand_hidden"] + cand_out_ln_scale: Float[Array, "d_cand_hidden"] + cand_w_out: Float[Array, "d_cand_hidden d_model"] + cand_b_out: Float[Array, "d_model"] + + pointer_q_w: Float[Array, "d_model d_score"] + pointer_k_w: Float[Array, "d_model d_score"] + d_in: int = eqx.field(static=True) + d_global: int = eqx.field(static=True) + d_edge: int = eqx.field(static=True) + d_model: int = eqx.field(static=True) + d_attn: int = eqx.field(static=True) + pointer_score_dim: int = eqx.field(static=True) + n_heads: int = eqx.field(static=True) + n_heads_kernel: int = eqx.field(static=True) + d_head: int = eqx.field(static=True) + max_n: int = eqx.field(static=True) + ffn_hidden: int = eqx.field(static=True) + msg_hidden: int = eqx.field(static=True) + cand_hidden: int = eqx.field(static=True) + rope_base: float = eqx.field(static=True) + rope_scaling: float = eqx.field(static=True) + ln_eps: float = eqx.field(static=True) + + def __init__( + self, + *, + d_in: int, + d_edge: int, + d_global: int, + d_model: int, + n_heads: int, + max_n: int, + key: PRNGKeyArray, + rope_base: float = 10000.0, + rope_scaling: float = 1.0, + attention_dim: int, + pointer_score_dim: int, + candidate_hidden: int, + summary_hidden: int, + ffn_hidden: int, + global_tap_dim: int, + score_init_scale: float = 1.0, + ): + ln_eps = 1.0e-5 + if d_in != d_model: + raise ValueError("router requires d_in == d_model") + if d_global < 1: + raise ValueError("route pointer d_global must be >= 1") + d_attn = int(attention_dim) + if d_attn < 1: + raise ValueError("route pointer attention_dim must be positive") + if d_attn % n_heads != 0: + raise ValueError("route pointer attention_dim must be divisible by n_heads") + d_head = d_attn // n_heads + if d_head % 2 != 0: + raise ValueError("route pointer RoPE requires an even per-head dim") + score_dim = int(pointer_score_dim) + if score_dim < 1: + raise ValueError("route pointer pointer_score_dim must be positive") + if max_n < 1: + raise ValueError("route pointer max_n must be >= 1") + if rope_base <= 0.0 or rope_scaling <= 0.0: + raise ValueError("route pointer RoPE base/scaling must be positive") + n_heads_kernel = 2 * n_heads + ffn_hidden = int(ffn_hidden) + msg_hidden = int(summary_hidden) + if ffn_hidden < 1 or msg_hidden < 1: + raise ValueError("route pointer FFN/summary widths must be positive") + n_virt_ratios = 3 + + _graw_dim = int(global_tap_dim) + if _graw_dim < 1 or _graw_dim >= int(d_global): + raise ValueError( + "global_tap_dim must be positive and smaller than d_global" + ) + cand_in = d_in + 5 * d_model + n_virt_ratios + _graw_dim + cand_hidden = int(candidate_hidden) + if cand_hidden < 1: + raise ValueError("route pointer candidate_hidden must be positive") + keys = jax.random.split(key, 22) + + def w(k, shape, fan_in): + return jax.random.normal(k, shape) * (fan_in**-0.5) + + k_in, k_global = jax.random.split(keys[0], 2) + del k_in + self.w_global = w(k_global, (int(d_global), d_model), int(d_global)) + self.b_global = jnp.zeros((d_model,)) + + self.cand_graw_ln_scale = jnp.ones((_graw_dim,)) + self.cand_g_tap_w = w( + jax.random.fold_in(k_global, 0x67AB), + (int(d_global), _graw_dim), + int(d_global), + ) + + self.pref_msg_ln_scale = jnp.ones((2 * d_edge,)) + self.pref_msg_w1 = w(keys[7], (2 * d_edge, msg_hidden), 2 * d_edge) + self.pref_msg_b1 = jnp.zeros((msg_hidden,)) + self.pref_msg_w2 = w(keys[8], (msg_hidden, d_model), msg_hidden) + self.pref_msg_b2 = jnp.zeros((d_model,)) + self.suff_msg_ln_scale = jnp.ones((2 * d_edge,)) + self.suff_msg_w1 = w(keys[9], (2 * d_edge, msg_hidden), 2 * d_edge) + self.suff_msg_b1 = jnp.zeros((msg_hidden,)) + self.suff_msg_w2 = w(keys[10], (msg_hidden, d_model), msg_hidden) + self.suff_msg_b2 = jnp.zeros((d_model,)) + + vkey = jax.random.fold_in(key, 0x5710C) + self.virt_emb = jax.random.normal(vkey, (1, d_model)) * (d_model**-0.5) + from .readout_leaf_context import lca_order_init_w_b + + self.order_decay_w, self.order_decay_b = lca_order_init_w_b(d_model) + self.virt_decay_w, self.virt_decay_b = lca_order_init_w_b(d_model) + self.cand_node_ln_scale = jnp.ones((d_in,)) + self.cand_global_ln_scale = jnp.ones((d_model,)) + self.cand_pref_ln_scale = jnp.ones((d_model,)) + self.cand_pref_order_ln_scale = jnp.ones((d_model,)) + self.cand_suff_ln_scale = jnp.ones((d_model,)) + self.cand_virt_pref_ln_scale = jnp.ones((d_model,)) + + compose_key = keys[15] + self.cand_node_w = w( + jax.random.fold_in(compose_key, 0), + (d_in, cand_hidden), + cand_in, + ) + self.cand_global_w = w( + jax.random.fold_in(compose_key, 1), + (d_model, cand_hidden), + cand_in, + ) + self.cand_graw_w = w( + jax.random.fold_in(compose_key, 2), + (_graw_dim, cand_hidden), + cand_in, + ) + self.cand_pref_w = w( + jax.random.fold_in(compose_key, 3), + (d_model, cand_hidden), + cand_in, + ) + self.cand_pref_order_w = w( + jax.random.fold_in(compose_key, 4), + (d_model, cand_hidden), + cand_in, + ) + self.cand_suff_w = w( + jax.random.fold_in(compose_key, 5), + (d_model, cand_hidden), + cand_in, + ) + self.cand_virt_pref_w = w( + jax.random.fold_in(compose_key, 6), + (d_model, cand_hidden), + cand_in, + ) + self.cand_virt_ratios_w = w( + jax.random.fold_in(compose_key, 10), + (n_virt_ratios, cand_hidden), + cand_in, + ) + self.cand_b_in = jnp.zeros((cand_hidden,)) + self.cand_block_ln_scale = jnp.ones((1, cand_hidden)) + self.cand_block_w1 = w( + keys[16], + (1, cand_hidden, cand_hidden), + cand_hidden, + ) + self.cand_block_b1 = jnp.zeros((1, cand_hidden)) + self.cand_block_w2 = w( + keys[17], + (1, cand_hidden, cand_hidden), + cand_hidden, + ) + self.cand_block_b2 = jnp.zeros((1, cand_hidden)) + cand_out_key = keys[18] + self.cand_out_ln_scale = jnp.ones((cand_hidden,)) + self.cand_w_out = w(cand_out_key, (cand_hidden, d_model), cand_hidden) + self.cand_b_out = jnp.zeros((d_model,)) + + pointer_q_key = keys[19] + pointer_k_key = keys[20] + + self.pointer_q_w = w(pointer_q_key, (d_model, score_dim), d_model) * float( + score_init_scale + ) + self.pointer_k_w = w(pointer_k_key, (d_model, score_dim), d_model) + del keys + + self.d_in = d_in + self.d_global = int(d_global) + self.d_edge = d_edge + self.d_model = d_model + self.d_attn = d_attn + self.pointer_score_dim = score_dim + self.n_heads = n_heads + self.n_heads_kernel = n_heads_kernel + self.d_head = d_head + self.max_n = max_n + self.ffn_hidden = ffn_hidden + self.msg_hidden = msg_hidden + self.cand_hidden = cand_hidden + self.rope_base = float(rope_base) + self.rope_scaling = float(rope_scaling) + self.ln_eps = float(ln_eps) + + def _ln( + self, + scale, + x, + *, + tag_id: str, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ): + from hamiltonzero.model.tree import _tagged_rms_eqx_style + + return _tagged_rms_eqx_style( + scale, + x, + eps=self.ln_eps, + tag_id=tag_id, + pathway="even", + var_floor=1e-2, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + def _cross_ln( + self, + scale, + shift, + x, + *, + tag_id: str, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ): + from hamiltonzero.model.tree import _tagged_ln_eqx_style + + return _tagged_ln_eqx_style( + scale, + shift, + x, + eps=self.ln_eps, + tag_id=tag_id, + pathway="even", + var_floor=1e-2, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + def _dense( + self, + weight, + bias, + x, + *, + tag_id: str, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ): + from hamiltonzero.model.tree import _tagged_dense + + return _tagged_dense( + weight, + bias, + x, + tag_id=tag_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + def _dense_no_bias( + self, + weight, + x, + *, + tag_id: str, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, + ): + from hamiltonzero.model.tree import _tagged_dense_no_bias + + return _tagged_dense_no_bias( + weight, + x, + tag_id=tag_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + def _project_nodes(self, h: Float[Array, "n d_in"], structural_mask=None): + del structural_mask + return h + + def _prepare_nodes( + self, + h: Float[Array, "n d_in"], + mask: Int[Array, "n"] | Array, + ): + + projected, node_mean = self._center_nodes( + self._project_nodes(h, mask.astype(bool)), + mask, + ) + return (h, projected), node_mean + + def _project_global( + self, + global_feat: Float[Array, "d_global"], + dtype, + structural_mask=None, + ): + + raw = global_feat.astype(dtype) + g_dm = self._dense( + self.w_global, + self.b_global, + raw, + tag_id="route.global_input", + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=0, + ) + return (raw, g_dm) + + def _center_nodes( + self, + node_state: Float[Array, "n d_model"], + mask: Int[Array, "n"] | Array, + ): + dtype = node_state.dtype + active = mask.astype(dtype).reshape(node_state.shape[0], 1) + denom = jnp.maximum(jnp.sum(active), jnp.asarray(1.0, dtype=dtype)) + global_state = jnp.sum(node_state * active, axis=0) / denom + return node_state, global_state + + def _message_mlp(self, edge_pair, *, prefix: bool, structural_mask=None): + if prefix: + ln_s = self.pref_msg_ln_scale + w1, b1, w2, b2 = ( + self.pref_msg_w1, + self.pref_msg_b1, + self.pref_msg_w2, + self.pref_msg_b2, + ) + name = "pref" + else: + ln_s = self.suff_msg_ln_scale + w1, b1, w2, b2 = ( + self.suff_msg_w1, + self.suff_msg_b1, + self.suff_msg_w2, + self.suff_msg_b2, + ) + name = "suff" + structural_mask = ( + jnp.ones(edge_pair.shape[:-1], dtype=bool) + if structural_mask is None + else jnp.broadcast_to( + jnp.asarray(structural_mask, dtype=bool), edge_pair.shape[:-1] + ) + ) + kfac_kwargs = dict( + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=structural_mask.ndim, + ) + x = self._ln( + ln_s, + edge_pair, + tag_id=f"route.candidate.{name}_msg_ln", + **kfac_kwargs, + ) + x = self._dense( + w1, + b1, + x, + tag_id=f"route.candidate.{name}_msg1", + **kfac_kwargs, + ) + x = fused_silu(x) + return self._dense( + w2, + b2, + x, + tag_id=f"route.candidate.{name}_msg2", + **kfac_kwargs, + ) + + def _edge_pair_for_source( + self, + edge: Float[Array, "n n d_edge"], + source: Int[Array, ""], + ) -> Float[Array, "n two_d_edge"]: + return jnp.concatenate( + [edge[:, source, :], edge[source, :, :]], + axis=-1, + ) + + def _ordered_edge_messages( + self, + edge: Float[Array, "n n d_edge"], + perm: Int[Array, "n"], + mask: Int[Array, "n"] | Array, + ): + n = edge.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + edge_i_p = edge[idx[None, :], perm[:, None], :] + edge_p_i = edge[perm[:, None], idx[None, :], :] + edge_pair = jnp.concatenate([edge_i_p, edge_p_i], axis=-1) + mask_bool = mask.astype(bool) + pair_structural_mask = mask_bool[perm][:, None] & mask_bool[None, :] + return ( + self._message_mlp( + edge_pair, + prefix=True, + structural_mask=pair_structural_mask, + ), + self._message_mlp( + edge_pair, + prefix=False, + structural_mask=pair_structural_mask, + ), + ) + + def _clock_root_center_from_mask(self, mask): + n_active = jnp.maximum(jnp.sum(jnp.asarray(mask, dtype=jnp.int32)), 1) + depth = jnp.ceil(jnp.log2(n_active.astype(jnp.float32))).astype(jnp.int32) + return _tree_clock_root_center_from_depth(depth, jnp.float32) + + def _route_position_embedding(self, pos, dtype, *, mask=None): + pos_f = jnp.asarray(pos, dtype=jnp.float32) + if mask is not None: + pos_f = pos_f - self._clock_root_center_from_mask(mask) + pos_f = pos_f / jnp.asarray( + self.rope_scaling, + dtype=jnp.float32, + ) + half = (self.d_model + 1) // 2 + band = jnp.arange(half, dtype=jnp.float32) + inv_freq = jnp.exp( + -jnp.log(jnp.asarray(self.rope_base, dtype=jnp.float32)) + * band + / jnp.asarray(max(half, 1), dtype=jnp.float32) + ) + angle = pos_f[..., None] * inv_freq + emb = jnp.concatenate([jnp.sin(angle), jnp.cos(angle)], axis=-1) + return emb[..., : self.d_model].astype(dtype) + + def _first_active_index(self, mask): + return jnp.argmax(mask.astype(jnp.int32)).astype(jnp.int32) + + def _compose_candidates( + self, + node_state, + global_state: Float[Array, "d_global"], + prefix_summary: Float[Array, "... n d_model"], + prefix_order_summary: Float[Array, "... n d_model"], + suffix_summary: Float[Array, "... n d_model"], + route_pos, + virt_pref_order_summary: Float[Array, "... n d_model"], + virt_ratios: Float[Array, "... n 3"], + clock_mask=None, + candidate_mask=None, + ) -> Float[Array, "... n d_model"]: + node_input, node_projected = node_state + g_raw, g_dm = global_state + candidate_structural_mask = ( + jnp.ones(prefix_summary.shape[:-1], dtype=bool) + if candidate_mask is None + else jnp.broadcast_to( + jnp.asarray(candidate_mask, dtype=bool), + prefix_summary.shape[:-1], + ) + ) + kfac_kwargs = dict( + kfac_structural_mask=candidate_structural_mask, + kfac_repeat_ndim=candidate_structural_mask.ndim, + ) + from hamiltonzero.model.tree import _tagged_dense_no_bias + + g_raw = _tagged_dense_no_bias( + self.cand_g_tap_w, + g_raw, + tag_id="route.candidate.gtap", + pathway="even", + kfac_structural_mask=jnp.any(candidate_structural_mask), + kfac_repeat_ndim=0, + ) + if prefix_summary.ndim == node_projected.ndim: + nodes = node_projected + node_inputs = node_input + global_nodes = jnp.broadcast_to(g_dm[None, :], nodes.shape) + graw_nodes = jnp.broadcast_to( + g_raw[None, :], nodes.shape[:-1] + (g_raw.shape[-1],) + ) + else: + nodes = jnp.broadcast_to( + node_projected, + prefix_summary.shape[:-1] + (self.d_model,), + ) + node_inputs = jnp.broadcast_to( + node_input, + prefix_summary.shape[:-1] + (self.d_in,), + ) + global_nodes = jnp.broadcast_to( + g_dm, + prefix_summary.shape[:-1] + (self.d_model,), + ) + graw_nodes = jnp.broadcast_to( + g_raw, + prefix_summary.shape[:-1] + (g_raw.shape[-1],), + ) + pos_nodes = self._route_position_embedding( + route_pos, + prefix_summary.dtype, + mask=clock_mask, + ) + while pos_nodes.ndim < global_nodes.ndim: + pos_nodes = pos_nodes[..., None, :] + global_nodes = global_nodes + jnp.broadcast_to(pos_nodes, global_nodes.shape) + node_in = self._ln( + self.cand_node_ln_scale, + node_inputs, + tag_id="route.candidate.node_ln", + **kfac_kwargs, + ) + global_in = self._ln( + self.cand_global_ln_scale, + global_nodes, + tag_id="route.candidate.global_ln", + **kfac_kwargs, + ) + graw_in = self._ln( + self.cand_graw_ln_scale, + graw_nodes, + tag_id="route.candidate.graw_ln", + **kfac_kwargs, + ) + pref_in = self._ln( + self.cand_pref_ln_scale, + prefix_summary, + tag_id="route.candidate.pref_ln", + **kfac_kwargs, + ) + pref_order_in = self._ln( + self.cand_pref_order_ln_scale, + prefix_order_summary, + tag_id="route.candidate.pref_order_ln", + **kfac_kwargs, + ) + suff_in = self._ln( + self.cand_suff_ln_scale, + suffix_summary, + tag_id="route.candidate.suff_ln", + **kfac_kwargs, + ) + + _vp = jnp.broadcast_to(virt_pref_order_summary, prefix_summary.shape) + virt_pref_in = self._ln( + self.cand_virt_pref_ln_scale, + _vp, + tag_id="route.candidate.virt_pref_ln", + **kfac_kwargs, + ) + virt_ratios_in = jnp.broadcast_to( + virt_ratios, + prefix_summary.shape[:-1] + (3,), + ).astype(prefix_summary.dtype) + + x = self._dense( + self.cand_node_w, + self.cand_b_in, + node_in, + tag_id="route.candidate.compose.node", + **kfac_kwargs, + ) + x = x + self._dense_no_bias( + self.cand_global_w, + global_in, + tag_id="route.candidate.compose.global", + **kfac_kwargs, + ) + x = x + self._dense_no_bias( + self.cand_graw_w, + graw_in, + tag_id="route.candidate.compose.graw", + **kfac_kwargs, + ) + x = x + self._dense_no_bias( + self.cand_pref_w, + pref_in, + tag_id="route.candidate.compose.pref", + **kfac_kwargs, + ) + x = x + self._dense_no_bias( + self.cand_pref_order_w, + pref_order_in, + tag_id="route.candidate.compose.pref_order", + **kfac_kwargs, + ) + x = x + self._dense_no_bias( + self.cand_suff_w, + suff_in, + tag_id="route.candidate.compose.suff", + **kfac_kwargs, + ) + x = x + self._dense_no_bias( + self.cand_virt_pref_w, + virt_pref_in, + tag_id="route.candidate.compose.virt_pref", + **kfac_kwargs, + ) + x = x + self._dense_no_bias( + self.cand_virt_ratios_w, + virt_ratios_in, + tag_id="route.candidate.compose.virt_ratios", + **kfac_kwargs, + ) + + def block_body(x, params): + ln_s, w1, b1, w2, b2 = params + y = self._ln( + ln_s, + x, + tag_id="route.candidate.block.ln", + **kfac_kwargs, + ) + y = self._dense( + w1, + b1, + y, + tag_id="route.candidate.block.ffn1", + **kfac_kwargs, + ) + y = fused_silu(y) + y = self._dense( + w2, + b2, + y, + tag_id="route.candidate.block.ffn2", + **kfac_kwargs, + ) + return x + y, None + + x, _ = jax.lax.scan( + block_body, + x, + ( + self.cand_block_ln_scale, + self.cand_block_w1, + self.cand_block_b1, + self.cand_block_w2, + self.cand_block_b2, + ), + ) + x = self._ln( + self.cand_out_ln_scale, + x, + tag_id="route.candidate.compose_out_ln", + **kfac_kwargs, + ) + delta = self._dense( + self.cand_w_out, + self.cand_b_out, + x, + tag_id="route.candidate.compose_out", + **kfac_kwargs, + ) + return nodes + delta + + def _teacher_candidate_states( + self, + node_state, + global_state: Float[Array, "d_global"], + edge: Float[Array, "n n d_edge"], + perm: Int[Array, "n"], + mask: Int[Array, "n"] | Array, + real_mask: Int[Array, "n"] | Array, + ) -> Float[Array, "n n d_model"]: + _node_input, node_projected = node_state + n = node_projected.shape[0] + dtype = node_projected.dtype + idx = jnp.arange(n, dtype=jnp.int32) + mask_bool = mask.astype(bool) + + rm_bool = real_mask.astype(bool) + virt_slot = mask_bool & (~rm_bool) + + virt_at_pos = virt_slot[perm].astype(dtype) + pref_msg, suff_msg = self._ordered_edge_messages(edge, perm, mask) + row_active = mask_bool.astype(dtype).reshape(n, 1, 1) + pref_msg = pref_msg * row_active + suff_msg = suff_msg * row_active + + pref_before = jnp.cumsum(pref_msg, axis=0) - pref_msg + + from .readout_leaf_context import lca_gaussian_decay, register_vector_as_dense + + _odw = register_vector_as_dense( + self.order_decay_w, tag_id="route.order_decay_w" + )[0] + _odb = register_vector_as_dense( + self.order_decay_b, tag_id="route.order_decay_b" + )[0] + _tri = (idx[:, None] > idx[None, :]).astype(dtype) + _odecay = lca_gaussian_decay(idx, idx, _odw, _odb) + pref_order_before = jnp.einsum("ts,tsd,sid->tid", _tri, _odecay, pref_msg) + suff_before = jnp.cumsum(suff_msg, axis=0) - suff_msg + suff_including_self = jnp.sum(suff_msg, axis=0)[None, :, :] - suff_before + + pos_of_node = jnp.zeros((n,), dtype=jnp.int32).at[perm].set(idx) + self_msg = suff_msg[pos_of_node, idx, :] + remaining = mask_bool[None, :] & (pos_of_node[None, :] >= idx[:, None]) + candidate_structural_mask = mask_bool[:, None] & remaining + suff_other = suff_including_self - jnp.where( + remaining[:, :, None], + self_msg[None, :, :], + 0.0, + ) + + virt_emb = register_vector_as_dense( + self.virt_emb, + tag_id="route.virt_emb", + )[0] + virt_msg = virt_at_pos[:, None] * virt_emb[None, :] + + _vdw = register_vector_as_dense(self.virt_decay_w, tag_id="route.virt_decay_w")[ + 0 + ] + _vdb = register_vector_as_dense(self.virt_decay_b, tag_id="route.virt_decay_b")[ + 0 + ] + _vdecay = lca_gaussian_decay(idx, idx, _vdw, _vdb) + virt_pref_order_before = jnp.einsum( + "ts,tsd,sd->td", + _tri, + _vdecay, + virt_msg, + ) + virt_cnt_prefix = jnp.cumsum(virt_at_pos) - virt_at_pos + total_empty = jnp.sum(virt_at_pos) + total_leafs = jnp.sum(mask_bool.astype(dtype)) + virt_cnt_suffix = total_empty - virt_cnt_prefix + virt_norm = jnp.sqrt(jnp.maximum(virt_cnt_prefix, 1.0))[:, None] + virt_ratios = jnp.stack( + [ + virt_cnt_suffix / jnp.maximum(total_empty, 1.0), + virt_cnt_suffix / jnp.maximum(total_leafs, 1.0), + jnp.log((virt_cnt_prefix + 1.0) / (virt_cnt_suffix + 1.0)), + ], + axis=-1, + ) + virt_pref_order_summary = (virt_pref_order_before / virt_norm)[:, None, :] + virt_ratios_summary = virt_ratios[:, None, :] + + pref_den = jnp.sqrt(jnp.maximum(idx, 1).astype(dtype)).reshape(n, 1, 1) + + n_active = jnp.sum(mask_bool.astype(jnp.int32)) + suff_den = jnp.sqrt(jnp.maximum(n_active - idx - 1, 1).astype(dtype)).reshape( + n, 1, 1 + ) + return self._compose_candidates( + node_state, + global_state, + pref_before / pref_den, + pref_order_before / pref_den, + suff_other / suff_den, + idx, + virt_pref_order_summary=virt_pref_order_summary, + virt_ratios=virt_ratios_summary, + clock_mask=mask, + candidate_mask=candidate_structural_mask, + ) + + def _initial_summaries_and_edge_messages( + self, + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + dtype, + ): + + n = edge.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + pref_msg, suff_msg = self._ordered_edge_messages(edge, idx, mask) + source_active = mask.astype(dtype).reshape(n, 1, 1) + not_self = (idx[:, None] != idx[None, :]).astype(dtype).reshape(n, n, 1) + suffix_raw = jnp.sum(suff_msg * source_active * not_self, axis=0) + zeros = jnp.zeros_like(suffix_raw) + virt_prefix_order0 = jnp.zeros((self.d_model,), dtype=suffix_raw.dtype) + virt_count0 = jnp.zeros((), dtype=suffix_raw.dtype) + summaries = ( + zeros, + zeros, + suffix_raw, + virt_prefix_order0, + virt_count0, + ) + return summaries, (pref_msg, suff_msg) + + def _initial_summaries( + self, + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + dtype, + ): + + summaries, _edge_messages = self._initial_summaries_and_edge_messages( + edge, mask, dtype + ) + return summaries + + def _initial_summaries_streamed( + self, + edge: Float[Array, "n n d_edge"], + edge_transpose: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + dtype, + *, + pair_tile_size: int | None = None, + sequence_axis_name: str | None = None, + sequence_mesh=None, + ): + + n = edge.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + + def _seq_constraint(value, *axes): + if sequence_axis_name is None: + return value + from jax.sharding import NamedSharding, PartitionSpec as P + + spec = P(*axes) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return jax.lax.with_sharding_constraint(value, spec) + + tile = n if pair_tile_size is None else min(int(pair_tile_size), n) + if tile < 1: + raise ValueError("pair_tile_size must be positive") + n_tiles = (n + tile - 1) // tile + padded_n = n_tiles * tile + source_pad = padded_n - n + edge_padded = jnp.pad(edge, ((0, 0), (0, source_pad), (0, 0))) + edge_transpose_padded = jnp.pad( + edge_transpose, + ((0, 0), (0, source_pad), (0, 0)), + ) + source_mask = jnp.pad(mask.astype(bool), ((0, source_pad),)) + candidate_mask = mask.astype(bool)[:, None] + candidate_ids = idx[:, None] + suffix0 = _seq_constraint( + jnp.zeros((n, self.d_model), dtype=dtype), + sequence_axis_name, + None, + ) + + def add_source_tile(tile_index, suffix_sum): + start = tile_index * tile + edge_tile = jax.lax.dynamic_slice_in_dim( + edge_padded, + start, + tile, + axis=1, + ) + edge_transpose_tile = jax.lax.dynamic_slice_in_dim( + edge_transpose_padded, + start, + tile, + axis=1, + ) + edge_pair = jnp.concatenate( + [edge_tile, edge_transpose_tile], + axis=-1, + ) + source_mask_tile = jax.lax.dynamic_slice_in_dim( + source_mask, + start, + tile, + axis=0, + ) + source_ids = start + jnp.arange(tile, dtype=jnp.int32) + pair_mask = candidate_mask & source_mask_tile[None, :] + suff_by_candidate = self._message_mlp( + edge_pair, + prefix=False, + structural_mask=pair_mask, + ) + source_weight = source_mask_tile.astype(dtype)[None, :, None] + not_self = (candidate_ids != source_ids[None, :]).astype(dtype) + suffix_sum = suffix_sum + jnp.sum( + suff_by_candidate * source_weight * not_self[..., None], + axis=1, + ) + return _seq_constraint( + suffix_sum, + sequence_axis_name, + None, + ) + + suffix_raw = jax.lax.fori_loop(0, n_tiles, add_source_tile, suffix0) + zeros = jnp.zeros_like(suffix_raw) + return ( + zeros, + zeros, + suffix_raw, + jnp.zeros((self.d_model,), dtype=suffix_raw.dtype), + jnp.zeros((), dtype=suffix_raw.dtype), + ) + + def _candidate_states_from_summaries( + self, + node_state, + global_state: Float[Array, "d_global"], + prefix_raw: Float[Array, "n d_model"], + prefix_order_raw: Float[Array, "n d_model"], + suffix_raw: Float[Array, "n d_model"], + route_pos: Int[Array, ""] | Float[Array, ""], + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + prefix_ids: Int[Array, "n"], + virt_prefix_order_raw: Float[Array, "d_model"], + virt_count: Float[Array, ""], + real_mask: Int[Array, "n"] | Array, + ) -> Float[Array, "n d_model"]: + _node_input, node_projected = node_state + dtype = node_projected.dtype + route_pos_i = jnp.asarray(route_pos, dtype=jnp.int32) + mask_bool = mask.astype(bool) + pref_den = jnp.sqrt(jnp.maximum(route_pos_i, 1).astype(dtype)) + n_active = jnp.sum(mask_bool.astype(jnp.int32)) + suff_den = jnp.sqrt(jnp.maximum(n_active - route_pos_i - 1, 1).astype(dtype)) + _cand_idx = jnp.arange(node_projected.shape[0], dtype=jnp.int32) + _placed_pos = jnp.arange(node_projected.shape[0], dtype=jnp.int32) + _already_picked = jnp.any( + (_placed_pos < route_pos_i)[:, None] + & (prefix_ids[:, None] == _cand_idx[None, :]), + axis=0, + ) + candidate_structural_mask = mask_bool & ~_already_picked + + _vnorm = jnp.sqrt(jnp.maximum(virt_count, 1.0)) + _vp = virt_prefix_order_raw / _vnorm + virt_slot = mask_bool & (~real_mask.astype(bool)) + total_empty = jnp.sum(virt_slot.astype(dtype)) + total_leafs = jnp.sum(mask.astype(dtype)) + virt_cnt_suffix = total_empty - virt_count + _vr = jnp.stack( + [ + virt_cnt_suffix / jnp.maximum(total_empty, 1.0), + virt_cnt_suffix / jnp.maximum(total_leafs, 1.0), + jnp.log((virt_count + 1.0) / (virt_cnt_suffix + 1.0)), + ], + ) + return self._compose_candidates( + node_state, + global_state, + prefix_raw / pref_den, + prefix_order_raw / pref_den, + suffix_raw / suff_den, + route_pos_i, + virt_pref_order_summary=_vp, + virt_ratios=_vr, + clock_mask=mask, + candidate_mask=candidate_structural_mask, + ) + + def _pointer_raw(self, hidden, candidate_state, structural_mask=None): + candidate_structural_mask = ( + jnp.ones(candidate_state.shape[:-1], dtype=bool) + if structural_mask is None + else jnp.broadcast_to( + jnp.asarray(structural_mask, dtype=bool), + candidate_state.shape[:-1], + ) + ) + if hidden.ndim == 1: + query_structural_mask = jnp.any(candidate_structural_mask) + q = self._dense_no_bias( + self.pointer_q_w, + hidden, + tag_id="route.pointer.q", + kfac_structural_mask=query_structural_mask, + kfac_repeat_ndim=0, + ) + k = self._dense_no_bias( + self.pointer_k_w, + candidate_state, + tag_id="route.pointer.k", + kfac_structural_mask=candidate_structural_mask, + kfac_repeat_ndim=1, + ) + raw = jnp.einsum("d,nd->n", q, k) + else: + query_structural_mask = jnp.any(candidate_structural_mask, axis=-1) + q = self._dense_no_bias( + self.pointer_q_w, + hidden, + tag_id="route.pointer.q", + kfac_structural_mask=query_structural_mask, + kfac_repeat_ndim=1, + ) + k = self._dense_no_bias( + self.pointer_k_w, + candidate_state, + tag_id="route.pointer.k", + kfac_structural_mask=candidate_structural_mask, + kfac_repeat_ndim=2, + ) + raw = jnp.einsum("td,tnd->tn", q, k) + scale = jax.lax.rsqrt(jnp.asarray(self.pointer_score_dim, dtype=raw.dtype)) + return raw * scale + + def _pointer_logits(self, hidden, candidate_state, picked, mask, tau): + active = mask.astype(bool) & (~picked) + raw = self._pointer_raw(hidden, candidate_state, structural_mask=active) + neg = jnp.asarray(-1.0e30, dtype=raw.dtype) + return jnp.where(active, raw / jnp.asarray(tau, dtype=raw.dtype), neg) + + def _learned_first_choice_mask(self, mask, real_mask): + mask_bool = mask.astype(bool) + if real_mask is None: + return mask_bool + real_bool = real_mask.astype(bool) & mask_bool + return jnp.where(jnp.any(real_bool), real_bool, mask_bool) + + def _step_choice_mask(self, first_step, mask, real_mask): + mask_bool = mask.astype(bool) + first_mask = self._learned_first_choice_mask(mask, real_mask) + return jnp.where(first_step, first_mask, mask_bool) + + def _step_pointer_hidden(self, first_step, global_state, hidden): + first_hidden = jnp.broadcast_to(global_state[1], hidden.shape) + return jnp.where(first_step, first_hidden, hidden) + + def _teacher_hidden_with_first(self, hidden, global_state, first_active_idx): + return hidden.at[first_active_idx].set(global_state[1]) + + def _score_step_for_logp(self, first_step, predict_step): + return predict_step | first_step + + def _logprob_contribute_mask(self, mask_bool, idx, first_active): + del idx, first_active + return mask_bool + + def _collapse_quotient_logits(self, logits, ids, valid): + n = logits.shape[-1] + dtype = logits.dtype + idx = jnp.arange(n, dtype=jnp.int32) + ids = jnp.asarray(ids, dtype=jnp.int32) + valid = valid.astype(bool) & (ids >= 0) + same = (ids[:, None] == ids[None, :]) & valid[:, None] & valid[None, :] + rep_idx = jnp.min( + jnp.where(same, idx[None, :], jnp.asarray(n, dtype=jnp.int32)), + axis=1, + ) + reps = valid & (idx == rep_idx) + neg = jnp.asarray(-1.0e30, dtype=dtype) + member_logits = jnp.where(same, logits[None, :], neg) + max_l = jnp.max(member_logits, axis=1) + max_l = jnp.where(jnp.isfinite(max_l), max_l, jnp.asarray(0.0, dtype=dtype)) + class_lse = max_l + jnp.log( + jnp.sum(jnp.exp(member_logits - max_l[:, None]), axis=1) + ) + class_size = jnp.maximum( + jnp.sum(same.astype(dtype), axis=1), + jnp.asarray(1.0, dtype=dtype), + ) + quotient_logits = class_lse - jnp.log(class_size) + return jnp.where(reps, quotient_logits, neg) + + def _apply_quotient_logits( + self, + logits, + first_orbit_ids, + valid_mask, + context_mask, + prefix_ids, + prefix_len, + ): + if len(first_orbit_ids) == 2: + node_key, edge_key = first_orbit_ids + ids = conditional_orbit_ids_from_keys( + node_key, + edge_key, + valid_mask, + context_mask, + prefix_ids, + prefix_len, + ) + elif len(first_orbit_ids) == 3: + node_key, edge_key, needs_fwl2 = first_orbit_ids + inputs = ( + node_key, + edge_key, + valid_mask, + context_mask, + prefix_ids, + prefix_len, + ) + ids = jax.lax.cond( + jnp.asarray(needs_fwl2, dtype=jnp.bool_), + lambda values: conditional_orbit_pair_ids_from_keys(*values), + lambda values: conditional_orbit_ids_from_keys(*values), + inputs, + ) + else: + raise ValueError( + "quotient carrier must contain node key, edge key, and " + "optionally needs_fwl2" + ) + return self._collapse_quotient_logits(logits, ids, valid_mask) + + def _stopgrad_logit_scale(self, logits, valid_mask): + del valid_mask + return logits + + def _append_token( + self, + token: Float[Array, "d_model"], + chosen: Int[Array, ""], + t: Int[Array, ""], + prefix_ids: Int[Array, "n"], + k_cache: Float[Array, "l n h d_head"], + v_cache: Float[Array, "l n h d_head"], + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + ): + del chosen, t, prefix_ids, edge, mask + return token, k_cache, v_cache + + def logprob_perm( + self, + h: Float[Array, "n d_in"], + edge: Float[Array, "n n d_edge"], + perm: Int[Array, "n"], + mask: Int[Array, "n"] | Array, + *, + global_feat: Float[Array, "d_global"] | None = None, + tau: float | Float[Array, ""] = 1.0, + real_mask: Int[Array, "n"] | Array | None = None, + first_orbit_ids: QuotientCarrier, + ) -> Float[Array, ""]: + scores = self._teacher_logits( + h, + edge, + perm, + mask, + global_feat=global_feat, + tau=tau, + real_mask=real_mask, + first_orbit_ids=first_orbit_ids, + ) + n = h.shape[0] + if n <= 1: + return jnp.asarray(0.0, dtype=h.dtype) + mask_bool = mask.astype(bool) + first_active = self._first_active_index(mask) + idx = jnp.arange(n, dtype=jnp.int32) + contribute = self._logprob_contribute_mask(mask_bool, idx, first_active) + neg = jnp.asarray(-1.0e30, dtype=scores.dtype) + scores = self._stopgrad_logit_scale( + scores, + scores > (neg * jnp.asarray(0.5, dtype=scores.dtype)), + ) + log_probs = jax.nn.log_softmax(scores.astype(jnp.float32), axis=-1) + chosen = jnp.take_along_axis(log_probs, perm[:, None], axis=-1)[:, 0] + return jnp.sum(jnp.where(contribute, chosen, 0.0)) + + def logprob_identity( + self, + h: Float[Array, "n d_in"], + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + *, + global_feat: Float[Array, "d_global"] | None = None, + tau: float | Float[Array, ""] = 1.0, + real_mask: Int[Array, "n"] | Array | None = None, + first_orbit_ids: QuotientCarrier, + ) -> Float[Array, ""]: + n = h.shape[0] + return self.logprob_perm( + h, + edge, + jnp.arange(n, dtype=jnp.int32), + mask, + global_feat=global_feat, + tau=tau, + real_mask=real_mask, + first_orbit_ids=first_orbit_ids, + ) + + +class _TreePrefixMerge(eqx.Module): + ln_scale: Float[Array, "d_in"] + + g_proj_w: Float[Array, "d_gstream d_gsec"] + w1: Float[Array, "d_in d_hidden"] + b1: Float[Array, "d_hidden"] + w2: Float[Array, "d_hidden d_model"] + b2: Float[Array, "d_model"] + + alpha_route: Float[Array, "d_model"] + + d_model: int = eqx.field(static=True) + d_hidden: int = eqx.field(static=True) + d_in: int = eqx.field(static=True) + max_depth: int = eqx.field(static=True) + ngpt_alpha_max: float = eqx.field(static=True) + ln_eps: float = eqx.field(static=True) + + def __init__( + self, + d_model: int, + *, + hidden: int, + max_depth: int, + key: PRNGKeyArray, + gladder_d_g: int, + alpha_init: float, + alpha_max: float, + ln_eps: float = 1e-5, + ): + d_hidden = int(hidden) + + d_in = 5 * int(d_model) + 64 + _TREE_NGPT_DEPTH_FEAT_DIM + k1, k2 = jax.random.split(key, 2) + self.ln_scale = jnp.ones((d_in,)) + self.w1 = jax.random.normal(k1, (d_in, d_hidden)) * (d_in**-0.5) + self.b1 = jnp.zeros((d_hidden,)) + self.w2 = jax.random.normal(k2, (d_hidden, d_model)) * (d_hidden**-0.5) + self.b2 = jnp.zeros((d_model,)) + kg = jax.random.fold_in(k2, 0x61B5) + self.g_proj_w = jax.random.normal(kg, (int(gladder_d_g), 64)) * ( + int(gladder_d_g) ** -0.5 + ) + self.alpha_route = float(alpha_init) * jnp.ones((int(d_model),)) + self.d_model = int(d_model) + self.d_hidden = int(d_hidden) + self.d_in = int(d_in) + self.max_depth = int(max_depth) + self.ngpt_alpha_max = float(alpha_max) + self.ln_eps = float(ln_eps) + + def project_global(self, g, structural_mask): + + from hamiltonzero.model.tree import _tagged_dense_no_bias + + return _tagged_dense_no_bias( + self.g_proj_w, + g, + tag_id="gladder.route.merge_gproj", + pathway="even", + kfac_structural_mask=jnp.asarray(structural_mask, dtype=bool), + kfac_scan_shared=False, + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=True, + ) + + def __call__( + self, + left, + right, + left_mask, + right_mask, + sibling_edge_lr, + sibling_edge_rl, + level_idx, + pair_idx=None, + pair_base=None, + clock_depth=None, + depth_feats=None, + g=None, + g_structural_mask=None, + g_projected=None, + ): + from hamiltonzero.model.tree import _tagged_dense, _tagged_rms_eqx_style + from hamiltonzero.model.tree import _rownorm_cols + + _we = _rownorm_cols + + dtype = left.dtype + left_mask = left_mask.astype(dtype) + right_mask = right_mask.astype(dtype) + out_mask = left_mask + right_mask - left_mask * right_mask + both = left_mask * right_mask + merge_structural_mask = both.astype(bool) + depth_i = ( + jnp.asarray(max(self.max_depth, 1), dtype=jnp.int32) + if clock_depth is None + else jnp.maximum(jnp.asarray(clock_depth, dtype=jnp.int32), 1) + ) + parts = [left, right, sibling_edge_lr, sibling_edge_rl] + g_active = ( + jnp.any(merge_structural_mask) + if g_structural_mask is None + else jnp.asarray(g_structural_mask, dtype=bool) + ) + gg = ( + g_projected if g_projected is not None else self.project_global(g, g_active) + ) + parts.append( + jnp.broadcast_to(gg[None, :], left.shape[:-1] + (gg.shape[-1],)).astype( + dtype + ) + ) + if depth_feats is None: + raise ValueError("tree prefix merge requires depth features") + parts.append(depth_feats.astype(dtype)) + clock = _route_merge_clock( + level_idx, + pair_idx, + pair_base, + self.d_model, + depth_i, + dtype, + root_centered=True, + ) + parts.append(jnp.broadcast_to(clock, left.shape)) + x = jnp.concatenate(parts, axis=-1) + x = _tagged_rms_eqx_style( + self.ln_scale, + x, + eps=self.ln_eps, + tag_id="route.tree_prefix.merge.ln", + pathway="even", + kfac_structural_mask=merge_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + h = _tagged_dense( + self.w1, + self.b1, + x, + tag_id="route.tree_prefix.merge.ffn1", + pathway="even", + weight_eff=_we(self.w1), + kfac_structural_mask=merge_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + h = fused_silu(h) + delta = _tagged_dense( + self.w2, + self.b2, + h, + tag_id="route.tree_prefix.merge.ffn2", + pathway="even", + weight_eff=_we(self.w2), + kfac_structural_mask=merge_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + raw = _tree_ngpt_residual( + 0.5 * (left + right), + delta, + self.alpha_route, + max_gain=self.ngpt_alpha_max, + tag_id="route.tree_prefix.merge.alpha_route", + pathway="even", + kfac_structural_mask=merge_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + carry = jnp.where(left_mask[:, None] > 0, left, right) + out = jnp.where(both[:, None] > 0, raw, carry) + out = jnp.where( + out_mask[:, None] > 0, + out, + jnp.zeros_like(out), + ) + return out, out_mask, both + + +class _TreePrefixSelfLayer(eqx.Module): + ln_scale: Float[Array, "d_model"] + w_qkv: Float[Array, "d_model three_qv"] + w_o: Float[Array, "d_o_in d_model"] + edge_ln_scale: Float[Array, "d_model"] + edge_w1: Float[Array, "d_model d_hidden"] + edge_b1: Float[Array, "d_hidden"] + edge_w2: Float[Array, "d_hidden h_kernel"] + edge_b2: Float[Array, "h_kernel"] + ffn_ln_scale: Float[Array, "d_model"] + ffn_w1: Float[Array, "d_model d_ffn"] + ffn_b1: Float[Array, "d_ffn"] + ffn_w2: Float[Array, "d_ffn d_model"] + ffn_b2: Float[Array, "d_model"] + + def as_tuple(self): + return ( + self.ln_scale, + self.w_qkv, + self.w_o, + self.edge_ln_scale, + self.edge_w1, + self.edge_b1, + self.edge_w2, + self.edge_b2, + self.ffn_ln_scale, + self.ffn_w1, + self.ffn_b1, + self.ffn_w2, + self.ffn_b2, + ) + + +class _TreePrefixCandidateLayer(eqx.Module): + cand_ln_scale: Float[Array, "d_model"] + prefix_ln_scale: Float[Array, "d_model"] + cand_w_qv: Float[Array, "d_model two_qv"] + prefix_w_kv: Float[Array, "d_model two_qv"] + w_o: Float[Array, "d_o_in d_model"] + edge_ln_scale: Float[Array, "d_model"] + edge_w1: Float[Array, "d_model d_hidden"] + edge_b1: Float[Array, "d_hidden"] + edge_w2: Float[Array, "d_hidden h_kernel"] + edge_b2: Float[Array, "h_kernel"] + ffn_ln_scale: Float[Array, "d_model"] + ffn_w1: Float[Array, "d_model d_ffn"] + ffn_b1: Float[Array, "d_ffn"] + ffn_w2: Float[Array, "d_ffn d_model"] + ffn_b2: Float[Array, "d_model"] + + def as_tuple(self): + return ( + self.cand_ln_scale, + self.prefix_ln_scale, + self.cand_w_qv, + self.prefix_w_kv, + self.w_o, + self.edge_ln_scale, + self.edge_w1, + self.edge_b1, + self.edge_w2, + self.edge_b2, + self.ffn_ln_scale, + self.ffn_w1, + self.ffn_b1, + self.ffn_w2, + self.ffn_b2, + ) + + +class _HeavyRouteLayer(eqx.Module): + cross_ln_scale: Float[Array, "d_model"] + cross_ln_shift: Float[Array, "d_model"] + cross_prefix_ln_scale: Float[Array, "d_model"] + cross_prefix_ln_shift: Float[Array, "d_model"] + cross_w_qv: Float[Array, "d_model two_qv"] + cross_w_kv: Float[Array, "d_model two_qv"] + cross_w_o: Float[Array, "d_o_in d_model"] + cross_edge_ln_scale: Float[Array, "two_d_edge"] + cross_edge_ln_shift: Float[Array, "two_d_edge"] + cross_edge_w1: Float[Array, "two_d_edge d_heavy_edge_hidden"] + cross_edge_b1: Float[Array, "d_heavy_edge_hidden"] + cross_edge_w2: Float[Array, "d_heavy_edge_hidden h_kernel"] + cross_edge_b2: Float[Array, "h_kernel"] + + self_ln_scale: Float[Array, "d_model"] + self_w_qkv: Float[Array, "d_model three_qv"] + self_w_o: Float[Array, "d_o_in d_model"] + self_edge_ln_scale: Float[Array, "two_d_edge"] + self_edge_w1: Float[Array, "two_d_edge d_heavy_edge_hidden"] + self_edge_b1: Float[Array, "d_heavy_edge_hidden"] + self_edge_w2: Float[Array, "d_heavy_edge_hidden h_kernel"] + self_edge_b2: Float[Array, "h_kernel"] + + ffn_ln_scale: Float[Array, "d_model"] + ffn_w1: Float[Array, "d_model d_ffn"] + ffn_b1: Float[Array, "d_ffn"] + ffn_w2: Float[Array, "d_ffn d_model"] + ffn_b2: Float[Array, "d_model"] + + def as_tuple(self): + return ( + self.cross_ln_scale, + self.cross_ln_shift, + self.cross_prefix_ln_scale, + self.cross_prefix_ln_shift, + self.cross_w_qv, + self.cross_w_kv, + self.cross_w_o, + self.cross_edge_ln_scale, + self.cross_edge_ln_shift, + self.cross_edge_w1, + self.cross_edge_b1, + self.cross_edge_w2, + self.cross_edge_b2, + self.self_ln_scale, + self.self_w_qkv, + self.self_w_o, + self.self_edge_ln_scale, + self.self_edge_w1, + self.self_edge_b1, + self.self_edge_w2, + self.self_edge_b2, + self.ffn_ln_scale, + self.ffn_w1, + self.ffn_b1, + self.ffn_w2, + self.ffn_b2, + ) + + +class _PrefixSuffixRouteBase(_RoutePointerBase): + heavy_layers: list[_HeavyRouteLayer] + route_prefix_suffix_layers: int = eqx.field(static=True) + route_decoder_attn_impl: str = eqx.field(static=True) + heavy_edge_hidden: int = eqx.field(static=True) + heavy_residual_gain: float = eqx.field(static=True) + + def __init__( + self, + *, + d_in: int, + d_edge: int, + d_global: int, + d_model: int, + n_heads: int, + max_n: int, + key: PRNGKeyArray, + route_prefix_suffix_layers: int = 1, + route_decoder_attn_impl: str = "mhsea_tuned", + score_init_scale: float = 1.0, + rope_base: float = 10000.0, + rope_scaling: float = 1.0, + attention_dim: int, + pointer_score_dim: int, + candidate_hidden: int, + summary_hidden: int, + ffn_hidden: int, + global_tap_dim: int, + ): + if route_prefix_suffix_layers < 0: + raise ValueError("route_prefix_suffix_layers must be >= 0") + allowed_impls = {"mhsea_tuned", "einsum"} + if route_decoder_attn_impl not in allowed_impls: + raise ValueError( + f"route_decoder_attn_impl must be one of {sorted(allowed_impls)}" + ) + + key_base, key_heavy = jax.random.split(key) + super().__init__( + d_in=d_in, + d_edge=d_edge, + d_global=d_global, + d_model=d_model, + n_heads=n_heads, + max_n=max_n, + key=key_base, + score_init_scale=score_init_scale, + rope_base=rope_base, + rope_scaling=rope_scaling, + attention_dim=attention_dim, + pointer_score_dim=pointer_score_dim, + candidate_hidden=candidate_hidden, + summary_hidden=summary_hidden, + ffn_hidden=ffn_hidden, + global_tap_dim=global_tap_dim, + ) + + layers = int(route_prefix_suffix_layers) + d_qv = self.n_heads_kernel * self.d_head + d_o_in = self.n_heads * self.d_head + pair_dim = 2 * self.d_edge + heavy_edge_hidden = max(32, 2 * self.n_heads_kernel, 4 * self.d_edge) + + def w(k, shape, fan_in): + return jax.random.normal(k, shape) * (fan_in**-0.5) + + layer_keys = jax.random.split(key_heavy, layers) + heavy_layers = [] + for li in range(layers): + keys = jax.random.split(layer_keys[li], 11) + heavy_layers.append( + _HeavyRouteLayer( + cross_ln_scale=jnp.ones((self.d_model,)), + cross_ln_shift=jnp.zeros((self.d_model,)), + cross_prefix_ln_scale=jnp.ones((self.d_model,)), + cross_prefix_ln_shift=jnp.zeros((self.d_model,)), + cross_w_qv=w(keys[0], (self.d_model, 2 * d_qv), self.d_model), + cross_w_kv=w(keys[1], (self.d_model, 2 * d_qv), self.d_model), + cross_w_o=w(keys[2], (d_o_in, self.d_model), d_o_in), + cross_edge_ln_scale=jnp.ones((pair_dim,)), + cross_edge_ln_shift=jnp.zeros((pair_dim,)), + cross_edge_w1=w(keys[3], (pair_dim, heavy_edge_hidden), pair_dim), + cross_edge_b1=jnp.zeros((heavy_edge_hidden,)), + cross_edge_w2=w( + keys[4], + (heavy_edge_hidden, self.n_heads_kernel), + heavy_edge_hidden, + ), + cross_edge_b2=jnp.zeros((self.n_heads_kernel,)), + self_ln_scale=jnp.ones((self.d_model,)), + self_w_qkv=w(keys[5], (self.d_model, 3 * d_qv), self.d_model), + self_w_o=w(keys[6], (d_o_in, self.d_model), d_o_in), + self_edge_ln_scale=jnp.ones((pair_dim,)), + self_edge_w1=w(keys[7], (pair_dim, heavy_edge_hidden), pair_dim), + self_edge_b1=jnp.zeros((heavy_edge_hidden,)), + self_edge_w2=w( + keys[8], + (heavy_edge_hidden, self.n_heads_kernel), + heavy_edge_hidden, + ), + self_edge_b2=jnp.zeros((self.n_heads_kernel,)), + ffn_ln_scale=jnp.ones((self.d_model,)), + ffn_w1=w(keys[9], (self.d_model, self.ffn_hidden), self.d_model), + ffn_b1=jnp.zeros((self.ffn_hidden,)), + ffn_w2=w( + keys[10], (self.ffn_hidden, self.d_model), self.ffn_hidden + ), + ffn_b2=jnp.zeros((self.d_model,)), + ) + ) + self.heavy_layers = heavy_layers + + self.route_prefix_suffix_layers = layers + self.route_decoder_attn_impl = str(route_decoder_attn_impl) + self.heavy_edge_hidden = heavy_edge_hidden + self.heavy_residual_gain = 0.0 if layers == 0 else float(layers) ** -0.5 + + def _heavy_layer_params(self): + return [layer.as_tuple() for layer in self.heavy_layers] + + def _resolve_heavy_attn_impl(self, n: int) -> str: + del n + return self.route_decoder_attn_impl + + def _heavy_edge_bias( + self, + edge_pair, + params, + *, + prefix: str, + structural_mask=None, + scan_shared: bool = False, + repeat_ndim: int | None = None, + context_primal_reused_over_walkers: bool = False, + ): + ln_s, w1, b1, w2, b2 = params + structural_mask = ( + jnp.ones(edge_pair.shape[:-1], dtype=bool) + if structural_mask is None + else jnp.broadcast_to( + jnp.asarray(structural_mask, dtype=bool), edge_pair.shape[:-1] + ) + ) + repeat_ndim = structural_mask.ndim if repeat_ndim is None else repeat_ndim + kfac_kwargs = dict( + kfac_structural_mask=structural_mask, + kfac_scan_shared=scan_shared, + kfac_repeat_ndim=repeat_ndim, + kfac_context_primal_reused_over_walkers=( + context_primal_reused_over_walkers + ), + ) + x = self._ln( + ln_s, + edge_pair, + tag_id=f"route.heavy.{prefix}.edge_ln", + **kfac_kwargs, + ) + x = self._dense( + w1, + b1, + x, + tag_id=f"route.heavy.{prefix}.edge_bias1", + **kfac_kwargs, + ) + x = fused_silu(x) + bias = self._dense( + w2, + b2, + x, + tag_id=f"route.heavy.{prefix}.edge_bias2", + **kfac_kwargs, + ) + return bias + + def _heavy_cross_edge_bias( + self, + edge_pair, + params, + *, + structural_mask=None, + scan_shared: bool = False, + repeat_ndim: int | None = None, + context_primal_reused_over_walkers: bool = False, + ): + ln_s, ln_b, w1, b1, w2, b2 = params + structural_mask = ( + jnp.ones(edge_pair.shape[:-1], dtype=bool) + if structural_mask is None + else jnp.broadcast_to( + jnp.asarray(structural_mask, dtype=bool), edge_pair.shape[:-1] + ) + ) + repeat_ndim = structural_mask.ndim if repeat_ndim is None else repeat_ndim + kfac_kwargs = dict( + kfac_structural_mask=structural_mask, + kfac_scan_shared=scan_shared, + kfac_repeat_ndim=repeat_ndim, + kfac_context_primal_reused_over_walkers=( + context_primal_reused_over_walkers + ), + ) + x = self._cross_ln( + ln_s, + ln_b, + edge_pair, + tag_id="route.heavy.cross.edge_ln", + **kfac_kwargs, + ) + x = self._dense( + w1, + b1, + x, + tag_id="route.heavy.cross.edge_bias1", + **kfac_kwargs, + ) + x = fused_silu(x) + return self._dense( + w2, + b2, + x, + tag_id="route.heavy.cross.edge_bias2", + **kfac_kwargs, + ) + + def _route_attention( + self, + q: Float[Array, "b n h d_head"], + k: Float[Array, "b n h d_head"], + v: Float[Array, "b n h d_head"], + edge_bias: Float[Array, "b n n h"], + key_mask: Int[Array, "b n"] | Array, + *, + impl: str, + key_mask_only: bool = False, + attention_mask: Array | None = None, + sequence_axis_name=None, + sequence_mesh=None, + ) -> Float[Array, "b n h d_head"]: + dtype = q.dtype + valid = key_mask.astype(bool)[:, None, :] + if attention_mask is not None: + valid = valid & attention_mask.astype(bool) + has_key = jnp.any(valid, axis=-1) + if impl == "einsum": + q_c = q + k_c = k + v_c = v + bias_c = edge_bias + if sequence_axis_name is not None: + from jax.sharding import NamedSharding, PartitionSpec as P + + def _sharding(*axes): + spec = P(*axes) + return ( + NamedSharding(sequence_mesh, spec) + if sequence_mesh is not None + else spec + ) + + q_c = jax.lax.with_sharding_constraint( + q_c, + _sharding(None, sequence_axis_name, None, None), + ) + k_c = jax.lax.with_sharding_constraint( + k_c, + _sharding(None, None, None, None), + ) + v_c = jax.lax.with_sharding_constraint( + v_c, + _sharding(None, None, None, None), + ) + bias_c = jax.lax.with_sharding_constraint( + bias_c, + _sharding(None, sequence_axis_name, None, None), + ) + logits = jnp.einsum("bihd,bjhd->bhij", q_c, k_c) + logits = logits / jnp.sqrt(jnp.asarray(self.d_head, dtype=dtype)) + logits = logits + jnp.transpose(bias_c, (0, 3, 1, 2)) + if sequence_axis_name is not None: + logits = jax.lax.with_sharding_constraint( + logits, + _sharding(None, None, sequence_axis_name, None), + ) + logits = jnp.where( + valid[:, None, :, :], + logits, + jnp.asarray(-1.0e30, dtype=dtype), + ) + if sequence_axis_name is not None: + logits = jax.lax.with_sharding_constraint( + logits, + _sharding(None, None, sequence_axis_name, None), + ) + alpha = jax.nn.softmax(logits, axis=-1) + if sequence_axis_name is not None: + alpha = jax.lax.with_sharding_constraint( + alpha, + _sharding(None, None, sequence_axis_name, None), + ) + out = jnp.einsum("bhij,bjhd->bihd", alpha, v_c) + if sequence_axis_name is not None: + out = jax.lax.with_sharding_constraint( + out, + _sharding(None, sequence_axis_name, None, None), + ) + elif impl == "mhsea_tuned": + from hamiltonzero.model.pallas_attention import mhsea_tuned_edge_attention + + if key_mask_only or attention_mask is not None: + edge_bias = jnp.where( + valid[..., None], + edge_bias, + jnp.asarray(-1.0e30, dtype=edge_bias.dtype), + ) + key_mask = jnp.ones_like(key_mask) + d_head_padded = max(16, self.d_head) + pad_amount = d_head_padded - self.d_head + if pad_amount: + scale = jnp.sqrt(jnp.asarray(d_head_padded / self.d_head, dtype=dtype)) + q = jnp.concatenate( + [ + q * scale, + jnp.zeros(q.shape[:-1] + (pad_amount,), dtype=q.dtype), + ], + axis=-1, + ) + k = jnp.concatenate( + [ + k, + jnp.zeros(k.shape[:-1] + (pad_amount,), dtype=k.dtype), + ], + axis=-1, + ) + v = jnp.concatenate( + [ + v, + jnp.zeros(v.shape[:-1] + (pad_amount,), dtype=v.dtype), + ], + axis=-1, + ) + out = jax.vmap( + lambda q_b, k_b, v_b, bias_b, mask_b: mhsea_tuned_edge_attention( + q_b, k_b, v_b, bias_b, mask_b.astype(jnp.int32) + ) + )(q, k, v, edge_bias, key_mask) + out = out[..., : self.d_head] + else: + raise ValueError("route attention must be 'einsum' or 'mhsea_tuned'") + return jnp.where(has_key[..., None, None], out, jnp.zeros_like(out)) + + def _collapse_heavy_heads(self, out): + gate_heads = out[..., : self.n_heads, :] + value_heads = out[..., self.n_heads :, :] + out = jax.nn.sigmoid(gate_heads) * value_heads + return out.reshape(out.shape[:-2] + (self.n_heads * self.d_head,)) + + def _heavy_prefix_pairs_teacher(self, edge, perm): + n = edge.shape[0] + edge_i_p = edge[:, perm, :] + edge_p_i = jnp.transpose(edge[perm, :, :], (1, 0, 2)) + pair = jnp.concatenate([edge_i_p, edge_p_i], axis=-1) + + return pair + + def _heavy_suffix_pairs_teacher(self, edge): + n = edge.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + edge_i_j = edge[idx[:, None], idx[None, :], :] + edge_j_i = edge[idx[None, :], idx[:, None], :] + pair = jnp.concatenate([edge_i_j, edge_j_i], axis=-1) + + return pair + + def _heavy_layer_teacher( + self, + cand: Float[Array, "n n d_model"], + z: Float[Array, "n d_model"], + edge: Float[Array, "n n d_edge"], + perm: Int[Array, "n"], + mask: Int[Array, "n"] | Array, + pos_of_node: Int[Array, "n"], + params, + *, + impl: str, + ) -> Float[Array, "n n d_model"]: + ( + cross_ln_s, + cross_ln_b, + cross_prefix_ln_s, + cross_prefix_ln_b, + cross_w_qv, + cross_w_kv, + cross_w_o, + cross_edge_ln_s, + cross_edge_ln_b, + cross_edge_w1, + cross_edge_b1, + cross_edge_w2, + cross_edge_b2, + self_ln_s, + self_w_qkv, + self_w_o, + self_edge_ln_s, + self_edge_w1, + self_edge_b1, + self_edge_w2, + self_edge_b2, + ffn_ln_s, + ffn_w1, + ffn_b1, + ffn_w2, + ffn_b2, + ) = params + n = cand.shape[0] + dtype = cand.dtype + idx = jnp.arange(n, dtype=jnp.int32) + mask_bool = mask.astype(bool) + candidate_structural_mask = ( + mask_bool[:, None] + & mask_bool[None, :] + & (pos_of_node[None, :] >= idx[:, None]) + ) + prefix_structural_mask = mask_bool + prefix_pair_structural_mask = ( + mask_bool[:, None] + & mask_bool[None, :] + & (idx[None, :] < pos_of_node[:, None]) + ) + suffix_pair_structural_mask = mask_bool[:, None] & mask_bool[None, :] + + context_reuse = True + candidate_kfac = dict( + kfac_structural_mask=candidate_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + prefix_kfac = dict( + kfac_structural_mask=prefix_structural_mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + query_mask = mask.astype(dtype)[None, :, None] + + x_ln = self._cross_ln( + cross_ln_s, + cross_ln_b, + cand, + tag_id="route.heavy.cross.ln", + **candidate_kfac, + ) + qv = self._dense_no_bias( + cross_w_qv, + x_ln, + tag_id="route.heavy.cross.qv", + **candidate_kfac, + ).reshape(n, n, 2, self.n_heads_kernel, self.d_head) + q = qv[:, :, 0] + v_self = qv[:, :, 1] + + z_ln = self._cross_ln( + cross_prefix_ln_s, + cross_prefix_ln_b, + z, + tag_id="route.heavy.cross.prefix_ln", + **prefix_kfac, + ) + kv = self._dense_no_bias( + cross_w_kv, + z_ln, + tag_id="route.heavy.cross.kv", + **prefix_kfac, + ).reshape(n, 2, self.n_heads_kernel, self.d_head) + k = kv[:, 0] + v = kv[:, 1] + k_b = jnp.broadcast_to(k[None, :, :, :], q.shape) + v_b = jnp.broadcast_to(v[None, :, :, :], q.shape) + prefix_pairs = self._heavy_prefix_pairs_teacher(edge, perm) + cross_bias = self._heavy_cross_edge_bias( + prefix_pairs, + ( + cross_edge_ln_s, + cross_edge_ln_b, + cross_edge_w1, + cross_edge_b1, + cross_edge_w2, + cross_edge_b2, + ), + structural_mask=prefix_pair_structural_mask, + repeat_ndim=2, + context_primal_reused_over_walkers=context_reuse, + ) + lca_tk = lca_alibi_bias( + idx, + idx, + lca_fixed_slopes(self.n_heads_kernel, dtype=cand.dtype), + ) + pos_bias = jnp.transpose(lca_tk, (1, 2, 0)) + cross_bias = cross_bias[None, :, :, :] + pos_bias[:, None, :, :] + key_mask = mask.astype(bool)[None, :] & (idx[None, :] < idx[:, None]) + cross_out = self._route_attention( + q, + k_b, + v_b, + cross_bias, + key_mask, + impl=impl, + key_mask_only=True, + ) + cross_flat = self._collapse_heavy_heads(cross_out) + delta = self._dense_no_bias( + cross_w_o, + cross_flat, + tag_id="route.heavy.cross.o", + **candidate_kfac, + ) + cand = cand + query_mask * self.heavy_residual_gain * delta + + x_ln = self._ln( + self_ln_s, + cand, + tag_id="route.heavy.self.ln", + **candidate_kfac, + ) + qkv = self._dense_no_bias( + self_w_qkv, + x_ln, + tag_id="route.heavy.self.qkv", + **candidate_kfac, + ).reshape(n, n, 3, self.n_heads_kernel, self.d_head) + q = qkv[:, :, 0] + k = qkv[:, :, 1] + v = qkv[:, :, 2] + suffix_pairs = self._heavy_suffix_pairs_teacher(edge) + suffix_bias = self._heavy_edge_bias( + suffix_pairs, + ( + self_edge_ln_s, + self_edge_w1, + self_edge_b1, + self_edge_w2, + self_edge_b2, + ), + prefix="self", + structural_mask=suffix_pair_structural_mask, + repeat_ndim=2, + context_primal_reused_over_walkers=context_reuse, + ) + + suffix_bias = jnp.broadcast_to( + suffix_bias[None, :, :, :], + (n,) + suffix_bias.shape, + ) + suffix_mask = mask.astype(bool)[None, :] & ( + pos_of_node[None, :] >= idx[:, None] + ) + self_out = self._route_attention( + q, + k, + v, + suffix_bias, + suffix_mask, + impl=impl, + ) + self_flat = self._collapse_heavy_heads(self_out) + delta = self._dense_no_bias( + self_w_o, + self_flat, + tag_id="route.heavy.self.o", + **candidate_kfac, + ) + cand = cand + query_mask * self.heavy_residual_gain * delta + + ffn_in = self._ln( + ffn_ln_s, + cand, + tag_id="route.heavy.ffn.ln", + **candidate_kfac, + ) + ffn = self._dense( + ffn_w1, + ffn_b1, + ffn_in, + tag_id="route.heavy.ffn1", + **candidate_kfac, + ) + ffn = fused_silu(ffn) + delta = self._dense( + ffn_w2, + ffn_b2, + ffn, + tag_id="route.heavy.ffn2", + **candidate_kfac, + ) + return cand + query_mask * self.heavy_residual_gain * delta + + def _apply_heavy_teacher(self, base, z, edge, perm, mask): + n = base.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + pos_of_node = jnp.zeros((n,), dtype=jnp.int32).at[perm].set(idx) + impl = self._resolve_heavy_attn_impl(n) + + def apply_one(cand, params): + return self._heavy_layer_teacher( + cand, + z, + edge, + perm, + mask, + pos_of_node, + params, + impl=impl, + ) + + params = self._heavy_layer_params() + cand = base + for layer in params: + cand = apply_one(cand, layer) + return cand + + def _heavy_prefix_pairs_step( + self, + edge, + prefix_ids, + *, + edge_transpose=None, + ): + edge_i_p = edge[:, prefix_ids, :] + edge_p_i = ( + jnp.transpose(edge[prefix_ids, :, :], (1, 0, 2)) + if edge_transpose is None + else edge_transpose[:, prefix_ids, :] + ) + return jnp.concatenate([edge_i_p, edge_p_i], axis=-1) + + def _heavy_suffix_pairs_step(self, edge, *, edge_transpose=None): + edge_i_j = edge + edge_j_i = ( + jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose + ) + return jnp.concatenate([edge_i_j, edge_j_i], axis=-1) + + def _heavy_cross_biases(self, edge): + all_pairs = self._heavy_suffix_pairs_step(edge) + params = self._heavy_layer_params() + return tuple( + self._heavy_cross_edge_bias( + all_pairs, + layer[7:13], + context_primal_reused_over_walkers=True, + ) + for layer in params + ) + + def _heavy_suffix_biases(self, edge): + suffix_pairs = self._heavy_suffix_pairs_step(edge) + params = self._heavy_layer_params() + return tuple( + self._heavy_edge_bias( + suffix_pairs, + layer[16:21], + prefix="self", + context_primal_reused_over_walkers=True, + ) + for layer in params + ) + + def _heavy_biases_tiled( + self, + edge, + edge_transpose, + *, + pair_tile_size: int, + sequence_axis_name: str | None = None, + sequence_mesh=None, + ): + + n = edge.shape[0] + tile = min(int(pair_tile_size), n) + if tile < 1: + raise ValueError("pair_tile_size must be positive") + n_tiles = (n + tile - 1) // tile + padded_n = n_tiles * tile + source_pad = padded_n - n + edge_transpose = ( + jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose + ) + edge_padded = jnp.pad(edge, ((0, 0), (0, source_pad), (0, 0))) + edge_transpose_padded = jnp.pad( + edge_transpose, + ((0, 0), (0, source_pad), (0, 0)), + ) + + def _seq_constraint(value, *axes): + if sequence_axis_name is None: + return value + from jax.sharding import NamedSharding, PartitionSpec as P + + spec = P(*axes) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return jax.lax.with_sharding_constraint(value, spec) + + params = self._heavy_layer_params() + layer_params = tuple(params) + + def project_layer(layer, param_slice, project_bias): + output0 = _seq_constraint( + jnp.zeros( + (n, padded_n, self.n_heads_kernel), + dtype=edge.dtype, + ), + sequence_axis_name, + None, + None, + ) + + def project_tile(tile_index, output): + start = tile_index * tile + edge_tile = jax.lax.dynamic_slice_in_dim( + edge_padded, + start, + tile, + axis=1, + ) + edge_transpose_tile = jax.lax.dynamic_slice_in_dim( + edge_transpose_padded, + start, + tile, + axis=1, + ) + edge_pair = jnp.concatenate( + [edge_tile, edge_transpose_tile], + axis=-1, + ) + edge_pair = _seq_constraint( + edge_pair, + sequence_axis_name, + None, + None, + ) + bias = project_bias( + edge_pair, + layer[param_slice], + context_primal_reused_over_walkers=True, + ) + output = jax.lax.dynamic_update_slice_in_dim( + output, + bias, + start, + axis=1, + ) + return _seq_constraint( + output, + sequence_axis_name, + None, + None, + ) + + output = jax.lax.fori_loop( + 0, + n_tiles, + project_tile, + output0, + ) + return output[:, :n, :] + + cross_biases = tuple( + project_layer(layer, slice(7, 13), self._heavy_cross_edge_bias) + for layer in layer_params + ) + suffix_biases = tuple( + project_layer( + layer, + slice(16, 21), + lambda edge_pair, params, **kwargs: self._heavy_edge_bias( + edge_pair, params, prefix="self", **kwargs + ), + ) + for layer in layer_params + ) + return cross_biases, suffix_biases + + def _pack_heavy_static_bias_tables(self, edge): + + return self._heavy_cross_biases(edge) + self._heavy_suffix_biases(edge) + + def _unpack_heavy_static_bias_tables(self, tables): + + layers = int(self.route_prefix_suffix_layers) + expected = 2 * layers + if len(tables) != expected: + raise ValueError( + "RouterStatic static_bias_tables has " + f"{len(tables)} leaves; expected {expected} for " + f"route_prefix_suffix_layers={layers}" + ) + return tables[:layers], tables[layers:] + + def _heavy_layer_step_candidates( + self, + cand: Float[Array, "n d_model"], + hidden_cache: Float[Array, "n d_model"], + edge: Float[Array, "n n d_edge"], + prefix_ids: Int[Array, "n"], + picked: Array, + mask: Int[Array, "n"] | Array, + t: Int[Array, ""], + params, + *, + impl: str, + cross_bias=None, + suffix_bias=None, + edge_transpose=None, + sequence_axis_name=None, + sequence_mesh=None, + ) -> Float[Array, "n d_model"]: + ( + cross_ln_s, + cross_ln_b, + cross_prefix_ln_s, + cross_prefix_ln_b, + cross_w_qv, + cross_w_kv, + cross_w_o, + cross_edge_ln_s, + cross_edge_ln_b, + cross_edge_w1, + cross_edge_b1, + cross_edge_w2, + cross_edge_b2, + self_ln_s, + self_w_qkv, + self_w_o, + self_edge_ln_s, + self_edge_w1, + self_edge_b1, + self_edge_w2, + self_edge_b2, + ffn_ln_s, + ffn_w1, + ffn_b1, + ffn_w2, + ffn_b2, + ) = params + n = cand.shape[0] + dtype = cand.dtype + idx = jnp.arange(n, dtype=jnp.int32) + mask_bool = mask.astype(bool) + row_active = mask_bool[t] + candidate_structural_mask = row_active & mask_bool & (~picked.astype(bool)) + prefix_structural_mask = row_active & mask_bool & (idx < t) + cross_pair_structural_mask = ( + candidate_structural_mask[:, None] & prefix_structural_mask[None, :] + ) + self_pair_structural_mask = ( + candidate_structural_mask[:, None] & candidate_structural_mask[None, :] + ) + context_reuse = True + candidate_kfac = dict( + kfac_structural_mask=candidate_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + prefix_kfac = dict( + kfac_structural_mask=prefix_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + query_mask = mask.astype(dtype).reshape(n, 1) + + def _seq_constraint(value, *axes): + if sequence_axis_name is None: + return value + from jax.sharding import NamedSharding, PartitionSpec as P + + spec = P(*axes) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return jax.lax.with_sharding_constraint(value, spec) + + cand = _seq_constraint(cand, sequence_axis_name, None) + + x_ln = self._cross_ln( + cross_ln_s, + cross_ln_b, + cand, + tag_id="route.heavy.cross.ln", + **candidate_kfac, + ) + qv = self._dense_no_bias( + cross_w_qv, + x_ln, + tag_id="route.heavy.cross.qv", + **candidate_kfac, + ).reshape(n, 2, self.n_heads_kernel, self.d_head) + qv = _seq_constraint( + qv, + sequence_axis_name, + None, + None, + None, + ) + q = qv[:, 0] + v_self = qv[:, 1] + + z_ln = self._cross_ln( + cross_prefix_ln_s, + cross_prefix_ln_b, + hidden_cache, + tag_id="route.heavy.cross.prefix_ln", + **prefix_kfac, + ) + kv = self._dense_no_bias( + cross_w_kv, + z_ln, + tag_id="route.heavy.cross.kv", + **prefix_kfac, + ).reshape(n, 2, self.n_heads_kernel, self.d_head) + kv = _seq_constraint(kv, None, None, None, None) + k = kv[:, 0] + v = kv[:, 1] + q = _seq_constraint(q, sequence_axis_name, None, None) + v_self = _seq_constraint( + v_self, + sequence_axis_name, + None, + None, + ) + + k = _seq_constraint(k, None, None, None) + v = _seq_constraint(v, None, None, None) + if cross_bias is None: + prefix_pairs = self._heavy_prefix_pairs_step( + edge, + prefix_ids, + edge_transpose=edge_transpose, + ) + cross_bias = self._heavy_cross_edge_bias( + prefix_pairs, + ( + cross_edge_ln_s, + cross_edge_ln_b, + cross_edge_w1, + cross_edge_b1, + cross_edge_w2, + cross_edge_b2, + ), + structural_mask=cross_pair_structural_mask, + scan_shared=True, + repeat_ndim=2, + context_primal_reused_over_walkers=context_reuse, + ) + else: + cross_bias = cross_bias[:, prefix_ids, :] + cross_bias = _seq_constraint( + cross_bias, + sequence_axis_name, + None, + None, + ) + pos_bias = lca_alibi_bias( + jnp.asarray([t], jnp.int32), + idx, + lca_fixed_slopes(self.n_heads_kernel, dtype=dtype), + )[:, 0, :] + cross_bias = cross_bias + jnp.transpose(pos_bias, (1, 0))[None, :, :] + key_mask = mask.astype(bool) & (idx < t) + cross_out = self._route_attention( + q[None, :, :, :], + k[None, :, :, :], + v[None, :, :, :], + cross_bias[None, :, :, :], + key_mask[None, :], + impl=impl, + key_mask_only=True, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + )[0] + cross_out = _seq_constraint( + cross_out, + sequence_axis_name, + None, + None, + ) + cross_flat = self._collapse_heavy_heads(cross_out) + delta = self._dense_no_bias( + cross_w_o, + cross_flat, + tag_id="route.heavy.cross.o", + **candidate_kfac, + ) + cand = cand + query_mask * self.heavy_residual_gain * delta + + x_ln = self._ln( + self_ln_s, + cand, + tag_id="route.heavy.self.ln", + **candidate_kfac, + ) + qkv = self._dense_no_bias( + self_w_qkv, + x_ln, + tag_id="route.heavy.self.qkv", + **candidate_kfac, + ).reshape(n, 3, self.n_heads_kernel, self.d_head) + qkv = _seq_constraint( + qkv, + sequence_axis_name, + None, + None, + None, + ) + q = qkv[:, 0] + k = qkv[:, 1] + v = qkv[:, 2] + q = _seq_constraint(q, sequence_axis_name, None, None) + k = _seq_constraint(k, None, None, None) + v = _seq_constraint(v, None, None, None) + if suffix_bias is None: + suffix_pairs = self._heavy_suffix_pairs_step( + edge, + edge_transpose=edge_transpose, + ) + suffix_bias = self._heavy_edge_bias( + suffix_pairs, + ( + self_edge_ln_s, + self_edge_w1, + self_edge_b1, + self_edge_w2, + self_edge_b2, + ), + prefix="self", + structural_mask=self_pair_structural_mask, + scan_shared=True, + repeat_ndim=2, + context_primal_reused_over_walkers=context_reuse, + ) + suffix_bias = _seq_constraint( + suffix_bias, + sequence_axis_name, + None, + None, + ) + suffix_mask = mask.astype(bool) & (~picked) + self_out = self._route_attention( + q[None, :, :, :], + k[None, :, :, :], + v[None, :, :, :], + suffix_bias[None, :, :, :], + suffix_mask[None, :], + impl=impl, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + )[0] + self_out = _seq_constraint( + self_out, + sequence_axis_name, + None, + None, + ) + self_flat = self._collapse_heavy_heads(self_out) + delta = self._dense_no_bias( + self_w_o, + self_flat, + tag_id="route.heavy.self.o", + **candidate_kfac, + ) + cand = cand + query_mask * self.heavy_residual_gain * delta + + ffn_in = self._ln( + ffn_ln_s, + cand, + tag_id="route.heavy.ffn.ln", + **candidate_kfac, + ) + ffn = self._dense( + ffn_w1, + ffn_b1, + ffn_in, + tag_id="route.heavy.ffn1", + **candidate_kfac, + ) + ffn = fused_silu(ffn) + delta = self._dense( + ffn_w2, + ffn_b2, + ffn, + tag_id="route.heavy.ffn2", + **candidate_kfac, + ) + return cand + query_mask * self.heavy_residual_gain * delta + + def _apply_heavy_step( + self, + base, + hidden_cache, + edge, + prefix_ids, + picked, + mask, + t, + *, + cross_biases=None, + suffix_biases=None, + edge_transpose=None, + sequence_axis_name=None, + sequence_mesh=None, + ): + impl = self._resolve_heavy_attn_impl(base.shape[0]) + layer_params = self._heavy_layer_params() + n_layers = int(self.route_prefix_suffix_layers) + + if cross_biases is None or len(cross_biases) == 0: + cross_biases = (None,) * n_layers + if suffix_biases is None: + suffix_biases = (None,) * n_layers + + cand = base + for params, cross_bias, suffix_bias in zip( + layer_params, + cross_biases, + suffix_biases, + ): + cand = self._heavy_layer_step_candidates( + cand, + hidden_cache, + edge, + prefix_ids, + picked, + mask, + t, + params, + impl=impl, + cross_bias=cross_bias, + suffix_bias=suffix_bias, + edge_transpose=edge_transpose, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + return cand + + +class TreePrefixPointerMHSEA(_PrefixSuffixRouteBase): + tree_merge: _TreePrefixMerge + tree_edge_merge: EdgeMergeOp + tree_level_layer: _TreePrefixSelfLayer + tree_level_fwl: CausalRouterEdgeFWLUpdate + alpha_tree_level_attn: Float[Array, "d_model"] + alpha_tree_level_ffn: Float[Array, "d_model"] + tree_prefix_layers: list[_TreePrefixSelfLayer] + tree_candidate_layers: list[_TreePrefixCandidateLayer] + + g_step_pool: "GDescriptorPool" + g_step_update: "TreeGlobalUpdate" + g_prefix_ffn_w: Float[Array, "d_gstream d_model"] + g_cand_ffn_w: Float[Array, "d_gstream d_model"] + + tree_edge_ln_scale: Float[Array, "two_d_edge"] + tree_edge_w1: Float[Array, "two_d_edge d_msg_hidden"] + tree_edge_b1: Float[Array, "d_msg_hidden"] + tree_edge_w2: Float[Array, "d_msg_hidden d_model"] + tree_edge_b2: Float[Array, "d_model"] + + route_tree_prefix_layers: int = eqx.field(static=True) + route_tree_prefix_candidate_layers: int = eqx.field(static=True) + tree_prefix_merge_hidden: int = eqx.field(static=True) + tree_prefix_edge_hidden: int = eqx.field(static=True) + tree_prefix_residual_gain: float = eqx.field(static=True) + tree_candidate_residual_gain: float = eqx.field(static=True) + tree_ngpt_alpha_max: float = eqx.field(static=True) + + def __init__( + self, + *, + d_in: int, + d_edge: int, + d_global: int, + d_model: int, + n_heads: int, + max_n: int, + key: PRNGKeyArray, + route_tree_prefix_layers: int = 1, + route_tree_prefix_candidate_layers: int = 1, + route_tree_prefix_merge_hidden: int, + route_tree_prefix_post_prefix_suffix_layers: int = 0, + score_init_scale: float = 1.0, + route_decoder_attn_impl: str = "mhsea_tuned", + rope_base: float = 10000.0, + rope_scaling: float = 1.0, + attention_dim: int, + pointer_score_dim: int, + candidate_hidden: int, + summary_hidden: int, + ffn_hidden: int, + global_tap_dim: int, + alpha_init: float, + alpha_max: float, + ): + if route_tree_prefix_layers < 0: + raise ValueError("route_tree_prefix_layers must be >= 0") + if route_tree_prefix_candidate_layers < 0: + raise ValueError("route_tree_prefix_candidate_layers must be >= 0") + if route_tree_prefix_post_prefix_suffix_layers < 0: + raise ValueError("route_tree_prefix_post_prefix_suffix_layers must be >= 0") + key_base, key_tree = jax.random.split(key) + super().__init__( + d_in=d_in, + d_edge=d_edge, + d_global=d_global, + d_model=d_model, + n_heads=n_heads, + max_n=max_n, + key=key_base, + score_init_scale=score_init_scale, + route_prefix_suffix_layers=int(route_tree_prefix_post_prefix_suffix_layers), + route_decoder_attn_impl=route_decoder_attn_impl, + rope_base=rope_base, + rope_scaling=rope_scaling, + attention_dim=attention_dim, + pointer_score_dim=pointer_score_dim, + candidate_hidden=candidate_hidden, + summary_hidden=summary_hidden, + ffn_hidden=ffn_hidden, + global_tap_dim=global_tap_dim, + ) + + layers = int(route_tree_prefix_layers) + cand_layers = int(route_tree_prefix_candidate_layers) + merge_hidden = int(route_tree_prefix_merge_hidden) + d_qv = self.n_heads_kernel * self.d_head + d_o_in = self.n_heads * self.d_head + edge_hidden = max(32, 2 * self.n_heads_kernel, self.msg_hidden) + msg_hidden = self.msg_hidden + k_edge, k_merge, k_level, k_prefix, k_cand = jax.random.split(key_tree, 5) + + def w(k, shape, fan_in): + return jax.random.normal(k, shape) * (fan_in**-0.5) + + self.tree_edge_ln_scale = jnp.ones((2 * self.d_edge,)) + ek1, ek2 = jax.random.split(k_edge) + self.tree_edge_w1 = w(ek1, (2 * self.d_edge, msg_hidden), 2 * self.d_edge) + self.tree_edge_b1 = jnp.zeros((msg_hidden,)) + self.tree_edge_w2 = w(ek2, (msg_hidden, self.d_model), msg_hidden) + self.tree_edge_b2 = jnp.zeros((self.d_model,)) + + self.tree_merge = _TreePrefixMerge( + self.d_model, + hidden=merge_hidden, + max_depth=max(1, default_tree_depth(max_n)), + key=k_merge, + ln_eps=1.0e-5, + gladder_d_g=int(self.d_global), + alpha_init=alpha_init, + alpha_max=alpha_max, + ) + self.tree_edge_merge = EdgeMergeOp( + d_edge=self.d_model, + d_c=self.d_model, + key=jax.random.fold_in(k_level, 0xE06E), + alpha_init=alpha_init, + alpha_max=alpha_max, + d_hidden=None, + n_blocks=2, + edge_node_ctx_dim=None, + ) + from .global_ladder import GDescriptorPool, TreeGlobalUpdate + + _k_rg = jax.random.split(jax.random.fold_in(k_merge, 0x61B6), 4) + self.g_step_pool = GDescriptorPool( + int(self.d_global), + self.d_model, + key=_k_rg[0], + tag="gladder.route.step.pool", + ) + self.g_step_update = TreeGlobalUpdate( + int(self.d_global), + self.g_step_pool.d_out, + key=_k_rg[1], + tag="gladder.route.step.upd", + tap_dim=global_tap_dim, + alpha_init=alpha_init, + alpha_max=alpha_max, + ) + self.g_prefix_ffn_w = jax.random.normal( + _k_rg[2], (int(self.d_global), self.d_model) + ) * (int(self.d_global) ** -0.5) + self.g_cand_ffn_w = jax.random.normal( + _k_rg[3], (int(self.d_global), self.d_model) + ) * (int(self.d_global) ** -0.5) + + ks = jax.random.split(k_level, 7) + self.tree_level_layer = _TreePrefixSelfLayer( + ln_scale=jnp.ones((self.d_model,)), + w_qkv=w(ks[0], (self.d_model, 3 * d_qv), self.d_model), + w_o=w(ks[1], (d_o_in, self.d_model), d_o_in), + edge_ln_scale=jnp.ones((self.d_model,)), + edge_w1=w(ks[2], (self.d_model, edge_hidden), self.d_model), + edge_b1=jnp.zeros((edge_hidden,)), + edge_w2=w(ks[3], (edge_hidden, self.n_heads_kernel), edge_hidden), + edge_b2=jnp.zeros((self.n_heads_kernel,)), + ffn_ln_scale=jnp.ones((self.d_model,)), + ffn_w1=w(ks[4], (self.d_model, self.ffn_hidden), self.d_model), + ffn_b1=jnp.zeros((self.ffn_hidden,)), + ffn_w2=w(ks[5], (self.ffn_hidden, self.d_model), self.ffn_hidden), + ffn_b2=jnp.zeros((self.d_model,)), + ) + self.tree_level_fwl = CausalRouterEdgeFWLUpdate( + d_c=self.d_model, + d_edge=self.d_model, + channels=max(32, self.d_model // 2), + alpha_init=alpha_init, + alpha_max=alpha_max, + key=ks[6], + ) + self.alpha_tree_level_attn = float(alpha_init) * jnp.ones((self.d_model,)) + self.alpha_tree_level_ffn = float(alpha_init) * jnp.ones((self.d_model,)) + self.tree_ngpt_alpha_max = float(alpha_max) + + prefix_keys = jax.random.split(k_prefix, max(layers, 1)) + prefix_layers = [] + for li in range(layers): + ks = jax.random.split(prefix_keys[li], 7) + prefix_layers.append( + _TreePrefixSelfLayer( + ln_scale=jnp.ones((self.d_model,)), + w_qkv=w(ks[0], (self.d_model, 3 * d_qv), self.d_model), + w_o=w(ks[1], (d_o_in, self.d_model), d_o_in), + edge_ln_scale=jnp.ones((self.d_model,)), + edge_w1=w(ks[2], (self.d_model, edge_hidden), self.d_model), + edge_b1=jnp.zeros((edge_hidden,)), + edge_w2=w(ks[3], (edge_hidden, self.n_heads_kernel), edge_hidden), + edge_b2=jnp.zeros((self.n_heads_kernel,)), + ffn_ln_scale=jnp.ones((self.d_model,)), + ffn_w1=w(ks[4], (self.d_model, self.ffn_hidden), self.d_model), + ffn_b1=jnp.zeros((self.ffn_hidden,)), + ffn_w2=w(ks[5], (self.ffn_hidden, self.d_model), self.ffn_hidden), + ffn_b2=jnp.zeros((self.d_model,)), + ) + ) + self.tree_prefix_layers = prefix_layers + + cand_keys = jax.random.split(k_cand, max(cand_layers, 1)) + tree_candidate_layers = [] + for li in range(cand_layers): + ks = jax.random.split(cand_keys[li], 8) + tree_candidate_layers.append( + _TreePrefixCandidateLayer( + cand_ln_scale=jnp.ones((self.d_model,)), + prefix_ln_scale=jnp.ones((self.d_model,)), + cand_w_qv=w(ks[0], (self.d_model, 2 * d_qv), self.d_model), + prefix_w_kv=w(ks[1], (self.d_model, 2 * d_qv), self.d_model), + w_o=w(ks[2], (d_o_in, self.d_model), d_o_in), + edge_ln_scale=jnp.ones((self.d_model,)), + edge_w1=w(ks[3], (self.d_model, edge_hidden), self.d_model), + edge_b1=jnp.zeros((edge_hidden,)), + edge_w2=w(ks[4], (edge_hidden, self.n_heads_kernel), edge_hidden), + edge_b2=jnp.zeros((self.n_heads_kernel,)), + ffn_ln_scale=jnp.ones((self.d_model,)), + ffn_w1=w(ks[5], (self.d_model, self.ffn_hidden), self.d_model), + ffn_b1=jnp.zeros((self.ffn_hidden,)), + ffn_w2=w(ks[6], (self.ffn_hidden, self.d_model), self.ffn_hidden), + ffn_b2=jnp.zeros((self.d_model,)), + ) + ) + self.tree_candidate_layers = tree_candidate_layers + + self.route_tree_prefix_layers = layers + self.route_tree_prefix_candidate_layers = cand_layers + self.tree_prefix_merge_hidden = int(merge_hidden) + self.tree_prefix_edge_hidden = int(edge_hidden) + self.tree_prefix_residual_gain = 0.0 if layers == 0 else float(layers) ** -0.5 + self.tree_candidate_residual_gain = ( + 0.0 if cand_layers == 0 else float(cand_layers) ** -0.5 + ) + + def _tree_prefix_layer_params(self): + return [layer.as_tuple() for layer in self.tree_prefix_layers] + + def _tree_candidate_layer_params(self): + return [layer.as_tuple() for layer in self.tree_candidate_layers] + + def _resolve_tree_attn_impl(self) -> str: + return self.route_decoder_attn_impl + + def _tree_edge_message_mlp(self, edge_pair, structural_mask): + x = self._ln( + self.tree_edge_ln_scale, + edge_pair, + tag_id="route.tree_prefix.edge_msg_ln", + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=2, + ) + x = self._dense( + self.tree_edge_w1, + self.tree_edge_b1, + x, + tag_id="route.tree_prefix.edge_msg1", + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=2, + ) + x = fused_silu(x) + return self._dense( + self.tree_edge_w2, + self.tree_edge_b2, + x, + tag_id="route.tree_prefix.edge_msg2", + kfac_structural_mask=structural_mask, + kfac_repeat_ndim=2, + ) + + def _tree_pair_messages(self, edge, mask): + edge_pair = jnp.concatenate([jnp.swapaxes(edge, 0, 1), edge], axis=-1) + mask_bool = mask.astype(bool) + structural_mask = mask_bool[:, None] & mask_bool[None, :] + return self._tree_edge_message_mlp(edge_pair, structural_mask) + + def _tree_pair_messages_for_route( + self, + edge, + route_ids, + mask, + *, + edge_transpose=None, + sequence_axis_name=None, + sequence_mesh=None, + row_permute_fn=None, + ): + + edge_transpose = ( + jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose + ) + edge_rows = ( + jnp.take(edge, route_ids, axis=0) + if row_permute_fn is None + else row_permute_fn(edge, route_ids) + ) + edge_transpose_rows = ( + jnp.take(edge_transpose, route_ids, axis=0) + if row_permute_fn is None + else row_permute_fn(edge_transpose, route_ids) + ) + edge_fwd = jnp.take(edge_rows, route_ids, axis=1) + edge_rev = jnp.take(edge_transpose_rows, route_ids, axis=1) + if sequence_axis_name is not None: + from jax.sharding import NamedSharding, PartitionSpec as P + + spec = P(sequence_axis_name, None, None) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + edge_fwd = jax.lax.with_sharding_constraint(edge_fwd, spec) + edge_rev = jax.lax.with_sharding_constraint(edge_rev, spec) + edge_pair = jnp.concatenate([edge_rev, edge_fwd], axis=-1) + route_mask = mask.astype(bool)[route_ids] + structural_mask = route_mask[:, None] & route_mask[None, :] + return self._tree_edge_message_mlp(edge_pair, structural_mask) + + def _tree_pair_message_row( + self, + edge, + source, + mask, + *, + edge_transpose=None, + ): + + edge_transpose = ( + jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose + ) + edge_pair = jnp.concatenate( + [edge[:, source, :], edge_transpose[:, source, :]], + axis=-1, + ) + structural_mask = mask.astype(bool)[source] & mask.astype(bool) + return self._tree_edge_message_mlp(edge_pair, structural_mask) + + def _tree_pair_message_column( + self, + edge, + destination, + mask, + *, + edge_transpose=None, + ): + + edge_transpose = ( + jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose + ) + edge_pair = jnp.concatenate( + [edge_transpose[:, destination, :], edge[:, destination, :]], + axis=-1, + ) + structural_mask = mask.astype(bool)[destination] & mask.astype(bool) + return self._tree_edge_message_mlp(edge_pair, structural_mask) + + def _tree_clock_depth_from_mask(self, mask): + n_active = jnp.maximum(jnp.sum(mask.astype(jnp.int32)), 1) + depth = jnp.ceil(jnp.log2(n_active.astype(jnp.float32))).astype(jnp.int32) + return jnp.maximum(depth, 1) + + def _apply_tree_level_attention(self, nodes, edges, mask, level_idx): + edges = self.tree_level_fwl.apply_residual( + edges, + nodes, + mask, + kfac_scan_shared=True, + ) + + ( + ln_s, + w_qkv, + w_o, + edge_ln_s, + edge_w1, + edge_b1, + edge_w2, + edge_b2, + ffn_ln_s, + ffn_w1, + ffn_b1, + ffn_w2, + ffn_b2, + ) = self.tree_level_layer.as_tuple() + del level_idx + + n = nodes.shape[0] + dtype = nodes.dtype + mask_bool = mask.astype(bool) + idx = jnp.arange(n, dtype=jnp.int32) + node_structural_mask = mask_bool + pair_structural_mask = ( + mask_bool[:, None] & mask_bool[None, :] & (idx[None, :] <= idx[:, None]) + ) + query_mask = mask.astype(dtype)[:, None] + x_ln = self._ln( + ln_s, + nodes, + tag_id="route.tree_prefix.level.ln", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + qkv = self._dense_no_bias( + w_qkv, + x_ln, + tag_id="route.tree_prefix.level.qkv", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ).reshape(n, 3, self.n_heads_kernel, self.d_head) + q = qkv[:, 0] + k = qkv[:, 1] + v = qkv[:, 2] + edge_bias = self._tree_prefix_edge_bias( + edges, + (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), + prefix="level", + kfac_structural_mask=pair_structural_mask, + kfac_scan_shared=True, + ) + edge_bias = edge_bias + jnp.transpose( + lca_alibi_bias( + idx, + idx, + lca_fixed_slopes(self.n_heads_kernel, dtype=dtype), + ), + (1, 2, 0), + ) + + compute_dtype = dtype + q_c = q.astype(compute_dtype) + k_c = k.astype(compute_dtype) + v_c = v.astype(compute_dtype) + logits = jnp.einsum("ihd,jhd->hij", q_c, k_c) + logits = logits / jnp.sqrt(jnp.asarray(self.d_head, dtype=compute_dtype)) + logits = logits + jnp.transpose(edge_bias.astype(compute_dtype), (2, 0, 1)) + valid = (idx[None, :] <= idx[:, None]) & mask_bool[None, :] + logits = jnp.where( + valid[None, :, :], + logits, + jnp.asarray(-1.0e30, dtype=compute_dtype), + ) + alpha = jax.nn.softmax(logits, axis=-1) + out = jnp.einsum("hij,jhd->ihd", alpha, v_c).astype(dtype) + delta = self._dense_no_bias( + w_o, + self._collapse_heavy_heads(out), + tag_id="route.tree_prefix.level.o", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + proposal_attn = query_mask * delta + x = _tree_ngpt_residual( + nodes, + proposal_attn, + self.alpha_tree_level_attn, + max_gain=self.tree_ngpt_alpha_max, + tag_id="", + update_mask=mask, + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + y = self._ln( + ffn_ln_s, + x, + tag_id="route.tree_prefix.level.ffn_ln", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + y = self._dense( + ffn_w1, + ffn_b1, + y, + tag_id="route.tree_prefix.level.ffn1", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + y = fused_silu(y) + y = self._dense( + ffn_w2, + ffn_b2, + y, + tag_id="route.tree_prefix.level.ffn2", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + proposal_ffn = query_mask * y + x = _tree_ngpt_residual( + x, + proposal_ffn, + self.alpha_tree_level_ffn, + max_gain=self.tree_ngpt_alpha_max, + tag_id="", + update_mask=mask, + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + return x, edges + + def _incremental_edge_parent_vector( + self, + e00, + e01, + e10, + e11, + c0, + c1, + d0, + d1, + active, + ): + + skip = jnp.asarray(0.25, dtype=e00.dtype) * (e00 + e01 + e10 + e11) + proposal = jax.vmap( + lambda x00, x01, x10, x11, xa, xb, ya, yb, keep: self.tree_edge_merge( + x00, + x01, + x10, + x11, + xa, + xb, + ya, + yb, + kfac_structural_mask=keep, + kfac_scan_shared=True, + ) + )(e00, e01, e10, e11, c0, c1, d0, d1, active) + merged = self.tree_edge_merge.apply_skip( + skip, + proposal, + kfac_structural_mask=active, + kfac_scan_shared=True, + ) + merged = _tree_sphere(merged) + return jnp.where(active[:, None], merged, jnp.zeros_like(merged)) + + def _apply_tree_level_attention_append( + self, + raw_nodes, + edge_pre, + active, + row, + b_cache, + *, + edge_row, + edge_col, + sequence_axis_name=None, + sequence_mesh=None, + ): + + edge_row, b_cache = self.tree_level_fwl.append_causal_row( + edge_pre, + raw_nodes, + active, + row, + b_cache, + edge_row=edge_row, + edge_col=edge_col, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + edge_row = jnp.where( + active[:, None], + _tree_sphere(edge_row), + edge_row, + ) + + ( + ln_s, + w_qkv, + w_o, + edge_ln_s, + edge_w1, + edge_b1, + edge_w2, + edge_b2, + ffn_ln_s, + ffn_w1, + ffn_b1, + ffn_w2, + ffn_b2, + ) = self.tree_level_layer.as_tuple() + n = raw_nodes.shape[0] + dtype = raw_nodes.dtype + idx = jnp.arange(n, dtype=jnp.int32) + x_ln = self._ln( + ln_s, + raw_nodes, + tag_id="route.tree_prefix.level.ln", + kfac_structural_mask=active, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + qkv = self._dense_no_bias( + w_qkv, + x_ln, + tag_id="route.tree_prefix.level.qkv", + kfac_structural_mask=active, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ).reshape(n, 3, self.n_heads_kernel, self.d_head) + q = qkv[row, 0] + k = qkv[:, 1] + v = qkv[:, 2] + edge_bias = self._tree_prefix_edge_bias( + edge_row, + (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), + prefix="level", + kfac_structural_mask=active, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + pos_bias = lca_alibi_bias( + jnp.asarray([row], dtype=jnp.int32), + idx, + lca_fixed_slopes(self.n_heads_kernel, dtype=dtype), + )[:, 0, :] + edge_bias = edge_bias + jnp.transpose(pos_bias, (1, 0)) + logits = jnp.einsum("hd,jhd->hj", q, k) + logits = logits / jnp.sqrt(jnp.asarray(self.d_head, dtype=dtype)) + logits = logits + jnp.transpose(edge_bias, (1, 0)) + logits = jnp.where( + active[None, :], + logits, + jnp.asarray(-1.0e30, dtype=dtype), + ) + alpha = jax.nn.softmax(logits, axis=-1) + out = jnp.einsum("hj,jhd->hd", alpha, v) + delta = self._dense_no_bias( + w_o, + self._collapse_heavy_heads(out), + tag_id="route.tree_prefix.level.o", + kfac_structural_mask=jnp.asarray(True), + kfac_scan_shared=True, + kfac_repeat_ndim=0, + ) + raw_row = raw_nodes[row] + x = _tree_ngpt_residual( + raw_row, + delta, + self.alpha_tree_level_attn, + max_gain=self.tree_ngpt_alpha_max, + tag_id="", + update_mask=jnp.asarray(True), + kfac_structural_mask=jnp.asarray(True), + kfac_scan_shared=True, + kfac_repeat_ndim=0, + ) + y = self._ln( + ffn_ln_s, + x, + tag_id="route.tree_prefix.level.ffn_ln", + kfac_structural_mask=jnp.asarray(True), + kfac_scan_shared=True, + kfac_repeat_ndim=0, + ) + y = self._dense( + ffn_w1, + ffn_b1, + y, + tag_id="route.tree_prefix.level.ffn1", + kfac_structural_mask=jnp.asarray(True), + kfac_scan_shared=True, + kfac_repeat_ndim=0, + ) + y = fused_silu(y) + y = self._dense( + ffn_w2, + ffn_b2, + y, + tag_id="route.tree_prefix.level.ffn2", + kfac_structural_mask=jnp.asarray(True), + kfac_scan_shared=True, + kfac_repeat_ndim=0, + ) + x = _tree_ngpt_residual( + x, + y, + self.alpha_tree_level_ffn, + max_gain=self.tree_ngpt_alpha_max, + tag_id="", + update_mask=jnp.asarray(True), + kfac_structural_mask=jnp.asarray(True), + kfac_scan_shared=True, + kfac_repeat_ndim=0, + ) + x = _tree_sphere(x) + return x, edge_row, edge_col, b_cache + + def _incremental_tree_append( + self, + state, + leaf, + chosen, + t, + prefix_ids, + edge, + mask, + *, + edge_transpose=None, + g=None, + sequence_axis_name=None, + sequence_mesh=None, + ): + + nodes_raw, nodes_post, level_states = state + n = mask.shape[0] + n_pad = nodes_raw.shape[1] + depth = nodes_raw.shape[0] - 1 + dtype = leaf.dtype + append = mask[t].astype(bool) + leaf_node = _tree_sphere(leaf) + nodes_raw = nodes_raw.at[0, t].set( + jnp.where(append, leaf_node, jnp.zeros_like(leaf_node)) + ) + nodes_post = nodes_post.at[0, t].set( + jnp.where(append, leaf_node, jnp.zeros_like(leaf_node)) + ) + if depth == 0: + return nodes_raw, nodes_post, level_states + + mask_pad = jnp.pad(mask.astype(dtype), (0, n_pad - n)) + route_pad = jnp.pad( + prefix_ids.astype(jnp.int32), + (0, n_pad - n), + ) + clock_depth = self._tree_clock_depth_from_mask(mask) + depth_features = _tree_ngpt_level_counts( + mask_pad, + n_pad // 2, + depth, + dtype, + feature_n_levels=clock_depth, + ) + clock_state = mask_pad + clock_pair_bases = [] + fixed_pairs = n_pad // 2 + for _level in range(depth): + clock_pairs = clock_state.reshape(fixed_pairs, 2) + clock_parent = ( + clock_pairs[:, 0] + + clock_pairs[:, 1] + - clock_pairs[:, 0] * clock_pairs[:, 1] + ) + clock_pair_bases.append( + jnp.maximum( + jnp.sum(clock_parent.astype(jnp.int32)), + jnp.asarray(2, dtype=jnp.int32), + ) + ) + clock_state = jnp.concatenate( + [clock_parent, jnp.zeros_like(clock_parent)], + axis=0, + ) + g_projected = self.tree_merge.project_global( + g, + append & (t > 0), + ) + + if sequence_axis_name is not None: + from jax.sharding import NamedSharding, PartitionSpec as P + + lanes = ( + int(sequence_mesh.shape[sequence_axis_name]) + if sequence_mesh is not None + else 1 + ) + + def constrain_nodes(value): + spec = P(None, sequence_axis_name, None) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return jax.lax.with_sharding_constraint(value, spec) + + def constrain_level(level_state): + width_local = level_state[0].shape[0] + row_axis = ( + sequence_axis_name + if width_local >= lanes and width_local % lanes == 0 + else None + ) + spec = P(row_axis, None, None) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return tuple( + jax.lax.with_sharding_constraint(value, spec) + for value in level_state + ) + + def constrain_row_value(value): + spec = P(None, None) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return jax.lax.with_sharding_constraint(value, spec) + + def constrain_column_value(value): + row_axis = ( + sequence_axis_name + if value.shape[0] >= lanes and value.shape[0] % lanes == 0 + else None + ) + spec = P(row_axis, None) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return jax.lax.with_sharding_constraint(value, spec) + else: + constrain_nodes = lambda value: value + constrain_level = lambda level_state: level_state + constrain_row_value = lambda value: value + constrain_column_value = lambda value: value + + levels_mut = list(level_states) + for level in range(depth): + width = n_pad >> (level + 1) + block = 1 << (level + 1) + create = append & (((t + 1) % block) == 0) + parent = (t + 1) // block - 1 + lower_post_edges = None if level == 0 else levels_mut[level - 1][1] + + def do_create(operand): + raw_all, post_all, level_state = operand + edge_pre, edge_post, b_cache = level_state + q = jnp.arange(width, dtype=jnp.int32) + q0 = 2 * q + q1 = q0 + 1 + p0 = 2 * parent + p1 = p0 + 1 + children = post_all[level] + left = children[p0] + right = children[p1] + + if level == 0: + source0 = route_pad[p0] + source1 = route_pad[p1] + row0 = self._tree_pair_message_row( + edge, + source0, + mask, + edge_transpose=edge_transpose, + )[route_pad] + row1 = self._tree_pair_message_row( + edge, + source1, + mask, + edge_transpose=edge_transpose, + )[route_pad] + col0 = self._tree_pair_message_column( + edge, + source0, + mask, + edge_transpose=edge_transpose, + )[route_pad] + col1 = self._tree_pair_message_column( + edge, + source1, + mask, + edge_transpose=edge_transpose, + )[route_pad] + row0 = _tree_sphere(row0) + row1 = _tree_sphere(row1) + col0 = _tree_sphere(col0) + col1 = _tree_sphere(col1) + row_cells = ( + row0[q0], + row0[q1], + row1[q0], + row1[q1], + ) + col_cells = ( + col0[q0], + col1[q0], + col0[q1], + col1[q1], + ) + sibling_lr = row0[p1] + sibling_rl = row1[p0] + else: + assert lower_post_edges is not None + lower_row0 = _square_row_by_reduction( + lower_post_edges, + p0, + ) + lower_row1 = _square_row_by_reduction( + lower_post_edges, + p1, + ) + lower_col0 = _square_column_local( + lower_post_edges, + p0, + ) + lower_col1 = _square_column_local( + lower_post_edges, + p1, + ) + row_cells = ( + lower_row0[q0], + lower_row0[q1], + lower_row1[q0], + lower_row1[q1], + ) + col_cells = ( + lower_col0[q0], + lower_col1[q0], + lower_col0[q1], + lower_col1[q1], + ) + sibling_lr = lower_row0[p1] + sibling_rl = lower_row1[p0] + + depth_row = ( + None + if depth_features is None + else depth_features[level, parent][None, :] + ) + merged, _valid, _genuine = self.tree_merge( + left[None, :], + right[None, :], + jnp.ones((1,), dtype=dtype), + jnp.ones((1,), dtype=dtype), + sibling_lr[None, :], + sibling_rl[None, :], + jnp.asarray(level, dtype=jnp.int32), + jnp.asarray([parent], dtype=jnp.int32), + clock_pair_bases[level], + clock_depth, + depth_feats=depth_row, + g=g, + g_structural_mask=jnp.asarray(True), + g_projected=g_projected, + ) + raw_parent = merged[0] + raw_all = raw_all.at[level + 1, parent].set(raw_parent) + + active = q <= parent + new_left = jnp.broadcast_to(left, (width, self.d_model)) + new_right = jnp.broadcast_to(right, (width, self.d_model)) + other_left = children[q0] + other_right = children[q1] + parent_row = self._incremental_edge_parent_vector( + *row_cells, + new_left, + new_right, + other_left, + other_right, + active, + ) + parent_col = self._incremental_edge_parent_vector( + *col_cells, + other_left, + other_right, + new_left, + new_right, + active, + ) + + parent_row = constrain_row_value(parent_row) + parent_col = constrain_column_value(parent_col) + edge_pre = _replace_square_row_column( + edge_pre, + parent, + parent_row, + parent_col, + ) + + post_parent, post_row, post_col, b_cache = ( + self._apply_tree_level_attention_append( + raw_all[level + 1, :width], + edge_pre, + active, + parent, + b_cache, + edge_row=parent_row, + edge_col=parent_col, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + ) + post_all = post_all.at[level + 1, parent].set(post_parent) + post_row = constrain_row_value(post_row) + post_col = constrain_column_value(post_col) + edge_post = _replace_square_row_column( + edge_post, + parent, + post_row, + post_col, + ) + return raw_all, post_all, (edge_pre, edge_post, b_cache) + + nodes_raw, nodes_post, levels_mut[level] = jax.lax.cond( + create, + do_create, + lambda operand: operand, + (nodes_raw, nodes_post, levels_mut[level]), + ) + + nodes_raw = constrain_nodes(nodes_raw) + nodes_post = constrain_nodes(nodes_post) + levels_mut[level] = constrain_level(levels_mut[level]) + return nodes_raw, nodes_post, tuple(levels_mut) + + def _tree_prefix_scan(self, seq, mask, pair_route, *, clock_mask=None, g=None): + n = seq.shape[0] + n_pad = _route_next_pow2(n) + depth = n_pad.bit_length() - 1 + dtype = seq.dtype + pad = n_pad - n + clock_mask = mask if clock_mask is None else clock_mask + clock_depth = self._tree_clock_depth_from_mask(clock_mask) + nodes0 = jnp.pad(seq, ((0, pad), (0, 0))) + valid0 = jnp.pad(mask.astype(dtype), (0, pad)) + nodes0 = jnp.where( + valid0.astype(bool)[:, None], + _tree_sphere(nodes0), + jnp.zeros_like(nodes0), + ) + clock_valid0 = jnp.pad(clock_mask.astype(dtype), (0, pad)) + tree_has_merge = jnp.sum(valid0.astype(jnp.int32)) > 1 + edge0 = jnp.pad(pair_route, ((0, pad), (0, pad), (0, 0))) + edge0_active = valid0.astype(bool)[:, None] & valid0.astype(bool)[None, :] + edge0 = jnp.where( + edge0_active[..., None], + _tree_sphere(edge0), + jnp.zeros_like(edge0), + ) + if depth == 0: + levels = nodes0[None, :, :] + valids = valid0[None, :] + edges = edge0[None, :, :, :] + scan_nodes = jnp.zeros((0,) + nodes0.shape, dtype=dtype) + return levels, valids, valids, edges, scan_nodes + n_pairs = n_pad // 2 + g_projected = self.tree_merge.project_global(g, tree_has_merge) + + def _split(x): + xr = x.reshape((n_pairs, 2) + x.shape[1:]) + return xr[:, 0], xr[:, 1] + + def _zpad(x): + return jnp.concatenate([x, jnp.zeros_like(x)], axis=0) + + def _zpad_edge(x): + pad_n = n_pad - n_pairs + return jnp.pad(x, ((0, pad_n), (0, pad_n), (0, 0))) + + pidx = jnp.arange(n_pairs, dtype=jnp.int32) + + def body(state, xs_lv): + level_idx, depth_feats_lv = xs_lv + nodes, valid, edge_state, clock_valid = state + left, right = _split(nodes) + left_m, right_m = _split(valid) + clock_left_m, clock_right_m = _split(clock_valid) + clock_pair_active = ( + clock_left_m + clock_right_m - clock_left_m * clock_right_m + ) + pair_base = jnp.maximum( + jnp.sum(clock_pair_active.astype(jnp.int32)), + jnp.asarray(2, dtype=jnp.int32), + ) + e_rs = edge_state.reshape(n_pairs, 2, n_pairs, 2, self.d_model) + merged, out_mask, genuine = self.tree_merge( + left, + right, + left_m, + right_m, + e_rs[pidx, 0, pidx, 1, :], + e_rs[pidx, 1, pidx, 0, :], + level_idx, + pidx, + pair_base, + clock_depth, + depth_feats=depth_feats_lv, + g=g, + g_structural_mask=tree_has_merge, + g_projected=g_projected, + ) + valid_pair = valid.reshape(n_pairs, 2).astype(dtype) + weights = valid_pair[:, :, None, None] * valid_pair[None, None, :, :] + cell_count = jnp.sum(weights, axis=(1, 3)) + denom = jnp.maximum( + cell_count, + jnp.asarray(1.0, dtype=dtype), + ) + edge_parent = ( + jnp.sum(e_rs * weights[..., None], axis=(1, 3)) / denom[..., None] + ) + e00 = e_rs[:, 0, :, 0, :] + e01 = e_rs[:, 0, :, 1, :] + e10 = e_rs[:, 1, :, 0, :] + e11 = e_rs[:, 1, :, 1, :] + edge_keep = (genuine[:, None] * genuine[None, :]).astype(bool) + + def _edge_row(e0, e1, e2, e3, c0, c1, keep_row): + return jax.vmap( + lambda x0, x1, x2, x3, d0, d1, keep: self.tree_edge_merge( + x0, + x1, + x2, + x3, + c0, + c1, + d0, + d1, + kfac_structural_mask=keep, + kfac_scan_shared=True, + ) + )(e0, e1, e2, e3, left, right, keep_row) + + edge_proposal = jax.vmap(_edge_row)( + e00, + e01, + e10, + e11, + left, + right, + edge_keep, + ) + edge_updated = self.tree_edge_merge.apply_skip( + edge_parent, + edge_proposal, + kfac_structural_mask=edge_keep, + kfac_scan_shared=True, + ) + edge_parent = jnp.where( + edge_keep[..., None], + edge_updated, + edge_parent, + ) + edge_parent = jnp.where( + (cell_count > 1)[..., None], + _tree_sphere(edge_parent), + edge_parent, + ) + merged_skip = merged + edge_skip = edge_parent + merged, edge_parent = self._apply_tree_level_attention( + merged, + edge_parent, + genuine, + level_idx, + ) + merged = jnp.where( + genuine.astype(bool)[:, None], + _tree_sphere(merged), + merged_skip, + ) + edge_update_mask = ( + genuine.astype(bool)[:, None] & genuine.astype(bool)[None, :] + ) + idx = jnp.arange(n_pairs, dtype=jnp.int32) + edge_update_mask = edge_update_mask & (idx[None, :] <= idx[:, None]) + edge_parent = jnp.where( + edge_update_mask[..., None], + _tree_sphere(edge_parent), + edge_skip, + ) + next_state = ( + _zpad(merged), + _zpad(out_mask), + _zpad_edge(edge_parent), + _zpad(clock_pair_active), + ) + ys = (next_state[0], next_state[1], _zpad(genuine), next_state[2]) + return next_state, ys + + depth_feat_levels = _tree_ngpt_level_counts( + clock_valid0, + n_pairs, + depth, + dtype, + ) + (_nodes, _valid, _edge, _clock_valid), ys = jax.lax.scan( + body, + (nodes0, valid0, edge0, clock_valid0), + (jnp.arange(depth, dtype=jnp.int32), depth_feat_levels), + ) + nodes_y, valid_y, genuine_y, edge_y = ys + tree_levels = jnp.concatenate([nodes0[None, :, :], nodes_y], axis=0) + valid_levels = jnp.concatenate([valid0[None, :], valid_y], axis=0) + genuine_levels = jnp.concatenate([valid0[None, :], genuine_y], axis=0) + edge_levels = jnp.concatenate([edge0[None, :, :, :], edge_y], axis=0) + return tree_levels, valid_levels, genuine_levels, edge_levels, nodes_y + + def _source_edge_levels(self, pair_msg, mask): + n = pair_msg.shape[0] + n_dst = pair_msg.shape[1] + n_pad = _route_next_pow2(n) + depth = n_pad.bit_length() - 1 + dtype = pair_msg.dtype + pad = n_pad - n + edge_state = jnp.pad(pair_msg, ((0, pad), (0, 0), (0, 0))) + valid = jnp.pad(mask.astype(dtype), (0, pad)) + levels = [edge_state] + n_pairs = n_pad // 2 + for _level in range(depth): + e_rs = edge_state.reshape(n_pairs, 2, n_dst, self.d_model) + v_rs = valid.reshape(n_pairs, 2) + weights = v_rs[:, :, None, None] + denom = jnp.maximum( + jnp.sum(v_rs, axis=1), + jnp.asarray(1.0, dtype=dtype), + ) + parent = jnp.sum(e_rs * weights, axis=1) / denom[:, None, None] + valid_parent = v_rs[:, 0] + v_rs[:, 1] - v_rs[:, 0] * v_rs[:, 1] + edge_state = jnp.concatenate([parent, jnp.zeros_like(parent)], axis=0) + valid = jnp.concatenate( + [valid_parent, jnp.zeros_like(valid_parent)], axis=0 + ) + levels.append(edge_state) + return jnp.stack(levels, axis=0) + + def _prefix_cover(self, n: int): + n_pad = _route_next_pow2(n) + depth = n_pad.bit_length() - 1 + if depth == 0: + return ( + jnp.zeros((n, 0), dtype=jnp.int32), + jnp.zeros((n, 0), dtype=jnp.int32), + jnp.zeros((n, 0), dtype=bool), + ) + t = jnp.arange(n, dtype=jnp.int32) + start = jnp.zeros((n,), dtype=jnp.int32) + levels = [] + nodes = [] + valids = [] + for bit in range(depth - 1, -1, -1): + take = (jnp.right_shift(t, bit) & 1) == 1 + node = jnp.right_shift(start, bit) + levels.append(jnp.where(take, jnp.asarray(bit, jnp.int32), 0)) + nodes.append(jnp.where(take, node, 0)) + valids.append(take) + start = start + jnp.where(take, jnp.asarray(1 << bit, jnp.int32), 0) + level_arr = jnp.stack(levels, axis=1) + node_arr = jnp.stack(nodes, axis=1) + valid_arr = jnp.stack(valids, axis=1) + prefix_width = _route_next_pow2(depth) + pad = prefix_width - depth + if pad: + level_arr = jnp.pad(level_arr, ((0, 0), (0, pad))) + node_arr = jnp.pad(node_arr, ((0, 0), (0, pad))) + valid_arr = jnp.pad(valid_arr, ((0, 0), (0, pad)), constant_values=False) + return level_arr, node_arr, valid_arr + + def _segment_weights(self, cover_level, cover_node, cover_valid, mask, dtype): + n = mask.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + leaf_node = jnp.right_shift(idx[None, None, :], cover_level[..., None]) + member = ( + (leaf_node == cover_node[..., None]) + & cover_valid[..., None] + & mask.astype(bool)[None, None, :] + ) + weights = member.astype(dtype) + denom = jnp.maximum( + jnp.sum(weights, axis=-1, keepdims=True), + jnp.asarray(1.0, dtype=dtype), + ) + return weights / denom + + def _tree_prefix_context_all( + self, + seq, + mask, + pair_route, + pair_to_nodes, + route_ids, + g=None, + ): + n = seq.shape[0] + dtype = seq.dtype + cover_level, cover_node, cover_valid = self._prefix_cover(n) + depth = cover_level.shape[1] + if depth == 0: + return ( + jnp.zeros((n, 0, self.d_model), dtype=dtype), + jnp.zeros((n, n, 0, self.d_model), dtype=dtype), + jnp.zeros((n, 0, 0, self.d_model), dtype=dtype), + jnp.zeros((n, 0), dtype=bool), + None, + ) + tree_levels, _valid_levels, genuine_levels, _edge_levels, _ys = ( + self._tree_prefix_scan(seq, mask, pair_route, g=g) + ) + source_to_nodes = self._source_edge_levels(pair_to_nodes, mask) + prefix_nodes = tree_levels[cover_level, cover_node] + + prefix_mask = ( + cover_valid + & (genuine_levels[cover_level, cover_node] > 0) + & mask.astype(bool)[:, None] + ) + source_nodes = source_to_nodes[cover_level, cover_node] + cand_prefix_edge = jnp.transpose(source_nodes, (0, 2, 1, 3)) + + source_route = jnp.take(source_nodes, route_ids, axis=-2) + dst_weights = self._segment_weights( + cover_level, + cover_node, + cover_valid, + mask, + dtype, + ) + prefix_prefix_edge = jnp.einsum("tlsd,tms->tlmd", source_route, dst_weights) + g_rows = self._causal_prefix_g(g, prefix_nodes, prefix_mask) + return (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_rows) + + def _causal_prefix_g(self, g, prefix_nodes, prefix_mask): + + if prefix_nodes.ndim == 2: + update_active = jnp.any(prefix_mask.astype(bool)) + return self.g_step_update( + g, + self.g_step_pool( + g, + prefix_nodes, + prefix_mask.astype(prefix_nodes.dtype), + kfac_structural_mask=prefix_mask.astype(bool), + kfac_update_mask=update_active, + kfac_repeat_ndim=1, + ), + update_mask=update_active, + kfac_structural_mask=update_active, + kfac_repeat_ndim=0, + ) + + pool_query_active = jnp.any(prefix_mask.astype(bool)) + return jax.vmap( + lambda pn, pm: self.g_step_update( + g, + self.g_step_pool( + g, + pn, + pm, + kfac_structural_mask=pm.astype(bool), + kfac_update_mask=pool_query_active, + kfac_repeat_ndim=2, + ), + update_mask=jnp.any(pm.astype(bool)), + kfac_structural_mask=jnp.any(pm.astype(bool)), + kfac_g_structural_mask=pool_query_active, + kfac_repeat_ndim=1, + ) + )(prefix_nodes, prefix_mask.astype(prefix_nodes.dtype)) + + def _tree_prefix_context_row( + self, + seq, + mask, + pair_route, + pair_to_nodes, + route_ids, + t, + *, + clock_mask=None, + g=None, + source_edge_frontier=None, + source_edge_counts=None, + sequence_axis_name=None, + sequence_mesh=None, + ): + n = seq.shape[0] + dtype = seq.dtype + cover_level, cover_node, cover_valid = self._prefix_cover(n) + depth = cover_level.shape[1] + if depth == 0: + return ( + jnp.zeros((0, self.d_model), dtype=dtype), + jnp.zeros((n, 0, self.d_model), dtype=dtype), + jnp.zeros((0, 0, self.d_model), dtype=dtype), + jnp.zeros((0,), dtype=bool), + None, + ) + tree_levels, _valid_levels, genuine_levels, _edge_levels, _ys = ( + self._tree_prefix_scan(seq, mask, pair_route, clock_mask=clock_mask, g=g) + ) + cl = cover_level[t] + cn = cover_node[t] + cv = cover_valid[t] + prefix_nodes = tree_levels[cl, cn] + row_active = ( + mask[t].astype(bool) if clock_mask is None else clock_mask[t].astype(bool) + ) + prefix_mask = cv & (genuine_levels[cl, cn] > 0) & row_active + if source_edge_frontier is None: + source_to_nodes = self._source_edge_levels(pair_to_nodes, mask) + source_nodes = source_to_nodes[cl, cn] + else: + if source_edge_counts is None: + raise ValueError( + "source_edge_counts is required with source_edge_frontier" + ) + source_sums = source_edge_frontier[cl] + source_counts = source_edge_counts[cl] + source_nodes = ( + source_sums + / jnp.maximum( + source_counts, + jnp.asarray(1.0, dtype=dtype), + )[:, None, None] + ) + source_nodes = jnp.where( + (cv & (source_counts > 0))[:, None, None], + source_nodes, + jnp.zeros_like(source_nodes), + ) + dst_weights = self._segment_weights( + cl[None, :], + cn[None, :], + cv[None, :], + mask, + dtype, + )[0] + + weights_by_candidate = ( + jnp.zeros_like(dst_weights).at[:, route_ids].add(dst_weights) + ) + if sequence_axis_name is not None: + from jax.sharding import NamedSharding, PartitionSpec as P + + def _sharding(*axes): + spec = P(*axes) + return ( + NamedSharding(sequence_mesh, spec) + if sequence_mesh is not None + else spec + ) + + source_nodes = jax.lax.with_sharding_constraint( + source_nodes, + _sharding(None, sequence_axis_name, None), + ) + weights_by_candidate = jax.lax.with_sharding_constraint( + weights_by_candidate, + _sharding(None, None), + ) + cand_prefix_edge = jnp.transpose(source_nodes, (1, 0, 2)) + prefix_prefix_edge = jnp.einsum( + "lcd,mc->lmd", + source_nodes, + weights_by_candidate, + ) + if sequence_axis_name is not None: + cand_prefix_edge = jax.lax.with_sharding_constraint( + cand_prefix_edge, + _sharding(sequence_axis_name, None, None), + ) + prefix_prefix_edge = jax.lax.with_sharding_constraint( + prefix_prefix_edge, + _sharding(None, None, None), + ) + g_row = self._causal_prefix_g(g, prefix_nodes, prefix_mask) + return (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_row) + + def _incremental_tree_state(self, n: int, dtype): + + n_pad = _route_next_pow2(n) + depth = n_pad.bit_length() - 1 + nodes_raw = jnp.zeros((depth + 1, n_pad, self.d_model), dtype=dtype) + nodes_post = jnp.zeros_like(nodes_raw) + channels = self.tree_level_fwl.two_hop_channels + levels = [] + for level in range(depth): + width = n_pad >> (level + 1) + edge_pre = jnp.zeros( + (width, width, self.d_model), + dtype=dtype, + ) + edge_post = jnp.zeros_like(edge_pre) + b_cache = jnp.zeros((width, width, channels), dtype=dtype) + levels.append((edge_pre, edge_post, b_cache)) + return nodes_raw, nodes_post, tuple(levels) + + def _tree_prefix_context_row_incremental( + self, + nodes_post, + mask, + route_ids, + t, + *, + g=None, + source_edge_frontier, + source_edge_counts, + sequence_axis_name=None, + sequence_mesh=None, + ): + + n = mask.shape[0] + dtype = nodes_post.dtype + cover_level, cover_node, cover_valid = self._prefix_cover(n) + cl = cover_level[t] + cn = cover_node[t] + cv = cover_valid[t] + prefix_nodes = nodes_post[cl, cn] + prefix_mask = cv & mask[t].astype(bool) + source_sums = source_edge_frontier[cl] + source_counts = source_edge_counts[cl] + source_nodes = ( + source_sums + / jnp.maximum( + source_counts, + jnp.asarray(1.0, dtype=dtype), + )[:, None, None] + ) + source_nodes = jnp.where( + (cv & (source_counts > 0))[:, None, None], + source_nodes, + jnp.zeros_like(source_nodes), + ) + dst_weights = self._segment_weights( + cl[None, :], + cn[None, :], + cv[None, :], + mask, + dtype, + )[0] + weights_by_candidate = ( + jnp.zeros_like(dst_weights).at[:, route_ids].add(dst_weights) + ) + if sequence_axis_name is not None: + from jax.sharding import NamedSharding, PartitionSpec as P + + def _sharding(*axes): + spec = P(*axes) + return ( + NamedSharding(sequence_mesh, spec) + if sequence_mesh is not None + else spec + ) + + source_nodes = jax.lax.with_sharding_constraint( + source_nodes, + _sharding(None, sequence_axis_name, None), + ) + weights_by_candidate = jax.lax.with_sharding_constraint( + weights_by_candidate, + _sharding(None, None), + ) + cand_prefix_edge = jnp.transpose(source_nodes, (1, 0, 2)) + prefix_prefix_edge = jnp.einsum( + "lcd,mc->lmd", + source_nodes, + weights_by_candidate, + ) + if sequence_axis_name is not None: + cand_prefix_edge = jax.lax.with_sharding_constraint( + cand_prefix_edge, + _sharding(sequence_axis_name, None, None), + ) + prefix_prefix_edge = jax.lax.with_sharding_constraint( + prefix_prefix_edge, + _sharding(None, None, None), + ) + g_row = self._causal_prefix_g(g, prefix_nodes, prefix_mask) + return ( + prefix_nodes, + cand_prefix_edge, + prefix_prefix_edge, + prefix_mask, + g_row, + ) + + def _tree_prefix_edge_bias( + self, + edge_msg, + params, + *, + prefix: str, + kfac_structural_mask, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 2, + kfac_context_primal_reused_over_walkers: bool = False, + ): + ln_s, w1, b1, w2, b2 = params + x = self._ln( + ln_s, + edge_msg, + tag_id=f"route.tree_prefix.{prefix}.edge_ln", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + x = self._dense( + w1, + b1, + x, + tag_id=f"route.tree_prefix.{prefix}.edge_bias1", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + x = fused_silu(x) + return self._dense( + w2, + b2, + x, + tag_id=f"route.tree_prefix.{prefix}.edge_bias2", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=( + kfac_context_primal_reused_over_walkers + ), + ) + + def _tree_query_seed(self, query_global, route_pos, mask, dtype): + + pos = jnp.asarray(route_pos, dtype=jnp.int32) + target_shape = pos.shape + (self.d_model,) + if query_global is None: + seed = jnp.zeros(target_shape, dtype=dtype) + else: + seed = jnp.broadcast_to( + jnp.asarray(query_global, dtype=dtype), + target_shape, + ) + seed = seed + self._route_position_embedding( + pos, + dtype, + mask=mask, + ) + return seed + + def _tree_prefix_layer( + self, + x, + prefix_edges, + token_mask, + attention_mask, + token_structural_mask, + pair_structural_mask, + params, + *, + impl, + g_projection, + ): + ( + ln_s, + w_qkv, + w_o, + edge_ln_s, + edge_w1, + edge_b1, + edge_w2, + edge_b2, + ffn_ln_s, + ffn_w1, + ffn_b1, + ffn_w2, + ffn_b2, + ) = params + context_reuse = True + token_kfac = dict( + kfac_structural_mask=token_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + bsz, n_tokens = x.shape[:2] + residual_mask = token_mask.astype(x.dtype)[..., None] + x_ln = self._ln( + ln_s, + x, + tag_id="route.tree_prefix.graph.ln", + **token_kfac, + ) + qkv = self._dense_no_bias( + w_qkv, + x_ln, + tag_id="route.tree_prefix.graph.qkv", + **token_kfac, + ).reshape(bsz, n_tokens, 3, self.n_heads_kernel, self.d_head) + q = qkv[:, :, 0] + k = qkv[:, :, 1] + v = qkv[:, :, 2] + edge_bias = self._tree_prefix_edge_bias( + prefix_edges, + (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), + prefix="graph", + kfac_structural_mask=pair_structural_mask, + kfac_repeat_ndim=3, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + out = self._route_attention( + q, + k, + v, + edge_bias, + token_mask, + impl=impl, + attention_mask=attention_mask, + ) + delta = self._dense_no_bias( + w_o, + self._collapse_heavy_heads(out), + tag_id="route.tree_prefix.graph.o", + **token_kfac, + ) + x = x + residual_mask * self.tree_prefix_residual_gain * delta + y = self._ln( + ffn_ln_s, + x, + tag_id="route.tree_prefix.graph.ffn_ln", + **token_kfac, + ) + if g_projection is not None: + y = y + g_projection + y = self._dense( + ffn_w1, + ffn_b1, + y, + tag_id="route.tree_prefix.graph.ffn1", + **token_kfac, + ) + y = fused_silu(y) + y = self._dense( + ffn_w2, + ffn_b2, + y, + tag_id="route.tree_prefix.graph.ffn2", + **token_kfac, + ) + return x + residual_mask * self.tree_prefix_residual_gain * y + + def _tree_candidate_layer( + self, + cand, + prefix_nodes, + cand_prefix_edge, + prefix_mask, + cand_mask, + candidate_structural_mask, + prefix_structural_mask, + cross_pair_structural_mask, + params, + *, + impl, + g_projection, + sequence_axis_name=None, + sequence_mesh=None, + ): + ( + cand_ln_s, + pref_ln_s, + cand_w_qv, + pref_w_kv, + w_o, + edge_ln_s, + edge_w1, + edge_b1, + edge_w2, + edge_b2, + ffn_ln_s, + ffn_w1, + ffn_b1, + ffn_w2, + ffn_b2, + ) = params + context_reuse = True + candidate_kfac = dict( + kfac_structural_mask=candidate_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + prefix_kfac = dict( + kfac_structural_mask=prefix_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + bsz, n_cand = cand.shape[:2] + n_pref = prefix_nodes.shape[1] + + def _seq_constraint(value, *axes): + if sequence_axis_name is None: + return value + from jax.sharding import NamedSharding, PartitionSpec as P + + spec = P(*axes) + if sequence_mesh is not None: + spec = NamedSharding(sequence_mesh, spec) + return jax.lax.with_sharding_constraint(value, spec) + + cand = _seq_constraint(cand, None, sequence_axis_name, None) + cand_ln = self._ln( + cand_ln_s, + cand, + tag_id="route.tree_prefix.candidate.ln", + **candidate_kfac, + ) + if sequence_axis_name is None: + qv = self._dense_no_bias( + cand_w_qv, + cand_ln, + tag_id="route.tree_prefix.candidate.qv", + **candidate_kfac, + ).reshape(bsz, n_cand, 2, self.n_heads_kernel, self.d_head) + q = qv[:, :, 0] + v_self = qv[:, :, 1] + else: + d_qv = self.n_heads_kernel * self.d_head + q = jnp.matmul(cand_ln, cand_w_qv[:, :d_qv]).reshape( + bsz, + n_cand, + self.n_heads_kernel, + self.d_head, + ) + v_self = jnp.matmul(cand_ln, cand_w_qv[:, d_qv:]).reshape( + bsz, + n_cand, + self.n_heads_kernel, + self.d_head, + ) + q = _seq_constraint(q, None, sequence_axis_name, None, None) + v_self = _seq_constraint( + v_self, + None, + sequence_axis_name, + None, + None, + ) + pref_ln = self._ln( + pref_ln_s, + prefix_nodes, + tag_id="route.tree_prefix.candidate.prefix_ln", + **prefix_kfac, + ) + kv = self._dense_no_bias( + pref_w_kv, + pref_ln, + tag_id="route.tree_prefix.candidate.kv", + **prefix_kfac, + ).reshape(bsz, n_pref, 2, self.n_heads_kernel, self.d_head) + kv = _seq_constraint(kv, None, None, None, None, None) + k = kv[:, :, 0] + v = kv[:, :, 1] + edge_bias = self._tree_prefix_edge_bias( + cand_prefix_edge, + (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), + prefix="candidate", + kfac_structural_mask=cross_pair_structural_mask, + kfac_repeat_ndim=3, + kfac_context_primal_reused_over_walkers=context_reuse, + ) + edge_bias = _seq_constraint( + edge_bias, + None, + sequence_axis_name, + None, + None, + ) + out = self._route_attention( + q, + k, + v, + edge_bias, + prefix_mask, + impl=impl, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + out = _seq_constraint(out, None, sequence_axis_name, None, None) + delta = self._dense_no_bias( + w_o, + self._collapse_heavy_heads(out), + tag_id="route.tree_prefix.candidate.o", + **candidate_kfac, + ) + cand = cand + cand_mask * self.tree_candidate_residual_gain * delta + y = self._ln( + ffn_ln_s, + cand, + tag_id="route.tree_prefix.candidate.ffn_ln", + **candidate_kfac, + ) + if g_projection is not None: + y = y + g_projection + y = self._dense( + ffn_w1, + ffn_b1, + y, + tag_id="route.tree_prefix.candidate.ffn1", + **candidate_kfac, + ) + y = fused_silu(y) + y = self._dense( + ffn_w2, + ffn_b2, + y, + tag_id="route.tree_prefix.candidate.ffn2", + **candidate_kfac, + ) + return cand + cand_mask * self.tree_candidate_residual_gain * y + + def _apply_tree_prefix_layers( + self, + prefix_nodes, + prefix_edges, + prefix_mask, + g=None, + query_seed=None, + row_mask=None, + ): + + wants_query = query_seed is not None + single = prefix_nodes.ndim == 2 + if single: + prefix_nodes = prefix_nodes[None, :, :] + prefix_edges = prefix_edges[None, :, :, :] + prefix_mask = prefix_mask[None, :] + if wants_query: + query_seed = query_seed[None, :] + if row_mask is not None: + row_mask = jnp.asarray(row_mask, dtype=bool).reshape(1) + if wants_query: + assert query_seed is not None + query_seed = jnp.broadcast_to( + query_seed, + prefix_nodes.shape[:-2] + (self.d_model,), + ) + n_pref = prefix_nodes.shape[1] + x = jnp.concatenate([prefix_nodes, query_seed[:, None, :]], axis=1) + prefix_edges = jnp.pad( + prefix_edges, + ((0, 0), (0, 1), (0, 1), (0, 0)), + ) + token_mask = jnp.concatenate( + [ + prefix_mask.astype(bool), + jnp.ones((prefix_mask.shape[0], 1), dtype=bool), + ], + axis=1, + ) + token_idx = jnp.arange(n_pref + 1, dtype=jnp.int32) + query_row = token_idx == n_pref + cover_key = token_idx < n_pref + attention_mask = token_mask[:, None, :] & ( + query_row[None, :, None] | cover_key[None, None, :] + ) + else: + n_pref = prefix_nodes.shape[1] + x = prefix_nodes + token_mask = prefix_mask.astype(bool) + attention_mask = None + row_structural_mask = ( + jnp.ones((x.shape[0],), dtype=bool) + if row_mask is None + else jnp.broadcast_to(jnp.asarray(row_mask, dtype=bool), (x.shape[0],)) + ) + token_structural_mask = token_mask.astype(bool) & row_structural_mask[:, None] + if attention_mask is None: + pair_structural_mask = ( + token_structural_mask[:, :, None] & token_structural_mask[:, None, :] + ) + else: + pair_structural_mask = ( + token_structural_mask[:, :, None] + & token_structural_mask[:, None, :] + & attention_mask.astype(bool) + ) + if self.route_tree_prefix_layers == 0: + cover = x[:, :n_pref] + if not wants_query: + return cover[0] if single else cover + query = x[:, n_pref] + return ( + cover[0] if single else cover, + query[0] if single else query, + ) + impl = self._resolve_tree_attn_impl() + from hamiltonzero.model.tree import _tagged_dense_no_bias as _tdnb + + _gg_pref = _tdnb( + self.g_prefix_ffn_w, + g, + tag_id="gladder.route.prefix_fproj", + pathway="even", + kfac_structural_mask=jnp.any(row_structural_mask), + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=True, + ).astype(prefix_nodes.dtype) + + params = self._tree_prefix_layer_params() + + def apply_one(state, layer): + return self._tree_prefix_layer( + state, + prefix_edges, + token_mask, + attention_mask, + token_structural_mask, + pair_structural_mask, + layer, + impl=impl, + g_projection=_gg_pref, + ) + + for layer in params: + x = apply_one(x, layer) + cover = x[:, :n_pref] + if not wants_query: + return cover[0] if single else cover + query = x[:, n_pref] + return ( + cover[0] if single else cover, + query[0] if single else query, + ) + + def _apply_tree_candidate_layers( + self, + base, + prefix_nodes, + cand_prefix_edge, + prefix_mask, + mask, + g_rows=None, + candidate_mask=None, + sequence_axis_name=None, + sequence_mesh=None, + ): + if self.route_tree_prefix_candidate_layers == 0 or prefix_nodes.shape[-2] == 0: + return base + single = base.ndim == 2 + if single: + base = base[None, :, :] + prefix_nodes = prefix_nodes[None, :, :] + cand_prefix_edge = cand_prefix_edge[None, :, :, :] + prefix_mask = prefix_mask[None, :] + if candidate_mask is not None: + candidate_mask = jnp.asarray(candidate_mask, dtype=bool)[None, :] + cand = base + impl = self._resolve_tree_attn_impl() + cand_mask = mask.astype(cand.dtype)[None, :, None] + candidate_structural_mask = ( + jnp.broadcast_to(mask.astype(bool), cand.shape[:2]) + if candidate_mask is None + else jnp.broadcast_to( + jnp.asarray(candidate_mask, dtype=bool), cand.shape[:2] + ) + ) + row_structural_mask = jnp.any(candidate_structural_mask, axis=-1) + prefix_structural_mask = prefix_mask.astype(bool) & row_structural_mask[:, None] + cross_pair_structural_mask = ( + candidate_structural_mask[:, :, None] & prefix_structural_mask[:, None, :] + ) + from hamiltonzero.model.tree import _tagged_dense_no_bias as _tdnb + + _gg = _tdnb( + self.g_cand_ffn_w, + g_rows, + tag_id="gladder.route.cand_fproj", + pathway="even", + kfac_structural_mask=( + row_structural_mask + if jnp.ndim(g_rows) > 1 + else jnp.any(row_structural_mask) + ), + kfac_repeat_ndim=(1 if jnp.ndim(g_rows) > 1 else 0), + kfac_context_primal_reused_over_walkers=True, + ).astype(cand.dtype) + _gg_cand = _gg[..., None, :] if _gg.ndim == cand.ndim - 1 else _gg + params = self._tree_candidate_layer_params() + + def apply_one(state, layer): + return self._tree_candidate_layer( + state, + prefix_nodes, + cand_prefix_edge, + prefix_mask, + cand_mask, + candidate_structural_mask, + prefix_structural_mask, + cross_pair_structural_mask, + layer, + impl=impl, + g_projection=_gg_cand, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + + for layer in params: + cand = apply_one(cand, layer) + return cand[0] if single else cand + + def _tree_enrich_teacher( + self, + base, + seq, + edge, + perm, + mask, + g=None, + query_global=None, + ): + pair_msg = self._tree_pair_messages(edge, mask) + pair_route = pair_msg[perm[:, None], perm[None, :]] + pair_to_nodes = pair_msg[perm, :] + (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_rows) = ( + self._tree_prefix_context_all( + seq, mask, pair_route, pair_to_nodes, perm, g=g + ) + ) + query_seed = self._tree_query_seed( + query_global, + jnp.arange(base.shape[0], dtype=jnp.int32), + mask, + base.dtype, + ) + idx = jnp.arange(base.shape[0], dtype=jnp.int32) + mask_bool = mask.astype(bool) + pos_of_node = jnp.zeros((base.shape[0],), dtype=jnp.int32).at[perm].set(idx) + candidate_structural_mask = ( + mask_bool[:, None] + & mask_bool[None, :] + & (pos_of_node[None, :] >= idx[:, None]) + ) + prefix_nodes, query = self._apply_tree_prefix_layers( + prefix_nodes, + prefix_prefix_edge, + prefix_mask, + g=g, + query_seed=query_seed, + row_mask=mask, + ) + candidate = self._apply_tree_candidate_layers( + base, + prefix_nodes, + cand_prefix_edge, + prefix_mask, + mask, + g_rows=g_rows, + candidate_mask=candidate_structural_mask, + ) + return candidate, query + + def _apply_tree_prefix_step( + self, + base, + base_cache, + prefix_ids, + mask, + t, + pair_msg, + g=None, + query_global=None, + picked=None, + source_edge_frontier=None, + source_edge_counts=None, + raw_edge_for_pair_messages=None, + raw_edge_transpose=None, + sequence_axis_name=None, + sequence_mesh=None, + row_permute_fn=None, + incremental_tree_state=None, + ): + idx = jnp.arange(base.shape[0], dtype=jnp.int32) + prefix_mask_positions = mask.astype(bool) & (idx < t) + if incremental_tree_state is None: + pair_route = ( + pair_msg[prefix_ids[:, None], prefix_ids[None, :]] + if raw_edge_for_pair_messages is None + else self._tree_pair_messages_for_route( + raw_edge_for_pair_messages, + prefix_ids, + mask, + edge_transpose=raw_edge_transpose, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + row_permute_fn=row_permute_fn, + ) + ) + pair_to_nodes = ( + pair_msg[prefix_ids, :] if source_edge_frontier is None else None + ) + (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_row) = ( + self._tree_prefix_context_row( + base_cache, + prefix_mask_positions, + pair_route, + pair_to_nodes, + prefix_ids, + t, + clock_mask=mask, + g=g, + source_edge_frontier=source_edge_frontier, + source_edge_counts=source_edge_counts, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + ) + else: + if source_edge_frontier is None or source_edge_counts is None: + raise ValueError( + "incremental tree context requires source-edge frontiers" + ) + (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_row) = ( + self._tree_prefix_context_row_incremental( + incremental_tree_state[1], + mask, + prefix_ids, + t, + g=g, + source_edge_frontier=source_edge_frontier, + source_edge_counts=source_edge_counts, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + ) + query_seed = self._tree_query_seed( + query_global, + t, + mask, + base.dtype, + ) + prefix_nodes, query = self._apply_tree_prefix_layers( + prefix_nodes, + prefix_prefix_edge, + prefix_mask, + g=g, + query_seed=query_seed, + row_mask=mask[t], + ) + candidate = self._apply_tree_candidate_layers( + base, + prefix_nodes, + cand_prefix_edge, + prefix_mask, + mask, + g_rows=g_row, + candidate_mask=( + mask.astype(bool) + & mask[t].astype(bool) + & ( + jnp.ones_like(mask, dtype=bool) + if picked is None + else ~picked.astype(bool) + ) + ), + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + return candidate, query + + def _teacher_logits( + self, + h: Float[Array, "n d_in"], + edge: Float[Array, "n n d_edge"], + perm: Int[Array, "n"], + mask: Int[Array, "n"] | Array, + *, + global_feat: Float[Array, "d_global"] | None = None, + tau: float | Float[Array, ""] = 1.0, + real_mask: Int[Array, "n"] | Array | None = None, + first_orbit_ids: QuotientCarrier, + ) -> Float[Array, "n n"]: + n = h.shape[0] + if n > self.max_n: + raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") + dtype = h.dtype + idx = jnp.arange(n, dtype=jnp.int32) + mask_bool = mask.astype(bool) + first_active_idx = self._first_active_index(mask) + node_state, _node_mean = self._prepare_nodes(h.astype(dtype), mask) + global_state = self._project_global( + global_feat, + dtype, + structural_mask=jnp.any(mask.astype(bool)), + ) + base = self._teacher_candidate_states( + node_state, + global_state, + edge, + perm, + mask, + real_mask=real_mask, + ) + seq = base[idx, perm, :] + tree_candidate_state, hidden = self._tree_enrich_teacher( + base, + seq, + edge, + perm, + mask, + g=global_state[0], + query_global=global_state[1], + ) + candidate_state = self._apply_heavy_teacher( + tree_candidate_state, + hidden, + edge, + perm, + mask, + ) + neg = jnp.asarray(-1.0e30, dtype=dtype) + + pos_of_node = jnp.zeros((n,), dtype=jnp.int32).at[perm].set(idx) + valid = mask_bool[None, :] & (pos_of_node[None, :] >= idx[:, None]) + first_choice_mask = self._learned_first_choice_mask(mask, real_mask) + valid = jnp.where( + (idx == first_active_idx)[:, None], + first_choice_mask[None, :], + valid, + ) + pointer_structural_mask = mask_bool[:, None] & valid + raw = self._pointer_raw( + hidden, + candidate_state, + structural_mask=pointer_structural_mask, + ) + raw = raw / jnp.asarray(tau, dtype=dtype) + + identity = jnp.where( + idx[None, :] == idx[:, None], + jnp.asarray(0.0, dtype=dtype), + neg, + ) + pointer = jnp.where(valid, raw, neg) + active_scores = pointer + active_scores = jax.vmap( + lambda row_i, row: self._apply_quotient_logits( + row, + first_orbit_ids, + row > (neg * jnp.asarray(0.5, dtype=dtype)), + mask, + perm, + row_i, + ) + )(idx, active_scores) + return jnp.where(mask_bool[:, None], active_scores, identity) + + def _decode( + self, + h: Float[Array, "n d_in"], + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + *, + tau: float | Float[Array, ""], + key: PRNGKeyArray, + real_mask: Int[Array, "n"] | Array | None = None, + first_orbit_ids: QuotientCarrier, + router_static, + ): + n = h.shape[0] + if n > self.max_n: + raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") + dtype = h.dtype + idx = jnp.arange(n, dtype=jnp.int32) + mask_bool = mask.astype(bool) + rm_bool = (real_mask if real_mask is not None else mask).astype(bool) + first_active = self._first_active_index(mask) + neg = jnp.asarray(-1.0e30, dtype=dtype) + node_state = (router_static.node_input, router_static.node_projected) + global_state = (router_static.global_input, router_static.global_projected) + suffix_raw0 = router_static.initial_suffix + prefix_raw0 = jnp.zeros_like(suffix_raw0) + virt_count0 = jnp.zeros((), dtype=dtype) + prefix_order_raw0 = jnp.zeros((n, n, self.d_model), dtype=dtype) + virt_prefix_order_raw0 = jnp.zeros((n, self.d_model), dtype=dtype) + order_decay = router_static.order_decay + virt_decay = router_static.virtual_decay + pair_msg = router_static.tree_pair_messages + cross_biases, suffix_biases = self._unpack_heavy_static_bias_tables( + router_static.static_bias_tables + ) + noise = jax.random.gumbel(key, (n, n), dtype=dtype) + + perm0 = idx + picked0 = jnp.zeros((n,), dtype=bool) + prefix_ids0 = jnp.zeros((n,), dtype=jnp.int32) + k_cache0 = jnp.zeros( + (0, n, self.n_heads_kernel, self.d_head), + dtype=dtype, + ) + v_cache0 = jnp.zeros_like(k_cache0) + hidden_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) + base_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) + last_hidden0 = jnp.zeros((self.d_model,), dtype=dtype) + + def body(carry, xs): + ( + perm, + picked, + prefix_raw, + prefix_order_raw, + suffix_raw, + virt_prefix_order_raw, + virt_count, + prefix_ids, + k_cache, + v_cache, + hidden_cache, + base_cache, + last_hidden, + ) = carry + t, noise_t = xs + pref_msg_buf = prefix_order_raw + virt_msg_buf = virt_prefix_order_raw + append_step = mask_bool[t] + first_step = append_step & (t == first_active) + tri_t = ((idx < t) & mask_bool).astype(dtype) + decay_t = order_decay[t] + prefix_order_row = jnp.einsum("s,sd,sid->id", tri_t, decay_t, pref_msg_buf) + vdecay_t = virt_decay[t] + virt_order_row = jnp.einsum("s,sd,sd->d", tri_t, vdecay_t, virt_msg_buf) + base = self._candidate_states_from_summaries( + node_state, + global_state, + prefix_raw, + prefix_order_row, + suffix_raw, + t, + edge, + mask, + prefix_ids, + virt_prefix_order_raw=virt_order_row, + virt_count=virt_count, + real_mask=real_mask, + ) + candidate_state, pointer_hidden = self._apply_tree_prefix_step( + base, + base_cache, + prefix_ids, + mask, + t, + pair_msg, + g=global_state[0], + query_global=global_state[1], + picked=picked, + ) + candidate_state = self._apply_heavy_step( + candidate_state, + hidden_cache, + edge, + prefix_ids, + picked, + mask, + t, + cross_biases=cross_biases, + suffix_biases=suffix_biases, + ) + active_logits = self._pointer_logits( + pointer_hidden, + candidate_state, + picked, + self._step_choice_mask(first_step, mask, real_mask), + tau, + ) + active_logits = self._apply_quotient_logits( + active_logits, + first_orbit_ids, + active_logits > (neg * jnp.asarray(0.5, dtype=dtype)), + mask, + prefix_ids, + t, + ) + identity_logits = jnp.where(idx == t, jnp.asarray(0.0, dtype=dtype), neg) + logits = jnp.where(append_step, active_logits, identity_logits) + select_scores = logits + noise_t + sampled = jnp.argmax(select_scores).astype(jnp.int32) + chosen = jnp.where(append_step, sampled, t) + + base_chosen = base[chosen] + token_in = jnp.where( + append_step, + base_chosen, + jnp.zeros((self.d_model,), dtype=dtype), + ) + token, k_new, v_new = self._append_token( + token_in, + chosen, + t, + prefix_ids, + k_cache, + v_cache, + edge, + mask, + ) + k_cache = k_new + v_cache = v_new + hidden_cache = hidden_cache.at[t].set(pointer_hidden) + base_cache = base_cache.at[t].set( + jnp.where(append_step, base_chosen, jnp.zeros_like(base_chosen)) + ) + last_hidden = pointer_hidden + pref_update = router_static.prefix_edge_messages[chosen] + suff_update = router_static.suffix_edge_messages[chosen] + update_mask = append_step.astype(dtype) + prefix_raw = prefix_raw + update_mask * pref_update + prefix_order_raw = pref_msg_buf.at[t].set(update_mask * pref_update) + suffix_raw = suffix_raw - update_mask * suff_update + virt_update = update_mask * (~rm_bool[chosen]).astype(dtype) + virt_prefix_order_raw = virt_msg_buf.at[t].set( + virt_update * self.virt_emb[0], + ) + virt_count = virt_count + virt_update + prefix_ids = prefix_ids.at[t].set(chosen) + picked = picked.at[chosen].set(jnp.where(append_step, True, picked[chosen])) + perm = perm.at[t].set(chosen) + return ( + perm, + picked, + prefix_raw, + prefix_order_raw, + suffix_raw, + virt_prefix_order_raw, + virt_count, + prefix_ids, + k_cache, + v_cache, + hidden_cache, + base_cache, + last_hidden, + ), None + + init = ( + perm0, + picked0, + prefix_raw0, + prefix_order_raw0, + suffix_raw0, + virt_prefix_order_raw0, + virt_count0, + prefix_ids0, + k_cache0, + v_cache0, + hidden_cache0, + base_cache0, + last_hidden0, + ) + final, _ = jax.lax.scan(body, init, (idx, noise)) + ( + perm, + _picked, + _prefix_raw, + _prefix_order_raw, + _suffix_raw, + _virt_po, + _virt_cnt, + _prefix_ids, + _k, + _v, + _hidden_cache, + _base_cache, + _hidden, + ) = final + return perm + + def _decode_greedy_compact( + self, + h: Float[Array, "n d_in"], + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + *, + global_feat: Float[Array, "d_global"], + tau: float | Float[Array, ""], + real_mask: Int[Array, "n"] | Array, + sequence_mesh, + pair_tile_size: int, + row_permute_fn, + ): + + n = h.shape[0] + if n > self.max_n: + raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") + if int(pair_tile_size) < 1: + raise ValueError("pair_tile_size must be positive") + sequence_axis_name = "seq" + dtype = h.dtype + idx = jnp.arange(n, dtype=jnp.int32) + mask_bool = mask.astype(bool) + rm_bool = real_mask.astype(bool) + first_active = self._first_active_index(mask) + neg = jnp.asarray(-1.0e30, dtype=dtype) + edge_transpose = jnp.swapaxes(edge, 0, 1) + from jax.sharding import NamedSharding, PartitionSpec as P + + edge_transpose = jax.lax.with_sharding_constraint( + edge_transpose, + NamedSharding(sequence_mesh, P(sequence_axis_name, None, None)), + ) + node_state, _node_mean = self._prepare_nodes(h.astype(dtype), mask) + global_state = self._project_global( + global_feat, + dtype, + structural_mask=jnp.any(mask_bool), + ) + ( + prefix_raw0, + _prefix_order_raw0_unused, + suffix_raw0, + _virt_po_unused, + virt_count0, + ) = self._initial_summaries_streamed( + edge, + edge_transpose, + mask, + dtype, + pair_tile_size=pair_tile_size, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + pair_msg = None + cross_biases, suffix_biases = self._heavy_biases_tiled( + edge, + edge_transpose, + pair_tile_size=int(pair_tile_size), + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + + frontier_depth = max(1, _route_next_pow2(n).bit_length() - 1) + prefix_order_frontier0 = jnp.zeros( + (frontier_depth, n, self.d_model), + dtype=dtype, + ) + virt_order_frontier0 = jnp.zeros( + (frontier_depth, self.d_model), + dtype=dtype, + ) + source_edge_frontier0 = jnp.zeros( + (frontier_depth, n, self.d_model), + dtype=dtype, + ) + source_edge_counts0 = jnp.zeros((frontier_depth,), dtype=dtype) + tree_state0 = self._incremental_tree_state(n, dtype) + + perm0 = idx + picked0 = jnp.zeros((n,), dtype=bool) + prefix_ids0 = jnp.zeros((n,), dtype=jnp.int32) + k_cache0 = jnp.zeros( + (0, n, self.n_heads_kernel, self.d_head), + dtype=dtype, + ) + v_cache0 = jnp.zeros_like(k_cache0) + hidden_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) + base_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) + logp0 = jnp.asarray(0.0, dtype=jnp.float32) + seq = sequence_axis_name + + def _seq_sharding(*axes): + return NamedSharding(sequence_mesh, P(*axes)) + + node_state = tuple( + jax.lax.with_sharding_constraint(x, _seq_sharding(seq, None)) + for x in node_state + ) + prefix_raw0 = jax.lax.with_sharding_constraint( + prefix_raw0, + _seq_sharding(seq, None), + ) + suffix_raw0 = jax.lax.with_sharding_constraint( + suffix_raw0, + _seq_sharding(seq, None), + ) + prefix_order_frontier0 = jax.lax.with_sharding_constraint( + prefix_order_frontier0, + _seq_sharding(None, seq, None), + ) + source_edge_frontier0 = jax.lax.with_sharding_constraint( + source_edge_frontier0, + _seq_sharding(None, seq, None), + ) + hidden_cache0 = jax.lax.with_sharding_constraint( + hidden_cache0, + _seq_sharding(seq, None), + ) + base_cache0 = jax.lax.with_sharding_constraint( + base_cache0, + _seq_sharding(seq, None), + ) + k_cache0 = jax.lax.with_sharding_constraint( + k_cache0, + _seq_sharding(None, seq, None, None), + ) + v_cache0 = jax.lax.with_sharding_constraint( + v_cache0, + _seq_sharding(None, seq, None, None), + ) + edge_transpose = jax.lax.with_sharding_constraint( + edge_transpose, + _seq_sharding(seq, None, None), + ) + tree_nodes_raw0, tree_nodes_post0, tree_levels0 = tree_state0 + tree_nodes_raw0 = jax.lax.with_sharding_constraint( + tree_nodes_raw0, + _seq_sharding(None, seq, None), + ) + tree_nodes_post0 = jax.lax.with_sharding_constraint( + tree_nodes_post0, + _seq_sharding(None, seq, None), + ) + lanes = int(sequence_mesh.shape[seq]) + constrained_levels = [] + for edge_pre0, edge_post0, b_cache0 in tree_levels0: + shard_rows = ( + seq + if edge_pre0.shape[0] >= lanes and edge_pre0.shape[0] % lanes == 0 + else None + ) + constrained_levels.append( + ( + jax.lax.with_sharding_constraint( + edge_pre0, + _seq_sharding(shard_rows, None, None), + ), + jax.lax.with_sharding_constraint( + edge_post0, + _seq_sharding(shard_rows, None, None), + ), + jax.lax.with_sharding_constraint( + b_cache0, + _seq_sharding(shard_rows, None, None), + ), + ) + ) + tree_state0 = ( + tree_nodes_raw0, + tree_nodes_post0, + tuple(constrained_levels), + ) + + def body(carry, t): + ( + perm, + picked, + prefix_raw, + prefix_order_frontier, + suffix_raw, + virt_order_frontier, + virt_count, + source_edge_frontier, + source_edge_counts, + tree_state, + prefix_ids, + k_cache, + v_cache, + hidden_cache, + base_cache, + total_logp, + ) = carry + append_step = mask_bool[t] + first_step = append_step & (t == first_active) + predict_step = append_step & (t != first_active) + + prefix_order_row = _dyadic_lca_frontier_sum( + prefix_order_frontier, + t, + self.order_decay_w[0], + self.order_decay_b[0], + ) + virt_order_row = _dyadic_lca_frontier_sum( + virt_order_frontier, + t, + self.virt_decay_w[0], + self.virt_decay_b[0], + ) + base = self._candidate_states_from_summaries( + node_state, + global_state, + prefix_raw, + prefix_order_row, + suffix_raw, + t, + edge, + mask, + prefix_ids, + virt_prefix_order_raw=virt_order_row, + virt_count=virt_count, + real_mask=real_mask, + ) + candidate_state, pointer_hidden = self._apply_tree_prefix_step( + base, + base_cache, + prefix_ids, + mask, + t, + pair_msg, + g=global_state[0], + query_global=global_state[1], + picked=picked, + source_edge_frontier=source_edge_frontier, + source_edge_counts=source_edge_counts, + raw_edge_for_pair_messages=edge, + raw_edge_transpose=edge_transpose, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + row_permute_fn=row_permute_fn, + incremental_tree_state=tree_state, + ) + candidate_state = self._apply_heavy_step( + candidate_state, + hidden_cache, + edge, + prefix_ids, + picked, + mask, + t, + cross_biases=cross_biases, + suffix_biases=suffix_biases, + edge_transpose=edge_transpose, + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + active_logits = self._pointer_logits( + pointer_hidden, + candidate_state, + picked, + self._step_choice_mask(first_step, mask, real_mask), + tau, + ) + identity_logits = jnp.where( + idx == t, + jnp.asarray(0.0, dtype=dtype), + neg, + ) + logits = jnp.where(append_step, active_logits, identity_logits) + sampled = jnp.argmax(logits).astype(jnp.int32) + chosen = jnp.where(append_step, sampled, t) + log_probs = jax.nn.log_softmax(logits.astype(jnp.float32), axis=-1) + score_step = self._score_step_for_logp(first_step, predict_step) + total_logp = total_logp + jnp.where( + score_step, + log_probs[chosen], + 0.0, + ) + + base_chosen = base[chosen] + token_in = jnp.where( + append_step, + base_chosen, + jnp.zeros((self.d_model,), dtype=dtype), + ) + _token, k_cache, v_cache = self._append_token( + token_in, + chosen, + t, + prefix_ids, + k_cache, + v_cache, + edge, + mask, + ) + hidden_cache = hidden_cache.at[t].set(pointer_hidden) + base_cache = base_cache.at[t].set( + jnp.where(append_step, base_chosen, jnp.zeros_like(base_chosen)) + ) + chosen_edge_pair = jnp.concatenate( + [edge[:, chosen, :], edge_transpose[:, chosen, :]], + axis=-1, + ) + pref_update = self._message_mlp(chosen_edge_pair, prefix=True) + suff_update = self._message_mlp(chosen_edge_pair, prefix=False) + + update_mask = append_step.astype(dtype) + prefix_raw = prefix_raw + update_mask * pref_update + prefix_order_frontier = _dyadic_frontier_add( + prefix_order_frontier, + update_mask * pref_update, + t, + ) + suffix_raw = suffix_raw - update_mask * suff_update + virt_update = update_mask * (~rm_bool[chosen]).astype(dtype) + virt_order_frontier = _dyadic_frontier_add( + virt_order_frontier, + virt_update * self.virt_emb[0], + t, + ) + virt_count = virt_count + virt_update + source_edge_update = self._tree_pair_message_row( + edge, + chosen, + mask, + edge_transpose=edge_transpose, + ) + source_edge_frontier = _dyadic_frontier_add( + source_edge_frontier, + update_mask * source_edge_update, + t, + ) + source_edge_counts = _dyadic_frontier_add( + source_edge_counts, + update_mask, + t, + ) + prefix_ids = prefix_ids.at[t].set(chosen) + tree_state = self._incremental_tree_append( + tree_state, + jnp.where( + append_step, + base_chosen, + jnp.zeros_like(base_chosen), + ), + chosen, + t, + prefix_ids, + edge, + mask, + edge_transpose=edge_transpose, + g=global_state[0], + sequence_axis_name=sequence_axis_name, + sequence_mesh=sequence_mesh, + ) + picked = picked.at[chosen].set(jnp.where(append_step, True, picked[chosen])) + perm = perm.at[t].set(chosen) + next_carry = ( + perm, + picked, + prefix_raw, + prefix_order_frontier, + suffix_raw, + virt_order_frontier, + virt_count, + source_edge_frontier, + source_edge_counts, + tree_state, + prefix_ids, + k_cache, + v_cache, + hidden_cache, + base_cache, + total_logp, + ) + return next_carry, None + + init = ( + perm0, + picked0, + prefix_raw0, + prefix_order_frontier0, + suffix_raw0, + virt_order_frontier0, + virt_count0, + source_edge_frontier0, + source_edge_counts0, + tree_state0, + prefix_ids0, + k_cache0, + v_cache0, + hidden_cache0, + base_cache0, + logp0, + ) + final, _ = jax.lax.scan(body, init, idx) + perm = final[0] + logp = final[-1] + return perm, logp + + def beam_search( + self, + h: Float[Array, "n d_in"], + edge: Float[Array, "n n d_edge"], + mask: Int[Array, "n"] | Array, + *, + global_feat: Float[Array, "d_global"] | None = None, + tau: float | Float[Array, ""] = 1.0, + beam_width: int = 4, + real_mask: Int[Array, "n"] | Array | None = None, + first_orbit_ids: QuotientCarrier, + router_static=None, + distributed_axis_name: str | None = None, + distributed_lanes: int | None = None, + ): + B = int(beam_width) + if B < 1: + raise ValueError("beam_width must be >= 1") + if distributed_axis_name is not None: + lanes = int(distributed_lanes if distributed_lanes is not None else 8) + if lanes < 1 or B % lanes: + raise ValueError( + f"distributed beam requires beam_width divisible by the " + f"lane count; got beam_width={B}, lanes={lanes}" + ) + if router_static is None: + raise ValueError("distributed audit beam requires RouterStatic") + n = h.shape[0] + if n > self.max_n: + raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") + dtype = h.dtype + idx = jnp.arange(n, dtype=jnp.int32) + mask_bool = mask.astype(bool) + rm_bool = (real_mask if real_mask is not None else mask).astype(bool) + first_active = self._first_active_index(mask) + neg = jnp.asarray(-1.0e30, dtype=dtype) + if router_static is None: + node_state, _node_mean = self._prepare_nodes(h.astype(dtype), mask) + global_state = self._project_global( + global_feat, + dtype, + structural_mask=jnp.any(mask.astype(bool)), + ) + ( + prefix_raw0, + _prefix_order_raw0_unused, + suffix_raw0, + _virt_po_unused, + virt_count0, + ) = self._initial_summaries(edge, mask, dtype) + else: + node_state = (router_static.node_input, router_static.node_projected) + global_state = (router_static.global_input, router_static.global_projected) + suffix_raw0 = router_static.initial_suffix + prefix_raw0 = jnp.zeros_like(suffix_raw0) + virt_count0 = jnp.zeros((), dtype=dtype) + prefix_order_raw0 = jnp.zeros((n, n, self.d_model), dtype=dtype) + virt_prefix_order_raw0 = jnp.zeros((n, self.d_model), dtype=dtype) + if router_static is None: + order_decay = lca_gaussian_decay( + idx, + idx, + self.order_decay_w[0], + self.order_decay_b[0], + ) + virt_decay = lca_gaussian_decay( + idx, + idx, + self.virt_decay_w[0], + self.virt_decay_b[0], + ) + pair_msg = self._tree_pair_messages(edge, mask) + else: + order_decay = router_static.order_decay + virt_decay = router_static.virtual_decay + pair_msg = router_static.tree_pair_messages + if router_static is None: + cross_biases = self._heavy_cross_biases(edge) + suffix_biases = self._heavy_suffix_biases(edge) + else: + cross_biases, suffix_biases = self._unpack_heavy_static_bias_tables( + router_static.static_bias_tables + ) + + def repeat(x): + return jnp.broadcast_to(x, (B,) + x.shape) + + perm0 = repeat(idx) + picked0 = jnp.zeros((B, n), dtype=bool) + prefix_ids0 = jnp.zeros((B, n), dtype=jnp.int32) + k_cache0 = jnp.zeros( + (B, 0, n, self.n_heads_kernel, self.d_head), + dtype=dtype, + ) + v_cache0 = jnp.zeros_like(k_cache0) + hidden_cache0 = jnp.zeros((B, n, self.d_model), dtype=dtype) + base_cache0 = jnp.zeros((B, n, self.d_model), dtype=dtype) + last_hidden0 = jnp.zeros((B, self.d_model), dtype=dtype) + logp0 = ( + jnp.full((B,), jnp.asarray(-1e9, dtype=jnp.float32), dtype=jnp.float32) + .at[0] + .set(0.0) + ) + beam_ids = jnp.arange(B, dtype=jnp.int32) + rows = jnp.arange(B, dtype=jnp.int32) + + def body(carry, t): + ( + perm, + picked, + prefix_raw, + prefix_order_raw, + suffix_raw, + virt_prefix_order_raw, + virt_count, + prefix_ids, + k_cache, + v_cache, + hidden_cache, + base_cache, + last_hidden, + total_logp, + ) = carry + append_step = mask_bool[t] + first_step = append_step & (t == first_active) + predict_step = append_step & (t != first_active) + + def states_one( + pr, por_buf, sr, vpo_buf, vcnt, pids, pk, bc, hcache, hidden + ): + tri_t = ((idx < t) & mask_bool).astype(dtype) + decay_t = order_decay[t] + por = jnp.einsum("s,sd,sid->id", tri_t, decay_t, por_buf) + vdecay_t = virt_decay[t] + virt_order_row = jnp.einsum("s,sd,sd->d", tri_t, vdecay_t, vpo_buf) + base = self._candidate_states_from_summaries( + node_state, + global_state, + pr, + por, + sr, + t, + edge, + mask, + pids, + virt_prefix_order_raw=virt_order_row, + virt_count=vcnt, + real_mask=real_mask, + ) + candidate_state, pointer_hidden = self._apply_tree_prefix_step( + base, + bc, + pids, + mask, + t, + pair_msg, + g=global_state[0], + query_global=global_state[1], + picked=pk, + ) + candidate_state = self._apply_heavy_step( + candidate_state, + hcache, + edge, + pids, + pk, + mask, + t, + cross_biases=cross_biases, + suffix_biases=suffix_biases, + ) + active_logits = self._pointer_logits( + pointer_hidden, + candidate_state, + pk, + self._step_choice_mask(first_step, mask, real_mask), + tau, + ) + active_logits = self._apply_quotient_logits( + active_logits, + first_orbit_ids, + active_logits > (neg * jnp.asarray(0.5, dtype=dtype)), + mask, + pids, + t, + ) + identity_logits = jnp.where( + idx == t, jnp.asarray(0.0, dtype=dtype), neg + ) + logits = jnp.where(append_step, active_logits, identity_logits) + return logits, candidate_state, base, pointer_hidden + + if distributed_axis_name is None: + parent_rows = rows + else: + lane = jax.lax.axis_index(distributed_axis_name) + _per_lane = B // lanes + parent_rows = lane * _per_lane + jnp.arange(_per_lane, dtype=jnp.int32) + ( + logits_local, + _candidate_state_local, + base_state_local, + query_state_local, + ) = jax.vmap(states_one)( + prefix_raw[parent_rows], + prefix_order_raw[parent_rows], + suffix_raw[parent_rows], + virt_prefix_order_raw[parent_rows], + virt_count[parent_rows], + prefix_ids[parent_rows], + picked[parent_rows], + base_cache[parent_rows], + hidden_cache[parent_rows], + last_hidden[parent_rows], + ) + score_step = self._score_step_for_logp(first_step, predict_step) + neg_f32 = jnp.asarray(-1e9, dtype=jnp.float32) + + def expansion_for(logits_arg, parent_total): + log_probs_arg = jax.nn.log_softmax( + logits_arg.astype(jnp.float32), axis=-1 + ) + step_logp_arg = jnp.where( + score_step, log_probs_arg, jnp.zeros_like(log_probs_arg) + ) + expansion_arg = parent_total[:, None] + step_logp_arg + forced_scores = jnp.where( + idx[None, :] == t.astype(jnp.int32), + expansion_arg, + neg_f32, + ) + return jnp.where(~append_step, forced_scores, expansion_arg) + + expansion_local = expansion_for(logits_local, total_logp[parent_rows]) + if distributed_axis_name is None: + expansion_scores = expansion_local + base_state = base_state_local + query_state = query_state_local + else: + _pl = B // lanes + base_shape = base_state_local.shape + query_shape = query_state_local.shape + payload_parts = [ + expansion_local.reshape((_pl, -1)), + ] + payload_parts.extend( + [ + base_state_local.astype(jnp.float32).reshape((_pl, -1)), + query_state_local.astype(jnp.float32).reshape((_pl, -1)), + ] + ) + payload = jnp.concatenate(payload_parts, axis=-1) + payload = jax.lax.all_gather( + payload, + distributed_axis_name, + axis=0, + tiled=True, + ) + cursor = 0 + expansion_scores = payload[:, cursor : cursor + n] + cursor += n + base_size = n * base_shape[-1] + base_state = ( + payload[:, cursor : cursor + base_size] + .reshape((B, n, base_shape[-1])) + .astype(dtype) + ) + cursor += base_size + query_state = payload[:, cursor : cursor + query_shape[-1]].astype( + dtype + ) + + rank_scores = expansion_scores + identity_distance = jnp.abs(idx - t.astype(jnp.int32)).astype(jnp.float32) + rank_scores = rank_scores - identity_distance[None, :] * 1.0e-6 + rank_scores = rank_scores - beam_ids[:, None].astype(jnp.float32) * 1.0e-9 + _rank_top, flat = jax.lax.top_k(rank_scores.reshape((-1,)), B) + parent = (flat // n).astype(jnp.int32) + chosen = (flat % n).astype(jnp.int32) + total_logp = expansion_scores.reshape((-1,))[flat] + + perm = perm[parent] + picked = picked[parent] + prefix_raw = prefix_raw[parent] + prefix_order_raw = prefix_order_raw[parent] + suffix_raw = suffix_raw[parent] + virt_prefix_order_raw = virt_prefix_order_raw[parent] + virt_count = virt_count[parent] + prefix_ids = prefix_ids[parent] + k_cache = k_cache[parent] + v_cache = v_cache[parent] + hidden_cache = hidden_cache[parent] + base_cache = base_cache[parent] + + base_chosen = base_state[parent, chosen] + query_chosen = query_state[parent] + token_in = jnp.where( + append_step, + base_chosen, + jnp.zeros_like(base_chosen), + ) + + token, k_cache, v_cache = jax.vmap( + lambda token_b, chosen_b, prefix_ids_b, k_b, v_b: self._append_token( + token_b, + chosen_b, + t, + prefix_ids_b, + k_b, + v_b, + edge, + mask, + ) + )(token_in, chosen, prefix_ids, k_cache, v_cache) + hidden_cache = hidden_cache.at[:, t, :].set(query_chosen) + base_cache = base_cache.at[:, t, :].set( + append_step.astype(dtype) * base_chosen + ) + last_hidden = query_chosen + + if router_static is None: + chosen_edge_pair = jax.vmap( + lambda chosen_b: self._edge_pair_for_source(edge, chosen_b) + )(chosen) + pref_update = jax.vmap( + lambda pair: self._message_mlp(pair, prefix=True) + )(chosen_edge_pair) + suff_update = jax.vmap( + lambda pair: self._message_mlp(pair, prefix=False) + )(chosen_edge_pair) + else: + pref_update = router_static.prefix_edge_messages[chosen] + suff_update = router_static.suffix_edge_messages[chosen] + update_mask = append_step.astype(dtype) + prefix_raw = prefix_raw + update_mask * pref_update + prefix_order_raw = prefix_order_raw.at[:, t].set(update_mask * pref_update) + suffix_raw = suffix_raw - update_mask * suff_update + virt_update = update_mask * (~rm_bool[chosen]).astype(dtype) + virt_prefix_order_raw = virt_prefix_order_raw.at[:, t].set( + virt_update[:, None] * self.virt_emb[0][None, :], + ) + virt_count = virt_count + virt_update + prefix_ids = prefix_ids.at[:, t].set(chosen) + old_picked = picked[rows, chosen] + picked = picked.at[rows, chosen].set( + jnp.where(append_step, True, old_picked) + ) + perm = perm.at[:, t].set(chosen) + return ( + perm, + picked, + prefix_raw, + prefix_order_raw, + suffix_raw, + virt_prefix_order_raw, + virt_count, + prefix_ids, + k_cache, + v_cache, + hidden_cache, + base_cache, + last_hidden, + total_logp, + ), None + + init = ( + perm0, + picked0, + repeat(prefix_raw0), + repeat(prefix_order_raw0), + repeat(suffix_raw0), + repeat(virt_prefix_order_raw0), + repeat(virt_count0), + prefix_ids0, + k_cache0, + v_cache0, + hidden_cache0, + base_cache0, + last_hidden0, + logp0, + ) + final, _ = jax.lax.scan(body, init, idx) + ( + perm, + _picked, + _prefix_raw, + _prefix_order_raw, + _suffix_raw, + _virt_po, + _virt_cnt, + _prefix_ids, + _k, + _v, + _hidden_cache, + _base_cache, + _hidden, + logp, + ) = final + return perm, logp + + +__all__ = ["TreePrefixPointerMHSEA"] diff --git a/src/hamiltonzero/model/route_quotient.py b/src/hamiltonzero/model/route_quotient.py new file mode 100644 index 0000000000000000000000000000000000000000..97542d30d158bd1897638b1bd439f1b9854aaece --- /dev/null +++ b/src/hamiltonzero/model/route_quotient.py @@ -0,0 +1,355 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import re + +import jax +import jax.numpy as jnp +import numpy as np +from jaxtyping import Array, Float, Int + + +def _edge_relation_tags(mask: Array, bmask: Array) -> Int[Array, "n n"]: + n = mask.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + ii = idx[:, None] + jj = idx[None, :] + real_i = mask.astype(bool)[:, None] + real_j = mask.astype(bool)[None, :] + ctx_i = bmask.astype(bool)[:, None] + ctx_j = bmask.astype(bool)[None, :] + ctx_pair = ctx_i & ctx_j + + real_real = real_i & real_j + empty_i = (~real_i) & ctx_i + empty_j = (~real_j) & ctx_j + self_pair = ii == jj + + tag = jnp.zeros((n, n), dtype=jnp.int32) + tag = jnp.where(real_real & self_pair, 1, tag) + tag = jnp.where(empty_i & empty_j & self_pair, 2, tag) + tag = jnp.where(empty_i & empty_j & (~self_pair), 3, tag) + tag = jnp.where(empty_i & real_j, 4, tag) + tag = jnp.where(real_i & empty_j, 5, tag) + return jnp.where(ctx_pair, tag, jnp.asarray(-1, dtype=jnp.int32)) + + +def _quantized_hash(x: Array, coeff: Array, tol: float) -> Int[Array, "..."]: + q = jnp.rint( + jnp.real(x).astype(jnp.float32) / jnp.asarray(tol, dtype=jnp.float32) + ).astype(jnp.int32) + return jnp.sum(q * coeff.astype(jnp.int32), axis=-1).astype(jnp.int32) + + +def _canonical_hermitian_j(J_full: Float[Array, "n n 3 3"]) -> Float[Array, "n n 3 3"]: + return 0.5 * (J_full + jnp.transpose(J_full, (1, 0, 3, 2))) + + +def _j_pair_hash( + J_full: Float[Array, "n n 3 3"], + *, + tol: float, +) -> Int[Array, "n n"]: + n = J_full.shape[0] + coeff = jnp.asarray( + [ + 1_000_003, + 1_009_003, + 1_021_009, + 1_033_013, + 1_049_009, + 1_061_009, + 1_073_003, + 1_087_009, + 1_093_013, + 1_109_009, + 1_117_001, + 1_123_003, + 1_129_009, + 1_151_003, + 1_159_013, + 1_171_009, + 1_181_003, + 1_187_009, + ], + dtype=jnp.int32, + ) + J_key = _canonical_hermitian_j(J_full) + pair = jnp.concatenate( + [ + J_key.reshape(n, n, 9), + jnp.transpose(J_key, (1, 0, 2, 3)).reshape(n, n, 9), + ], + axis=-1, + ) + return _quantized_hash(pair, coeff, tol) + + +def _j_diag_hash( + J_full: Float[Array, "n n 3 3"], + *, + tol: float, +) -> Int[Array, "n"]: + + n = J_full.shape[0] + coeff = jnp.asarray( + [ + 1_000_003, + 1_009_003, + 1_021_009, + 1_033_013, + 1_049_009, + 1_061_009, + 1_073_003, + 1_087_009, + 1_093_013, + ], + dtype=jnp.int32, + ) + J_key = _canonical_hermitian_j(J_full) + idx = jnp.arange(n) + diag = J_key[idx, idx].reshape(n, 9) + return _quantized_hash(diag, coeff, tol) + + +def route_quotient_keys( + J_full: Float[Array, "n n 3 3"], + h: Float[Array, "n 3"], + mask: Int[Array, "n"] | Array, + bmask: Int[Array, "n"] | Array, + *, + tol: float = 1e-6, +) -> tuple[Int[Array, "n"], Int[Array, "n n"]]: + + real = mask.astype(bool) + context = bmask.astype(bool) + h_coeff = jnp.asarray([1_000_003, 1_009_003, 1_021_009], dtype=jnp.int32) + h_key = _quantized_hash(h, h_coeff, tol) + + node_key = jnp.where( + real, + h_key + + _j_diag_hash(J_full, tol=tol) * jnp.asarray(131_063, dtype=jnp.int32) + + jnp.asarray(17_071, dtype=jnp.int32), + jnp.asarray(-313_037, dtype=jnp.int32), + ) + node_key = jnp.where(context, node_key, jnp.asarray(-1, dtype=jnp.int32)) + + edge_hash = _j_pair_hash(J_full, tol=tol) + tags = _edge_relation_tags(mask, bmask) + edge_hash = jnp.where(tags == 1, jnp.asarray(0, dtype=jnp.int32), edge_hash) + edge_key = edge_hash + tags * jnp.asarray(131_071, dtype=jnp.int32) + edge_key = jnp.where(context[:, None] & context[None, :], edge_key, 0) + return node_key, edge_key + + +def conditional_orbit_ids_from_keys( + node_key: Int[Array, "n"], + edge_key: Int[Array, "n n"], + valid_mask: Int[Array, "n"] | Array, + context_mask: Int[Array, "n"] | Array, + prefix_ids: Int[Array, "n"], + prefix_len, + *, + max_rounds: int | None = None, +) -> Int[Array, "n"]: + + n = node_key.shape[0] + if max_rounds is None: + max_rounds = n + idx = jnp.arange(n, dtype=jnp.int32) + prefix_len_i = jnp.asarray(prefix_len, dtype=jnp.int32) + context = context_mask.astype(bool) + valid = valid_mask.astype(bool) & context + + prefix_active = (idx < prefix_len_i) & context[prefix_ids] + prefix_pos = jnp.max( + jnp.where( + prefix_active[:, None] & (prefix_ids[:, None] == idx[None, :]), + idx[:, None], + jnp.asarray(-1, dtype=jnp.int32), + ), + axis=0, + ) + is_prefix = prefix_pos >= 0 + valid = valid & (~is_prefix) + prefix_color = jnp.asarray( + 2_000_000_000, dtype=jnp.int32 + ) - prefix_pos * jnp.asarray(1_000_003, dtype=jnp.int32) + colors0 = jnp.where(is_prefix, prefix_color, node_key) + colors0 = jnp.where(context, colors0, jnp.asarray(-1, dtype=jnp.int32)) + big = jnp.asarray(n + 1, dtype=jnp.int32) + + def body(colors, _): + pair_key = edge_key + colors[None, :] * jnp.asarray(1_310_719, dtype=jnp.int32) + pair_key = jnp.where( + context[None, :], pair_key, jnp.asarray(0, dtype=jnp.int32) + ) + sorted_keys = jnp.sort(pair_key, axis=1) + same_sig = ( + context[:, None] + & context[None, :] + & (colors[:, None] == colors[None, :]) + & jnp.all(sorted_keys[:, None, :] == sorted_keys[None, :, :], axis=-1) + ) + new_colors = jnp.min(jnp.where(same_sig, idx[None, :], big), axis=1) + new_colors = jnp.where(is_prefix, prefix_color, new_colors) + return jnp.where(context, new_colors, jnp.asarray(-1, dtype=jnp.int32)), None + + colors, _ = jax.lax.scan(body, colors0, None, length=int(max_rounds)) + same_valid = valid[:, None] & valid[None, :] & (colors[:, None] == colors[None, :]) + reps = jnp.min(jnp.where(same_valid, idx[None, :], big), axis=1) + return jnp.where(valid, reps, jnp.asarray(-1, dtype=jnp.int32)) + + +def conditional_orbit_pair_ids_from_keys( + node_key: Int[Array, "n"], + edge_key: Int[Array, "n n"], + valid_mask: Int[Array, "n"] | Array, + context_mask: Int[Array, "n"] | Array, + prefix_ids: Int[Array, "n"], + prefix_len, +) -> Int[Array, "n"]: + + n = int(node_key.shape[0]) + idx = jnp.arange(n, dtype=jnp.int32) + pidx = jnp.arange(n * n, dtype=jnp.int32) + nn = jnp.asarray(n * n, dtype=jnp.int32) + big = jnp.asarray(n * n, dtype=jnp.int32) + prefix_len_i = jnp.asarray(prefix_len, dtype=jnp.int32) + context = context_mask.astype(bool) + valid = valid_mask.astype(bool) & context + prefix_active = (idx < prefix_len_i) & context[prefix_ids] + prefix_pos = jnp.max( + jnp.where( + prefix_active[:, None] & (prefix_ids[:, None] == idx[None, :]), + idx[:, None], + jnp.asarray(-1, dtype=jnp.int32), + ), + axis=0, + ) + is_prefix = prefix_pos >= 0 + valid = valid & (~is_prefix) + prefix_color = jnp.asarray( + 2_000_000_000, dtype=jnp.int32 + ) - prefix_pos * jnp.asarray(1_000_003, dtype=jnp.int32) + node_colors = jnp.where(is_prefix, prefix_color, node_key) + node_colors = jnp.where(context, node_colors, jnp.asarray(-1, dtype=jnp.int32)) + cc = (context[:, None] & context[None, :]).reshape(-1) + + def _canon(components): + + same = cc[:, None] & cc[None, :] + for s in components: + sf = s.reshape(-1) + same = same & (sf[:, None] == sf[None, :]) + reps = jnp.min(jnp.where(same, pidx[None, :], big), axis=1) + return jnp.where(cc, reps, jnp.asarray(-1, dtype=jnp.int32)).reshape(n, n) + + ni = jnp.broadcast_to(node_colors[:, None], (n, n)) + nj = jnp.broadcast_to(node_colors[None, :], (n, n)) + pc0 = _canon([ni, nj, edge_key, jnp.transpose(edge_key)]) + + base_a = jnp.asarray(1_000_003, dtype=jnp.int32) + base_b = jnp.asarray(1_300_021, dtype=jnp.int32) + if n > 1: + pow_a = jnp.concatenate( + [ + jnp.ones((1,), jnp.int32), + jnp.cumprod(jnp.full((n - 1,), base_a, jnp.int32)), + ] + ) + pow_b = jnp.concatenate( + [ + jnp.ones((1,), jnp.int32), + jnp.cumprod(jnp.full((n - 1,), base_b, jnp.int32)), + ] + ) + else: + pow_a = jnp.ones((1,), jnp.int32) + pow_b = jnp.ones((1,), jnp.int32) + ctx_k = context[None, None, :] + off = jnp.asarray(7, dtype=jnp.int32) + neutral = jnp.asarray(-1, dtype=jnp.int32) + + def _body(pc, _): + a = pc[:, None, :] + b = jnp.transpose(pc)[None, :, :] + code = jnp.where(ctx_k, a * nn + b, neutral) + sc = jnp.sort(code, axis=-1) + + ph_a = jnp.sum((sc + off) * pow_a[None, None, :], axis=-1) + ph_b = jnp.sum((sc + off) * pow_b[None, None, :], axis=-1) + return _canon([pc, ph_a, ph_b]), None + + pc, _ = jax.lax.scan(_body, pc0, None, length=n) + diag = jnp.diagonal(pc) + same_valid = valid[:, None] & valid[None, :] & (diag[:, None] == diag[None, :]) + reps = jnp.min(jnp.where(same_valid, idx[None, :], big), axis=1) + return jnp.where(valid, reps, jnp.asarray(-1, dtype=jnp.int32)) + + +_WL1_BREAKER_PATTERN = re.compile(r"wl1|srg|paley|shrikhande|rook", re.IGNORECASE) +_WL1_BREAKER_CATEGORY = "13_wl1_breaking" + + +def _tag_forces_fwl2( + *, + category: str | None = None, + tag: str | None = None, + topology_class: str | None = None, + j_class: str | None = None, +) -> bool: + if category is not None and str(category) == _WL1_BREAKER_CATEGORY: + return True + for value in (tag, topology_class, j_class): + if value is not None and _WL1_BREAKER_PATTERN.search(str(value)): + return True + return False + + +def system_needs_fwl2( + J_full, + h, + n_spins: int, + *, + category: str | None = None, + tag: str | None = None, + topology_class: str | None = None, + j_class: str | None = None, + tol: float = 1e-6, +) -> bool: + if _tag_forces_fwl2( + category=category, + tag=tag, + topology_class=topology_class, + j_class=j_class, + ): + return True + + n = int(n_spins) + J_arr = jnp.asarray(np.asarray(J_full, dtype=np.float64)) + h_arr = jnp.asarray(np.asarray(h, dtype=np.float64)) + mask = jnp.ones((n,), dtype=jnp.int32) + bmask = jnp.ones((n,), dtype=jnp.int32) + idx = jnp.arange(n, dtype=jnp.int32) + prefix_len = jnp.int32(0) + + node_key, edge_key = route_quotient_keys(J_arr, h_arr, mask, bmask, tol=tol) + wl1 = conditional_orbit_ids_from_keys( + node_key, edge_key, mask, bmask, idx, prefix_len + ) + fwl2 = conditional_orbit_pair_ids_from_keys( + node_key, edge_key, mask, bmask, idx, prefix_len + ) + return bool(not jnp.array_equal(wl1, fwl2)) + + +__all__ = [ + "conditional_orbit_ids_from_keys", + "conditional_orbit_pair_ids_from_keys", + "route_quotient_keys", + "system_needs_fwl2", +] diff --git a/src/hamiltonzero/model/tree.py b/src/hamiltonzero/model/tree.py new file mode 100644 index 0000000000000000000000000000000000000000..29448b44e87f95233ccdfcbd8dc241a59f4fd929 --- /dev/null +++ b/src/hamiltonzero/model/tree.py @@ -0,0 +1,2387 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +from functools import partial +import equinox as eqx +import jax +import jax.numpy as jnp +import kfac_jax +from jaxtyping import Array, Float, Int, PRNGKeyArray + + +def _kfac_name_kw(tag_id: str) -> dict: + return {"name": tag_id} if tag_id else {} + + +_TREE_DYADIC_CLOCK_BASE = 10000.0 + + +def _tree_coord_clock( + pos, width: int, dtype, *, base: float = _TREE_DYADIC_CLOCK_BASE, scale=None +): + if int(width) <= 0: + pos_arr = jnp.asarray(pos) + return jnp.zeros(pos_arr.shape + (0,), dtype=dtype) + pos_f = jnp.asarray(pos, dtype=jnp.float32) + if scale is not None: + denom = jnp.maximum( + jnp.asarray(scale, dtype=jnp.float32) - jnp.asarray(1.0, dtype=jnp.float32), + jnp.asarray(1.0, dtype=jnp.float32), + ) + pos_f = pos_f / denom + half = (int(width) + 1) // 2 + band = jnp.arange(half, dtype=jnp.float32) + base_f = jnp.maximum( + jnp.asarray(base, dtype=jnp.float32), jnp.asarray(2.0, dtype=jnp.float32) + ) + inv_freq = jnp.exp( + -jnp.log(base_f) * band / jnp.asarray(max(half, 1), dtype=jnp.float32) + ) + phase = pos_f[..., None] * inv_freq + enc = jnp.concatenate([jnp.sin(phase), jnp.cos(phase)], axis=-1) + return enc[..., : int(width)].astype(dtype) + + +def _tree_level_clock(pos, width: int, max_depth, dtype): + del max_depth + return _tree_coord_clock(pos, width, dtype, base=_TREE_DYADIC_CLOCK_BASE) + + +def _tree_clock_root_center_from_depth(depth, dtype): + depth_f = jnp.maximum(jnp.asarray(depth, dtype=jnp.float32), 0.0) + span = jnp.power(jnp.asarray(2.0, dtype=jnp.float32), depth_f) + return (span - jnp.asarray(1.0, dtype=jnp.float32)) * 0.5 + + +def _tree_dyadic_segment_clock( + level_idx, pair_idx, width: int, dtype, *, root_center=None +): + if int(width) <= 0: + level_arr = jnp.asarray(level_idx) + if pair_idx is None: + return jnp.zeros(level_arr.shape + (0,), dtype=dtype) + pair_arr = jnp.asarray(pair_idx) + return jnp.zeros(pair_arr.shape + (0,), dtype=dtype) + center_width = max(1, int(width) - 2) + scale_width = int(width) - center_width + level_f = jnp.asarray(level_idx, dtype=jnp.float32) + span = jnp.power( + jnp.asarray(2.0, dtype=jnp.float32), + level_f + jnp.asarray(1.0, dtype=jnp.float32), + ) + scale_pos = span + if pair_idx is None: + center_pos = (span - jnp.asarray(1.0, dtype=jnp.float32)) * 0.5 + else: + pair_f = jnp.asarray(pair_idx, dtype=jnp.float32) + center_pos = pair_f * span + jnp.asarray(0.5, dtype=jnp.float32) * ( + span - jnp.asarray(1.0, dtype=jnp.float32) + ) + if root_center is not None: + center_pos = center_pos - jnp.asarray(root_center, dtype=jnp.float32) + scale_pos = jnp.broadcast_to(scale_pos, center_pos.shape) + center_clock = _tree_coord_clock( + center_pos, center_width, dtype, base=_TREE_DYADIC_CLOCK_BASE + ) + scale_clock = _tree_coord_clock( + scale_pos, scale_width, dtype, base=_TREE_DYADIC_CLOCK_BASE + ) + return jnp.concatenate([center_clock, scale_clock], axis=-1) + + +def _tree_merge_clock(level_idx, pair_idx, pair_base, width: int, max_depth, dtype): + del pair_base + root_center = _tree_clock_root_center_from_depth(max_depth, dtype) + return _tree_dyadic_segment_clock( + level_idx, pair_idx, width, dtype, root_center=root_center + ) + + +_TREE_NGPT_DEPTH_FEAT_DIM = 32 + + +def _tree_sphere(x, axis=-1): + ms = jnp.mean(jnp.square(x), axis=axis, keepdims=True) + return x * jax.lax.rsqrt(jnp.maximum(ms, 0.0001)) + + +def _tree_depth_count_features(cnt_a, cnt_b, n_total, level, n_levels, dtype): + a = cnt_a.astype(jnp.float32) + b = cnt_b.astype(jnp.float32) + nt = jnp.asarray(n_total, jnp.float32) + la = jnp.log2(1.0 + a) + lb = jnp.log2(1.0 + b) + rem = jnp.log2(1.0 + jnp.maximum(nt - a - b, 0.0)) + counts = jnp.stack([la, lb, rem], axis=-1) + omg_c = jnp.pi / 2.0 ** jnp.arange(4, dtype=jnp.float32) + ang_c = counts[..., :, None] * omg_c + f_c = jnp.concatenate([jnp.sin(ang_c), jnp.cos(ang_c)], axis=-1) + f_c = f_c.reshape(f_c.shape[:-2] + (24,)) + lv = jnp.asarray(level, jnp.float32) + lv_rem = jnp.asarray(n_levels, jnp.float32) - 1.0 - lv + levels = jnp.stack( + [jnp.broadcast_to(lv, a.shape), jnp.broadcast_to(lv_rem, a.shape)], axis=-1 + ) + omg_l = jnp.pi / 2.0 ** jnp.arange(2, dtype=jnp.float32) + ang_l = levels[..., :, None] * omg_l + f_l = jnp.concatenate([jnp.sin(ang_l), jnp.cos(ang_l)], axis=-1) + f_l = f_l.reshape(f_l.shape[:-2] + (8,)) + return jnp.concatenate([f_c, f_l], axis=-1).astype(dtype) + + +def _tree_ngpt_level_counts(m0, n_pairs, n_levels, dtype, *, feature_n_levels=None): + cnt = m0.astype(jnp.float32) + n_total = jnp.sum(cnt) + if feature_n_levels is None: + feature_n_levels = _tree_active_clock_depth(m0) + feats = [] + for lv in range(n_levels): + pairs = cnt.reshape(n_pairs, 2) + feats.append( + _tree_depth_count_features( + pairs[:, 0], pairs[:, 1], n_total, lv, feature_n_levels, dtype + ) + ) + parents = pairs.sum(axis=1) + cnt = jnp.concatenate([parents, jnp.zeros_like(parents)], axis=0) + return jnp.stack(feats, axis=0) + + +from .odd_ops import BiasFreeLinear, HypernetMatrix, Linear, MLP, _RMS +from .readout_leaf_context import lca_alibi_bias, lca_fixed_slopes +from .fused_silu import fused_silu + + +def _replace_square_row_column(matrix, index, row_value, column_value): + idx = jnp.arange(matrix.shape[0], dtype=jnp.int32) + select = idx == jnp.asarray(index, dtype=jnp.int32) + diagonal = row_value[index] + column_value = jnp.where(select[:, None], diagonal, column_value) + matrix = jnp.where(select[:, None, None], row_value[None, :, :], matrix) + return jnp.where(select[None, :, None], column_value[:, None, :], matrix) + + +def _quadrilinear_merge( + T, + u_a, + u_b, + *, + tag_id: str = "", + pathway: str | None = None, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, +): + from ._custom_lap_primitives import custom_lap_active, quadrilinear_merge_p + + if custom_lap_active(): + return quadrilinear_merge_p.bind(T, u_a, u_b) + G, d_r = (T.shape[0], T.shape[1]) + _odt = jnp.float32 + T_param = T + T = T if T.dtype == _odt else T.astype(_odt) + u_a = u_a if u_a.dtype == _odt else u_a.astype(_odt) + u_b = u_b if u_b.dtype == _odt else u_b.astype(_odt) + leading = u_a.shape[:-1] + u_a_2d = u_a.reshape(*leading, G, d_r) + u_b_2d = u_b.reshape(*leading, G, d_r) + Tu_a = jnp.einsum("ijkl,...ik->...ijl", T, u_a_2d) + y_2d = jnp.einsum("...ijl,...il->...ij", Tu_a, u_b_2d) + y = y_2d.reshape(*leading, G * d_r) + if kfac_structural_mask is None: + return y + from hamiltonzero.optim.spin_blocks import register_structural_quadrilinear_merge + + return register_structural_quadrilinear_merge( + y, + u_a, + u_b, + T_param, + kfac_structural_mask, + scan_shared=kfac_scan_shared, + repeat_ndim=kfac_repeat_ndim, + **_kfac_name_kw(tag_id), + ) + + +def _rownorm_cols(weight): + nsq = jnp.sum(jnp.square(weight), axis=0, keepdims=True) + return weight * jax.lax.rsqrt(jnp.maximum(nsq, 0.0001)) + + +def _tagged_dense( + weight, + bias, + x, + *, + tag_id: str = "", + pathway: str, + weight_eff=None, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + cdtype = _compute_dtype() if pathway in ("even", "hypernet_eside") else jnp.float32 + w_src = weight if weight_eff is None else weight_eff + w_compute = w_src.astype(cdtype) if w_src.dtype != cdtype else w_src + b_compute = bias.astype(cdtype) if bias.dtype != cdtype else bias + x_compute = x.astype(cdtype) if x.dtype != cdtype else x + y = x_compute @ w_compute + b_compute + if kfac_structural_mask is not None: + from hamiltonzero.optim.blocks import register_structural_dense + + return register_structural_dense( + y, + x, + kfac_structural_mask, + weight, + bias, + scan_shared=kfac_scan_shared, + repeat_ndim=kfac_repeat_ndim, + context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + **_kfac_name_kw(tag_id), + ) + return kfac_jax.register_dense(y, x, weight, bias, **_kfac_name_kw(tag_id)) + + +def _tagged_dense_no_bias( + weight, + x, + *, + tag_id: str = "", + pathway: str, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + cdtype = _compute_dtype() if pathway in ("even", "hypernet_eside") else jnp.float32 + w_compute = weight.astype(cdtype) if weight.dtype != cdtype else weight + x_compute = x.astype(cdtype) if x.dtype != cdtype else x + y = x_compute @ w_compute + if kfac_structural_mask is not None: + from hamiltonzero.optim.blocks import register_structural_dense + + return register_structural_dense( + y, + x, + kfac_structural_mask, + weight, + scan_shared=kfac_scan_shared, + repeat_ndim=kfac_repeat_ndim, + context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + **_kfac_name_kw(tag_id), + ) + return kfac_jax.register_dense(y, x, weight, **_kfac_name_kw(tag_id)) + + +def _tagged_ln_eqx_style( + scale, + shift, + x, + eps: float = 1e-05, + *, + tag_id: str = "", + pathway: str, + var_floor: float | None = None, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + out_cdtype = ( + _compute_dtype() if pathway in ("even", "hypernet_eside") else jnp.float32 + ) + in_dtype = x.dtype + stats_dtype = jnp.promote_types(jnp.float32, in_dtype) + x_hi = x.astype(stats_dtype) if in_dtype != stats_dtype else x + mean = jnp.mean(x_hi, axis=-1, keepdims=True) + centered_hi = x_hi - mean + var = jnp.mean(centered_hi * centered_hi, axis=-1, keepdims=True) + if var_floor is not None: + var = jnp.maximum(var, var_floor) + normalized_hi = centered_hi * jax.lax.rsqrt(var + eps) + normalized = ( + normalized_hi.astype(out_cdtype) + if normalized_hi.dtype != out_cdtype + else normalized_hi + ) + scale_compute = scale.astype(out_cdtype) if scale.dtype != out_cdtype else scale + shift_compute = shift.astype(out_cdtype) if shift.dtype != out_cdtype else shift + y = normalized * scale_compute + shift_compute + if kfac_structural_mask is not None: + from hamiltonzero.optim.blocks import register_structural_scale_and_shift + + return register_structural_scale_and_shift( + y, + normalized, + kfac_structural_mask, + scale=scale, + shift=shift, + scan_shared=kfac_scan_shared, + repeat_ndim=kfac_repeat_ndim, + context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + **_kfac_name_kw(tag_id), + ) + return kfac_jax.register_scale_and_shift( + y, normalized, scale, shift, **_kfac_name_kw(tag_id) + ) + + +def _tagged_rms_eqx_style( + scale, + x, + eps: float = 1e-05, + *, + tag_id: str = "", + pathway: str, + var_floor: float | None = None, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + out_cdtype = ( + _compute_dtype() if pathway in ("even", "hypernet_eside") else jnp.float32 + ) + in_dtype = x.dtype + _stats_dtype = jnp.promote_types(jnp.float32, in_dtype) + x_hi = x.astype(_stats_dtype) if in_dtype != _stats_dtype else x + mean_sq = jnp.mean(x_hi * x_hi, axis=-1, keepdims=True) + if var_floor is not None: + rsqrt_hi = jax.lax.rsqrt(jnp.maximum(mean_sq, var_floor)) + else: + rsqrt_hi = jax.lax.rsqrt(mean_sq + eps) + normalized_hi = x_hi * rsqrt_hi + normalized = ( + normalized_hi.astype(out_cdtype) + if normalized_hi.dtype != out_cdtype + else normalized_hi + ) + rsqrt = rsqrt_hi.astype(out_cdtype) if rsqrt_hi.dtype != out_cdtype else rsqrt_hi + s_compute = scale.astype(out_cdtype) if scale.dtype != out_cdtype else scale + x_compute = x.astype(out_cdtype) if x.dtype != out_cdtype else x + inv = s_compute * rsqrt + y = x_compute * inv + if kfac_structural_mask is not None: + from hamiltonzero.optim.blocks import register_structural_scale_and_shift + + return register_structural_scale_and_shift( + y, + normalized, + kfac_structural_mask, + scale=scale, + scan_shared=kfac_scan_shared, + repeat_ndim=kfac_repeat_ndim, + context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + **_kfac_name_kw(tag_id), + ) + return kfac_jax.register_scale_and_shift( + y, normalized, scale=s_compute, shift=None, **_kfac_name_kw(tag_id) + ) + + +def _tagged_lerp_alpha( + alpha, + d, + *, + tag_id: str = "", + pathway: str, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + out_cdtype = ( + _compute_dtype() if pathway in ("even", "hypernet_eside") else jnp.float32 + ) + d_compute = d.astype(out_cdtype) if d.dtype != out_cdtype else d + a_compute = alpha.astype(out_cdtype) if alpha.dtype != out_cdtype else alpha + y = d_compute * a_compute + if kfac_structural_mask is not None: + from hamiltonzero.optim.blocks import register_structural_scale_and_shift + + return register_structural_scale_and_shift( + y, + d_compute, + kfac_structural_mask, + scale=alpha, + scan_shared=kfac_scan_shared, + repeat_ndim=kfac_repeat_ndim, + context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + **_kfac_name_kw(tag_id), + ) + return kfac_jax.register_scale_and_shift( + y, d_compute, scale=alpha, shift=None, **_kfac_name_kw(tag_id) + ) + + +def _tagged_bounded_ngpt_gain( + alpha, + like, + *, + max_gain: float = 0.5, + tag_id: str = "", + pathway: str, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + ones = jax.lax.stop_gradient(like) * jnp.asarray( + 0.0, dtype=like.dtype + ) + jnp.asarray(1.0, dtype=like.dtype) + tagged_alpha = _tagged_lerp_alpha( + alpha, + ones, + tag_id=tag_id, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + return jnp.asarray(max_gain, dtype=tagged_alpha.dtype) * jax.nn.sigmoid( + tagged_alpha + ) + + +def _tree_ngpt_residual( + skip, + proposal, + alpha, + *, + max_gain: float, + tag_id: str, + pathway: str = "even", + update_mask=None, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + skip_n = _tree_sphere(skip) + proposal_n = _tree_sphere(proposal) + direction = proposal_n - skip_n + gain = _tagged_bounded_ngpt_gain( + alpha, + direction, + max_gain=max_gain, + tag_id=tag_id, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + updated = _tree_sphere(skip_n + gain * direction) + if update_mask is None: + return updated + active = update_mask.astype(bool) + while active.ndim < updated.ndim: + active = active[..., None] + return jnp.where(active, updated, skip) + + +def _inline_norm_forward( + nrm, + x, + *, + pathway: str, + tag_id=None, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 0, + kfac_context_primal_reused_over_walkers: bool = False, +): + tid = nrm._use_id if tag_id is None else tag_id + return _tagged_rms_eqx_style( + nrm.weight, + x, + eps=nrm.eps, + tag_id=tid, + pathway=pathway, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + kfac_context_primal_reused_over_walkers=kfac_context_primal_reused_over_walkers, + ) + + +class LeafBuilder(eqx.Module): + P_c: Linear + P_u: HypernetMatrix + + def __init__( + self, + d_e: int, + d_o: int, + d_c: int, + d_r: int, + rank: int, + *, + key: PRNGKeyArray, + d_g: int, + leaf_hypernet_rank: int | None = None, + d_m_merge: int | None = None, + ): + keys = jax.random.split(key, 5) + ctx_dim = d_e + d_g + p_u_rank = leaf_hypernet_rank if leaf_hypernet_rank is not None else rank + d_m_eff = d_m_merge if d_m_merge is not None else d_r + if d_m_eff % d_r != 0: + raise ValueError( + f"LeafBuilder: d_m_merge={d_m_merge} must be divisible by d_r={d_r} (HT carrier reshape requires G = d_m_eff // d_r)." + ) + self.P_c = Linear(d_e, d_c, key=keys[0]) + self.P_u = HypernetMatrix(d_o, d_m_eff, ctx_dim, p_u_rank, key=keys[1]) + + def conditioner_context( + self, e: Float[Array, "n d_e"], g_emb: Float[Array, "d_g"] + ) -> Float[Array, "n d_ctx"]: + n = e.shape[0] + g_emb_b = jnp.broadcast_to(g_emb[None, :], (n, g_emb.shape[0])) + return jnp.concatenate([e, g_emb_b], axis=-1) + + def __call__(self, e, z, *, g_emb, kfac_structural_mask, kfac_odd_structural_mask): + n = e.shape[0] + ctx = self.conditioner_context(e, g_emb) + c = _tagged_dense( + self.P_c.weight, + self.P_c.bias, + e, + tag_id=self.P_c._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + ) + u = self.P_u.apply( + ctx, + z, + e_pathway="even", + kfac_structural_mask=kfac_odd_structural_mask, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + return (_tree_sphere(c), u, jnp.zeros((n,), dtype=jnp.float32)) + + +class EdgeMergeOp(eqx.Module): + mlp: MLP + node_ctx_proj: BiasFreeLinear | None + alpha: Float[Array, "d_edge"] + ngpt_alpha_max: float = eqx.field(static=True, default=0.5) + + def __init__( + self, + d_edge: int, + d_c: int, + *, + key: PRNGKeyArray, + alpha_init: float, + alpha_max: float, + d_hidden: int | None = None, + n_blocks: int = 2, + edge_node_ctx_dim: int | None = None, + ): + node_ctx_dim = int(d_c) if edge_node_ctx_dim is None else int(edge_node_ctx_dim) + if node_ctx_dim < 1: + raise ValueError( + f"tree edge_node_ctx_dim must be positive or None, got {edge_node_ctx_dim}" + ) + self.node_ctx_proj = ( + None + if node_ctx_dim == int(d_c) + else BiasFreeLinear(d_c, node_ctx_dim, key=jax.random.fold_in(key, 60782)) + ) + d_in = 4 * d_edge + 4 * node_ctx_dim + d_hidden_eff = d_hidden if d_hidden is not None else max(d_edge * 2, 64) + self.mlp = MLP(d_in, d_hidden_eff, d_edge, key=key, n_blocks=n_blocks) + self.ngpt_alpha_max = float(alpha_max) + self.alpha = float(alpha_init) * jnp.ones((int(d_edge),)) + + def apply_skip( + self, + skip, + proposal, + *, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + ): + return _tree_ngpt_residual( + skip, + proposal, + self.alpha, + max_gain=self.ngpt_alpha_max, + tag_id="", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + + def __call__( + self, + e_2i_2j: Float[Array, "d_edge"], + e_2i_2j1: Float[Array, "d_edge"], + e_2i1_2j: Float[Array, "d_edge"], + e_2i1_2j1: Float[Array, "d_edge"], + c_2i: Float[Array, "d_c"], + c_2i1: Float[Array, "d_c"], + c_2j: Float[Array, "d_c"], + c_2j1: Float[Array, "d_c"], + *, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + ) -> Float[Array, "d_edge"]: + child_ctx = jnp.stack([c_2i, c_2i1, c_2j, c_2j1], axis=0) + if self.node_ctx_proj is not None: + child_structural_mask = ( + None + if kfac_structural_mask is None + else jnp.broadcast_to(kfac_structural_mask, (4,)) + ) + child_ctx = _tagged_dense_no_bias( + self.node_ctx_proj.weight, + child_ctx, + tag_id=self.node_ctx_proj._use_id, + pathway="even", + kfac_structural_mask=child_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=3, + ) + c_2i, c_2i1, c_2j, c_2j1 = child_ctx + mlp_in = jnp.concatenate( + [e_2i_2j, e_2i_2j1, e_2i1_2j, e_2i1_2j1, c_2i, c_2i1, c_2j, c_2j1] + ) + mlp = self.mlp + x = _tagged_dense( + mlp.in_proj.weight, + mlp.in_proj.bias, + mlp_in, + tag_id=mlp.in_proj._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + for nrm, l1, l2 in zip(mlp.block_norms, mlp.block_l1s, mlp.block_l2s): + normed = _inline_norm_forward( + nrm, + x, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + inner = _tagged_dense( + l1.weight, + l1.bias, + normed, + tag_id=l1._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + inner_act = mlp._act(inner) + inner_out = _tagged_dense( + l2.weight, + l2.bias, + inner_act, + tag_id=l2._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + x = x + mlp.inner_gain * inner_out + out_normed = _inline_norm_forward( + mlp.out_norm, + x, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + return _tagged_dense( + mlp.out_proj.weight, + mlp.out_proj.bias, + out_normed, + tag_id=mlp.out_proj._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + + +class EdgeFWLUpdate(eqx.Module): + ln_edge: _RMS + ln_c: _RMS + node_ctx_proj: BiasFreeLinear | None + psi_L_in: Linear + psi_L_out: Linear + psi_R_in: Linear + psi_R_out: Linear + ln_path: _RMS + ffn_in: Linear + ffn_out: Linear + alpha: Float[Array, "d_edge"] + ngpt_alpha_max: float = eqx.field(static=True, default=0.5) + + def __init__( + self, + d_c: int, + d_edge: int, + *, + key: PRNGKeyArray, + alpha_init: float, + alpha_max: float, + channels: int = 64, + edge_node_ctx_dim: int | None = None, + ): + node_ctx_dim = int(d_c) if edge_node_ctx_dim is None else int(edge_node_ctx_dim) + if node_ctx_dim < 1: + raise ValueError( + f"tree edge_node_ctx_dim must be positive or None, got {edge_node_ctx_dim}" + ) + d_pair = d_edge + 2 * node_ctx_dim + d_psi_hidden = 2 * channels + d_ffn_hidden = max(d_edge, 2 * channels) + ( + k_psi_L_in, + k_psi_L_out, + k_psi_L_gate, + k_psi_R_in, + k_psi_R_out, + k_psi_R_gate, + k_ffn_in, + k_ffn_out, + ) = jax.random.split(key, 8) + self.ln_edge = _RMS(d_edge) + self.ln_c = _RMS(d_c) + self.node_ctx_proj = ( + None + if node_ctx_dim == int(d_c) + else BiasFreeLinear(d_c, node_ctx_dim, key=jax.random.fold_in(key, 63262)) + ) + self.psi_L_in = Linear(d_pair, d_psi_hidden, key=k_psi_L_in) + self.psi_L_out = Linear(d_psi_hidden, channels, key=k_psi_L_out) + self.psi_R_in = Linear(d_pair, d_psi_hidden, key=k_psi_R_in) + self.psi_R_out = Linear(d_psi_hidden, channels, key=k_psi_R_out) + self.ln_path = _RMS(channels) + self.ffn_in = Linear(d_pair + channels, d_ffn_hidden, key=k_ffn_in) + self.ffn_out = Linear(d_ffn_hidden, d_edge, key=k_ffn_out) + self.ngpt_alpha_max = float(alpha_max) + self.alpha = float(alpha_init) * jnp.ones((int(d_edge),)) + + def _psi_apply( + self, + pair_ij: Float[Array, "n n d_pair"], + which: str, + *, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + kfac_repeat_ndim: int = 2, + ) -> Float[Array, "n n C"]: + if which == "L": + l_in, l_out = (self.psi_L_in, self.psi_L_out) + else: + l_in, l_out = (self.psi_R_in, self.psi_R_out) + hidden = fused_silu( + _tagged_dense( + l_in.weight, + l_in.bias, + pair_ij, + tag_id=l_in._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + ) + ) + return _tagged_dense( + l_out.weight, + l_out.bias, + hidden, + tag_id=l_out._use_id, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=kfac_repeat_ndim, + ) + + def __call__( + self, + edge: Float[Array, "n n d_edge"], + c_level: Float[Array, "n d_c"], + mask: Float[Array, "n"] | None = None, + *, + kfac_scan_shared: bool = False, + ) -> Float[Array, "n n d_edge"]: + n = c_level.shape[0] + d_c_dim = c_level.shape[-1] + node_structural_mask = mask + full_pair_structural_mask = ( + None if mask is None else mask[:, None] * mask[None, :] + ) + edge_ln = _inline_norm_forward( + self.ln_edge, + edge, + pathway="even", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + c_ln = _inline_norm_forward( + self.ln_c, + c_level, + pathway="even", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + c_ctx = ( + c_ln + if self.node_ctx_proj is None + else _tagged_dense_no_bias( + self.node_ctx_proj.weight, + c_ln, + tag_id=self.node_ctx_proj._use_id, + pathway="even", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + ) + d_c_dim = c_ctx.shape[-1] + c_i_b = jnp.broadcast_to(c_ctx[:, None, :], (n, n, d_c_dim)) + c_j_b = jnp.broadcast_to(c_ctx[None, :, :], (n, n, d_c_dim)) + pair_ij = jnp.concatenate([edge_ln, c_i_b, c_j_b], axis=-1) + A = self._psi_apply( + pair_ij, + "L", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + ) + B = self._psi_apply( + pair_ij, + "R", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + ) + if mask is not None: + m = mask.astype(A.dtype) + A = A * (m[:, None, None] * m[None, :, None]) + B = B * m[None, :, None] + n_eff = jnp.maximum(jnp.sum(m), 1.0).astype(A.dtype) + else: + n_eff = jnp.asarray(float(n), dtype=A.dtype) + P = jnp.einsum("ikc,kjc->ijc", A, B) + P = P / jnp.sqrt(n_eff) + p_ij = _inline_norm_forward( + self.ln_path, + P, + pathway="even", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + cat = jnp.concatenate([pair_ij, p_ij], axis=-1) + hidden = fused_silu( + _tagged_dense( + self.ffn_in.weight, + self.ffn_in.bias, + cat, + tag_id=self.ffn_in._use_id, + pathway="even", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + ) + delta = _tagged_dense( + self.ffn_out.weight, + self.ffn_out.bias, + hidden, + tag_id=self.ffn_out._use_id, + pathway="even", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + return delta + + def apply_residual( + self, + edge: Float[Array, "n n d_edge"], + c_level: Float[Array, "n d_c"], + mask: Float[Array, "n"] | None = None, + *, + kfac_scan_shared: bool = False, + ) -> Float[Array, "n n d_edge"]: + delta = self(edge, c_level, mask, kfac_scan_shared=kfac_scan_shared) + update_mask = None + if mask is not None: + m = mask.astype(bool) + update_mask = m[:, None] & m[None, :] + return _tree_ngpt_residual( + edge, + delta, + self.alpha, + max_gain=self.ngpt_alpha_max, + tag_id="", + update_mask=update_mask, + kfac_structural_mask=update_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + + +class CausalRouterEdgeFWLUpdate(EdgeFWLUpdate): + two_hop_channels: int = eqx.field(static=True) + + def __init__( + self, + d_c: int, + d_edge: int, + *, + key: PRNGKeyArray, + alpha_init: float, + alpha_max: float, + channels: int = 64, + edge_node_ctx_dim: int | None = None, + ): + super().__init__( + d_c, + d_edge, + key=key, + alpha_init=alpha_init, + alpha_max=alpha_max, + channels=channels, + edge_node_ctx_dim=edge_node_ctx_dim, + ) + self.two_hop_channels = int(channels) + + def __call__( + self, + edge: Float[Array, "n n d_edge"], + c_level: Float[Array, "n d_c"], + mask: Float[Array, "n"] | None = None, + *, + kfac_scan_shared: bool = False, + ) -> Float[Array, "n n d_edge"]: + n = c_level.shape[0] + node_structural_mask = mask + full_pair_structural_mask = ( + None if mask is None else mask[:, None] * mask[None, :] + ) + idx = jnp.arange(n, dtype=jnp.int32) + causal_pair = idx[None, :] <= idx[:, None] + if full_pair_structural_mask is None: + full_pair_structural_mask = jnp.ones((n, n), dtype=bool) + causal_pair_structural_mask = ( + full_pair_structural_mask.astype(bool) & causal_pair + ) + edge_ln = _inline_norm_forward( + self.ln_edge, + edge, + pathway="even", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + c_ln = _inline_norm_forward( + self.ln_c, + c_level, + pathway="even", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + c_ctx = ( + c_ln + if self.node_ctx_proj is None + else _tagged_dense_no_bias( + self.node_ctx_proj.weight, + c_ln, + tag_id=self.node_ctx_proj._use_id, + pathway="even", + kfac_structural_mask=node_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + ) + d_c_dim = c_ctx.shape[-1] + c_i_b = jnp.broadcast_to(c_ctx[:, None, :], (n, n, d_c_dim)) + c_j_b = jnp.broadcast_to(c_ctx[None, :, :], (n, n, d_c_dim)) + pair_ij = jnp.concatenate([edge_ln, c_i_b, c_j_b], axis=-1) + A = self._psi_apply( + pair_ij, + "L", + kfac_structural_mask=causal_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + ) + B = self._psi_apply( + pair_ij, + "R", + kfac_structural_mask=full_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + ) + if mask is not None: + m = mask.astype(A.dtype) + A = A * (m[:, None, None] * m[None, :, None]) + B = B * m[None, :, None] + allowed = causal_pair.astype(A.dtype) * m[None, :] + else: + allowed = causal_pair.astype(A.dtype) + A = A * causal_pair[..., None].astype(A.dtype) + n_eff = jnp.maximum(jnp.sum(allowed, axis=1), 1.0).astype(A.dtype) + P = jnp.einsum("ikc,kjc->ijc", A, B) + P = P / jnp.sqrt(n_eff)[:, None, None] + P = P * causal_pair[..., None].astype(P.dtype) + p_ij = _inline_norm_forward( + self.ln_path, + P, + pathway="even", + kfac_structural_mask=causal_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + cat = jnp.concatenate([pair_ij, p_ij], axis=-1) + hidden = fused_silu( + _tagged_dense( + self.ffn_in.weight, + self.ffn_in.bias, + cat, + tag_id=self.ffn_in._use_id, + pathway="even", + kfac_structural_mask=causal_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + ) + delta = _tagged_dense( + self.ffn_out.weight, + self.ffn_out.bias, + hidden, + tag_id=self.ffn_out._use_id, + pathway="even", + kfac_structural_mask=causal_pair_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + delta_mask = causal_pair.astype(delta.dtype) + if mask is not None: + delta_mask = delta_mask * (m[:, None] * m[None, :]) + return delta * delta_mask[..., None] + + def append_causal_row( + self, + edge: Float[Array, "n n d_edge"], + c_level: Float[Array, "n d_c"], + mask: Float[Array, "n"], + row: Int[Array, ""], + b_cache: Float[Array, "n n channels"], + *, + edge_row: Float[Array, "n d_edge"] | None = None, + edge_col: Float[Array, "n d_edge"] | None = None, + sequence_axis_name=None, + sequence_mesh=None, + ): + n = c_level.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + row = jnp.asarray(row, dtype=jnp.int32) + active = mask.astype(bool) + allowed = active & (idx <= row) + if edge_row is None: + edge_row = edge[row] + if edge_col is None: + edge_col = edge[:, row] + edge_row_ln = _inline_norm_forward( + self.ln_edge, + edge_row, + pathway="even", + kfac_structural_mask=allowed, + kfac_repeat_ndim=1, + ) + edge_col_ln = _inline_norm_forward( + self.ln_edge, + edge_col, + pathway="even", + kfac_structural_mask=active, + kfac_repeat_ndim=1, + ) + c_ln = _inline_norm_forward( + self.ln_c, + c_level, + pathway="even", + kfac_structural_mask=active, + kfac_repeat_ndim=1, + ) + c_ctx = ( + c_ln + if self.node_ctx_proj is None + else _tagged_dense_no_bias( + self.node_ctx_proj.weight, + c_ln, + tag_id=self.node_ctx_proj._use_id, + pathway="even", + kfac_structural_mask=active, + kfac_repeat_ndim=1, + ) + ) + c_row = c_ctx[row] + c_row_b = jnp.broadcast_to(c_row, c_ctx.shape) + pair_row = jnp.concatenate([edge_row_ln, c_row_b, c_ctx], axis=-1) + pair_col = jnp.concatenate([edge_col_ln, c_ctx, c_row_b], axis=-1) + a_row = self._psi_apply( + pair_row, "L", kfac_structural_mask=allowed, kfac_repeat_ndim=1 + ) + b_row = self._psi_apply( + pair_row, "R", kfac_structural_mask=active, kfac_repeat_ndim=1 + ) + b_col = self._psi_apply( + pair_col, "R", kfac_structural_mask=active, kfac_repeat_ndim=1 + ) + if sequence_axis_name is not None: + from jax.sharding import NamedSharding, PartitionSpec as P + + lanes = ( + int(sequence_mesh.shape[sequence_axis_name]) + if sequence_mesh is not None + else 1 + ) + row_spec = P(None, None) + col_spec = P( + sequence_axis_name if n >= lanes and n % lanes == 0 else None, None + ) + if sequence_mesh is not None: + row_spec = NamedSharding(sequence_mesh, row_spec) + col_spec = NamedSharding(sequence_mesh, col_spec) + b_row = jax.lax.with_sharding_constraint(b_row, row_spec) + b_col = jax.lax.with_sharding_constraint(b_col, col_spec) + b_cache = _replace_square_row_column(b_cache, row, b_row, b_col) + a_row = a_row * allowed[:, None].astype(a_row.dtype) + path = jnp.einsum("kc,kjc->jc", a_row, b_cache) + n_eff = jnp.maximum( + jnp.sum(allowed.astype(path.dtype)), jnp.asarray(1.0, dtype=path.dtype) + ) + path = path / jnp.sqrt(n_eff) + path = path * allowed[:, None].astype(path.dtype) + path_ln = _inline_norm_forward( + self.ln_path, + path, + pathway="even", + kfac_structural_mask=allowed, + kfac_repeat_ndim=1, + ) + cat = jnp.concatenate([pair_row, path_ln], axis=-1) + hidden = fused_silu( + _tagged_dense( + self.ffn_in.weight, + self.ffn_in.bias, + cat, + tag_id=self.ffn_in._use_id, + pathway="even", + kfac_structural_mask=allowed, + kfac_repeat_ndim=1, + ) + ) + delta = _tagged_dense( + self.ffn_out.weight, + self.ffn_out.bias, + hidden, + tag_id=self.ffn_out._use_id, + pathway="even", + kfac_structural_mask=allowed, + kfac_repeat_ndim=1, + ) + delta = delta * allowed[:, None].astype(delta.dtype) + updated = _tree_ngpt_residual( + edge_row, + delta, + self.alpha, + max_gain=self.ngpt_alpha_max, + tag_id="", + update_mask=allowed, + kfac_structural_mask=allowed, + kfac_repeat_ndim=1, + ) + return updated, b_cache + + def apply_residual( + self, + edge: Float[Array, "n n d_edge"], + c_level: Float[Array, "n d_c"], + mask: Float[Array, "n"] | None = None, + *, + kfac_scan_shared: bool = False, + ) -> Float[Array, "n n d_edge"]: + delta = self(edge, c_level, mask, kfac_scan_shared=kfac_scan_shared) + idx = jnp.arange(edge.shape[0], dtype=jnp.int32) + update_mask = idx[None, :] <= idx[:, None] + if mask is not None: + m = mask.astype(bool) + update_mask = update_mask & m[:, None] & m[None, :] + return _tree_ngpt_residual( + edge, + delta, + self.alpha, + max_gain=self.ngpt_alpha_max, + tag_id="", + update_mask=update_mask, + kfac_structural_mask=update_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + + +class LevelEdgeAttn(eqx.Module): + ln_scale: Float[Array, "d_c"] + w_qkv: Float[Array, "d_c d_qkv"] + w_o: Float[Array, "d_o_in d_c"] + bias_mlp: MLP + ffn_ln_scale: Float[Array, "d_c"] + ffn_w1: Float[Array, "d_c d_ffn_hidden"] + ffn_b1: Float[Array, "d_ffn_hidden"] + ffn_w2: Float[Array, "d_ffn_hidden d_c"] + ffn_b2: Float[Array, "d_c"] + alpha_attn: Float[Array, "d_c"] + alpha_ffn: Float[Array, "d_c"] + n_heads: int = eqx.field(static=True) + n_heads_kernel: int = eqx.field(static=True) + d_attn: int = eqx.field(static=True) + d_head: int = eqx.field(static=True) + d_ffn_hidden: int = eqx.field(static=True) + attn_impl: str = eqx.field(static=True) + ln_eps: float = eqx.field(static=True) + max_n: int = eqx.field(static=True) + rope_base: float = eqx.field(static=True) + rope_scaling: float = eqx.field(static=True) + ngpt_alpha_max: float = eqx.field(static=True, default=0.5) + _use_id_ln: str = eqx.field(static=True, default="") + _use_id_qkv: str = eqx.field(static=True, default="") + _use_id_o: str = eqx.field(static=True, default="") + _use_id_ffn_ln: str = eqx.field(static=True, default="") + _use_id_ffn1: str = eqx.field(static=True, default="") + _use_id_ffn2: str = eqx.field(static=True, default="") + + def __init__( + self, + d_c: int, + d_edge: int, + *, + key: PRNGKeyArray, + alpha_init: float, + alpha_max: float, + n_heads: int = 4, + attn_dim: int | None = None, + attn_impl: str = "mhsea_tuned", + bias_mlp_hidden: int | None = None, + bias_mlp_n_blocks: int = 1, + ffn_d_hidden: int | None = None, + ln_eps: float = 1e-05, + max_n: int = 128, + rope_base: float = 10000.0, + rope_scaling: float = 1.0, + ): + d_attn = int(d_c) if attn_dim is None else int(attn_dim) + if d_attn < 1: + raise ValueError( + f"LevelEdgeAttn attn_dim must be positive or None, got {attn_dim}" + ) + assert d_attn % n_heads == 0, ( + f"attn_dim ({d_attn}) must be divisible by n_heads ({n_heads})" + ) + if attn_impl not in ("einsum", "mhsea_tuned"): + raise ValueError("attn_impl must be 'einsum' or 'mhsea_tuned'") + n_heads_kernel = 2 * n_heads + d_head = d_attn // n_heads + if d_head % 2 != 0: + raise ValueError( + f"LevelEdgeAttn requires even d_head for RoPE, got d_head={d_head} (= attn_dim={d_attn} / n_heads={n_heads})" + ) + if max_n < 2: + raise ValueError(f"LevelEdgeAttn max_n must be >= 2, got {max_n}") + if rope_base <= 0.0: + raise ValueError("LevelEdgeAttn rope_base must be positive") + if rope_scaling <= 0.0: + raise ValueError("LevelEdgeAttn rope_scaling must be positive") + d_qkv_out = n_heads_kernel * d_head + d_o_in = n_heads * d_head + k_qkv, k_b, k_f1, k_o, k_f2 = jax.random.split(key, 5) + self.w_qkv = jax.random.normal(k_qkv, (d_c, 3 * d_qkv_out)) * d_c ** (-0.5) + self.w_o = jax.random.normal(k_o, (d_o_in, d_c)) * d_o_in ** (-0.5) + if bias_mlp_hidden is None: + bias_mlp_hidden = max(32, n_heads_kernel * 2) + self.bias_mlp = MLP( + d_edge, bias_mlp_hidden, n_heads_kernel, key=k_b, n_blocks=bias_mlp_n_blocks + ) + self.ln_scale = jnp.ones((d_c,)) + d_ffn_eff = ffn_d_hidden if ffn_d_hidden is not None else 4 * d_c + self.ffn_ln_scale = jnp.ones((d_c,)) + self.ffn_w1 = jax.random.normal(k_f1, (d_c, d_ffn_eff)) * d_c ** (-0.5) + self.ffn_b1 = jnp.zeros((d_ffn_eff,)) + self.ffn_w2 = jax.random.normal(k_f2, (d_ffn_eff, d_c)) * d_ffn_eff ** (-0.5) + self.ffn_b2 = jnp.zeros((d_c,)) + self.n_heads = n_heads + self.n_heads_kernel = n_heads_kernel + self.d_attn = d_attn + self.d_head = d_head + self.d_ffn_hidden = d_ffn_eff + self.attn_impl = attn_impl + self.ln_eps = ln_eps + self.max_n = int(max_n) + self.rope_base = float(rope_base) + self.rope_scaling = float(rope_scaling) + self.ngpt_alpha_max = float(alpha_max) + self.alpha_attn = float(alpha_init) * jnp.ones((int(d_c),)) + self.alpha_ffn = float(alpha_init) * jnp.ones((int(d_c),)) + + def __call__( + self, + c_level: Float[Array, "n d_c"], + edge: Float[Array, "n n d_edge"], + mask: Float[Array, "n"], + level_idx=None, + *, + kfac_scan_shared: bool = False, + ) -> Float[Array, "n d_c"]: + n = c_level.shape[0] + H_k = self.n_heads_kernel + d_h = self.d_head + pair_mask = mask[:, None] * mask[None, :] + x = _tagged_rms_eqx_style( + self.ln_scale, + c_level, + eps=self.ln_eps, + tag_id=self._use_id_ln, + pathway="even", + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + qkv = _tagged_dense_no_bias( + self.w_qkv, + x, + tag_id=self._use_id_qkv, + pathway="even", + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + qkv = qkv.reshape(n, 3, H_k, d_h) + Q = qkv[:, 0] + K = qkv[:, 1] + V = qkv[:, 2] + bmlp = self.bias_mlp + b = _tagged_dense( + bmlp.in_proj.weight, + bmlp.in_proj.bias, + edge, + tag_id=bmlp.in_proj._use_id, + pathway="even", + kfac_structural_mask=pair_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + for nrm, l1, l2 in zip(bmlp.block_norms, bmlp.block_l1s, bmlp.block_l2s): + normed = _inline_norm_forward( + nrm, + b, + pathway="even", + kfac_structural_mask=pair_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + inner = _tagged_dense( + l1.weight, + l1.bias, + normed, + tag_id=l1._use_id, + pathway="even", + kfac_structural_mask=pair_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + inner_act = bmlp._act(inner) + inner_out = _tagged_dense( + l2.weight, + l2.bias, + inner_act, + tag_id=l2._use_id, + pathway="even", + kfac_structural_mask=pair_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + b = b + bmlp.inner_gain * inner_out + b_normed = _inline_norm_forward( + bmlp.out_norm, + b, + pathway="even", + kfac_structural_mask=pair_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + coup_bias = _tagged_dense( + bmlp.out_proj.weight, + bmlp.out_proj.bias, + b_normed, + tag_id=bmlp.out_proj._use_id, + pathway="even", + kfac_structural_mask=pair_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=2, + ) + coup_bias = coup_bias / jnp.sqrt(d_h) + if level_idx is not None: + _idx = jnp.arange(n, dtype=jnp.int32) + _lca = lca_alibi_bias(_idx, _idx, lca_fixed_slopes(H_k, dtype=x.dtype)) + coup_bias = coup_bias + jnp.transpose(_lca, (1, 2, 0)) + impl = self.attn_impl + from .pallas_attention import ( + mhsea_tuned_edge_attention, + reference_edge_attention, + ) + + if impl == "einsum": + out = reference_edge_attention(Q, K, V, coup_bias, mask) + elif impl == "mhsea_tuned": + d_head_padded = max(16, d_h) + pad_amount = d_head_padded - d_h + scale = jnp.sqrt(jnp.float32(d_head_padded / d_h)) + pad_shape = (n, H_k, pad_amount) + Q_p = jnp.concatenate([Q * scale, jnp.zeros(pad_shape, Q.dtype)], axis=-1) + K_p = jnp.concatenate([K, jnp.zeros(pad_shape, K.dtype)], axis=-1) + V_p = jnp.concatenate([V, jnp.zeros(pad_shape, V.dtype)], axis=-1) + out = mhsea_tuned_edge_attention(Q_p, K_p, V_p, coup_bias, mask)[..., :d_h] + else: + raise ValueError("attn_impl must be 'einsum' or 'mhsea_tuned'") + gate_heads = out[:, : self.n_heads, :] + value_heads = out[:, self.n_heads :, :] + out = jax.nn.sigmoid(gate_heads) * value_heads + out_flat = out.reshape(n, self.n_heads * d_h) + delta = _tagged_dense_no_bias( + self.w_o, + out_flat, + tag_id=self._use_id_o, + pathway="even", + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + mask_q = mask.reshape(-1, 1).astype(delta.dtype) + proposal_attn = mask_q * delta + c_attn = _tree_ngpt_residual( + c_level, + proposal_attn, + self.alpha_attn, + max_gain=self.ngpt_alpha_max, + tag_id="", + update_mask=mask, + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + x_ffn = _tagged_rms_eqx_style( + self.ffn_ln_scale, + c_attn, + eps=self.ln_eps, + tag_id=self._use_id_ffn_ln, + pathway="even", + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + h = _tagged_dense( + self.ffn_w1, + self.ffn_b1, + x_ffn, + tag_id=self._use_id_ffn1, + pathway="even", + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + h = fused_silu(h) + delta_ffn = _tagged_dense( + self.ffn_w2, + self.ffn_b2, + h, + tag_id=self._use_id_ffn2, + pathway="even", + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + proposal_ffn = mask_q * delta_ffn + return _tree_ngpt_residual( + c_attn, + proposal_ffn, + self.alpha_ffn, + max_gain=self.ngpt_alpha_max, + tag_id="", + update_mask=mask, + kfac_structural_mask=mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + + +class MergeOp(eqx.Module): + T: Float[Array, "G d_r d_r d_r"] + _use_id_T: str = eqx.field(static=True, default="") + mlp_c: MLP + alpha_c: Float[Array, "d_c"] + _use_id_alpha_c: str = eqx.field(static=True, default="") + ngpt_alpha_max: float = eqx.field(static=True, default=0.5) + eps: float = eqx.field(static=True) + output_hypernet: HypernetMatrix + edge_merge: EdgeMergeOp + level_edge_attn: LevelEdgeAttn + tree_edge_fwl: EdgeFWLUpdate + + def __init__( + self, + d_r: int, + d_c: int, + *, + key: PRNGKeyArray, + d_g: int, + alpha_init: float, + alpha_max: float, + eps: float = 1e-06, + merge_output_hypernet_rank: int = 128, + d_m_merge: int | None = None, + level_edge_attn_d_edge: int = 64, + level_edge_attn_n_heads: int = 4, + level_edge_attn_attn_dim: int | None = None, + tree_edge_node_ctx_dim: int | None = None, + level_edge_attn_attn_impl: str = "mhsea_tuned", + level_edge_attn_edge_mlp_hidden: int | None = None, + level_edge_attn_edge_mlp_n_blocks: int = 2, + level_edge_attn_ffn_d_hidden: int | None = None, + level_edge_attn_max_n: int = 128, + level_edge_attn_rope_base: float = 10000.0, + level_edge_attn_rope_scaling: float = 1.0, + tree_edge_fwl_channels: int = 64, + level_edge_attn_bias_mlp_hidden: int | None = None, + level_edge_attn_bias_mlp_n_blocks: int = 1, + merge_c_mlp_hidden: int | None = None, + ): + keys = jax.random.split(key, 11) + k_t, k_c, k_h = keys[:3] + k_em, k_lea, k_fwl = keys[4], keys[5], keys[10] + d_m_eff = d_m_merge if d_m_merge is not None else d_r + if d_m_eff % d_r != 0: + raise ValueError( + f"MergeOp: d_m_merge={d_m_merge} must be divisible by d_r={d_r} (HT carrier T is reshape-indexed as [G, d_r, d_r, d_r] with G = d_m_eff // d_r)." + ) + G_merge = d_m_eff // d_r + merge_edge_dim = 2 * int(level_edge_attn_d_edge) + merge_clock_dim = int(d_c) + merge_ngpt_dim = _TREE_NGPT_DEPTH_FEAT_DIM + merge_extra_dim = int(d_g) + merge_edge_dim + merge_clock_dim + merge_ngpt_dim + std = 1.0 / d_r + self.T = jax.random.normal(k_t, (G_merge, d_r, d_r, d_r)) * std + _c_hidden = ( + int(merge_c_mlp_hidden) if merge_c_mlp_hidden is not None else max(d_c, 16) + ) + self.mlp_c = MLP(2 * d_c + merge_extra_dim, _c_hidden, d_c, key=k_c) + self.ngpt_alpha_max = float(alpha_max) + self.alpha_c = float(alpha_init) * jnp.ones((int(d_c),)) + self.output_hypernet = HypernetMatrix( + d_in=d_m_eff, + d_out=d_m_eff, + d_e=d_c + merge_ngpt_dim, + rank=merge_output_hypernet_rank, + key=k_h, + ) + self.eps = eps + self.edge_merge = EdgeMergeOp( + d_edge=level_edge_attn_d_edge, + d_c=d_c, + key=k_em, + alpha_init=alpha_init, + alpha_max=alpha_max, + d_hidden=level_edge_attn_edge_mlp_hidden, + n_blocks=level_edge_attn_edge_mlp_n_blocks, + edge_node_ctx_dim=tree_edge_node_ctx_dim, + ) + self.level_edge_attn = LevelEdgeAttn( + d_c=d_c, + d_edge=level_edge_attn_d_edge, + key=k_lea, + alpha_init=alpha_init, + alpha_max=alpha_max, + n_heads=level_edge_attn_n_heads, + attn_dim=level_edge_attn_attn_dim, + attn_impl=level_edge_attn_attn_impl, + ffn_d_hidden=level_edge_attn_ffn_d_hidden, + max_n=level_edge_attn_max_n, + rope_base=level_edge_attn_rope_base, + rope_scaling=level_edge_attn_rope_scaling, + bias_mlp_hidden=level_edge_attn_bias_mlp_hidden, + bias_mlp_n_blocks=level_edge_attn_bias_mlp_n_blocks, + ) + self.tree_edge_fwl = EdgeFWLUpdate( + d_c=d_c, + d_edge=level_edge_attn_d_edge, + key=k_fwl, + alpha_init=alpha_init, + alpha_max=alpha_max, + channels=tree_edge_fwl_channels, + edge_node_ctx_dim=tree_edge_node_ctx_dim, + ) + + def _merge_extra_inputs( + self, + c_a, + c_b, + g_emb, + sibling_edge_lr, + sibling_edge_rl, + level_idx, + pair_idx, + pair_base, + clock_depth, + depth_feats, + ): + parts = [g_emb] + parts.append( + jnp.concatenate( + [sibling_edge_lr.astype(c_a.dtype), sibling_edge_rl.astype(c_a.dtype)] + ) + ) + parts.append( + _tree_merge_clock( + level_idx, pair_idx, pair_base, c_a.shape[-1], clock_depth, c_a.dtype + ) + ) + parts.append(depth_feats.astype(c_a.dtype)) + return parts + + def _apply_c_skip( + self, + c_a, + c_b, + c_delta, + *, + kfac_structural_mask=None, + kfac_scan_shared: bool = False, + ): + return _tree_ngpt_residual( + 0.5 * (c_a + c_b), + c_delta, + self.alpha_c, + max_gain=self.ngpt_alpha_max, + tag_id=self._use_id_alpha_c, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + + def context_candidate( + self, + c_a: Float[Array, "d_c"], + c_b: Float[Array, "d_c"], + g_emb: Float[Array, "d_g"], + *, + sibling_edge_lr: Float[Array, "d_edge"], + sibling_edge_rl: Float[Array, "d_edge"], + level_idx: Array, + pair_idx: Array, + pair_base: Array, + clock_depth: Array, + depth_feats: Array, + kfac_structural_mask=None, + kfac_g_structural_mask=None, + kfac_scan_shared: bool = False, + ): + ffn_in = jnp.concatenate( + [ + c_a, + c_b, + *self._merge_extra_inputs( + c_a, + c_b, + g_emb, + sibling_edge_lr, + sibling_edge_rl, + level_idx, + pair_idx, + pair_base, + clock_depth, + depth_feats, + ), + ] + ) + mlp = self.mlp_c + _we = _rownorm_cols + x = _tagged_dense( + mlp.in_proj.weight, + mlp.in_proj.bias, + ffn_in, + tag_id=mlp.in_proj._use_id, + pathway="even", + weight_eff=_we(mlp.in_proj.weight), + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + for nrm, l1, l2 in zip(mlp.block_norms, mlp.block_l1s, mlp.block_l2s): + normed = _inline_norm_forward( + nrm, + x, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + inner = _tagged_dense( + l1.weight, + l1.bias, + normed, + tag_id=l1._use_id, + pathway="even", + weight_eff=_we(l1.weight), + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + inner_act = mlp._act(inner) + inner_out = _tagged_dense( + l2.weight, + l2.bias, + inner_act, + tag_id=l2._use_id, + pathway="even", + weight_eff=_we(l2.weight), + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + x = x + mlp.inner_gain * inner_out + out_normed = _inline_norm_forward( + mlp.out_norm, + x, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + c_delta = _tagged_dense( + mlp.out_proj.weight, + mlp.out_proj.bias, + out_normed, + tag_id=mlp.out_proj._use_id, + pathway="even", + weight_eff=_we(mlp.out_proj.weight), + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + return self._apply_c_skip( + c_a, + c_b, + c_delta, + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=kfac_scan_shared, + ) + + def __call__( + self, + c_a: Float[Array, "d_c"], + u_a: Float[Array, "d_m_eff"], + s_a: Float[Array, ""], + c_b: Float[Array, "d_c"], + u_b: Float[Array, "d_m_eff"], + s_b: Float[Array, ""], + g_emb: Float[Array, "d_g"], + *, + sibling_edge_lr: Float[Array, "d_edge"], + sibling_edge_rl: Float[Array, "d_edge"], + level_idx: Array, + pair_idx: Array, + pair_base: Array, + clock_depth: Array, + depth_feats: Array, + kfac_context_mask=None, + kfac_g_context_mask=None, + kfac_odd_mask=None, + kfac_scan_shared: bool = False, + ): + c_p = self.context_candidate( + c_a, + c_b, + g_emb, + sibling_edge_lr=sibling_edge_lr, + sibling_edge_rl=sibling_edge_rl, + level_idx=level_idx, + pair_idx=pair_idx, + pair_base=pair_base, + clock_depth=clock_depth, + depth_feats=depth_feats, + kfac_structural_mask=kfac_context_mask, + kfac_g_structural_mask=kfac_g_context_mask, + kfac_scan_shared=kfac_scan_shared, + ) + raw = _quadrilinear_merge( + self.T, + u_a, + u_b, + tag_id=self._use_id_T, + pathway="odd", + kfac_structural_mask=kfac_odd_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + H = self.output_hypernet + h_ctx = jnp.concatenate([c_p, depth_feats.astype(c_p.dtype)]) + h_p = _tagged_dense_no_bias( + H.W_h, + h_ctx, + tag_id=H._use_id_W_h, + pathway="hypernet_eside", + kfac_structural_mask=kfac_odd_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + V_x = _tagged_dense_no_bias( + H.V, + raw, + tag_id=H._use_id_V, + pathway="odd", + kfac_structural_mask=kfac_odd_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + m_p = h_p * V_x + H_out = _tagged_dense_no_bias( + H.U, + m_p, + tag_id=H._use_id_U, + pathway="odd", + kfac_structural_mask=kfac_odd_mask, + kfac_scan_shared=kfac_scan_shared, + kfac_repeat_ndim=1, + ) + out = raw + H_out + scale_sq = jnp.mean(out * out) + scale = jnp.sqrt(scale_sq + self.eps) + u_p = out / scale + s_p = s_a + s_b + jnp.log(scale) + return (c_p, u_p, s_p) + + +def project_tree_ngpt_rownorm(model): + import equinox as _eqx + + def _collect(m): + mlp = m.merge.mlp_c + leaves = [mlp.in_proj.weight] + for l1, l2 in zip(mlp.block_l1s, mlp.block_l2s): + leaves.append(l1.weight) + leaves.append(l2.weight) + leaves.append(mlp.out_proj.weight) + leaves.append(m.route_decoder.tree_merge.w1) + leaves.append(m.route_decoder.tree_merge.w2) + return tuple(leaves) + + targets = _collect(model) + return _eqx.tree_at(_collect, model, _rownorm_project_jit(targets)) + + +@partial(jax.jit, donate_argnums=0) +def _rownorm_project_jit(ws): + return tuple((_rownorm_cols(w) for w in ws)) + + +def merge_masked( + c_a, + u_a, + s_a, + m_a, + c_b, + u_b, + s_b, + m_b, + merge, + g_emb, + *, + k_a, + k_b, + sibling_edge_lr, + sibling_edge_rl, + level_idx, + pair_idx, + pair_base, + clock_depth, + depth_feats, + kfac_g_context_mask=None, + kfac_scan_shared=False, +): + both = m_a * m_b + only_a = m_a * (1.0 - m_b) + only_b = (1.0 - m_a) * m_b + both_k = k_a * k_b + only_a_k = k_a * (1.0 - k_b) + only_b_k = (1.0 - k_a) * k_b + + def gate(value, left, right): + return jnp.where( + both.astype(bool), + value, + jnp.where( + only_a.astype(bool), + left, + jnp.where(only_b.astype(bool), right, jnp.zeros_like(value)), + ), + ) + + def structural_gate(value, left, right): + return jnp.where( + both_k.astype(bool), + value, + jnp.where( + only_a_k.astype(bool), + left, + jnp.where(only_b_k.astype(bool), right, jnp.zeros_like(value)), + ), + ) + + c_new, u_new, s_new = merge( + c_a, + u_a, + s_a, + c_b, + u_b, + s_b, + g_emb, + sibling_edge_lr=sibling_edge_lr, + sibling_edge_rl=sibling_edge_rl, + level_idx=level_idx, + pair_idx=pair_idx, + pair_base=pair_base, + clock_depth=clock_depth, + depth_feats=depth_feats, + kfac_context_mask=both_k, + kfac_g_context_mask=kfac_g_context_mask, + kfac_odd_mask=both, + kfac_scan_shared=kfac_scan_shared, + ) + c_out = structural_gate(c_new, c_a, c_b) + out = (c_out, gate(u_new, u_a, u_b), gate(s_new, s_a, s_b), m_a + m_b - m_a * m_b) + return out + + +def edge_merge_masked( + e_2i_2j: Float[Array, "d_edge"], + e_2i_2j1: Float[Array, "d_edge"], + e_2i1_2j: Float[Array, "d_edge"], + e_2i1_2j1: Float[Array, "d_edge"], + m_2i: Float[Array, ""], + m_2i1: Float[Array, ""], + m_2j: Float[Array, ""], + m_2j1: Float[Array, ""], + c_2i: Float[Array, "d_c"], + c_2i1: Float[Array, "d_c"], + c_2j: Float[Array, "d_c"], + c_2j1: Float[Array, "d_c"], + edge_merge: EdgeMergeOp, + *, + k_2i: Float[Array, ""], + k_2i1: Float[Array, ""], + k_2j: Float[Array, ""], + k_2j1: Float[Array, ""], + kfac_scan_shared: bool = False, +) -> tuple[Float[Array, "d_edge"], Float[Array, ""]]: + m_p = m_2i + m_2i1 - m_2i * m_2i1 + m_q = m_2j + m_2j1 - m_2j * m_2j1 + m_pq = m_p * m_q + both_p = k_2i * k_2i1 + both_q = k_2j * k_2j1 + out_mask = both_p * both_q + proposal = edge_merge( + e_2i_2j, + e_2i_2j1, + e_2i1_2j, + e_2i1_2j1, + c_2i, + c_2i1, + c_2j, + c_2j1, + kfac_structural_mask=out_mask, + kfac_scan_shared=kfac_scan_shared, + ) + cell_weights = jnp.stack( + [k_2i * k_2j, k_2i * k_2j1, k_2i1 * k_2j, k_2i1 * k_2j1] + ).astype(proposal.dtype) + child_edges = jnp.stack([e_2i_2j, e_2i_2j1, e_2i1_2j, e_2i1_2j1], axis=0) + mean_denom = jnp.maximum( + jnp.sum(cell_weights), jnp.asarray(1.0, dtype=proposal.dtype) + ) + masked_mean = jnp.sum(child_edges * cell_weights[:, None], axis=0) / mean_denom + e_pq = edge_merge.apply_skip( + masked_mean, + proposal, + kfac_structural_mask=out_mask, + kfac_scan_shared=kfac_scan_shared, + ) + return ( + jnp.where(jnp.asarray(out_mask).astype(bool), e_pq, jnp.zeros_like(e_pq)), + m_pq, + ) + + +def _next_pow2(n: int) -> int: + return 1 if n <= 1 else 1 << (n - 1).bit_length() + + +def _balanced_subtree_mask(m, n_pad: int): + N = jnp.sum(m.astype(jnp.int32)) + powers = 2 ** jnp.arange(max(1, n_pad.bit_length()), dtype=jnp.int32) + big = jnp.asarray(1 << 30, dtype=jnp.int32) + next_p = jnp.min(jnp.where(powers >= jnp.maximum(N, 1), powers, big)) + return (jnp.arange(n_pad, dtype=jnp.int32) < next_p).astype(m.dtype) + + +def _tree_active_clock_depth(mask): + n_active = jnp.maximum(jnp.sum(mask.astype(jnp.int32)), 1) + depth = jnp.ceil(jnp.log2(n_active.astype(jnp.float32))).astype(jnp.int32) + return jnp.maximum(depth, 1) + + +def balanced_tree_reduce_masked_scan( + c: Float[Array, "n d_c"], + u: Float[Array, "n d_m_eff"], + s: Float[Array, "n"], + m: Float[Array, "n"], + merge: MergeOp, + g_emb: Float[Array, "d_g"], + *, + edges_init: Float[Array, "n n d_edge"], + gladder, + g_stream0, +): + from hamiltonzero.model.fp32 import compute_dtype as _compute_dtype + + _cdtype = _compute_dtype() + if c.dtype != _cdtype: + c = c.astype(_cdtype) + if edges_init.dtype != _cdtype: + edges_init = edges_init.astype(_cdtype) + n = c.shape[0] + if n == 1: + return (c[0], u[0], s[0], m[0], edges_init[0, 0], g_stream0) + n_pad = _next_pow2(n) + n_levels = n_pad.bit_length() - 1 + pad_amount = n_pad - n + if pad_amount > 0: + c = jnp.pad(c, ((0, pad_amount),) + ((0, 0),) * (c.ndim - 1)) + u = jnp.pad(u, ((0, pad_amount),) + ((0, 0),) * (u.ndim - 1)) + s = jnp.pad(s, (0, pad_amount)) + m = jnp.pad(m, (0, pad_amount)) + k = _balanced_subtree_mask(m, n_pad) + clock_depth = _tree_active_clock_depth(k) + d_edge = edges_init.shape[-1] + edges_padded = _tree_sphere( + jnp.pad(edges_init, ((0, pad_amount), (0, pad_amount), (0, 0))) + ) + initial_state = (c, u, s, m, k, edges_padded, g_stream0) + n_pairs = n_pad // 2 + pidx = jnp.arange(n_pairs, dtype=jnp.int32) + + def _merge_one_pair( + c_a_, + u_a_, + s_a_, + c_b_, + u_b_, + s_b_, + m_a_, + m_b_, + k_a_, + k_b_, + sibling_edge_lr_, + sibling_edge_rl_, + depth_feats_, + pair_idx_, + pair_base_, + clock_depth_, + level_idx_, + g_emb_, + g_structural_mask_, + ): + return merge_masked( + c_a_, + u_a_, + s_a_, + m_a_, + c_b_, + u_b_, + s_b_, + m_b_, + merge, + g_emb_, + k_a=k_a_, + k_b=k_b_, + sibling_edge_lr=sibling_edge_lr_, + sibling_edge_rl=sibling_edge_rl_, + level_idx=level_idx_, + pair_idx=pair_idx_, + pair_base=pair_base_, + clock_depth=clock_depth_, + depth_feats=depth_feats_, + kfac_g_context_mask=g_structural_mask_, + kfac_scan_shared=True, + ) + + _vmap_in_axes = (0,) * 14 + (None,) * 5 + _vmapped_merge = jax.vmap( + _merge_one_pair, in_axes=_vmap_in_axes, axis_name="tree_pair" + ) + + def _edge_one( + e0, + e1, + e2, + e3, + m_p_a, + m_p_b, + m_q_a, + m_q_b, + k_p_a, + k_p_b, + k_q_a, + k_q_b, + c_p_a, + c_p_b, + c_q_a, + c_q_b, + ): + return edge_merge_masked( + e0, + e1, + e2, + e3, + m_p_a, + m_p_b, + m_q_a, + m_q_b, + c_p_a, + c_p_b, + c_q_a, + c_q_b, + merge.edge_merge, + k_2i=k_p_a, + k_2i1=k_p_b, + k_2j=k_q_a, + k_2j1=k_q_b, + kfac_scan_shared=True, + ) + + _edge_inner = jax.vmap( + _edge_one, + in_axes=(0, 0, 0, 0, None, None, 0, 0, None, None, 0, 0, None, None, 0, 0), + axis_name="tree_edge_q", + ) + _edge_outer = jax.vmap( + _edge_inner, + in_axes=(0, 0, 0, 0, 0, 0, None, None, 0, 0, None, None, 0, 0, None, None), + axis_name="tree_edge_p", + ) + + def body(state, xs_lv): + level_idx, depth_feats_lv = xs_lv + c, u, s, m, k, E, g_carry = state + + def _split(x): + xr = x.reshape((n_pairs, 2) + x.shape[1:]) + return (xr[:, 0], xr[:, 1]) + + m_a, m_b = _split(m) + k_a, k_b = _split(k) + level_active = jnp.any((k_a * k_b).astype(bool)) + c_a, c_b = _split(c) + u_a, u_b = _split(u) + s_a, s_b = _split(s) + pair_args: list = [c_a, u_a, s_a, c_b, u_b, s_b] + pair_args.extend([m_a, m_b]) + pair_args.extend([k_a, k_b]) + E_rs_for_merge = E.reshape(n_pairs, 2, n_pairs, 2, d_edge) + pair_args.extend( + [E_rs_for_merge[pidx, 0, pidx, 1, :], E_rs_for_merge[pidx, 1, pidx, 0, :]] + ) + g_emb_lvl = _tagged_dense( + gladder[2], + gladder[3], + g_carry, + tag_id="gladder.tree.proj", + pathway="even", + kfac_structural_mask=level_active, + kfac_scan_shared=True, + kfac_repeat_ndim=0, + ) + pair_args.append(depth_feats_lv) + clock_pair_active = k_a + k_b - k_a * k_b + pair_base = jnp.maximum( + jnp.sum(clock_pair_active.astype(jnp.int32)), + jnp.asarray(2, dtype=jnp.int32), + ) + pair_args.extend( + [pidx, pair_base, clock_depth, level_idx, g_emb_lvl, level_active] + ) + merged = _vmapped_merge(*pair_args) + c_p, u_p, s_p, m_p = merged + k_p = k_a + k_b - k_a * k_b + both_struct = k_a * k_b + attn_mask = both_struct + E_rs = E.reshape(n_pairs, 2, n_pairs, 2, d_edge) + E_00 = E_rs[:, 0, :, 0, :] + E_01 = E_rs[:, 0, :, 1, :] + E_10 = E_rs[:, 1, :, 0, :] + E_11 = E_rs[:, 1, :, 1, :] + E_new, _m_edge_new = _edge_outer( + E_00, + E_01, + E_10, + E_11, + m_a, + m_b, + m_a, + m_b, + k_a, + k_b, + k_a, + k_b, + c_a, + c_b, + c_a, + c_b, + ) + E_new = merge.tree_edge_fwl.apply_residual( + E_new, c_p, attn_mask, kfac_scan_shared=True + ) + edge_keep = (both_struct[:, None] * both_struct[None, :]).astype(bool) + E_new = jnp.where(edge_keep[..., None], E_new, E_00) + E_new = jnp.where(edge_keep[..., None], _tree_sphere(E_new), E_00) + c_skip = c_p + c_p = merge.level_edge_attn( + c_p, E_new, attn_mask, level_idx=level_idx, kfac_scan_shared=True + ) + c_p = jnp.where(attn_mask.astype(bool)[:, None], _tree_sphere(c_p), c_skip) + _lvl_mask = k_p.astype(c_p.dtype) + update_active = jnp.any(attn_mask.astype(bool)) + pool_structural_mask = _lvl_mask * update_active.astype(_lvl_mask.dtype) + pooled = gladder[0]( + g_carry, + c_p, + _lvl_mask, + kfac_structural_mask=pool_structural_mask, + kfac_update_mask=update_active, + kfac_scan_shared=True, + kfac_repeat_ndim=1, + ) + g_carry = gladder[1]( + g_carry, + pooled, + update_mask=update_active, + kfac_structural_mask=update_active, + kfac_scan_shared=True, + ) + + def _zpad(x_half): + return jnp.concatenate([x_half, jnp.zeros_like(x_half)], axis=0) + + pad_amt = n_pad - n_pairs + E_padded = jnp.pad(E_new, ((0, pad_amt), (0, pad_amt), (0, 0))) + return ( + ( + _zpad(c_p), + _zpad(u_p), + _zpad(s_p), + _zpad(m_p), + _zpad(k_p), + E_padded, + g_carry, + ), + None, + ) + + depth_feat_levels = _tree_ngpt_level_counts(m, n_pairs, n_levels, c.dtype) + final_state, _ = jax.lax.scan( + body, initial_state, (jnp.arange(n_levels), depth_feat_levels) + ) + c_f, u_f, s_f, m_f, _k_f, E_f, g_final = final_state + return (c_f[0], u_f[0], s_f[0], m_f[0], E_f[0, 0], g_final) + + +class RootReadout(eqx.Module): + output_hypernet: HypernetMatrix + ln_e: _RMS + + def __init__( + self, + d_r: int, + *, + key: PRNGKeyArray, + d_m_merge: int | None = None, + d_edge: int, + edge_rank: int = 64, + d_g: int = 0, + d_c: int, + ): + d_m_eff = d_m_merge if d_m_merge is not None else d_r + if d_m_eff % d_r != 0: + raise ValueError( + f"RootReadout: d_m_merge={d_m_merge} must be divisible by d_r={d_r} (HT carrier dim must match MergeOp's d_m_eff)." + ) + keys = jax.random.split(key, 5) + d_e_ctx = int(d_edge) + int(d_c) + int(d_g) + self.output_hypernet = HypernetMatrix( + d_in=d_m_eff, d_out=2, d_e=d_e_ctx, rank=edge_rank, key=keys[0] + ) + self.ln_e = _RMS(d_edge) + + def __call__(self, u_r, s_r, *, e_root, g_emb, c_root, kfac_structural_mask=None): + e_norm = _inline_norm_forward( + self.ln_e, + e_root, + pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + ) + h_ctx = jnp.concatenate( + [e_norm, c_root.astype(e_norm.dtype), g_emb.astype(e_norm.dtype)] + ) + psi = self.output_hypernet.apply( + h_ctx, + u_r, + e_pathway="even", + kfac_structural_mask=kfac_structural_mask, + kfac_scan_shared=False, + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=True, + ) + re = 0.5 * jnp.log(psi[0] * psi[0] + psi[1] * psi[1]) + s_r + return (re, jnp.arctan2(psi[1], psi[0])) diff --git a/src/hamiltonzero/model/trunk.py b/src/hamiltonzero/model/trunk.py new file mode 100644 index 0000000000000000000000000000000000000000..68a1ca120a597c405d951026841a2bd3945f9923 --- /dev/null +++ b/src/hamiltonzero/model/trunk.py @@ -0,0 +1,525 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int, PRNGKeyArray +from .context import SpinContext +from .fused_silu import fused_silu +from .odd_ops import BiasFreeLinear, Linear, MLP, UnnormalizedMLP, _RMS + + +class MultiHeadEvenAttention(eqx.Module): + W_QKV: BiasFreeLinear + W_O: BiasFreeLinear + bias_mlp: UnnormalizedMLP + ln_edge: _RMS + n_heads: int = eqx.field(static=True) + n_heads_kernel: int = eqx.field(static=True) + d_head: int = eqx.field(static=True) + d_attn: int = eqx.field(static=True) + attn_impl: str = eqx.field(static=True) + + def __init__( + self, + d_e: int, + n_heads: int, + n_edge: int, + *, + key: PRNGKeyArray, + attn_impl: str, + n_layers: int, + attn_dim: int, + bias_hidden_dim: int, + ): + if n_heads < 1: + raise ValueError(f"n_heads must be >= 1, got {n_heads}") + d_attn = int(attn_dim) + if d_attn < 1: + raise ValueError(f"attn_dim must be positive or None, got {attn_dim}") + if d_attn % n_heads != 0: + raise ValueError( + f"attention inner width must be divisible by n_heads: attn_dim={d_attn}, n_heads={n_heads}" + ) + if attn_impl not in ("einsum", "mhsea_tuned"): + raise ValueError("attn_impl must be 'einsum' or 'mhsea_tuned'") + k_qkv, k_o, k_b, k_ln_edge = jax.random.split(key, 4) + del k_ln_edge + d_head = d_attn // n_heads + n_heads_kernel = 2 * n_heads + d_qkv_out = n_heads_kernel * d_head + self.W_QKV = BiasFreeLinear(d_e, 3 * d_qkv_out, key=k_qkv) + self.W_O = BiasFreeLinear(d_attn, d_e, key=k_o) + self.bias_mlp = UnnormalizedMLP( + n_edge, + bias_hidden_dim, + n_heads_kernel, + key=k_b, + n_blocks=1, + inner_gain=float(n_layers) ** (-0.5), + ) + self.ln_edge = _RMS(n_edge) + self.n_heads = n_heads + self.n_heads_kernel = n_heads_kernel + self.d_head = d_head + self.d_attn = d_attn + self.attn_impl = attn_impl + + def __call__( + self, + e: Float[Array, "n d_e"], + edge: Float[Array, "n n n_edge"], + mask: Int[Array, "n"], + ) -> Float[Array, "n d_e"]: + n = e.shape[0] + node_structural_mask = mask.astype(bool) + pair_structural_mask = ( + node_structural_mask[:, None] & node_structural_mask[None, :] + ) + qkv = self.W_QKV( + e, + pathway="even", + kfac_structural_mask=node_structural_mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ).reshape(n, 3, self.n_heads_kernel, self.d_head) + Q, K, V = (qkv[:, 0], qkv[:, 1], qkv[:, 2]) + edge_pre = self.ln_edge( + edge, + pathway="even", + kfac_structural_mask=pair_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + coup_bias = self.bias_mlp( + edge_pre, + pathway="even", + kfac_structural_mask=pair_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + coup_bias = coup_bias / jnp.sqrt( + jnp.asarray(self.d_head, dtype=coup_bias.dtype) + ) + from .pallas_attention import ( + mhsea_tuned_edge_attention, + reference_edge_attention, + ) + + if self.attn_impl == "einsum": + out = reference_edge_attention(Q, K, V, coup_bias, mask) + else: + d_head_padded = max(16, self.d_head) + pad_amount = d_head_padded - self.d_head + scale = jnp.sqrt(jnp.asarray(d_head_padded / self.d_head, dtype=Q.dtype)) + Q_pad = jnp.concatenate( + [ + Q * scale, + jnp.zeros((n, self.n_heads_kernel, pad_amount), dtype=Q.dtype), + ], + axis=-1, + ) + K_pad = jnp.concatenate( + [K, jnp.zeros((n, self.n_heads_kernel, pad_amount), dtype=K.dtype)], + axis=-1, + ) + V_pad = jnp.concatenate( + [V, jnp.zeros((n, self.n_heads_kernel, pad_amount), dtype=V.dtype)], + axis=-1, + ) + out = mhsea_tuned_edge_attention(Q_pad, K_pad, V_pad, coup_bias, mask) + out = out[..., : self.d_head] + gate_heads = out[:, : self.n_heads, :] + value_heads = out[:, self.n_heads :, :] + out = jax.nn.sigmoid(gate_heads) * value_heads + out = out.reshape(n, -1) + return self.W_O( + out, + pathway="even", + kfac_structural_mask=node_structural_mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + + +class EvenFFN(eqx.Module): + l1: Linear + l2: Linear + + def __init__(self, d_e: int, d_hidden: int, *, key: PRNGKeyArray): + k1, k2 = jax.random.split(key, 2) + self.l1 = Linear(d_e, d_hidden, key=k1) + self.l2 = Linear(d_hidden, d_e, key=k2) + + def __call__( + self, e: Float[Array, "... d_e"], *, kfac_structural_mask=None + ) -> Float[Array, "... d_e"]: + kwargs = dict( + kfac_structural_mask=kfac_structural_mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + return self.l2( + fused_silu(self.l1(e, pathway="even", **kwargs)), pathway="even", **kwargs + ) + + +class EdgeUpdateContextAware(eqx.Module): + ln_edge: _RMS + ln_even: _RMS + node_ctx_proj: BiasFreeLinear + ffn: MLP + psi_L_in: Linear + psi_L_out: Linear + psi_R_in: Linear + psi_R_out: Linear + ln_path: _RMS + two_hop_channels: int = eqx.field(static=True, default=0) + edge_node_ctx_dim: int = eqx.field(static=True, default=0) + + def __init__( + self, + d_e: int, + n_edge: int, + *, + key: PRNGKeyArray, + d_hidden: int, + n_layers: int, + edge_node_ctx_dim: int, + two_hop_channels: int, + two_hop_hidden_dim: int, + ): + node_ctx_dim = int(edge_node_ctx_dim) + if node_ctx_dim < 1: + raise ValueError( + f"edge_node_ctx_dim must be positive, got {edge_node_ctx_dim}" + ) + self.ln_edge = _RMS(n_edge) + self.ln_even = _RMS(d_e) + self.node_ctx_proj = BiasFreeLinear( + d_e, node_ctx_dim, key=jax.random.fold_in(key, 60782) + ) + d_pair = n_edge + 2 * node_ctx_dim + d_in = d_pair + two_hop_channels + ( + k_ffn, + k_psi_L_in, + k_psi_L_out, + k_psi_L_gate, + k_psi_R_in, + k_psi_R_out, + k_psi_R_gate, + ) = jax.random.split(key, 7) + del k_psi_L_gate, k_psi_R_gate + self.ffn = MLP( + d_in, + d_hidden, + n_edge, + key=k_ffn, + n_blocks=1, + inner_gain=float(n_layers) ** (-0.5), + ) + self.edge_node_ctx_dim = int(node_ctx_dim) + self.psi_L_in = Linear(d_pair, two_hop_hidden_dim, key=k_psi_L_in) + self.psi_L_out = Linear(two_hop_hidden_dim, two_hop_channels, key=k_psi_L_out) + self.psi_R_in = Linear(d_pair, two_hop_hidden_dim, key=k_psi_R_in) + self.psi_R_out = Linear(two_hop_hidden_dim, two_hop_channels, key=k_psi_R_out) + self.ln_path = _RMS(two_hop_channels) + self.two_hop_channels = int(two_hop_channels) + + def __call__( + self, + edge: Float[Array, "n n n_edge"], + even: Float[Array, "n d_e"], + mask: Int[Array, "n"] | Float[Array, "n"], + ) -> Float[Array, "n n n_edge"]: + n = even.shape[0] + node_structural_mask = mask.astype(bool) + pair_structural_mask = ( + node_structural_mask[:, None] & node_structural_mask[None, :] + ) + node_kfac = dict( + kfac_structural_mask=node_structural_mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + pair_kfac = dict( + kfac_structural_mask=pair_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + edge_ln = self.ln_edge(edge, pathway="even", **pair_kfac) + even_ln = self.ln_even(even, pathway="even", **node_kfac) + even_ctx = self.node_ctx_proj(even_ln, pathway="even", **node_kfac) + d_ctx = even_ctx.shape[-1] + even_i_b = jnp.broadcast_to(even_ctx[:, None, :], (n, n, d_ctx)) + even_j_b = jnp.broadcast_to(even_ctx[None, :, :], (n, n, d_ctx)) + pair_ij = jnp.concatenate([edge_ln, even_i_b, even_j_b], axis=-1) + A = self._psi_apply( + pair_ij, + self.psi_L_in, + self.psi_L_out, + kfac_structural_mask=pair_structural_mask, + ) + B = self._psi_apply( + pair_ij, + self.psi_R_in, + self.psi_R_out, + kfac_structural_mask=pair_structural_mask, + ) + m = mask.astype(A.dtype) + A = A * (m[:, None, None] * m[None, :, None]) + B = B * m[None, :, None] + n_eff = jnp.maximum(jnp.sum(m), 1.0).astype(A.dtype) + P = jnp.einsum("ikc,kjc->ijc", A, B) / jnp.sqrt(n_eff) + p_ij = self.ln_path(P, pathway="even", **pair_kfac) + cat = jnp.concatenate([pair_ij, p_ij], axis=-1) + return self.ffn(cat, pathway="even", **pair_kfac) + + def _psi_apply( + self, + pair_ij: Float[Array, "n n d_pair"], + l_in: Linear, + l_out: Linear, + *, + kfac_structural_mask=None, + ) -> Float[Array, "n n C"]: + kfac_kwargs = dict( + kfac_structural_mask=kfac_structural_mask, + kfac_repeat_ndim=2, + kfac_context_primal_reused_over_walkers=True, + ) + hidden = fused_silu(l_in(pair_ij, pathway="even", **kfac_kwargs)) + return l_out(hidden, pathway="even", **kfac_kwargs) + + +class TransformerBlock(eqx.Module): + edge_update_ctx: EdgeUpdateContextAware + ln_attn: _RMS + attn: MultiHeadEvenAttention + ln_ffn: _RMS + ffn: EvenFFN + g_pool: "GDescriptorPool" + g_update: "ResidualGlobalUpdate" + g_ffn_proj_w: Float[Array, "d_gstream d_e"] + residual_gain: float = eqx.field(static=True) + + def __init__( + self, + d_e: int, + n_heads: int, + n_edge: int, + gladder_d_g: int, + *, + key: PRNGKeyArray, + global_tap_dim: int, + n_layers: int, + attn_impl: str, + attn_dim: int, + attn_bias_hidden_dim: int, + ffn_hidden_dim: int, + edge_hidden_dim: int, + edge_node_ctx_dim: int, + two_hop_channels: int, + two_hop_hidden_dim: int, + ): + k_e, k_a, k_f, k_o = jax.random.split(key, 4) + del k_o + self.edge_update_ctx = EdgeUpdateContextAware( + d_e=d_e, + n_edge=n_edge, + key=k_e, + d_hidden=edge_hidden_dim, + n_layers=n_layers, + edge_node_ctx_dim=edge_node_ctx_dim, + two_hop_channels=two_hop_channels, + two_hop_hidden_dim=two_hop_hidden_dim, + ) + self.ln_attn = _RMS(d_e) + self.attn = MultiHeadEvenAttention( + d_e, + n_heads, + n_edge, + key=k_a, + attn_impl=attn_impl, + n_layers=n_layers, + attn_dim=attn_dim, + bias_hidden_dim=attn_bias_hidden_dim, + ) + self.ln_ffn = _RMS(d_e) + self.ffn = EvenFFN(d_e, ffn_hidden_dim, key=k_f) + from .global_ladder import GDescriptorPool, ResidualGlobalUpdate + + k_global = jax.random.split(jax.random.fold_in(key, 25009), 3) + self.g_pool = GDescriptorPool( + gladder_d_g, d_e, key=k_global[0], tag="gladder.trunk.pool" + ) + self.g_update = ResidualGlobalUpdate( + gladder_d_g, + self.g_pool.d_out, + key=k_global[1], + tap_dim=global_tap_dim, + tag="gladder.trunk.upd", + residual_gain=float(n_layers) ** (-0.5), + ) + self.g_ffn_proj_w = jax.random.normal( + k_global[2], (gladder_d_g, d_e) + ) * gladder_d_g ** (-0.5) + self.residual_gain = float(n_layers) ** (-0.5) + + def _even_edge_step( + self, + e: Float[Array, "n d_e"], + edge: Float[Array, "n n n_edge"], + mask: Int[Array, "n"], + g: Float[Array, "d_gstream"], + ): + node_structural_mask = mask.astype(bool) + node_kfac = dict( + kfac_structural_mask=node_structural_mask, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + edge_update = self.residual_gain * self.edge_update_ctx(edge, e, mask) + edge = edge + edge_update + e_pre = self.ln_attn(e, pathway="even", **node_kfac) + e = e + self.residual_gain * self.attn(e_pre, edge, mask) + e_pre = self.ln_ffn(e, pathway="even", **node_kfac) + from .tree import _tagged_dense_no_bias + + gg = _tagged_dense_no_bias( + self.g_ffn_proj_w, + g, + tag_id="gladder.trunk.fproj", + pathway="even", + kfac_structural_mask=jnp.any(mask.astype(bool)), + kfac_scan_shared=False, + kfac_repeat_ndim=0, + kfac_context_primal_reused_over_walkers=True, + ) + e_pre = e_pre + gg[None, :].astype(e_pre.dtype) + e = e + self.residual_gain * self.ffn( + e_pre, kfac_structural_mask=node_structural_mask + ) + system_active = jnp.any(mask.astype(bool)) + pooled = self.g_pool( + g, + e, + mask, + kfac_structural_mask=mask, + kfac_update_mask=system_active, + kfac_scan_shared=False, + kfac_repeat_ndim=1, + kfac_context_primal_reused_over_walkers=True, + ) + g = self.g_update( + g, + pooled, + kfac_structural_mask=system_active, + kfac_scan_shared=False, + kfac_context_primal_reused_over_walkers=True, + ) + return (e, edge, g) + + def __call__( + self, + e: Float[Array, "n d_e"], + edge: Float[Array, "n n n_edge"], + mask: Int[Array, "n"], + g: Float[Array, "d_gstream"], + ): + return self._even_edge_step(e, edge, mask, g=g) + + +class Trunk(eqx.Module): + blocks: TransformerBlock + + def __init__( + self, + d_e: int, + n_heads: int, + n_layers: int, + n_edge: int, + d_local_in: int, + d_edge_in: int, + *, + key: PRNGKeyArray, + gladder_d_g: int, + global_tap_dim: int, + attn_impl: str, + attn_dim: int, + attn_bias_hidden_dim: int, + ffn_hidden_dim: int, + edge_hidden_dim: int, + edge_node_ctx_dim: int, + two_hop_channels: int, + two_hop_hidden_dim: int, + ): + if d_local_in != d_e: + raise ValueError( + f"Trunk requires d_local_in (= feat_n_heads*feat_head_dim = {d_local_in}) == d_e (= {d_e}). Adjust featurizer config so the widths match." + ) + if d_edge_in != n_edge: + raise ValueError( + f"Trunk requires d_edge_in (= feat_d_edge = {d_edge_in}) == n_edge (= {n_edge})." + ) + k_odd, k_blocks = jax.random.split(key, 2) + del k_odd + block_keys = jax.random.split(k_blocks, n_layers) + + def make_block(k: PRNGKeyArray) -> TransformerBlock: + return TransformerBlock( + d_e, + n_heads, + n_edge, + gladder_d_g, + key=k, + global_tap_dim=global_tap_dim, + n_layers=n_layers, + attn_impl=attn_impl, + attn_dim=attn_dim, + attn_bias_hidden_dim=attn_bias_hidden_dim, + ffn_hidden_dim=ffn_hidden_dim, + edge_hidden_dim=edge_hidden_dim, + edge_node_ctx_dim=edge_node_ctx_dim, + two_hop_channels=two_hop_channels, + two_hop_hidden_dim=two_hop_hidden_dim, + ) + + block_list = [make_block(k) for k in block_keys] + dynamic_static = [eqx.partition(block, eqx.is_array) for block in block_list] + dynamic = [part for part, _ in dynamic_static] + _, static_template = dynamic_static[0] + stacked_dynamic = jax.tree.map(lambda *xs: jnp.stack(xs, axis=0), *dynamic) + self.blocks = eqx.combine(stacked_dynamic, static_template) + + def __call__( + self, + ctx: SpinContext, + edge_feat: Float[Array, "n n d_edge_in"], + local_feat: Float[Array, "n d_local_in"], + g: Float[Array, "d_gstream"], + ): + e = local_feat.astype(jnp.float32) + edge = edge_feat.astype(jnp.float32) + dynamic, static = eqx.partition(self.blocks, eqx.is_array) + + def scan_body(carry, layer_dynamic): + e_carry, edge_carry, g_carry = carry + block = eqx.combine(layer_dynamic, static) + e_carry, edge_carry, g_carry = block( + e_carry, + edge_carry, + ctx.mask, + g_carry, + ) + return (e_carry, edge_carry, g_carry), None + + (e, edge, g), _ = jax.lax.scan( + scan_body, + (e, edge, g), + dynamic, + ) + return (e, edge, g) diff --git a/src/hamiltonzero/modes/__init__.py b/src/hamiltonzero/modes/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..73bee8e197a6b25cc9314d020272060e398a5d38 --- /dev/null +++ b/src/hamiltonzero/modes/__init__.py @@ -0,0 +1,2 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 diff --git a/src/hamiltonzero/modes/eval.py b/src/hamiltonzero/modes/eval.py new file mode 100644 index 0000000000000000000000000000000000000000..a77a5505fe750e6f34bd0ead69aca8cf6750856e --- /dev/null +++ b/src/hamiltonzero/modes/eval.py @@ -0,0 +1,55 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import dataclasses +import json +import math +from pathlib import Path + +from hamiltonzero.config import EvalConfig +from hamiltonzero.evaluation import EvalBackend, EvalMetric, EvalResult, evaluate + + +def _strict_json(value): + if isinstance(value, float): + return value if math.isfinite(value) else None + if isinstance(value, dict): + return {key: _strict_json(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return [_strict_json(item) for item in value] + return value + + +def run(config: EvalConfig, backend: EvalBackend | None = None) -> EvalResult: + if backend is None: + from hamiltonzero.evaluation.runtime import build_eval_backend + + backend = build_eval_backend() + output = Path(config.output) + output.mkdir(parents=True, exist_ok=True) + metrics = output / "eval.metrics.jsonl" + metrics.write_text("") + + def write_metric(metric: EvalMetric) -> None: + with metrics.open("a", encoding="utf-8") as stream: + stream.write( + json.dumps( + _strict_json(dataclasses.asdict(metric)), + separators=(",", ":"), + ) + + "\n" + ) + + result = evaluate(config, backend, metric_sink=write_metric) + destination = output / "eval.json" + temporary = output / "eval.json.tmp" + temporary.write_text( + json.dumps(_strict_json(result.as_dict()), indent=2, sort_keys=True) + "\n" + ) + temporary.replace(destination) + return result + + +__all__ = ["run"] diff --git a/src/hamiltonzero/modes/finetune.py b/src/hamiltonzero/modes/finetune.py new file mode 100644 index 0000000000000000000000000000000000000000..e443d8d237b9e9b45547c2de31735a4d9750bbc6 --- /dev/null +++ b/src/hamiltonzero/modes/finetune.py @@ -0,0 +1,325 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import functools +import time +from dataclasses import dataclass +from typing import Callable + +import equinox as eqx +import jax +import jax.numpy as jnp +import numpy as np +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +from hamiltonzero.checkpoint import load_mcmc, load_model, save_model +from hamiltonzero.compiled.model import ( + CompiledFinetuneWaveFunction, + compile_finetune_model, +) +from hamiltonzero.config import FineTuneConfig +from hamiltonzero.data import build_context_and_energy, load_system +from hamiltonzero.energy import vmc_energy_custom_lap_finetune +from hamiltonzero.energy.frame import compile_energy_frame +from hamiltonzero.mcmc import ( + REState, + adapt_batched, + cold_samples, + init_batched_state, + run_batched, +) +from hamiltonzero.model import build_model +from hamiltonzero.optim import ( + KFACBundle, + apply_finetune_kfac_step, + init_finetune_kfac_state, + learning_rate, + process_finetune_targets, +) +from hamiltonzero.router import ( + batch_context, + route_context, + route_state, + select_frozen_route, + strip_router, +) + + +@dataclass(frozen=True, slots=True) +class FineTuneMetric: + step: int + energy: float + energy_std: float + step_walltime: float + walltime: float + + +@dataclass(frozen=True, slots=True) +class FineTuneResult: + model: CompiledFinetuneWaveFunction + route_perm: jax.Array + mcmc_state: REState + kfac: KFACBundle + last_metric: FineTuneMetric | None + + +def _adapt(state: REState, config: FineTuneConfig) -> REState: + return adapt_batched( + state, + beta_history_weight=config.mcmc.beta_history_weight, + sigma_target=config.mcmc.langevin_target_acceptance, + sigma_scale=config.mcmc.sigma_scale, + haar_target=config.mcmc.haar_target_acceptance, + ) + + +def _burn_in( + model: CompiledFinetuneWaveFunction, + context, + state: REState, + config: FineTuneConfig, +) -> REState: + replica_steps = config.mcmc.burn_in_replica_steps + step_mcmc = eqx.filter_jit( + functools.partial( + run_batched, + n_steps=replica_steps, + walker_chunk_size=config.mcmc.walker_chunk_size, + ) + ) + adapt = eqx.filter_jit(functools.partial(_adapt, config=config)) + for iteration in range(config.mcmc.burn_in): + state = step_mcmc( + model, + context, + state, + ) + if iteration > 0 and iteration % config.mcmc.adapt_every == 0: + state = adapt(state) + return state + + +def _replicate(value, sharding: NamedSharding): + return jax.device_put( + value, + jax.tree_util.tree_map(lambda _leaf: sharding, value), + ) + + +def _place_state(state: REState, mesh: Mesh) -> REState: + replicated = NamedSharding(mesh, P()) + batched = NamedSharding(mesh, P("batch")) + shardings = REState( + q=batched, + log_p=batched, + grad_log_p=batched, + beta=replicated, + sigma=replicated, + step=replicated, + key=batched, + n_local_accept=batched, + n_local=batched, + n_swap_accept=batched, + n_swap=batched, + mask=replicated, + m=replicated, + n_haar_accept=batched, + n_haar=batched, + ) + return jax.device_put(state, shardings) + + +def _metric( + step: int, + total, + step_started: float, + run_started: float, +) -> FineTuneMetric: + jax.block_until_ready(total) + return FineTuneMetric( + step=step, + energy=float(jax.device_get(jnp.mean(total.real))), + energy_std=float(jax.device_get(jnp.std(total.real))), + step_walltime=time.perf_counter() - step_started, + walltime=time.perf_counter() - run_started, + ) + + +def run_finetune( + config: FineTuneConfig, + *, + metric_sink: Callable[[FineTuneMetric], None] | None = None, +) -> FineTuneResult: + key = jax.random.PRNGKey(config.seed) + key_model, key_mcmc = jax.random.split(key) + system = load_system(config.system) + context, energy_inputs = build_context_and_energy( + system, + n_max=None, + mu=config.energy.mu, + eps=config.energy.eps, + ) + template = build_model( + config.model, + key_model, + n_max=int(context.mask.shape[-1]), + ) + eager_model = load_model(config.checkpoint, template) + state = init_batched_state( + key_mcmc, + context, + batch_size=config.mcmc.batch_size, + n_replicas=config.mcmc.replicas, + initial_m=config.mcmc.initial_haar_sites, + initial_sigma=config.mcmc.initial_sigma, + ) + reused = False + if config.mcmc.reuse_mcmc is not None: + state = load_mcmc(config.mcmc.reuse_mcmc, state) + reused = True + key, _route_key = jax.random.split(key) + freeze_route = eqx.filter_jit( + functools.partial( + select_frozen_route, + tau=config.route_temperature, + ) + ) + route_perm = freeze_route( + eager_model, + context, + ) + energy_frame = compile_energy_frame( + energy_inputs, + context.mask, + context.bmask, + route_perm, + ) + context = route_context(context, route_perm) + state = route_state(state, route_perm) + eager_model = strip_router(eager_model) + key, key_expand = jax.random.split(key) + compile_model = eqx.filter_jit( + functools.partial( + compile_finetune_model, + leaf_rank=config.leaf_rank, + merge_rank=config.merge_rank, + ) + ) + model = compile_model( + eager_model, + context, + physical_perm=route_perm, + key=key_expand, + ) + del eager_model, template, _route_key + devices = tuple(jax.devices()) + if config.mcmc.batch_size % len(devices): + raise ValueError( + f"batch_size={config.mcmc.batch_size} must be divisible by " + f"the {len(devices)} visible devices" + ) + mesh = Mesh(np.asarray(devices, dtype=object), ("batch",)) + replicated = NamedSharding(mesh, P()) + model = _replicate(model, replicated) + context = _replicate(context, replicated) + energy_frame = _replicate(energy_frame, replicated) + state = _place_state(state, mesh) + context_batch = _replicate(batch_context(context), replicated) + q_cold = cold_samples(state) + kfac_data = NamedSharding(mesh, P(None, "batch")) + q_kfac = jax.device_put(q_cold[None], kfac_data) + energy_seed = jax.device_put( + jnp.zeros((1, config.mcmc.batch_size), dtype=jnp.complex64), + kfac_data, + ) + kfac = init_finetune_kfac_state( + config.kfac, + model, + q_kfac, + energy_seed, + context_batch, + t=0.0, + key=jax.random.fold_in(key, 0xCAFE), + multi_device=mesh.size > 1, + ) + if not reused: + state = _burn_in(model, context, state, config) + step_mcmc = eqx.filter_jit( + functools.partial( + run_batched, + n_steps=config.mcmc.steps, + walker_chunk_size=config.mcmc.walker_chunk_size, + ) + ) + adapt = eqx.filter_jit(functools.partial(_adapt, config=config)) + local_energy = eqx.filter_jit( + functools.partial( + vmc_energy_custom_lap_finetune, + chunk_size=config.energy.chunk_size, + ) + ) + run_started = time.perf_counter() + last_metric = None + for step in range(config.steps): + step_started = time.perf_counter() + state = step_mcmc( + model, + context, + state, + ) + if step > 0 and step % config.mcmc.adapt_every == 0: + state = adapt(state) + q_cold = cold_samples(state) + total, _exchange, _casimir, _field = local_energy( + model, + energy_frame, + q_cold, + ) + target = process_finetune_targets( + total[None], + context_batch.s_norm, + mad_width=config.kfac.mad_clip_width, + ) + key, key_kfac = jax.random.split(key) + model, kfac = apply_finetune_kfac_step( + kfac, + model, + jax.device_put(q_cold[None], kfac_data), + jax.device_put(target, kfac_data), + context_batch, + t=0.0, + key=key_kfac, + momentum=config.kfac.momentum, + learning_rate=learning_rate(config.kfac, step), + damping=config.kfac.damping, + ) + jax.block_until_ready(model) + last_metric = _metric(step, total, step_started, run_started) + if metric_sink is not None: + metric_sink(last_metric) + save_model( + config.output, + model, + kind="compiled_finetune", + metadata={ + "leaf_rank": int(config.leaf_rank), + "merge_rank": int(config.merge_rank), + "n_max": int(context.mask.shape[-1]), + }, + ) + return FineTuneResult( + model=model, + route_perm=route_perm, + mcmc_state=state, + kfac=kfac, + last_metric=last_metric, + ) + + +__all__ = [ + "FineTuneMetric", + "FineTuneResult", + "run_finetune", +] diff --git a/src/hamiltonzero/modes/train.py b/src/hamiltonzero/modes/train.py new file mode 100644 index 0000000000000000000000000000000000000000..1a44bf4cae7c1c6db8d0377c19f955ebc55a8798 --- /dev/null +++ b/src/hamiltonzero/modes/train.py @@ -0,0 +1,516 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import time +from dataclasses import dataclass +from typing import Callable + +import equinox as eqx +import jax +import jax.numpy as jnp +import numpy as np +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +from hamiltonzero.checkpoint import load_mcmc, save_model +from hamiltonzero.compiled.tree import compile_physical_tree_reference +from hamiltonzero.compiled.trunk import bind_shared_kernel, compile_shared_trunk +from hamiltonzero.compiled.types import ( + CompiledWaveFunction, +) +from hamiltonzero.config import TrainConfig +from hamiltonzero.data import build_context_and_energy, load_systems +from hamiltonzero.energy import vmc_energy_custom_lap_compiled +from hamiltonzero.energy.frame import compile_energy_frame +from hamiltonzero.mcmc import ( + REState, + adapt_batched, + cold_samples, + init_batched_state, + run_batched, +) +from hamiltonzero.model import MultiSystemContext, build_model +from hamiltonzero.model.tree import project_tree_ngpt_rownorm +from hamiltonzero.optim import ( + KFACBundle, + apply_router_kfac_step, + init_router_kfac_state, + learning_rate, + process_route_targets, +) +from hamiltonzero.router import ( + ROUTE_SAMPLES, + bind_router_kernel, + build_beam16, + build_route_sampler, + compile_router_static, + rebase_cold_samples, + reframe_state_context, + snis_mode_baseline, +) + + +@dataclass(frozen=True, slots=True) +class TrainMetric: + step: int + system: int + energy: float + energy_std: float + step_walltime: float + walltime: float + + +@dataclass(frozen=True, slots=True) +class TrainResult: + model: object + kfac: KFACBundle + mcmc_states: tuple[REState | None, ...] + last_metric: TrainMetric | None + + +@dataclass(slots=True) +class _SystemState: + sampler: REState + context: MultiSystemContext + perms: jax.Array + + +def _identity_perms(n_max: int): + return jnp.broadcast_to( + jnp.arange(n_max, dtype=jnp.int32), + (ROUTE_SAMPLES, n_max), + ) + + +def _route_sharding(mesh: Mesh, value): + return jax.tree_util.tree_map( + lambda x: NamedSharding( + mesh, + P("systems", *([None] * (x.ndim - 1))) + if x.ndim and x.shape[0] == ROUTE_SAMPLES + else P(), + ), + value, + ) + + +def _replicate(mesh: Mesh, value): + replicated = NamedSharding(mesh, P()) + return jax.device_put(value, jax.tree_util.tree_map(lambda _: replicated, value)) + + +def _place_routes(mesh: Mesh, value): + return jax.device_put(value, _route_sharding(mesh, value)) + + +def _host_pool(value): + def materialize(x): + if not eqx.is_array(x): + return x + result = np.asarray(jax.device_get(x)) + if not result.flags.writeable: + result = result.copy() + return result + + return jax.tree_util.tree_map(materialize, value) + + +def _host_system_state(state: _SystemState) -> _SystemState: + return _SystemState( + sampler=_host_pool(state.sampler), + context=_host_pool(state.context), + perms=_host_pool(state.perms), + ) + + +def _activate_system(mesh: Mesh, state: _SystemState) -> _SystemState: + return _SystemState( + sampler=_place_routes(mesh, state.sampler), + context=_place_routes(mesh, state.context), + perms=_place_routes(mesh, state.perms), + ) + + +def _compile_trees(model, trunk, perms): + return jax.vmap(lambda perm: compile_physical_tree_reference(model, trunk, perm))( + perms + ) + + +def _compile_frames(inputs, mask, bmask, perms): + return jax.vmap( + lambda perm: compile_energy_frame( + inputs, + mask, + bmask, + perm, + ) + )(perms) + + +def _run_routes(kernel, trees, contexts, state, n_steps, chunk_size): + def one(tree, context, sampler): + return run_batched( + CompiledWaveFunction(kernel=kernel, tree=tree), + context, + sampler, + n_steps, + walker_chunk_size=chunk_size, + ) + + return jax.vmap(one)(trees, contexts, state) + + +def _adapt_routes(state, config): + return jax.vmap( + lambda sampler: adapt_batched( + sampler, + beta_history_weight=config.mcmc.beta_history_weight, + sigma_target=config.mcmc.langevin_target_acceptance, + sigma_scale=config.mcmc.sigma_scale, + haar_target=config.mcmc.haar_target_acceptance, + ) + )(state) + + +def _sampled_energy(kernel, trees, frames, q, chunk_size): + def one(tree, frame, q_row): + return vmc_energy_custom_lap_compiled( + kernel, + tree, + frame, + q_row, + chunk_size=chunk_size, + ) + + return jax.vmap(one)(trees, frames, q) + + +def _mode_energy(kernel, tree, frame, q_canonical, mode_perm, chunk_size): + q_mode = jnp.take(q_canonical, mode_perm, axis=-2) + + def one(q_row): + with jax.default_matmul_precision("default"): + log_p = 2.0 * jax.vmap( + lambda walker: CompiledWaveFunction(kernel, tree)(walker)[0] + )(q_row) + total, exchange, casimir, field = vmc_energy_custom_lap_compiled( + kernel, + tree, + frame, + q_row, + chunk_size=chunk_size, + ) + return total, exchange, casimir, field, log_p + + return jax.vmap(one)(q_mode) + + +def _initial_system_state( + model, + context, + config, + mcmc_key, + system_index, + *, + n_systems, + mesh, + compile_plan, + compile_trees, + run_routes, +): + walkers = config.mcmc.batch_size // ROUTE_SAMPLES + cpu = jax.devices("cpu")[0] + mcmc_key = jax.device_put(mcmc_key, cpu) + with jax.default_device(cpu): + lane_indices = system_index * ROUTE_SAMPLES + jnp.arange( + ROUTE_SAMPLES, dtype=jnp.int32 + ) + keys = jax.vmap(lambda index: jax.random.fold_in(mcmc_key, index))(lane_indices) + sampler = jax.vmap( + lambda lane_key: init_batched_state( + lane_key, + context, + batch_size=walkers, + n_replicas=config.mcmc.replicas, + initial_m=config.mcmc.initial_haar_sites, + initial_sigma=config.mcmc.initial_sigma, + ) + )(keys) + contexts = MultiSystemContext.stack([context] * ROUTE_SAMPLES) + perms = _identity_perms(config.n_max) + if config.mcmc.reuse_mcmc is not None: + source = config.mcmc.reuse_mcmc + if source.is_dir(): + source = source / f"{system_index}.eqx" + elif n_systems != 1: + raise ValueError( + "multisystem --reuse-mcmc must point to a directory of " + ".eqx files" + ) + sampler = load_mcmc(source, sampler) + sampler = _place_routes(mesh, sampler) + contexts = _place_routes(mesh, contexts) + perms = _place_routes(mesh, perms) + trunk = compile_plan(model, context)[0] + kernel = bind_shared_kernel(model) + trees = compile_trees(model, trunk, perms) + if config.mcmc.reuse_mcmc is None: + for iteration in range(config.mcmc.burn_in): + sampler = run_routes( + kernel, + trees, + contexts, + sampler, + config.mcmc.burn_in_replica_steps, + config.mcmc.walker_chunk_size, + ) + if iteration and iteration % config.mcmc.adapt_every == 0: + sampler = _adapt_routes(sampler, config) + return _SystemState(sampler=sampler, context=contexts, perms=perms) + + +def _metric(step, system_index, total, step_started, run_started): + jax.block_until_ready(total) + return TrainMetric( + step=step, + system=system_index, + energy=float(jax.device_get(jnp.mean(total.real))), + energy_std=float(jax.device_get(jnp.std(total.real))), + step_walltime=time.perf_counter() - step_started, + walltime=time.perf_counter() - run_started, + ) + + +def run_train( + config: TrainConfig, + *, + metric_sink: Callable[[TrainMetric], None] | None = None, +) -> TrainResult: + if config.mcmc.batch_size % ROUTE_SAMPLES: + raise ValueError("mcmc.batch_size must be divisible by K=8") + systems = load_systems(config.systems) + if not systems: + raise ValueError("training requires at least one system") + systems_data = [ + _host_pool( + build_context_and_energy( + system, + n_max=config.n_max, + mu=config.energy.mu, + eps=config.energy.eps, + ) + ) + for system in systems + ] + contexts = [context for context, _energy in systems_data] + energy_inputs = [energy for _context, energy in systems_data] + key = jax.random.PRNGKey(config.seed) + key_model, key_mcmc = jax.random.split(key) + model = build_model(config.model, key_model, n_max=config.n_max) + devices = tuple(jax.devices()) + if len(devices) < ROUTE_SAMPLES: + raise ValueError( + "learned-router train requires eight devices for the K=8 systems mesh" + ) + mesh = Mesh(np.asarray(devices[:ROUTE_SAMPLES], dtype=object), ("systems",)) + model = _replicate(mesh, model) + compile_plan = jax.jit( + lambda model_value, context_value: ( + compile_shared_trunk(model_value, context_value), + ) + ) + compile_trees = jax.jit(_compile_trees) + compile_frames = jax.jit(_compile_frames) + run_routes = jax.jit( + _run_routes, + static_argnums=(4, 5), + donate_argnums=(3,), + ) + sampled_energy = jax.jit(_sampled_energy, static_argnums=(4,)) + mode_energy = jax.jit(_mode_energy, static_argnums=(5,)) + reframe = jax.jit(reframe_state_context, donate_argnums=(0, 1)) + system_states: list[_SystemState | None] = [None] * len(systems) + + def get_system(index: int): + cached = system_states[index] + if cached is None: + return _initial_system_state( + model, + contexts[index], + config, + key_mcmc, + index, + n_systems=len(systems), + mesh=mesh, + compile_plan=compile_plan, + compile_trees=compile_trees, + run_routes=run_routes, + ) + return _activate_system(mesh, cached) + + first = get_system(0) + q_seed = jax.vmap(cold_samples)(first.sampler) + energy_seed = _place_routes(mesh, np.zeros(q_seed.shape[:2], dtype=np.complex64)) + kfac = init_router_kfac_state( + config.kfac, + model, + q_seed, + energy_seed, + first.context, + t=0.0, + key=jax.random.fold_in(key, 0xCAFE), + multi_device=True, + route_tau=config.router.temperature, + route_loss_weight=config.router.loss_weight, + ) + system_states[0] = _host_system_state(first) + del first, q_seed, energy_seed + order_rng = np.random.default_rng(config.seed) + order = np.arange(len(systems), dtype=np.int32) + order_rng.shuffle(order) + run_started = time.perf_counter() + last_metric = None + for step in range(config.steps): + step_started = time.perf_counter() + if step and step % len(order) == 0: + order_rng.shuffle(order) + system_index = int(order[step % len(order)]) + state = get_system(system_index) + trunk = compile_plan(model, contexts[system_index])[0] + router_kernel = bind_router_kernel(model) + router_static = compile_router_static( + router_kernel, + trunk, + contexts[system_index].route_quotient_node_key, + contexts[system_index].route_quotient_edge_key, + contexts[system_index].needs_fwl2, + ) + tau = jnp.asarray(config.router.temperature, dtype=jnp.float32) + router_kernel = _replicate(mesh, router_kernel) + router_static = _replicate(mesh, router_static) + key, key_route = jax.random.split(key) + new_perms = build_route_sampler(mesh, router_kernel.decoder, router_static)( + router_kernel.decoder, router_static, key_route, tau + ) + mode_perm = build_beam16(mesh, router_kernel.decoder, router_static)( + router_kernel.decoder, router_static, tau + ) + state.sampler, state.context = reframe( + state.sampler, state.context, state.perms, new_perms + ) + state.perms = new_perms + kernel = bind_shared_kernel(model) + trees = compile_trees(model, trunk, new_perms) + frames = compile_frames( + energy_inputs[system_index], + contexts[system_index].mask, + contexts[system_index].bmask, + new_perms, + ) + state.sampler = run_routes( + kernel, + trees, + state.context, + state.sampler, + config.mcmc.steps, + config.mcmc.walker_chunk_size, + ) + if step and step % config.mcmc.adapt_every == 0: + state.sampler = _adapt_routes(state.sampler, config) + q_cold = jax.vmap(cold_samples)(state.sampler) + total, _exchange, _casimir, _field = sampled_energy( + kernel, + trees, + frames, + q_cold, + config.energy.chunk_size, + ) + baseline_is_sampled = bool( + np.asarray( + jax.device_get( + jnp.all( + new_perms.astype(jnp.int32) + == mode_perm[None, :].astype(jnp.int32) + ) + ) + ) + ) + if baseline_is_sampled: + baseline_total = total + baseline_weights = jnp.full( + total.shape, + 1.0 / total.shape[-1], + dtype=total.real.dtype, + ) + else: + mode_tree = compile_physical_tree_reference(model, trunk, mode_perm) + mode_frame = compile_energy_frame( + energy_inputs[system_index], + contexts[system_index].mask, + contexts[system_index].bmask, + mode_perm, + ) + q_canonical = rebase_cold_samples(q_cold, new_perms) + baseline_total, _bx, _bc, _bf, candidate_log_p = mode_energy( + kernel, + mode_tree, + mode_frame, + q_canonical, + mode_perm, + config.energy.chunk_size, + ) + sampled_log_p = state.sampler.log_p[..., -1] + baseline_weights = snis_mode_baseline( + baseline_total, candidate_log_p, sampled_log_p + ) + target, advantage = process_route_targets( + total, + baseline_total, + state.context.s_norm, + baseline_weights, + mad_width=config.kfac.mad_clip_width, + ) + key, key_kfac = jax.random.split(key) + model, kfac = apply_router_kfac_step( + kfac, + model, + q_cold, + target, + state.context, + t=0.0, + key=key_kfac, + momentum=config.kfac.momentum, + learning_rate=learning_rate(config.kfac, step), + damping=config.kfac.damping, + route_advantage=advantage, + route_tau=config.router.temperature, + ) + model = project_tree_ngpt_rownorm(model) + jax.block_until_ready(model) + system_states[system_index] = _host_system_state(state) + last_metric = _metric(step, system_index, total, step_started, run_started) + if metric_sink is not None: + metric_sink(last_metric) + save_model( + config.output, + model, + kind="router", + metadata={"n_max": config.n_max}, + ) + return TrainResult( + model=model, + kfac=kfac, + mcmc_states=tuple( + None if value is None else value.sampler for value in system_states + ), + last_metric=last_metric, + ) + + +__all__ = [ + "TrainMetric", + "TrainResult", + "run_train", +] diff --git a/src/hamiltonzero/observables.py b/src/hamiltonzero/observables.py new file mode 100644 index 0000000000000000000000000000000000000000..95a831965fafa68d687bdc0d098d1aefaf4ac8f1 --- /dev/null +++ b/src/hamiltonzero/observables.py @@ -0,0 +1,64 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import math +from typing import Any + +import jax +import jax.numpy as jnp + +from hamiltonzero.energy.kernel import _right_su2_chart_jet + + +def _local_spin_single( + wavefunction: Any, + q_routed: jax.Array, +) -> jax.Array: + n_sites = q_routed.shape[0] + + def f_entry(z): + q_perturbed = _right_su2_chart_jet(q_routed, z) + real, imaginary = wavefunction(q_perturbed, None, 0.0) + return jnp.stack([real, imaginary]) + + jac_pair = jax.jacrev(f_entry)(jnp.zeros((n_sites, 3), dtype=q_routed.dtype)) + g_lie = jac_pair[0] + 1j * jac_pair[1] + return -0.5j * g_lie + + +def local_spin( + wavefunction: Any, + context: Any, + q_routed: jax.Array, + *, + chunk_size: int | None = 512, +) -> jax.Array: + q_routed = jnp.asarray(q_routed) + if q_routed.ndim < 2 or q_routed.shape[-1] != 4: + raise ValueError("q_routed must have shape [..., N, 4]") + if context.mask.shape[-1] != q_routed.shape[-2]: + raise ValueError("context mask and q_routed must have the same site width") + lead = q_routed.shape[:-2] + n_items = math.prod(lead) if lead else 1 + flat = q_routed.reshape((n_items,) + q_routed.shape[-2:]) + + with jax.default_matmul_precision("highest"): + if chunk_size is None or chunk_size >= n_items: + values = jax.vmap(lambda q: _local_spin_single(wavefunction, q))(flat) + else: + if chunk_size < 1: + raise ValueError("chunk_size must be positive or None") + values = jax.lax.map( + lambda q: _local_spin_single(wavefunction, q), + flat, + batch_size=int(chunk_size), + ) + values = values.reshape(lead + q_routed.shape[-2:-1] + (3,)) + return values * jnp.asarray(context.mask, dtype=values.real.dtype)[..., None] + + +__all__ = [ + "local_spin", +] diff --git a/src/hamiltonzero/optim/__init__.py b/src/hamiltonzero/optim/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..2c6290efe401465ab21eb2badbde65110c33267f --- /dev/null +++ b/src/hamiltonzero/optim/__init__.py @@ -0,0 +1,26 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .api import ( + KFACBundle, + apply_finetune_kfac_step, + apply_router_kfac_step, + init_finetune_kfac_state, + init_router_kfac_state, + learning_rate, + process_finetune_targets, + process_route_targets, + register_scale_and_shift, +) + +__all__ = [ + "KFACBundle", + "apply_finetune_kfac_step", + "apply_router_kfac_step", + "init_finetune_kfac_state", + "init_router_kfac_state", + "learning_rate", + "process_finetune_targets", + "process_route_targets", + "register_scale_and_shift", +] diff --git a/src/hamiltonzero/optim/api.py b/src/hamiltonzero/optim/api.py new file mode 100644 index 0000000000000000000000000000000000000000..65887478fa49a5edda3fc7b9e6d1fa61b40686b8 --- /dev/null +++ b/src/hamiltonzero/optim/api.py @@ -0,0 +1,47 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import kfac_jax + +from .production import ( + KFACBundle, + apply_finetune_kfac_step, + apply_router_kfac_step, + init_finetune_kfac_state, + init_router_kfac_state, +) +from .targets import process_finetune_targets, process_route_targets + + +def learning_rate(config, step: int) -> float: + return float(config.learning_rate_numerator) / ( + float(config.learning_rate_offset) + + float(step) / float(config.learning_rate_decay_steps) + ) + + +def register_scale_and_shift(y, x, scale, tag_id: str): + from hamiltonzero.model.tree import _kfac_name_kw + + return kfac_jax.register_scale_and_shift( + y, + x, + scale=scale, + shift=None, + **_kfac_name_kw(tag_id), + ) + + +__all__ = [ + "KFACBundle", + "apply_finetune_kfac_step", + "apply_router_kfac_step", + "init_finetune_kfac_state", + "init_router_kfac_state", + "learning_rate", + "process_finetune_targets", + "process_route_targets", + "register_scale_and_shift", +] diff --git a/src/hamiltonzero/optim/blocks.py b/src/hamiltonzero/optim/blocks.py new file mode 100644 index 0000000000000000000000000000000000000000..97e0f7e7f9dc90a7d6781ddb942e3e2a9ae42f8c --- /dev/null +++ b/src/hamiltonzero/optim/blocks.py @@ -0,0 +1,1681 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + +import math + +import jax +import jax.numpy as jnp +import kfac_jax +from kfac_jax._src import utils as kfac_utils +from kfac_jax._src.curvature_blocks import utils as cb_utils +from kfac_jax._src.layers_and_loss_tags import LayerMetaData, layer_tag + + +STRUCTURAL_DENSE_TAG_VARIANT = "structural_repeated_dense" +STRUCTURAL_SCALE_SHIFT_TAG_VARIANT = "structural_scale_and_shift" +STRUCTURAL_STACKED_DENSE_TAG_VARIANT = "structural_stacked_repeated_dense" +STRUCTURAL_STACKED_SCALE_SHIFT_TAG_VARIANT = "structural_stacked_scale_and_shift" +STRUCTURAL_TRAILING_STACKED_SCALE_SHIFT_TAG_VARIANT = ( + "structural_trailing_stacked_scale_and_shift" +) + + +def _optional_name_kw(name: str | None) -> dict[str, str]: + return {} if name is None else {"name": name} + + +def _validate_structural_registration( + x, + structural_mask, + *, + repeat_ndim: int, + feature_ndim: int, +) -> None: + if int(repeat_ndim) < 0: + raise ValueError(f"repeat_ndim must be non-negative, got {repeat_ndim}") + expected_mask_shape = x.shape if feature_ndim == 0 else x.shape[:-feature_ndim] + if tuple(structural_mask.shape) != tuple(expected_mask_shape): + raise ValueError( + "structural_mask must exactly cover the local repeat axes: " + f"mask={structural_mask.shape}, expected={expected_mask_shape}, " + f"x={x.shape}, feature_ndim={feature_ndim}" + ) + + +def register_structural_dense( + y, + x, + structural_mask, + weight, + bias=None, + *, + scan_shared: bool, + repeat_ndim: int, + context_primal_reused_over_walkers: bool = False, + name: str | None = None, +): + + _validate_structural_registration( + x, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + args = ( + (y, x, structural_mask, weight) + if bias is None + else (y, x, structural_mask, weight, bias) + ) + return layer_tag.bind( + *args, + meta=LayerMetaData( + variant=STRUCTURAL_DENSE_TAG_VARIANT, + outputs_index=(0,), + inputs_index=(1, 2), + params_index=tuple(range(3, len(args))), + ), + scan_shared=bool(scan_shared), + repeat_ndim=int(repeat_ndim), + context_primal_reused_over_walkers=bool(context_primal_reused_over_walkers), + **_optional_name_kw(name), + ) + + +def register_structural_scale_and_shift( + y, + x, + structural_mask, + scale=None, + shift=None, + *, + scan_shared: bool, + repeat_ndim: int, + context_primal_reused_over_walkers: bool = False, + name: str | None = None, +): + + params = tuple(value for value in (scale, shift) if value is not None) + if not params: + raise ValueError("At least one of scale and shift must be provided") + feature_ndim = params[0].ndim + if any(tuple(param.shape) != tuple(params[0].shape) for param in params[1:]): + raise ValueError("structural scale and shift shapes must match") + _validate_structural_registration( + x, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=feature_ndim, + ) + args = (y, x, structural_mask, *params) + return layer_tag.bind( + *args, + meta=LayerMetaData( + variant=STRUCTURAL_SCALE_SHIFT_TAG_VARIANT, + outputs_index=(0,), + inputs_index=(1, 2), + params_index=tuple(range(3, len(args))), + ), + has_scale=scale is not None, + has_shift=shift is not None, + scan_shared=bool(scan_shared), + repeat_ndim=int(repeat_ndim), + context_primal_reused_over_walkers=bool(context_primal_reused_over_walkers), + **_optional_name_kw(name), + ) + + +def register_structural_trailing_stacked_scale_and_shift( + y, + x, + structural_mask, + scale, + *, + repeat_ndim: int, + context_primal_reused_over_walkers: bool = False, + name: str | None = None, +): + + if scale.ndim != 2: + raise ValueError( + "trailing stacked scale/shift parameters must have shape [K,d]; " + f"got {scale.shape}" + ) + _validate_structural_registration( + x, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=2, + ) + args = (y, x, structural_mask, scale) + return layer_tag.bind( + *args, + meta=LayerMetaData( + variant=STRUCTURAL_TRAILING_STACKED_SCALE_SHIFT_TAG_VARIANT, + outputs_index=(0,), + inputs_index=(1, 2), + params_index=(3,), + ), + has_scale=True, + has_shift=False, + scan_shared=False, + repeat_ndim=int(repeat_ndim), + context_primal_reused_over_walkers=bool(context_primal_reused_over_walkers), + **_optional_name_kw(name), + ) + + +def _structural_tag_contract(layer_tag_eq): + params = layer_tag_eq.params + return ( + bool(params["scan_shared"]), + int(params["repeat_ndim"]), + bool(params.get("context_primal_reused_over_walkers", False)), + ) + + +def _align_structural_primal_and_mask( + x, + dy, + structural_mask, + *, + repeat_ndim: int, + feature_ndim: int, + context_primal_reused_over_walkers: bool, +): + + structural_mask = jnp.asarray(structural_mask, dtype=bool) + x_leading = tuple(x.shape[:-feature_ndim]) if feature_ndim else tuple(x.shape) + dy_leading = tuple(dy.shape[:-feature_ndim]) if feature_ndim else tuple(dy.shape) + + def _missing_walker_axis(source_leading, target_leading, *, what): + if source_leading == target_leading: + return None + if len(target_leading) != len(source_leading) + 1: + raise ValueError( + f"{what} supports exactly one missing walker sample axis: " + f"source={source_leading}, target={target_leading}" + ) + insert_axis = len(source_leading) - int(repeat_ndim) + if insert_axis < 0 or ( + source_leading[:insert_axis] != target_leading[:insert_axis] + or source_leading[insert_axis:] != target_leading[insert_axis + 1 :] + ): + raise ValueError( + f"{what} walker axis must be the final logical-sample axis " + f"before the {repeat_ndim} repeat axes: " + f"source={source_leading}, target={target_leading}" + ) + return insert_axis + + x_insert_axis = _missing_walker_axis( + x_leading, + dy_leading, + what="context primal reuse", + ) + if x_insert_axis is not None: + if not context_primal_reused_over_walkers: + raise ValueError( + "x/dy structural layouts differ without " + "context_primal_reused_over_walkers: " + f"x={x.shape}, dy={dy.shape}, mask={structural_mask.shape}" + ) + x = jnp.expand_dims(x, axis=x_insert_axis) + x = jnp.broadcast_to( + x, + (*dy_leading, *x.shape[-feature_ndim:]) if feature_ndim else dy_leading, + ) + + structural_mask = align_structural_mask_to_leading( + structural_mask, + dy_leading, + repeat_ndim=repeat_ndim, + ) + return x, dy, structural_mask + + +def align_structural_mask_to_leading( + structural_mask, + target_leading, + *, + repeat_ndim: int, +): + + structural_mask = jnp.asarray(structural_mask, dtype=bool) + source_leading = tuple(structural_mask.shape) + target_leading = tuple(target_leading) + if source_leading == target_leading: + return structural_mask + if len(target_leading) != len(source_leading) + 1: + raise ValueError( + "structural mask supports exactly one missing walker sample " + f"axis: source={source_leading}, target={target_leading}" + ) + insert_axis = len(source_leading) - int(repeat_ndim) + if insert_axis < 0 or ( + source_leading[:insert_axis] != target_leading[:insert_axis] + or source_leading[insert_axis:] != target_leading[insert_axis + 1 :] + ): + raise ValueError( + "structural mask walker axis must be the final logical-sample " + f"axis before the {repeat_ndim} repeat axes: " + f"source={source_leading}, target={target_leading}" + ) + structural_mask = jnp.expand_dims(structural_mask, axis=insert_axis) + return jnp.broadcast_to(structural_mask, target_leading) + + +def structural_group_repeats( + value, + structural_mask, + *, + scan_shared: bool, + repeat_ndim: int, + feature_ndim: int, +): + + feature_shape = tuple(value.shape[-feature_ndim:]) if feature_ndim else () + leading_shape = ( + tuple(value.shape[:-feature_ndim]) if feature_ndim else tuple(value.shape) + ) + if tuple(structural_mask.shape) != leading_shape: + raise ValueError( + f"mask/value leading mismatch: {structural_mask.shape} vs {leading_shape}" + ) + scan_ndim = 1 if scan_shared else 0 + if len(leading_shape) < scan_ndim + int(repeat_ndim): + raise ValueError( + "not enough leading axes for structural layout: " + f"shape={value.shape}, scan_shared={scan_shared}, " + f"repeat_ndim={repeat_ndim}" + ) + sample_end = len(leading_shape) - int(repeat_ndim) + sample_axes = tuple(range(scan_ndim, sample_end)) + repeat_axes = ((0,) if scan_shared else ()) + tuple( + range(sample_end, len(leading_shape)) + ) + feature_axes = tuple(range(len(leading_shape), value.ndim)) + permutation = (*sample_axes, *repeat_axes, *feature_axes) + mask_permutation = (*sample_axes, *repeat_axes) + value = jnp.transpose(value, permutation) if permutation else value + structural_mask = ( + jnp.transpose(structural_mask, mask_permutation) + if mask_permutation + else structural_mask + ) + logical_batch = int(math.prod(leading_shape[i] for i in sample_axes)) or 1 + repeats = int(math.prod(leading_shape[i] for i in repeat_axes)) or 1 + return ( + value.reshape(logical_batch, repeats, *feature_shape), + structural_mask.reshape(logical_batch, repeats), + logical_batch, + repeats, + ) + + +def _floor_matrix_avg_diag(mat, eps: float): + + d = mat.shape[-1] + eps_arr = jnp.asarray(eps, dtype=mat.dtype) + avg_diag = jnp.trace(mat) / d + shift = jnp.maximum(eps_arr, eps_arr - avg_diag) + return mat + shift * jnp.eye(d, dtype=mat.dtype) + + +def _floor_diag_avg(vec, eps: float): + + eps_arr = jnp.asarray(eps, dtype=vec.dtype) + avg_diag = jnp.mean(vec) + shift = jnp.maximum(eps_arr, eps_arr - avg_diag) + return vec + shift + + +def _iter_factor_update(raw_update, n_iter: int, eps: float, dtype): + + del eps + if n_iter == 1: + return jnp.ones((1, 1), dtype=dtype) + return raw_update + + +def _iter_factor_for_inverse(raw_update, n_iter: int, eps: float, dtype): + + return _floor_matrix_avg_diag( + _iter_factor_update(raw_update, n_iter, eps, dtype), + eps, + ) + + +def _validate_approx_inverse_cache_request( + exact_powers_to_cache, + approx_powers_to_cache, +): + + if exact_powers_to_cache: + raise NotImplementedError( + "Custom Kronecker blocks do not implement exact cached powers." + ) + unsupported = set(approx_powers_to_cache) - {-1} + if unsupported: + raise NotImplementedError( + f"Unsupported approximate cached powers: {sorted(unsupported)}." + ) + + +def _identity_factor(shape, dtype): + + shape = tuple(shape) + if len(shape) == 1: + return jnp.ones(shape, dtype=dtype) + if len(shape) == 2 and shape[0] == shape[1]: + return jnp.eye(shape[0], dtype=dtype) + raise ValueError(f"Unsupported Kronecker factor shape: {shape}.") + + +def _init_factor_inverse_cache( + factor_shapes, + dtype, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, + eigenvalue_count, +): + + _validate_approx_inverse_cache_request( + exact_powers_to_cache, + approx_powers_to_cache, + ) + cache = {} + if -1 in approx_powers_to_cache: + cache["-1"] = { + f"{i}_factor": _identity_factor(shape, dtype) + for i, shape in enumerate(factor_shapes) + } + if cache_eigenvalues: + cache["eigenvalues"] = jnp.zeros((eigenvalue_count,), dtype=dtype) + return cache + + +@kfac_utils.register_state_class +class _StackedRepeatedDenseState(kfac_jax.CurvatureBlock.State): + K_iter: kfac_utils.WeightedMovingAverage + A: kfac_utils.WeightedMovingAverage + G: kfac_utils.WeightedMovingAverage + average_repeats: kfac_utils.WeightedMovingAverage + + +class _StackedRepeatedDense(kfac_jax.CurvatureBlock): + State = _StackedRepeatedDenseState + + _MATPOWER_EPSILON_FLOOR: float = 1e-6 + + @property + def n_iter(self) -> int: + + return int(self.parameters_shapes[0][0]) + + @property + def in_dim(self) -> int: + + wshape = tuple(self.parameters_shapes[0][1:]) + if len(wshape) == 0: + return 1 + if len(wshape) == 1: + return 1 + return int(math.prod(wshape[:-1])) + + @property + def out_dim(self) -> int: + + wshape = tuple(self.parameters_shapes[0][1:]) + if len(wshape) == 0: + return 1 + return int(wshape[-1]) + + @property + def in_dim_aug(self) -> int: + + return self.in_dim + (1 if self.number_of_parameters == 2 else 0) + + def _init( + self, + rng, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, + ): + del rng + K = self.n_iter + cache = _init_factor_inverse_cache( + ( + (K, K), + (self.in_dim_aug, self.in_dim_aug), + (self.out_dim, self.out_dim), + ), + self.dtype, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, + self.dim, + ) + return self.State( + cache=cache, + K_iter=kfac_utils.WeightedMovingAverage( + value=jnp.eye(K, dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + A=kfac_utils.WeightedMovingAverage( + value=jnp.eye(self.in_dim_aug, dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + G=kfac_utils.WeightedMovingAverage( + value=jnp.eye(self.out_dim, dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + average_repeats=kfac_utils.WeightedMovingAverage( + value=jnp.ones((K,), dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + ) + + def sync(self, state, pmap_axis_name): + state = state.copy() + state.K_iter.sync(pmap_axis_name) + state.A.sync(pmap_axis_name) + state.G.sync(pmap_axis_name) + state.average_repeats.sync(pmap_axis_name) + return state + + def _locate_iter_axis(self, arr_shape) -> int: + + n_iter = self.n_iter + candidates = [i for i, s in enumerate(arr_shape) if s == n_iter] + if not candidates: + raise ValueError( + f"{type(self).__name__}: no axis of size n_iter={n_iter} " + f"in shape {arr_shape}. Hoist contract drifted. " + f"parameters_shapes={self.parameters_shapes!r}" + ) + return 0 if 0 in candidates else candidates[0] + + def _iter_axis_tensors(self, x, dy): + ax_x = self._locate_iter_axis(x.shape) + ax_dy = self._locate_iter_axis(dy.shape) + return jnp.moveaxis(x, ax_x, 0), jnp.moveaxis(dy, ax_dy, 0) + + def state_dependent_scale(self, state): + + repeats = jnp.mean(state.average_repeats.value) + return 1.0 / jnp.where(repeats > 0, repeats, 1.0) + + def _multiply_matpower_unscaled( + self, + state, + vector, + identity_weight, + power, + exact_power, + use_cached, + ): + if exact_power and power != 1: + raise NotImplementedError( + "StackedRepeatedDense implements approximate inverse powers only." + ) + + grad_aug = self._params_list_to_aug_array(vector) + + if power == 1: + factors = ( + _iter_factor_update( + state.K_iter.value, + self.n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ), + state.A.value, + state.G.value, + ) + scale = self.state_dependent_scale(state) if use_cached else 1.0 + new_grad_aug = kfac_utils.kronecker_product_axis_mul_v( + factors, + grad_aug, + axis_groups=[(0,), (1,), (2,)], + ) + new_grad_aug = scale * new_grad_aug + identity_weight * grad_aug + elif power == -1: + if use_cached: + inv_factors = tuple(state.cache["-1"][f"{i}_factor"] for i in range(3)) + else: + eps = self._MATPOWER_EPSILON_FLOOR + inv_factors = kfac_utils.pi_adjusted_kronecker_inverse( + _iter_factor_for_inverse( + state.K_iter.value, + self.n_iter, + eps, + self.dtype, + ), + _floor_matrix_avg_diag(state.A.value, eps), + _floor_matrix_avg_diag(state.G.value, eps), + damping=identity_weight, + ) + new_grad_aug = kfac_utils.kronecker_product_axis_mul_v( + inv_factors, + grad_aug, + axis_groups=[(0,), (1,), (2,)], + ) + else: + raise NotImplementedError( + f"StackedRepeatedDense: power={power} not implemented " + f"(only ±1 supported)." + ) + + return self._aug_array_to_params_list(new_grad_aug) + + def _params_list_to_aug_array(self, parameters_list): + + W = parameters_list[0] + W_arr = W.reshape(self.n_iter, self.in_dim, self.out_dim) + if self.number_of_parameters == 2: + b = parameters_list[1] + b_aug = b.reshape(self.n_iter, 1, self.out_dim) + return jnp.concatenate([W_arr, b_aug], axis=1) + return W_arr + + def _aug_array_to_params_list(self, arr): + + W_shape = self.parameters_shapes[0] + W = arr[:, : self.in_dim, :].reshape(W_shape) + if self.number_of_parameters == 2: + b_shape = self.parameters_shapes[1] + b = arr[:, self.in_dim :, :].reshape(b_shape) + return [W, b] + return [W] + + def _eigenvalues_unscaled(self, state, use_cached): + if use_cached: + return state.cache["eigenvalues"] + s_K, _ = kfac_utils.safe_psd_eigh( + _iter_factor_update( + state.K_iter.value, + self.n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ) + ) + s_A, _ = kfac_utils.safe_psd_eigh(state.A.value) + s_G, _ = kfac_utils.safe_psd_eigh(state.G.value) + + return jnp.einsum("k,a,o->kao", s_K, s_A, s_G).reshape(-1) + + def _update_cache( + self, + state, + identity_weight, + exact_powers, + approx_powers, + eigenvalues, + ): + _validate_approx_inverse_cache_request(exact_powers, approx_powers) + state = state.copy() + eps = self._MATPOWER_EPSILON_FLOOR + factors = ( + _iter_factor_for_inverse( + state.K_iter.value, + self.n_iter, + eps, + self.dtype, + ), + _floor_matrix_avg_diag(state.A.value, eps), + _floor_matrix_avg_diag(state.G.value, eps), + ) + scale = self.state_dependent_scale(state) + + if eigenvalues: + state.cache["eigenvalues"] = scale * self._eigenvalues_unscaled( + state, use_cached=False + ) + + if -1 in approx_powers: + inv_factors = kfac_utils.pi_adjusted_kronecker_inverse( + *factors, + damping=identity_weight, + ) + factor_scale = jnp.power(scale, 1.0 / len(factors)) + for i, inv_factor in enumerate(inv_factors): + state.cache["-1"][f"{i}_factor"] = inv_factor / factor_scale + + return state + + def _to_dense_unscaled(self, state): + + F_KA = jnp.kron( + _iter_factor_update( + state.K_iter.value, + self.n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ), + state.A.value, + ) + return jnp.kron(F_KA, state.G.value) + + def _norm_unscaled(self, state, norm_type): + n_K = kfac_utils.psd_matrix_norm( + _iter_factor_update( + state.K_iter.value, + self.n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ), + norm_type=norm_type, + ) + n_A = kfac_utils.psd_matrix_norm( + state.A.value, + norm_type=norm_type, + ) + n_G = kfac_utils.psd_matrix_norm( + state.G.value, + norm_type=norm_type, + ) + return n_K * n_A * n_G + + +STACKED_SCALE_SHIFT_TAG_VARIANT = "stacked_scale_and_shift" + + +@kfac_utils.register_state_class +class _StackedScaleAndShiftState(kfac_jax.CurvatureBlock.State): + K_iter_factors: tuple[kfac_utils.WeightedMovingAverage, ...] + D_shared_factors: tuple[kfac_utils.WeightedMovingAverage, ...] + + +class _StackedScaleAndShiftDiagonal(kfac_jax.CurvatureBlock): + State = _StackedScaleAndShiftState + _MATPOWER_EPSILON_FLOOR: float = 1e-6 + + @property + def n_iter(self) -> int: + return int(self.parameters_shapes[0][0]) + + @property + def _per_iter_shapes(self) -> tuple[tuple[int, ...], ...]: + + return (tuple(self.parameters_shapes[0][1:]),) + + @property + def _per_iter_d_flats(self) -> tuple[int, ...]: + + shape = self._per_iter_shapes[0] + return (int(math.prod(shape)) if shape else 1,) + + def _locate_iter_axis(self, arr_shape) -> int: + n_iter = self.n_iter + candidates = [i for i, s in enumerate(arr_shape) if s == n_iter] + if not candidates: + raise ValueError( + f"{type(self).__name__}: no axis of size n_iter={n_iter} " + f"in shape {arr_shape}." + ) + + return 0 if 0 in candidates else candidates[0] + + def _iter_axis_tensors(self, x, dy): + ax_x = self._locate_iter_axis(x.shape) + ax_dy = self._locate_iter_axis(dy.shape) + return jnp.moveaxis(x, ax_x, 0), jnp.moveaxis(dy, ax_dy, 0) + + def _init( + self, + rng, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, + ): + del rng + K = self.n_iter + d = self._per_iter_d_flats[0] + cache = _init_factor_inverse_cache( + ((K, K), (d,)), + self.dtype, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, + self.dim, + ) + return self.State( + cache=cache, + K_iter_factors=( + kfac_utils.WeightedMovingAverage( + value=jnp.eye(K, dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + ), + D_shared_factors=( + kfac_utils.WeightedMovingAverage( + value=jnp.ones((d,), dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + ), + ) + + def sync(self, state, pmap_axis_name): + state = state.copy() + state.K_iter_factors[0].sync(pmap_axis_name) + state.D_shared_factors[0].sync(pmap_axis_name) + return state + + @kfac_utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight, batch_size + state = state.copy() + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + x_iter, dy_iter = self._iter_axis_tensors(x, dy) + mask_iter = 1.0 - jnp.all( + dy_iter == 0.0, + axis=-1, + keepdims=True, + ) + n_iter = self.n_iter + per_iter_shape = self._per_iter_shapes[0] + d_flat = self._per_iter_d_flats[0] + + def _per_iter(arr_i): + return cb_utils.compatible_sum( + arr_i, + per_iter_shape, + skip_axes=[0], + ) + + d_grad = jax.vmap(_per_iter)(x_iter * dy_iter).reshape( + n_iter, + -1, + d_flat, + ) + mask = jnp.any( + mask_iter.reshape(n_iter, mask_iter.shape[1], -1) > 0, + axis=-1, + ).astype(self.dtype) + d_grad = d_grad * mask[..., None] + n_active = jnp.maximum(jnp.sum(mask), 1.0).astype(self.dtype) + D_update = jnp.einsum("kbi,kbi->i", d_grad, d_grad) / n_active + if n_iter == 1: + K_update = jnp.ones((1, 1), dtype=self.dtype) + else: + weighted = d_grad * jnp.sqrt(jnp.maximum(D_update, 0.0))[None, None, :] + numerator = jnp.einsum("kbi,lbi->kl", weighted, weighted) + per_iter_active = jnp.sum(mask, axis=-1) + active_norm = jnp.sqrt( + jnp.maximum( + per_iter_active[:, None] * per_iter_active[None, :], + 1.0, + ) + ).astype(self.dtype) + D_frob2 = jnp.maximum(jnp.sum(D_update * D_update), 1e-12) + K_update = numerator / (active_norm * D_frob2) + K_update = _iter_factor_update( + 0.5 * (K_update + K_update.T), + n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ) + state.K_iter_factors[0].update(K_update, ema_old, ema_new) + state.D_shared_factors[0].update(D_update, ema_old, ema_new) + return state + + def _multiply_matpower_unscaled( + self, + state, + vector, + identity_weight, + power, + exact_power, + use_cached, + ): + if exact_power and power != 1: + raise NotImplementedError( + "StackedScaleAndShiftDiagonal implements approximate " + "inverse powers only." + ) + n_iter = self.n_iter + v = vector[0] + v_flat = v.reshape(n_iter, -1) + if power == -1 and use_cached: + K_iter_inv = state.cache["-1"]["0_factor"] + D_shared_inv = state.cache["-1"]["1_factor"] + Kv = jnp.einsum("kl,li->ki", K_iter_inv, v_flat) + result_flat = D_shared_inv[None, :] * Kv + elif power == 1: + K_factor = _iter_factor_update( + state.K_iter_factors[0].value, + n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ) + D_factor = state.D_shared_factors[0].value + Kv = jnp.einsum("kl,li->ki", K_factor, v_flat) + result_flat = D_factor[None, :] * Kv + identity_weight * v_flat + elif power == -1: + eps = self._MATPOWER_EPSILON_FLOOR + K_floored = _iter_factor_for_inverse( + state.K_iter_factors[0].value, + n_iter, + eps, + self.dtype, + ) + D_floored = _floor_diag_avg( + state.D_shared_factors[0].value, + eps, + ) + shrink = jnp.maximum( + 1.0, + jnp.mean(D_floored) / identity_weight, + ) + D_floored = D_floored / shrink + K_iter_inv, D_shared_inv = kfac_utils.pi_adjusted_kronecker_inverse( + K_floored, + D_floored, + damping=identity_weight, + ) + Kv = jnp.einsum("kl,li->ki", K_iter_inv, v_flat) + result_flat = D_shared_inv[None, :] * Kv + else: + raise NotImplementedError( + f"StackedScaleAndShiftDiagonal: power={power} not " + f"implemented (only ±1 supported)." + ) + return (result_flat.reshape(v.shape),) + + def _eigenvalues_unscaled(self, state, use_cached): + if use_cached: + return state.cache["eigenvalues"] + s_K, _ = kfac_utils.safe_psd_eigh( + _iter_factor_update( + state.K_iter_factors[0].value, + self.n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ) + ) + return jnp.einsum( + "k,i->ki", + s_K, + state.D_shared_factors[0].value, + ).reshape(-1) + + def _update_cache( + self, + state, + identity_weight, + exact_powers, + approx_powers, + eigenvalues, + ): + _validate_approx_inverse_cache_request(exact_powers, approx_powers) + state = state.copy() + + if eigenvalues: + state.cache["eigenvalues"] = self._eigenvalues_unscaled( + state, + use_cached=False, + ) + + if -1 in approx_powers: + eps = self._MATPOWER_EPSILON_FLOOR + K_floored = _iter_factor_for_inverse( + state.K_iter_factors[0].value, + self.n_iter, + eps, + self.dtype, + ) + D_floored = _floor_diag_avg( + state.D_shared_factors[0].value, + eps, + ) + shrink = jnp.maximum( + 1.0, + jnp.mean(D_floored) / identity_weight, + ) + D_floored = D_floored / shrink + K_inv, D_inv = kfac_utils.pi_adjusted_kronecker_inverse( + K_floored, + D_floored, + damping=identity_weight, + ) + state.cache["-1"]["0_factor"] = K_inv + state.cache["-1"]["1_factor"] = D_inv + + return state + + def _to_dense_unscaled(self, state): + return jnp.kron( + _iter_factor_update( + state.K_iter_factors[0].value, + self.n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ), + jnp.diag(state.D_shared_factors[0].value), + ) + + def _norm_unscaled(self, state, norm_type): + if norm_type in ("trace", "avg_diag"): + component_norm = "trace" + elif norm_type in ("fro", "avg_fro"): + component_norm = "fro" + else: + component_norm = norm_type + + K = _iter_factor_update( + state.K_iter_factors[0].value, + self.n_iter, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ) + D = state.D_shared_factors[0].value + if component_norm == "trace": + norm = jnp.trace(K) * jnp.sum(D) + elif component_norm == "fro": + norm = jnp.linalg.norm(K) * jnp.linalg.norm(D) + elif component_norm == "2_norm": + norm = jnp.max(jnp.linalg.eigvalsh(K)) * jnp.max(D) + elif component_norm == "1_norm": + norm = jnp.max(jnp.sum(jnp.abs(K), axis=0)) * jnp.max(jnp.abs(D)) + elif component_norm == "one_over_dim": + norm = jnp.asarray(1.0, dtype=self.dtype) + else: + raise NotImplementedError( + f"Kronecker norm {norm_type!r} is not needed by KFAC stats" + ) + total_dim = self.n_iter * self._per_iter_d_flats[0] + if norm_type == "trace": + return norm + if norm_type == "avg_diag": + return norm / total_dim + if norm_type == "one_over_dim": + return jnp.asarray(1.0 / total_dim, dtype=self.dtype) + if norm_type in ("2_norm", "1_norm"): + return norm + if norm_type in ("fro", "avg_fro"): + return norm if norm_type == "fro" else norm / jnp.sqrt(total_dim) + raise NotImplementedError( + f"direct-sum norm {norm_type!r} is not needed by KFAC stats" + ) + + +class _ScaleAndShiftDiagonal(kfac_jax.ScaleAndShiftDiagonal): + @kfac_utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight, batch_size + state = state.copy() + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + mask = 1.0 - jnp.all(dy == 0.0, axis=-1, keepdims=True) + x_masked = x * mask + n_active = jnp.maximum(jnp.sum(mask), 1.0).astype(x.dtype) + if self.has_scale: + scale_shape = estimation_data.primals.params[0].shape + n_param_dims = len(scale_shape) + x_flat = x_masked.reshape( + (-1,) + tuple(x_masked.shape[-n_param_dims:]) if n_param_dims else (-1,) + ) + dy_flat = dy.reshape( + (-1,) + tuple(dy.shape[-n_param_dims:]) if n_param_dims else (-1,) + ) + d_scale = cb_utils.compatible_sum( + x_flat * dy_flat, + scale_shape, + skip_axes=[0], + ) + scale_diag_update = ( + jnp.sum( + d_scale * d_scale, + axis=0, + keepdims=d_scale.ndim == len(scale_shape), + ) + / n_active + ) + state.diagonal_factors[0].update( + scale_diag_update, + ema_old, + ema_new, + ) + if self.has_shift: + shift_shape = estimation_data.primals.params[-1].shape + n_param_dims = len(shift_shape) + dy_flat = dy.reshape( + (-1,) + tuple(dy.shape[-n_param_dims:]) if n_param_dims else (-1,) + ) + d_shift = cb_utils.compatible_sum( + dy_flat, + shift_shape, + skip_axes=[0], + ) + shift_diag_update = ( + jnp.sum( + d_shift * d_shift, + axis=0, + keepdims=d_shift.ndim == len(shift_shape), + ) + / n_active + ) + state.diagonal_factors[-1].update( + shift_diag_update, + ema_old, + ema_new, + ) + return state + + def _norm_unscaled(self, state, norm_type): + diagonal = jnp.concatenate( + [factor.value.flatten() for factor in state.diagonal_factors], + axis=0, + ) + return kfac_utils.psd_matrix_norm( + diagonal, + norm_type=norm_type, + ) + + def _multiply_matpower_unscaled( + self, + state, + vector, + identity_weight, + power, + exact_power, + use_cached, + ): + scale = self.state_dependent_scale(state) if use_cached else 1.0 + factors = [] + for diagonal_factor in state.diagonal_factors: + value = scale * diagonal_factor.value + shrink = jnp.maximum(1.0, jnp.mean(value) / identity_weight) + factors.append(value / shrink + identity_weight) + assert len(factors) == len(vector) + if power == 1: + return tuple(factor * value for factor, value in zip(factors, vector)) + elif power == -1: + return tuple(value / factor for factor, value in zip(factors, vector)) + return tuple( + jnp.power(factor, power) * value for factor, value in zip(factors, vector) + ) + + +class StructuralRepeatedDenseKroneckerFactored( + kfac_jax.RepeatedDenseKroneckerFactored, +): + def state_dependent_scale(self, state): + repeats = state.average_repeats.value + return 1.0 / jnp.where(repeats > 0, repeats, 1.0) + + @kfac_utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight + state = state.copy() + x, structural_mask = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + scan_shared, repeat_ndim, context_reuse = _structural_tag_contract( + self._layer_tag_eq + ) + try: + x, dy, structural_mask = _align_structural_primal_and_mask( + x, + dy, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=1, + context_primal_reused_over_walkers=context_reuse, + ) + except ValueError as error: + meta = self._layer_tag_eq.params.get("meta") + raise ValueError( + f"{error}; structural dense tag=" + f"{getattr(meta, 'name', None)!r}, scan_shared={scan_shared}, " + f"repeat_ndim={repeat_ndim}, context_reuse={context_reuse}, " + f"x={x.shape}, dy={dy.shape}, mask={structural_mask.shape}" + ) from error + xg, mg, logical_batch, _ = structural_group_repeats( + x, + structural_mask, + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + dyg, _, dy_batch, _ = structural_group_repeats( + dy, + structural_mask, + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + if logical_batch != dy_batch: + raise ValueError( + f"x/dy logical batch mismatch: {logical_batch} vs {dy_batch}" + ) + mask = mg.astype(xg.dtype)[..., None] + xg = xg * mask + dyg = dyg * mask.astype(dyg.dtype) + x_flat = xg.reshape((-1, xg.shape[-1])) + dy_flat = dyg.reshape((-1, dyg.shape[-1])) + if self.number_of_parameters == 2: + x_flat = jnp.concatenate( + [x_flat, mask.reshape((-1, 1))], + axis=-1, + ) + logical_divisor = jnp.asarray(logical_batch, dtype=x_flat.dtype) + global_divisor = jnp.asarray(batch_size, dtype=dy_flat.dtype) + input_stats = jnp.einsum("ai,aj->ij", x_flat, x_flat) / logical_divisor + output_stats = jnp.einsum("ao,ap->op", dy_flat, dy_flat) / global_divisor + average_repeats = jnp.sum(mask) / logical_divisor + state.factors[0].update(input_stats, ema_old, ema_new) + state.factors[1].update(output_stats, ema_old, ema_new) + + state.average_repeats.update( + average_repeats, + ema_old, + ema_new, + ) + return state + + +class StructuralScaleAndShiftDiagonal(_ScaleAndShiftDiagonal): + @kfac_utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight + state = state.copy() + x, structural_mask = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + scan_shared, repeat_ndim, context_reuse = _structural_tag_contract( + self._layer_tag_eq + ) + reference_param = estimation_data.primals.params[0] + feature_ndim = reference_param.ndim + try: + x, dy, structural_mask = _align_structural_primal_and_mask( + x, + dy, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=feature_ndim, + context_primal_reused_over_walkers=context_reuse, + ) + except ValueError as error: + meta = self._layer_tag_eq.params.get("meta") + raise ValueError( + f"{error}; structural scale/shift tag=" + f"{getattr(meta, 'name', None)!r}, scan_shared={scan_shared}, " + f"repeat_ndim={repeat_ndim}, context_reuse={context_reuse}, " + f"x={x.shape}, dy={dy.shape}, mask={structural_mask.shape}" + ) from error + xg, mg, logical_batch, _ = structural_group_repeats( + x, + structural_mask, + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=feature_ndim, + ) + dyg, _, dy_batch, _ = structural_group_repeats( + dy, + structural_mask, + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=feature_ndim, + ) + if logical_batch != dy_batch: + raise ValueError( + f"x/dy logical batch mismatch: {logical_batch} vs {dy_batch}" + ) + mask = mg.astype(dyg.dtype) + mask = mask.reshape((*mask.shape, *(1,) * feature_ndim)) + xg = xg * mask.astype(xg.dtype) + dyg = dyg * mask + divisor = jnp.asarray(batch_size, dtype=dyg.dtype) + param_index = 0 + if self.has_scale: + d_scale = jnp.sum(xg * dyg, axis=1) + scale_update = jnp.sum(d_scale * d_scale, axis=0) / divisor + state.diagonal_factors[param_index].update( + scale_update, + ema_old, + ema_new, + ) + param_index += 1 + if self.has_shift: + d_shift = jnp.sum(dyg, axis=1) + shift_update = jnp.sum(d_shift * d_shift, axis=0) / divisor + state.diagonal_factors[param_index].update( + shift_update, + ema_old, + ema_new, + ) + return state + + +class StructuralStackedRepeatedDense(_StackedRepeatedDense): + def state_dependent_scale(self, state): + repeats = jnp.mean(state.average_repeats.value) + return 1.0 / jnp.where(repeats > 0, repeats, 1.0) + + @kfac_utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight + state = state.copy() + x, structural_mask = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + scan_shared, repeat_ndim, context_reuse = _structural_tag_contract( + self._layer_tag_eq + ) + x, dy, structural_mask = _align_structural_primal_and_mask( + x, + dy, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=1, + context_primal_reused_over_walkers=context_reuse, + ) + ax_x = self._locate_iter_axis(x.shape) + ax_dy = self._locate_iter_axis(dy.shape) + ax_mask = self._locate_iter_axis(structural_mask.shape) + x_iter = jnp.moveaxis(x, ax_x, 0) + dy_iter = jnp.moveaxis(dy, ax_dy, 0) + mask_iter = jnp.moveaxis(structural_mask, ax_mask, 0) + + x_groups = [] + dy_groups = [] + mask_groups = [] + logical_batch = None + for k in range(self.n_iter): + xg, mg, B, _ = structural_group_repeats( + x_iter[k], + mask_iter[k], + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + dyg, _, B_dy, _ = structural_group_repeats( + dy_iter[k], + mask_iter[k], + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + if B != B_dy or (logical_batch is not None and B != logical_batch): + raise ValueError("stacked structural logical batches differ") + logical_batch = B + x_groups.append(xg) + dy_groups.append(dyg) + mask_groups.append(mg) + x_group = jnp.stack(x_groups, axis=0) + dy_group = jnp.stack(dy_groups, axis=0) + mask_group = jnp.stack(mask_groups, axis=0).astype(x_group.dtype) + x_group = x_group * mask_group[..., None] + dy_group = dy_group * mask_group[..., None].astype(dy_group.dtype) + if self.number_of_parameters == 2: + x_group = jnp.concatenate( + [x_group, mask_group[..., None]], + axis=-1, + ) + K = self.n_iter + logical_batch = int(logical_batch or 1) + logical_divisor = jnp.asarray(K * logical_batch, dtype=self.dtype) + global_divisor = jnp.asarray(K * batch_size, dtype=self.dtype) + x_flat = x_group.reshape(K, -1, x_group.shape[-1]) + dy_flat = dy_group.reshape(K, -1, dy_group.shape[-1]) + A_update = jnp.einsum("kbi,kbj->ij", x_flat, x_flat) / logical_divisor + G_update = jnp.einsum("kbo,kbp->op", dy_flat, dy_flat) / global_divisor + per_iter_active = jnp.sum(mask_group, axis=(1, 2)) + if K == 1: + K_iter_update = jnp.ones((1, 1), dtype=self.dtype) + else: + if context_reuse: + A_projection = state.A.value + G_projection = state.G.value + else: + A_projection = A_update + G_projection = G_update + xA = jnp.einsum("kbi,ij->kbj", x_flat, A_projection) + dyG = jnp.einsum("kbo,op->kbp", dy_flat, G_projection) + numerator = jnp.einsum( + "klb,klb->kl", + jnp.einsum("kbi,lbi->klb", xA, x_flat), + jnp.einsum("kbo,lbo->klb", dyG, dy_flat), + ) + mean_repeats = jnp.mean(per_iter_active) / jnp.asarray( + logical_batch, self.dtype + ) + denominator = ( + jnp.asarray(batch_size, self.dtype) + * jnp.maximum( + jnp.sum(A_projection * A_projection), + 1e-12, + ) + * jnp.maximum( + jnp.sum(G_projection * G_projection), + 1e-12, + ) + ) + K_iter_update = mean_repeats * numerator / denominator + K_iter_update = _iter_factor_update( + 0.5 * (K_iter_update + K_iter_update.T), + K, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ) + state.K_iter.update(K_iter_update, ema_old, ema_new) + state.A.update(A_update, ema_old, ema_new) + state.G.update(G_update, ema_old, ema_new) + state.average_repeats.update( + per_iter_active / jnp.asarray(logical_batch, self.dtype), + ema_old, + ema_new, + ) + return state + + +class StructuralStackedScaleAndShiftDiagonal(_StackedScaleAndShiftDiagonal): + def _structural_iter_axis(self, shape) -> int: + if shape and int(shape[0]) == self.n_iter: + return 0 + return self._locate_iter_axis(shape) + + @kfac_utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight + state = state.copy() + x, structural_mask = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + scan_shared, repeat_ndim, context_reuse = _structural_tag_contract( + self._layer_tag_eq + ) + + feature_ndim = len(self._per_iter_shapes[0]) + x, dy, structural_mask = _align_structural_primal_and_mask( + x, + dy, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=feature_ndim, + context_primal_reused_over_walkers=context_reuse, + ) + x_iter = jnp.moveaxis(x, self._structural_iter_axis(x.shape), 0) + dy_iter = jnp.moveaxis(dy, self._structural_iter_axis(dy.shape), 0) + mask_iter = jnp.moveaxis( + structural_mask, + self._structural_iter_axis(structural_mask.shape), + 0, + ) + self._update_structural_scale( + state, + x_iter, + dy_iter, + mask_iter, + self._per_iter_shapes[0], + scan_shared, + repeat_ndim, + context_reuse, + batch_size, + ema_old, + ema_new, + ) + return state + + def _update_structural_scale( + self, + state, + x_iter, + dy_iter, + mask_iter, + per_iter_shape, + scan_shared, + repeat_ndim, + context_reuse, + batch_size, + ema_old, + ema_new, + ): + K = self.n_iter + feature_ndim = len(per_iter_shape) + grads = [] + logical_batch = None + for k in range(K): + xg, mg, B, _ = structural_group_repeats( + x_iter[k], + mask_iter[k], + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=feature_ndim, + ) + dyg, _, B_dy, _ = structural_group_repeats( + dy_iter[k], + mask_iter[k], + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=feature_ndim, + ) + if B != B_dy or (logical_batch is not None and B != logical_batch): + raise ValueError("stacked scale logical batches differ") + logical_batch = B + mask = mg.astype(dyg.dtype).reshape((*mg.shape, *(1,) * feature_ndim)) + row_grad = xg * dyg * mask + grads.append(jnp.sum(row_grad, axis=1).reshape(B, -1)) + grad = jnp.stack(grads, axis=0) + logical_batch = int(logical_batch or 1) + D_update = jnp.einsum("kbi,kbi->i", grad, grad) / jnp.asarray( + K * batch_size, + self.dtype, + ) + if K == 1: + K_update = jnp.ones((1, 1), dtype=self.dtype) + else: + D_projection = ( + state.D_shared_factors[0].value if context_reuse else D_update + ) + weighted = grad * jnp.sqrt(jnp.maximum(D_projection, 0.0))[None, None, :] + numerator = jnp.einsum("kbi,lbi->kl", weighted, weighted) + denom = jnp.asarray(batch_size, self.dtype) * jnp.maximum( + jnp.sum(D_projection * D_projection), + 1e-12, + ) + K_update = _iter_factor_update( + 0.5 * (numerator / denom + (numerator / denom).T), + K, + self._MATPOWER_EPSILON_FLOOR, + self.dtype, + ) + state.K_iter_factors[0].update(K_update, ema_old, ema_new) + state.D_shared_factors[0].update(D_update, ema_old, ema_new) + + +class StructuralTrailingStackedScaleAndShiftDiagonal( + StructuralStackedScaleAndShiftDiagonal, +): + @kfac_utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight + state = state.copy() + x, structural_mask = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + _, repeat_ndim, context_reuse = _structural_tag_contract(self._layer_tag_eq) + x, dy, structural_mask = _align_structural_primal_and_mask( + x, + dy, + structural_mask, + repeat_ndim=repeat_ndim, + feature_ndim=2, + context_primal_reused_over_walkers=context_reuse, + ) + K = self.n_iter + if int(x.shape[-2]) != K or int(dy.shape[-2]) != K: + raise ValueError( + f"{type(self).__name__}: expected trailing K={K} axis, " + f"got x={x.shape}, dy={dy.shape}" + ) + x_iter = jnp.moveaxis(x, -2, 0) + dy_iter = jnp.moveaxis(dy, -2, 0) + mask_with_groups = jnp.broadcast_to( + structural_mask[..., None], + x.shape[:-1], + ) + mask_iter = jnp.moveaxis(mask_with_groups, -1, 0) + + self._update_structural_scale( + state, + x_iter, + dy_iter, + mask_iter, + self._per_iter_shapes[0], + False, + repeat_ndim, + context_reuse, + batch_size, + ema_old, + ema_new, + ) + return state + + +class _DenseBlock(kfac_jax.DenseTwoKroneckerFactored): + def update_curvature_matrix_estimate( + self, + state, + estimation_data, + ema_old, + ema_new, + identity_weight, + batch_size, + ): + del identity_weight + state = state.copy() + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + if not kfac_jax.utils.first_dim_is_size(batch_size, x, dy): + x, dy = ( + jnp.tile(a[None], (batch_size, *(1 for _ in a.shape))).reshape( + (-1, a.shape[-1]) + ) + for a in (x, dy) + ) + batch_size = x.size // x.shape[-1] + assert kfac_jax.utils.first_dim_is_size(batch_size, x, dy) + mask = 1.0 - jnp.all(dy == 0.0, axis=-1, keepdims=True) + x = x * mask + n_active = jnp.maximum(jnp.sum(mask), 1.0).astype(x.dtype) + x = x.reshape((-1, x.shape[-1])) + dy = dy.reshape((-1, dy.shape[-1])) + input_stats = jnp.einsum("ay,az->yz", x, x) / n_active + output_stats = jnp.einsum("ay,az->yz", dy, dy) / n_active + state.factors[0].update(input_stats, ema_old, ema_new) + state.factors[1].update(output_stats, ema_old, ema_new) + return state + + +kfac_jax.set_default_tag_to_block_ctor("dense", _DenseBlock) +kfac_jax.set_default_tag_to_block_ctor( + "scale_and_shift", + _ScaleAndShiftDiagonal, +) +kfac_jax.set_default_tag_to_block_ctor( + STACKED_SCALE_SHIFT_TAG_VARIANT, + _StackedScaleAndShiftDiagonal, +) +kfac_jax.set_default_tag_to_block_ctor( + STRUCTURAL_DENSE_TAG_VARIANT, + StructuralRepeatedDenseKroneckerFactored, +) +kfac_jax.set_default_tag_to_block_ctor( + STRUCTURAL_SCALE_SHIFT_TAG_VARIANT, + StructuralScaleAndShiftDiagonal, +) +kfac_jax.set_default_tag_to_block_ctor( + STRUCTURAL_STACKED_DENSE_TAG_VARIANT, + StructuralStackedRepeatedDense, +) +kfac_jax.set_default_tag_to_block_ctor( + STRUCTURAL_STACKED_SCALE_SHIFT_TAG_VARIANT, + StructuralStackedScaleAndShiftDiagonal, +) +kfac_jax.set_default_tag_to_block_ctor( + STRUCTURAL_TRAILING_STACKED_SCALE_SHIFT_TAG_VARIANT, + StructuralTrailingStackedScaleAndShiftDiagonal, +) + + +def make_graph_patterns(): + return () + + +__all__ = [ + "STRUCTURAL_DENSE_TAG_VARIANT", + "STRUCTURAL_SCALE_SHIFT_TAG_VARIANT", + "STRUCTURAL_STACKED_DENSE_TAG_VARIANT", + "STRUCTURAL_STACKED_SCALE_SHIFT_TAG_VARIANT", + "STRUCTURAL_TRAILING_STACKED_SCALE_SHIFT_TAG_VARIANT", + "StructuralRepeatedDenseKroneckerFactored", + "StructuralScaleAndShiftDiagonal", + "StructuralStackedRepeatedDense", + "StructuralStackedScaleAndShiftDiagonal", + "StructuralTrailingStackedScaleAndShiftDiagonal", + "make_graph_patterns", + "register_structural_dense", + "register_structural_scale_and_shift", + "register_structural_trailing_stacked_scale_and_shift", + "structural_group_repeats", +] diff --git a/src/hamiltonzero/optim/compat.py b/src/hamiltonzero/optim/compat.py new file mode 100644 index 0000000000000000000000000000000000000000..ff3c12bae2832a65c8b5a54d18c6c0afc5e5f027 --- /dev/null +++ b/src/hamiltonzero/optim/compat.py @@ -0,0 +1,926 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + + +from __future__ import annotations + + +def _scan_partition_sizes(eqn) -> tuple[int, int, int]: + + consts, carry, xs = eqn.params["ft_in"].unpack() + return len(consts), len(carry), len(xs) + + +def _extend_scan_flat_trees( + params: dict, + *, + extra_xs: int, + extra_ys: int, +) -> None: + + from jax._src import flattree as _ft + + consts, carry, xs = params["ft_in"].unpack() + carry_out, ys = params["ft_out"].unpack() + if extra_xs: + xs = _ft.pack((xs, _ft.nones(extra_xs))) + if extra_ys: + ys = _ft.pack((ys, _ft.nones(extra_ys))) + params["ft_in"] = _ft.pack((consts, carry, xs)) + params["ft_out"] = _ft.pack((carry_out, ys)) + + +def _patch() -> None: + + from jax._src import source_info_util as jex_source_info_util + from kfac_jax._src import tag_graph_matcher as tgm + + if getattr(tgm.eval_jaxpr_eqn, "__hamiltonzero_patched__", False): + return + + def eval_jaxpr_eqn(eqn, in_values): + bind_params = eqn.primitive.get_bind_params(eqn.params) + user_context = jex_source_info_util.user_context + with user_context(eqn.source_info.traceback): + output = eqn.primitive.bind(*in_values, **bind_params) + return [output] if not isinstance(output, list) else output + + eval_jaxpr_eqn.__hamiltonzero_patched__ = True + tgm.eval_jaxpr_eqn = eval_jaxpr_eqn + + +def _patch_allow_multiple_registrations() -> None: + + import threading + + from kfac_jax._src import tag_graph_matcher as tgm + + if getattr(tgm.auto_register_tags, "__hamiltonzero_allow_multi__", False): + return + + _orig_auto = tgm.auto_register_tags + _orig_check = tgm.TaggedFunction.check_multiple_registrations + _state = threading.local() + + def auto_register_tags( + func, + func_args, + *, + allow_multiple_registrations: bool = False, + **kwargs, + ): + prev = getattr(_state, "allow", False) + _state.allow = bool(allow_multiple_registrations) + try: + return _orig_auto(func, func_args, **kwargs) + finally: + _state.allow = prev + + def check_multiple_registrations(self): + if getattr(_state, "allow", False): + return + return _orig_check(self) + + auto_register_tags.__hamiltonzero_allow_multi__ = True + tgm.auto_register_tags = auto_register_tags + tgm.TaggedFunction.check_multiple_registrations = check_multiple_registrations + + +def _patch_orphan_registration_in_sub_graphs() -> None: + + from kfac_jax._src import tag_graph_matcher as tgm + + if getattr(tgm._auto_register_tags, "__hamiltonzero_orphan_subgraph__", False): + return + + _orig = tgm._auto_register_tags + + def _patched(graph, *args, register_orphans=True, **kwargs): + + return _orig(graph, *args, register_orphans=True, **kwargs) + + _patched.__hamiltonzero_orphan_subgraph__ = True + tgm._auto_register_tags = _patched + + +def _patch_manual_tag_outputs_that_are_graph_inputs() -> None: + + from kfac_jax._src import tag_graph_matcher as tgm + import jax.extend as jex + + graph_cls = tgm.JaxprGraph + if getattr( + graph_cls.sub_graph_eqns, + "__hamiltonzero_graph_input_tag_output__", + False, + ): + return + original = graph_cls.sub_graph_eqns + + def sub_graph_eqns(self, root_vars, leaf_vars): + kept = [] + for value in leaf_vars: + if ( + isinstance(value, jex.core.Literal) + or value in self.params_vars + or value in self.var_to_creation_op + ): + kept.append(value) + elif value in self.jaxpr.invars: + continue + else: + raise KeyError(value) + return original(self, root_vars, tuple(kept)) + + sub_graph_eqns.__hamiltonzero_graph_input_tag_output__ = True + graph_cls.sub_graph_eqns = sub_graph_eqns + + +def _patch_hoist_layer_tags_from_scan() -> None: + + from kfac_jax._src import tag_graph_matcher as tgm + from kfac_jax._src import layers_and_loss_tags as tags + import jax + from jax._src import core as _jcore + from kfac_jax._src.tag_graph_matcher import ( + ClosedJaxpr, + HIGHER_ORDER_NAMES, + to_closed_jaxpr, + to_jaxpr_or_closed_jaxpr, + ) + + from jax.extend.core import gensym, new_jaxpr_eqn + + if getattr(tgm.clean_layer_tags_jaxpr, "__hamiltonzero_hoist_tags__", False): + return + + _orig_clean_layer = tgm.clean_layer_tags_jaxpr + + def _make_transpose_swap01_eqn(scan_outvar, make_var_func): + + ndim = len(scan_outvar.aval.shape) + if ndim < 2: + return None, scan_outvar + from jax._src.lax import lax as _jlax + + permutation = (1, 0) + tuple(range(2, ndim)) + new_shape = tuple(scan_outvar.aval.shape[p] for p in permutation) + new_aval = _jcore.ShapedArray(new_shape, scan_outvar.aval.dtype) + new_outvar = make_var_func(new_aval) + eqn = new_jaxpr_eqn( + invars=[scan_outvar], + outvars=[new_outvar], + primitive=_jlax.transpose_p, + params={"permutation": permutation}, + effects=frozenset(), + ) + return eqn, new_outvar + + def _hoist_tags_recursive(closed_jaxpr, make_var_func): + + new_eqns = [] + hoisted_tags = [] + + for eqn in closed_jaxpr.jaxpr.eqns: + if eqn.primitive.name not in HIGHER_ORDER_NAMES: + new_eqns.append(eqn) + continue + if eqn.primitive.name == "cond": + new_eqns.append(eqn) + continue + if eqn.primitive.name == "while": + body_jaxpr = eqn.params["body_jaxpr"] + key = "body_jaxpr" + supports_extension = False + elif eqn.primitive.name == "scan": + body_jaxpr = eqn.params["jaxpr"] + key = "jaxpr" + supports_extension = True + elif eqn.primitive.name == "pjit": + body_jaxpr = eqn.params["jaxpr"] + key = "jaxpr" + supports_extension = False + elif eqn.primitive.name in ("xla_call", "xla_pmap"): + body_jaxpr = eqn.params["call_jaxpr"] + key = "call_jaxpr" + supports_extension = False + else: + new_eqns.append(eqn) + continue + + body_closed = to_closed_jaxpr(body_jaxpr) + new_body_closed, nested_hoisted = _hoist_tags_recursive( + body_closed, make_var_func + ) + + body_invars = body_jaxpr.jaxpr.invars + body_eqns_no_tags = [] + tag_var_map = {} + + new_body_captures: list = [] + new_body_capture_id_to_idx: dict[int, int] = {} + + output_var_to_aux_xs: dict[int, tuple] = {} + + scan_length = eqn.params["length"] if eqn.primitive.name == "scan" else None + + deferred_specs: list = [] + + import jax.extend as _jex_chain + + def _resolve_tag_chain(w): + + while not isinstance(w, _jex_chain.core.Literal) and w in tag_var_map: + w = tag_var_map[w] + return w + + for body_eqn in new_body_closed.jaxpr.eqns: + if not isinstance(body_eqn.primitive, tags.LayerTag): + body_eqns_no_tags.append(body_eqn) + continue + + meta = body_eqn.params["meta"] + for ind1, ind2 in enumerate(meta.outputs_index): + tag_var_map[body_eqn.outvars[ind1]] = body_eqn.invars[ind2] + + params_index_set = set(meta.params_index) + partial_invars: list = [None] * len(body_eqn.invars) + deferred: list = [] + hoistable = True + + ( + _scan_num_consts, + _scan_num_carry, + _, + ) = _scan_partition_sizes(eqn) + _xs_threshold = _scan_num_consts + _scan_num_carry + + tag_has_xs_iterated_params = False + output_idx_set = set(meta.outputs_index) + + for _arg_idx_pre, _v_pre in enumerate(body_eqn.invars): + if _arg_idx_pre not in params_index_set: + continue + _v_pre_resolved = _resolve_tag_chain(_v_pre) + if _v_pre_resolved in body_invars: + _idx = body_invars.index(_v_pre_resolved) + if eqn.primitive.name == "scan" and _idx >= _xs_threshold: + tag_has_xs_iterated_params = True + break + + use_aux_xs_for_outputs = supports_extension and scan_length is not None + + use_accumulating_aux_base = False + if ( + use_aux_xs_for_outputs + and not tag_has_xs_iterated_params + and getattr(meta, "variant", None) == "dense" + and len(meta.inputs_index) == 1 + and len(meta.outputs_index) == 1 + and len(meta.params_index) >= 1 + ): + _const_indices = [] + for _candidate_idx in ( + *meta.inputs_index, + *meta.params_index, + ): + _candidate = _resolve_tag_chain(body_eqn.invars[_candidate_idx]) + if _candidate not in body_invars: + _const_indices = [] + break + _candidate_body_idx = body_invars.index(_candidate) + if _candidate_body_idx >= _scan_num_consts: + _const_indices = [] + break + _const_indices.append(_candidate_body_idx) + _output_candidate = _resolve_tag_chain( + body_eqn.invars[meta.outputs_index[0]] + ) + if ( + len(_const_indices) + == len(meta.inputs_index) + len(meta.params_index) + and _output_candidate not in body_invars + ): + _outer_input = eqn.invars[_const_indices[0]] + use_accumulating_aux_base = ( + _outer_input.aval.shape[:-1] + == _output_candidate.aval.shape[:-1] + ) + + for arg_idx, v in enumerate(body_eqn.invars): + v_resolved = _resolve_tag_chain(v) + if v_resolved in body_invars: + idx = body_invars.index(v_resolved) + is_xs_iterated = ( + eqn.primitive.name == "scan" and idx >= _xs_threshold + ) + is_param = arg_idx in params_index_set + if is_xs_iterated and is_param: + tag_has_xs_iterated_params = True + + _slot_is_scan_carry = ( + eqn.primitive.name == "scan" + and idx >= _scan_num_consts + and idx < _xs_threshold + ) + if ( + (tag_has_xs_iterated_params or _slot_is_scan_carry) + and not is_param + and not is_xs_iterated + and supports_extension + and scan_length is not None + ): + cap_id = id(v_resolved) + if cap_id not in new_body_capture_id_to_idx: + new_body_capture_id_to_idx[cap_id] = len( + new_body_captures, + ) + new_body_captures.append(v_resolved) + deferred.append( + (arg_idx, new_body_capture_id_to_idx[cap_id]), + ) + continue + + partial_invars[arg_idx] = eqn.invars[idx] + elif arg_idx in params_index_set: + hoistable = False + break + elif ( + arg_idx in output_idx_set + and use_aux_xs_for_outputs + and supports_extension + and scan_length is not None + ): + body_v_id = id(v_resolved) + if body_v_id not in output_var_to_aux_xs: + aux_body_invar = make_var_func(v_resolved.aval) + aux_outer_aval = _jcore.ShapedArray( + (scan_length, *v_resolved.aval.shape), + v_resolved.aval.dtype, + ) + aux_outer_var = make_var_func(aux_outer_aval) + aux_tag_var = ( + make_var_func(v_resolved.aval) + if use_accumulating_aux_base + else aux_outer_var + ) + output_var_to_aux_xs[body_v_id] = ( + v_resolved, + aux_body_invar, + aux_outer_var, + aux_tag_var, + ) + ( + _, + _, + aux_outer_var, + aux_tag_var, + ) = output_var_to_aux_xs[body_v_id] + if ( + aux_tag_var is not aux_outer_var + ) != use_accumulating_aux_base: + raise ValueError( + "Conflicting scan accumulation contracts for " + "the same hoisted layer output." + ) + partial_invars[arg_idx] = aux_tag_var + elif supports_extension: + cap_id = id(v_resolved) + if cap_id not in new_body_capture_id_to_idx: + new_body_capture_id_to_idx[cap_id] = len( + new_body_captures, + ) + new_body_captures.append(v_resolved) + deferred.append( + (arg_idx, new_body_capture_id_to_idx[cap_id]), + ) + else: + import jax.extend as _jex + import numpy as _np + + zero_val = _np.zeros( + v_resolved.aval.shape, + dtype=v_resolved.aval.dtype, + ) + partial_invars[arg_idx] = _jex.core.Literal( + zero_val, + v_resolved.aval, + ) + + if not hoistable: + body_eqns_no_tags.append(body_eqn) + continue + + deferred_specs.append( + ( + body_eqn, + partial_invars, + deferred, + tag_has_xs_iterated_params, + use_aux_xs_for_outputs, + ) + ) + + import jax.extend as _jex + + def _remap_invars(eqns): + + out = [] + for e in eqns: + new_invars = [ + _resolve_tag_chain(w) + if not isinstance(w, _jex.core.Literal) + else w + for w in e.invars + ] + out.append(e.replace(invars=new_invars)) + return out + + body_eqns_no_tags = _remap_invars(body_eqns_no_tags) + new_body_outvars = [ + _resolve_tag_chain(v) if not isinstance(v, _jex.core.Literal) else v + for v in new_body_closed.jaxpr.outvars + ] + + if output_var_to_aux_xs: + from jax._src.lax import lax as _jlax + + aug_for_id: dict[int, tuple] = {} + for body_v_id, ( + body_v, + aux_body_invar, + _, + _, + ) in output_var_to_aux_xs.items(): + aug_var = make_var_func(body_v.aval) + aug_for_id[body_v_id] = (aug_var, aux_body_invar) + + seen_ids: set[int] = set() + + def _retarget_to_aug(w): + if isinstance(w, _jex.core.Literal): + return w + wid = id(w) + if wid in aug_for_id and wid in seen_ids: + return aug_for_id[wid][0] + return w + + augmented_eqns = [] + for body_eqn_clean in body_eqns_no_tags: + augmented_eqns.append( + body_eqn_clean.replace( + invars=[_retarget_to_aug(w) for w in body_eqn_clean.invars] + ) + ) + for o in body_eqn_clean.outvars: + oid = id(o) + if oid in aug_for_id and oid not in seen_ids: + aug_var, aux_body_invar = aug_for_id[oid] + aug_eqn = new_jaxpr_eqn( + invars=[o, aux_body_invar], + outvars=[aug_var], + primitive=_jlax.add_p, + params={}, + effects=frozenset(), + ) + augmented_eqns.append(aug_eqn) + seen_ids.add(oid) + + assert seen_ids == set(aug_for_id.keys()), ( + "aux-xs injection: some output Vars not encountered as " + "body-eqn outvars" + ) + body_eqns_no_tags = augmented_eqns + new_body_outvars = [ + _retarget_to_aug(v) if not isinstance(v, _jex.core.Literal) else v + for v in new_body_outvars + ] + + new_body_outvars = list(new_body_outvars) + list(new_body_captures) + + new_body_invars_list = list(new_body_closed.jaxpr.invars) + [ + aux_body_invar + for _, aux_body_invar, _, _ in output_var_to_aux_xs.values() + ] + + new_body_jaxpr = new_body_closed.jaxpr.replace( + eqns=body_eqns_no_tags, + outvars=new_body_outvars, + invars=new_body_invars_list, + ) + new_body_closed_clean = ClosedJaxpr( + new_body_jaxpr, + new_body_closed.consts, + ) + + params_dict = dict(**eqn.params) + params_dict[key] = to_jaxpr_or_closed_jaxpr( + new_body_closed_clean, + body_jaxpr, + ) + if eqn.primitive.name == "scan": + _extend_scan_flat_trees( + params_dict, + extra_xs=len(output_var_to_aux_xs), + extra_ys=len(new_body_captures), + ) + + if output_var_to_aux_xs: + from jax._src.lax import lax as _jlax + import jax.extend as _jex + import numpy as _np + + for ( + body_v, + _aux_body_invar, + aux_outer_var, + aux_tag_var, + ) in output_var_to_aux_xs.values(): + zero_scalar_aval = _jcore.ShapedArray( + (), + body_v.aval.dtype, + ) + zero_scalar_literal = _jex.core.Literal( + _np.array(0.0, dtype=body_v.aval.dtype), + zero_scalar_aval, + ) + first_bcast_outvar = aux_tag_var + first_bcast_shape = aux_tag_var.aval.shape + bcast_eqn = new_jaxpr_eqn( + invars=[zero_scalar_literal], + outvars=[first_bcast_outvar], + primitive=_jlax.broadcast_in_dim_p, + params={ + "shape": first_bcast_shape, + "broadcast_dimensions": (), + "sharding": None, + }, + effects=frozenset(), + ) + new_eqns.append(bcast_eqn) + if aux_tag_var is not aux_outer_var: + expand_eqn = new_jaxpr_eqn( + invars=[aux_tag_var], + outvars=[aux_outer_var], + primitive=_jlax.broadcast_in_dim_p, + params={ + "shape": ( + scan_length, + *body_v.aval.shape, + ), + "broadcast_dimensions": tuple( + range(1, body_v.aval.ndim + 1) + ), + "sharding": None, + }, + effects=frozenset(), + ) + new_eqns.append(expand_eqn) + + new_capture_outvars = [] + for cap_v in new_body_captures: + cap_aval = _jcore.ShapedArray( + (scan_length, *cap_v.aval.shape), + cap_v.aval.dtype, + ) + new_capture_outvars.append(make_var_func(cap_aval)) + + new_scan_invars = list(eqn.invars) + [ + aux_outer_var + for _, _, aux_outer_var, _ in output_var_to_aux_xs.values() + ] + new_eqn = eqn.replace( + params=params_dict, + invars=new_scan_invars, + outvars=list(eqn.outvars) + list(new_capture_outvars), + ) + new_eqns.append(new_eqn) + + aux_xs_capture_ids: set[int] = set() + for _be, _pi, _df, _has_xs, _use_aux in deferred_specs: + if _use_aux: + for _aidx, _cidx in _df: + aux_xs_capture_ids.add(_cidx) + transposed_outvars: list = [] + for _cap_idx, cap_outvar in enumerate(new_capture_outvars): + if _cap_idx in aux_xs_capture_ids: + transposed_outvars.append(cap_outvar) + continue + t_eqn, t_outvar = _make_transpose_swap01_eqn( + cap_outvar, + make_var_func, + ) + if t_eqn is not None: + new_eqns.append(t_eqn) + transposed_outvars.append(t_outvar) + + for ( + body_eqn, + partial_invars, + deferred, + has_xs_iter_params, + _use_aux_xs, + ) in deferred_specs: + final_invars = list(partial_invars) + for arg_idx, cap_idx in deferred: + final_invars[arg_idx] = transposed_outvars[cap_idx] + new_outvars = [make_var_func(v.aval) for v in body_eqn.outvars] + + hoisted_params = body_eqn.params + if has_xs_iter_params: + import dataclasses as _dc + + orig_meta = body_eqn.params["meta"] + + _v = orig_meta.variant or "" + new_variant = None + if _v == "scale_and_shift": + new_variant = "stacked_scale_and_shift" + elif _v == "structural_repeated_dense": + new_variant = "structural_stacked_repeated_dense" + elif _v == "structural_scale_and_shift": + new_variant = "structural_stacked_scale_and_shift" + if new_variant is not None: + new_meta = _dc.replace(orig_meta, variant=new_variant) + hoisted_params = {**body_eqn.params, "meta": new_meta} + hoisted_tags.append( + new_jaxpr_eqn( + invars=final_invars, + outvars=new_outvars, + primitive=body_eqn.primitive, + params=hoisted_params, + effects=body_eqn.effects, + ) + ) + + for nh_eqn in nested_hoisted: + outer_remapped = [] + ok = True + for v in nh_eqn.invars: + if v in body_jaxpr.jaxpr.invars: + idx = body_jaxpr.jaxpr.invars.index(v) + outer_remapped.append(eqn.invars[idx]) + else: + ok = False + break + if ok: + new_outvars2 = [make_var_func(v.aval) for v in nh_eqn.outvars] + hoisted_tags.append( + new_jaxpr_eqn( + invars=outer_remapped, + outvars=new_outvars2, + primitive=nh_eqn.primitive, + params=nh_eqn.params, + effects=nh_eqn.effects, + ) + ) + + new_closed = ClosedJaxpr( + closed_jaxpr.jaxpr.replace(eqns=new_eqns), + closed_jaxpr.consts, + ) + return new_closed, hoisted_tags + + def clean_layer_tags_jaxpr_patched(jaxpr, only_remove_auto_tags=False): + + closed = to_closed_jaxpr(jaxpr) + make_var_func = gensym() + closed, hoisted = _hoist_tags_recursive(closed, make_var_func) + + seen_param_keys = set() + deduped = [] + for h in hoisted: + meta = h.params["meta"] + key = tuple(id(h.invars[i]) for i in meta.params_index) + if key in seen_param_keys: + continue + seen_param_keys.add(key) + deduped.append(h) + hoisted = deduped + + if hoisted: + new_eqns = list(closed.jaxpr.eqns) + list(hoisted) + closed = ClosedJaxpr( + closed.jaxpr.replace(eqns=new_eqns), + closed.consts, + ) + return _orig_clean_layer( + to_jaxpr_or_closed_jaxpr(closed, jaxpr), + only_remove_auto_tags=only_remove_auto_tags, + ) + + clean_layer_tags_jaxpr_patched.__hamiltonzero_hoist_tags__ = True + tgm.clean_layer_tags_jaxpr = clean_layer_tags_jaxpr_patched + + +def _patch_kfactor_identity_init() -> None: + + import jax.numpy as _jnp + + from kfac_jax._src.curvature_blocks import ( + kronecker_factored as _kf, + ) + from kfac_jax._src import utils as _kfac_utils + + if getattr(_kf.KroneckerFactored._init, "__hamiltonzero_kfactor_identity__", False): + return + + _orig_kf_init = _kf.KroneckerFactored._init + + def _patched_kf_init( + self, rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues + ): + + cache = {} + factors = [] + for i, d in enumerate(self.array_shape): + eye = _jnp.eye(d, dtype=self.dtype) * _jnp.asarray(1.0, dtype=self.dtype) + + wma = _kfac_utils.WeightedMovingAverage( + value=eye, + weight=_jnp.asarray(1.0, dtype=self.dtype), + ) + factors.append(wma) + if cache_eigenvalues or exact_powers_to_cache: + cache[f"{i}_factor_eigenvalues"] = _jnp.ones((d,), dtype=self.dtype) + if exact_powers_to_cache: + cache[f"{i}_factor_eigen_vectors"] = _jnp.eye(d, dtype=self.dtype) + for power in approx_powers_to_cache: + if power != -1: + raise NotImplementedError( + f"Approximations for power {power} not implemented." + ) + if str(power) not in cache: + cache[str(power)] = {} + cache[str(power)][f"{i}_factor"] = _jnp.eye(d, dtype=self.dtype) + return _kf.KroneckerFactored.State( + cache=cache, + factors=tuple(factors), + ) + + _patched_kf_init.__hamiltonzero_kfactor_identity__ = True + _kf.KroneckerFactored._init = _patched_kf_init + + if getattr( + _kf.RepeatedDenseKroneckerFactored._init, + "__hamiltonzero_avg_repeats_one__", + False, + ): + return + + def _patched_rd_init( + self, rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues + ): + super_state = _kf.KroneckerFactored._init( + self, + rng, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, + ) + avg = _kfac_utils.WeightedMovingAverage( + value=_jnp.asarray(1.0, dtype=self.dtype), + weight=_jnp.asarray(1.0, dtype=self.dtype), + ) + return _kf.RepeatedDenseKroneckerFactored.State( + average_repeats=avg, + **super_state.__dict__, + ) + + _patched_rd_init.__hamiltonzero_avg_repeats_one__ = True + _kf.RepeatedDenseKroneckerFactored._init = _patched_rd_init + + +def _patch_pi_adjusted_kronecker_factors_floor() -> None: + + import jax.numpy as _jnp + from kfac_jax._src.utils import math as _kfac_math + + if getattr( + _kfac_math.pi_adjusted_kronecker_factors, + "__hamiltonzero_kron_floor__", + False, + ): + return + + _orig = _kfac_math.pi_adjusted_kronecker_factors + EPS_FLOOR = 1e-6 + EPS_REL = 1e-4 + + def _shift_from_avg_diag(avg_diag, scale): + eps_abs = _jnp.asarray(EPS_FLOOR, dtype=avg_diag.dtype) + eps_rel = _jnp.asarray(EPS_REL, dtype=avg_diag.dtype) + floor = _jnp.maximum(eps_abs, eps_rel * scale) + return _jnp.maximum(floor, floor - avg_diag) + + def _floor_factor(f): + if f.ndim == 0 or f.size == 1: + return f + _shift_from_avg_diag(f, _jnp.abs(f)) + if f.ndim == 1: + avg_diag = _jnp.mean(f) + scale = _jnp.max(_jnp.abs(f)) + return f + _shift_from_avg_diag(avg_diag, scale) + if f.ndim == 2: + d = f.shape[-1] + diag = _jnp.diagonal(f) + avg_diag = _jnp.sum(diag) / d + scale = _jnp.max(diag) + shift = _shift_from_avg_diag(avg_diag, scale) + return f + shift * _jnp.eye(d, dtype=f.dtype) + + if f.ndim >= 3 and f.shape[-1] == f.shape[-2]: + d = f.shape[-1] + eye = _jnp.eye(d, dtype=f.dtype) + for _ in range(f.ndim - 2): + eye = eye[None, ...] + diag = _jnp.diagonal(f, axis1=-2, axis2=-1) + avg_diag = _jnp.mean(diag, axis=-1) + scale = _jnp.max(diag, axis=-1) + shift = _shift_from_avg_diag(avg_diag, scale) + return f + shift[..., None, None] * eye + return f + + def patched(*factors, damping): + floored = tuple(_floor_factor(f) for f in factors) + return _orig(*floored, damping=damping) + + patched.__hamiltonzero_kron_floor__ = True + _kfac_math.pi_adjusted_kronecker_factors = patched + + from kfac_jax._src import utils as _kfac_utils_pkg + + if hasattr(_kfac_utils_pkg, "pi_adjusted_kronecker_factors"): + _kfac_utils_pkg.pi_adjusted_kronecker_factors = patched + + +def _patch_nested_scan_parent_walk() -> None: + + from kfac_jax._src import tag_graph_matcher as tgm + + _TagLocation = tgm.TagLocation + if getattr(_TagLocation, "__hamiltonzero_nested_parent_walk__", False): + return + + def _invars_of(eqn): + nm = eqn.primitive.name + if nm in ("scan", "pjit"): + return eqn.params["jaxpr"].jaxpr.invars + if nm == "while": + return eqn.params["body_jaxpr"].jaxpr.invars + if nm in ("xla_call", "xla_pmap"): + return eqn.params["call_jaxpr"].invars + raise NotImplementedError(f"higher-order primitive {nm!r}") + + def _walk(param_vars, eqns_in_order): + for eqn, _ in eqns_in_order: + invars = _invars_of(eqn) + p_indexes = [invars.index(p) for p in param_vars] + param_vars = tuple(eqn.invars[pi] for pi in p_indexes) + return param_vars + + def _top_level_parameters(self): + pv = self.bottom_level_parameters + return _walk(pv, list(self.parent_equations)) + + def _full_name_ordered(self, eqns_in_order): + param_vars = self.bottom_level_parameters + parts = [] + for eqn, n in eqns_in_order: + nm = eqn.primitive.name + invars = _invars_of(eqn) + p_indexes = [invars.index(p) for p in param_vars] + piece = f"{nm}_{n}/" + if nm == "scan": + num_consts, _, _ = _scan_partition_sizes(eqn) + checks = [pi < num_consts for pi in p_indexes] + if not (all(checks) or all(not ci for ci in checks)): + raise ValueError( + "Parameters inside scan of the same tag are not both " + "carry or const." + ) + piece = piece + ("const/" if all(checks) else "carry/") + parts.append(piece) + param_vars = [eqn.invars[pi] for pi in p_indexes] + + prefix = "".join(reversed(parts)) + return prefix + self.base_name + + def _full_name(self): + return _full_name_ordered(self, list(self.parent_equations)) + + _TagLocation.top_level_parameters = property(_top_level_parameters) + _TagLocation.full_name = property(_full_name) + _TagLocation.__hamiltonzero_nested_parent_walk__ = True + + +_patch() +_patch_allow_multiple_registrations() +_patch_orphan_registration_in_sub_graphs() +_patch_manual_tag_outputs_that_are_graph_inputs() +_patch_hoist_layer_tags_from_scan() +_patch_kfactor_identity_init() +_patch_pi_adjusted_kronecker_factors_floor() +_patch_nested_scan_parent_walk() + + +__all__: list[str] = [] diff --git a/src/hamiltonzero/optim/production.py b/src/hamiltonzero/optim/production.py new file mode 100644 index 0000000000000000000000000000000000000000..b7fe9b68ab69c0eefa0803679152e8ae5e647859 --- /dev/null +++ b/src/hamiltonzero/optim/production.py @@ -0,0 +1,460 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any, Callable, NamedTuple + +import equinox as eqx +import jax +import jax.numpy as jnp +import kfac_jax + +from hamiltonzero.optim import compat as _kfac_compat +from hamiltonzero.optim import spin_blocks as _spin_blocks +from hamiltonzero.optim.blocks import ( + make_graph_patterns, +) + + +_GRAPH_PATTERNS = make_graph_patterns() +_FISHER_SIGN_STREAM = 1262895427 +_FISHER_CHANNELS = 3 +_ROUTE_SAMPLES = 8 + + +class KFACBundle(NamedTuple): + optimizer: Any + loss_fn: Callable + state: Any + + +def _make_fisher_signs(key, q_cold): + sign_key = jax.random.fold_in(key, _FISHER_SIGN_STREAM) + signs = jax.random.rademacher( + sign_key, + q_cold.shape[:2] + (_FISHER_CHANNELS,), + dtype=q_cold.dtype, + ) + q_sharding = getattr(q_cold, "sharding", None) + if isinstance(q_sharding, jax.sharding.NamedSharding): + q_spec = tuple(q_sharding.spec) + sharding = jax.sharding.NamedSharding( + q_sharding.mesh, + jax.sharding.PartitionSpec(*q_spec[:2], None), + ) + signs = jax.device_put(signs, sharding) + elif isinstance(q_sharding, jax.sharding.SingleDeviceSharding): + signs = jax.device_put(signs, q_sharding) + return signs + + +def _signed_identity(value, signs): + stopped = jax.lax.stop_gradient(value) + return stopped + signs.astype(value.dtype) * (value - stopped) + + +def _register_fisher_output(value, signs): + kfac_jax.register_normal_predictive_distribution( + _signed_identity(value, signs).reshape(-1, 1) + ) + + +def _center_preprocessed_energy(energy, pmap_axis_name): + if pmap_axis_name is None: + energy_sum = jnp.sum(energy, axis=1) + count = jnp.asarray(energy.shape[1], dtype=energy.real.dtype) + else: + from kfac_jax._src.utils import parallel as kfac_parallel + + energy_sum = kfac_parallel.psum_if_pmap( + jnp.sum(energy, axis=1), + pmap_axis_name, + ) + count = kfac_parallel.psum_if_pmap( + jnp.asarray(energy.shape[1], dtype=energy.real.dtype), + pmap_axis_name, + ) + mean = energy_sum / jnp.maximum(count, 1.0) + delta = energy - mean[:, None] + return jax.lax.stop_gradient(delta), mean + + +def _router_loss(apply_fn, route_loss_weight: float): + def apply_walkers(params, q_cold, context, t, tau): + systems, walkers = q_cold.shape[:2] + if systems == 1: + context_single = jax.tree.map( + lambda x: x[0] if isinstance(x, jnp.ndarray) and x.ndim > 0 else x, + context, + ) + re, im, route_logp = jax.vmap( + lambda p, q, time: apply_fn( + p, + q, + context_single, + time, + tau, + ), + in_axes=(None, 0, None), + )(params, q_cold[0], t) + return ( + re.reshape(1, walkers), + im.reshape(1, walkers), + route_logp.reshape(1, walkers), + ) + + def apply_system(params_value, context_value, q_value, t_value): + return jax.vmap( + lambda p, q, time: apply_fn( + p, + q, + context_value, + time, + tau, + ), + in_axes=(None, 0, None), + )(params_value, q_value, t_value) + + return jax.vmap( + apply_system, + in_axes=(None, 0, 0, None), + )(params, context, q_cold, t) + + @jax.custom_jvp + def total_energy(params, batch): + _q, energy, _context, _t, _tau, _advantage, _signs = batch + _delta, mean = _center_preprocessed_energy(energy, None) + return jnp.mean(mean.real) + + @total_energy.defjvp + def total_energy_jvp(primals, tangents): + params, batch = primals + params_t, _batch_t = tangents + q_cold, energy, context, t, tau, advantage, fisher_signs = batch + (re, im, route_logp), (tan_re, tan_im, tan_route_logp) = jax.jvp( + lambda p: apply_walkers(p, q_cold, context, t, tau), + (params,), + (params_t,), + ) + _register_fisher_output(re, fisher_signs[..., 0]) + _register_fisher_output(im, fisher_signs[..., 1]) + _register_fisher_output(route_logp, fisher_signs[..., 2]) + delta, mean = _center_preprocessed_energy(energy, None) + real_num = jnp.sum(tan_re * delta.real) + imag_num = jnp.sum(tan_im * delta.imag) + n_eff = jnp.maximum( + jnp.sum(jnp.abs(delta) > 0).astype(tan_re.dtype), + 1.0, + ) + loss_tangent = 2.0 * (real_num + imag_num) / n_eff + advantage = jax.lax.stop_gradient(advantage.reshape((-1,))) + route_tangent = jnp.mean( + advantage * jnp.mean(tan_route_logp, axis=1).reshape((-1,)) + ) + loss_tangent = ( + loss_tangent + + jnp.asarray( + route_loss_weight, + dtype=loss_tangent.dtype, + ) + * route_tangent + ) + return jnp.mean(mean.real), loss_tangent + + return total_energy + + +def _finetune_loss(apply_fn, pmap_axis_name): + def apply_walkers(params, q_cold, context, t): + context = jax.tree.map( + lambda x: x[0] if isinstance(x, jnp.ndarray) and x.ndim > 0 else x, + context, + ) + re, im = jax.vmap( + lambda p, q, time: apply_fn(p, q, context, time), + in_axes=(None, 0, None), + )(params, q_cold[0], t) + batch_size = q_cold.shape[1] + return re.reshape(1, batch_size), im.reshape(1, batch_size) + + @jax.custom_jvp + def total_energy(params, batch): + _q, energy, _context, _t, _signs = batch + _delta, mean = _center_preprocessed_energy( + energy, + pmap_axis_name, + ) + return jnp.mean(mean.real) + + @total_energy.defjvp + def total_energy_jvp(primals, tangents): + params, batch = primals + params_t, _batch_t = tangents + q_cold, energy, context, t, fisher_signs = batch + (re, im), (tan_re, tan_im) = jax.jvp( + lambda p: apply_walkers(p, q_cold, context, t), + (params,), + (params_t,), + ) + _register_fisher_output(re, fisher_signs[..., 0]) + _register_fisher_output(im, fisher_signs[..., 1]) + delta, mean = _center_preprocessed_energy(energy, pmap_axis_name) + local_real_num = jnp.sum(tan_re * delta.real) + local_imag_num = jnp.sum(tan_im * delta.imag) + local_n_eff = jnp.sum(jnp.abs(delta) > 0).astype(tan_re.dtype) + if pmap_axis_name is None: + real_num = local_real_num + imag_num = local_imag_num + n_eff = local_n_eff + else: + from kfac_jax._src.utils import parallel as kfac_parallel + + real_num = kfac_parallel.psum_if_pmap( + local_real_num, + pmap_axis_name, + ) + imag_num = kfac_parallel.psum_if_pmap( + local_imag_num, + pmap_axis_name, + ) + n_eff = kfac_parallel.psum_if_pmap( + local_n_eff, + pmap_axis_name, + ) + loss_tangent = 2.0 * (real_num + imag_num) / jnp.maximum(n_eff, 1.0) + return jnp.mean(mean.real), loss_tangent + + return total_energy + + +def _configure_kfac(): + kfac_jax.utils.set_use_cholesky_inversion(True) + + +def _new_optimizer(config, loss_fn, *, multi_device: bool, axis_name): + _configure_kfac() + return kfac_jax.Optimizer( + jax.value_and_grad(loss_fn), + learning_rate_schedule=None, + damping_schedule=None, + norm_constraint=float(config.norm_constraint), + multi_device=multi_device, + pmap_axis_name=axis_name if multi_device else None, + value_func_has_aux=False, + value_func_has_rng=False, + register_only_generic=False, + auto_register_kwargs={ + "graph_patterns": _GRAPH_PATTERNS, + "allow_multiple_registrations": True, + }, + include_norms_in_stats=False, + estimation_mode="fisher_exact", + share_curvature_and_grad_forward=False, + num_burnin_steps=0, + batch_size_extractor=lambda batch, *_: batch[0].shape[0] * batch[0].shape[1], + min_damping=float(config.minimum_damping), + inverse_update_period=int(config.inverse_update_period), + curvature_update_period=int(config.curvature_update_period), + curvature_ema=float(config.curvature_ema), + l2_reg=float(config.l2_regularization), + ) + + +def _partition(model): + return eqx.partition(model, jax.tree.map(eqx.is_inexact_array, model)) + + +def _assert_no_naive_full(optimizer, state): + blocks = list(enumerate(getattr(state, "blocks_states", []) or [])) + if not blocks: + try: + blocks = list(enumerate(optimizer._estimator.blocks)) + except AttributeError: + blocks = [] + bad = [] + for index, block in blocks: + name = type(block).__name__ + if "NaiveFull" in name: + bad.append((index, name, getattr(block, "parameters_shapes", None))) + if bad: + details = "; ".join( + f"block[{index}] {name} shapes={shapes}" for index, name, shapes in bad + ) + raise RuntimeError(f"KFAC produced unsupported NaiveFull blocks: {details}") + + +def _router_initial_advantage(energy): + rewards = jnp.mean(energy.real, axis=1) + grouped = rewards.reshape((-1, _ROUTE_SAMPLES)) + centered = grouped - jnp.mean(grouped, axis=1, keepdims=True) + return jax.lax.stop_gradient( + (float(_ROUTE_SAMPLES) / float(_ROUTE_SAMPLES - 1) * centered).reshape((-1,)) + ) + + +def init_router_kfac_state( + config, + model, + q_cold, + energy, + context, + *, + t: float, + key, + multi_device: bool, + route_tau, + route_loss_weight: float, +): + params, static = _partition(model) + + def apply_fn(params_value, q, context_value, t_value, tau_value): + combined = eqx.combine(params_value, static) + return combined.call_with_route_logprob( + q, + context_value, + t_value, + tau=tau_value, + ) + + loss_fn = _router_loss(apply_fn, route_loss_weight) + optimizer = _new_optimizer( + config, + loss_fn, + multi_device=multi_device, + axis_name="systems", + ) + fisher_signs = _make_fisher_signs(key, q_cold) + batch = ( + q_cold, + energy, + context, + jnp.asarray(t, dtype=q_cold.dtype), + jnp.asarray(route_tau, dtype=q_cold.dtype), + _router_initial_advantage(energy), + fisher_signs, + ) + _configure_kfac() + state = optimizer.init(params, key, batch) + _assert_no_naive_full(optimizer, state) + return KFACBundle(optimizer=optimizer, loss_fn=loss_fn, state=state) + + +def init_finetune_kfac_state( + config, + model, + q_cold, + energy, + context, + *, + t: float, + key, + multi_device: bool, +): + params, static = _partition(model) + + def apply_fn(params_value, q, context_value, t_value): + combined = eqx.combine(params_value, static) + return combined.call_tagged(q, context_value, t_value) + + loss_axis = "batch" if multi_device else None + loss_fn = _finetune_loss(apply_fn, loss_axis) + optimizer = _new_optimizer( + config, + loss_fn, + multi_device=multi_device, + axis_name="batch", + ) + fisher_signs = _make_fisher_signs(key, q_cold) + batch = ( + q_cold, + energy, + context, + jnp.asarray(t, dtype=q_cold.dtype), + fisher_signs, + ) + _configure_kfac() + state = optimizer.init(params, key, batch) + _assert_no_naive_full(optimizer, state) + return KFACBundle(optimizer=optimizer, loss_fn=loss_fn, state=state) + + +def apply_router_kfac_step( + bundle, + model, + q_cold, + energy, + context, + *, + t: float, + key, + momentum, + learning_rate, + damping, + route_advantage, + route_tau, +): + _configure_kfac() + params, static = _partition(model) + batch = ( + q_cold, + energy, + context, + jnp.asarray(t, dtype=q_cold.dtype), + jnp.asarray(route_tau, dtype=q_cold.dtype), + route_advantage, + _make_fisher_signs(key, q_cold), + ) + new_params, state, _stats = bundle.optimizer.step( + params, + bundle.state, + key, + batch=batch, + momentum=jnp.asarray(momentum, dtype=jnp.float32), + learning_rate=jnp.asarray(learning_rate, dtype=jnp.float32), + damping=jnp.asarray(damping, dtype=jnp.float32), + ) + return eqx.combine(new_params, static), bundle._replace(state=state) + + +def apply_finetune_kfac_step( + bundle, + model, + q_cold, + energy, + context, + *, + t: float, + key, + momentum, + learning_rate, + damping, +): + _configure_kfac() + params, static = _partition(model) + batch = ( + q_cold, + energy, + context, + jnp.asarray(t, dtype=q_cold.dtype), + _make_fisher_signs(key, q_cold), + ) + new_params, state, _stats = bundle.optimizer.step( + params, + bundle.state, + key, + batch=batch, + momentum=jnp.asarray(momentum, dtype=jnp.float32), + learning_rate=jnp.asarray(learning_rate, dtype=jnp.float32), + damping=jnp.asarray(damping, dtype=jnp.float32), + ) + return eqx.combine(new_params, static), bundle._replace(state=state) + + +__all__ = [ + "KFACBundle", + "apply_finetune_kfac_step", + "apply_router_kfac_step", + "init_finetune_kfac_state", + "init_router_kfac_state", +] diff --git a/src/hamiltonzero/optim/spin_blocks.py b/src/hamiltonzero/optim/spin_blocks.py new file mode 100644 index 0000000000000000000000000000000000000000..395bf4863f3ab299d19d8d01efa70cefdc1db8ff --- /dev/null +++ b/src/hamiltonzero/optim/spin_blocks.py @@ -0,0 +1,726 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations +import jax +import jax.numpy as jnp +import kfac_jax +from kfac_jax._src import utils as kfac_utils +from kfac_jax._src.layers_and_loss_tags import LayerMetaData, layer_tag + +_FEATURIZER_OUTPUT_SPLIT_IDS = frozenset( + {"featurizer.global_w1", "featurizer.combine_w1"} +) +_FEATURIZER_INPUT_SPLIT_IDS = frozenset( + {"featurizer.global_w2", "featurizer.combine_w2"} +) + + +def _floor_matrix_avg_diag(mat, eps: float): + d = mat.shape[-1] + eps_arr = jnp.asarray(eps, dtype=mat.dtype) + avg_diag = jnp.trace(mat) / d + shift = jnp.maximum(eps_arr, eps_arr - avg_diag) + return mat + shift * jnp.eye(d, dtype=mat.dtype) + + +def _balanced_axis_partition(shape: tuple[int, ...]): + n = len(shape) + full_mask = (1 << n) - 1 + best = None + for mask in range(1, full_mask): + if not mask & 1: + continue + left_axes = tuple((i for i in range(n) if mask & 1 << i)) + right_axes = tuple((i for i in range(n) if not mask & 1 << i)) + left_prod = _prod_int((shape[i] for i in left_axes)) + right_prod = _prod_int((shape[i] for i in right_axes)) + score = (max(left_prod, right_prod), abs(left_prod - right_prod)) + if best is None or score < best[0]: + best = (score, left_axes, right_axes) + assert best is not None + return (best[1], best[2]) + + +def _prod_int(vals) -> int: + out = 1 + for v in vals: + out *= int(v) + return out + + +def _matricize(x, left_axes, right_axes): + shape = tuple(x.shape) + perm = tuple(left_axes) + tuple(right_axes) + left_dim = _prod_int((shape[i] for i in left_axes)) + right_dim = _prod_int((shape[i] for i in right_axes)) + return jnp.transpose(x, perm).reshape(left_dim, right_dim) + + +def _unmatricize(x_mat, shape, left_axes, right_axes): + left_shape = tuple((shape[i] for i in left_axes)) + right_shape = tuple((shape[i] for i in right_axes)) + perm = tuple(left_axes) + tuple(right_axes) + inv_perm_list = [0] * len(perm) + for pos, axis in enumerate(perm): + inv_perm_list[axis] = pos + inv_perm = tuple(inv_perm_list) + x_perm = x_mat.reshape(left_shape + right_shape) + return jnp.transpose(x_perm, inv_perm) + + +def _validate_approx_inverse_cache_request( + exact_powers_to_cache, approx_powers_to_cache +): + if exact_powers_to_cache: + raise NotImplementedError( + "Custom merge blocks do not implement exact cached powers." + ) + unsupported = set(approx_powers_to_cache) - {-1} + if unsupported: + raise NotImplementedError( + f"Unsupported approximate cached powers: {sorted(unsupported)}." + ) + + +def _init_two_kron_cache( + left_dim, + right_dim, + dtype, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, +): + _validate_approx_inverse_cache_request( + exact_powers_to_cache, approx_powers_to_cache + ) + cache = {} + if -1 in approx_powers_to_cache: + cache["-1"] = { + "left_factor": jnp.eye(left_dim, dtype=dtype), + "right_factor": jnp.eye(right_dim, dtype=dtype), + } + if cache_eigenvalues: + cache["eigenvalues"] = jnp.zeros((left_dim * right_dim,), dtype=dtype) + return cache + + +def _update_two_kron_cache( + state, + left_factor, + right_factor, + identity_weight, + exact_powers, + approx_powers, + eigenvalues, + *, + inverse_epsilon=None, +): + _validate_approx_inverse_cache_request(exact_powers, approx_powers) + state = state.copy() + if eigenvalues: + s_left, _ = kfac_utils.safe_psd_eigh(left_factor) + s_right, _ = kfac_utils.safe_psd_eigh(right_factor) + state.cache["eigenvalues"] = jnp.einsum("p,q->pq", s_left, s_right).reshape(-1) + if -1 in approx_powers: + if inverse_epsilon is not None: + left_for_inverse = _floor_matrix_avg_diag(left_factor, inverse_epsilon) + right_for_inverse = _floor_matrix_avg_diag(right_factor, inverse_epsilon) + else: + left_for_inverse = left_factor + right_for_inverse = right_factor + inv_left, inv_right = kfac_utils.pi_adjusted_kronecker_inverse( + left_for_inverse, right_for_inverse, damping=identity_weight + ) + state.cache["-1"]["left_factor"] = inv_left + state.cache["-1"]["right_factor"] = inv_right + return state + + +def _two_kron_marginal_from_merge(dy_m, uA_m, uB_m, group_axes): + lower = ("i", "j", "k", "l") + upper = ("I", "J", "K", "L") + group_axes = tuple(group_axes) + group_set = set(group_axes) + + def _labels(axes, primed: bool): + out = ["n"] + for ax in axes: + out.append(upper[ax] if primed and ax in group_set else lower[ax]) + return "".join(out) + + dy1 = _labels((0, 1), primed=False) + uA1 = _labels((0, 2), primed=False) + uB1 = _labels((0, 3), primed=False) + dy2 = _labels((0, 1), primed=True) + uA2 = _labels((0, 2), primed=True) + uB2 = _labels((0, 3), primed=True) + out = "".join((lower[ax] for ax in group_axes)) + out += "".join((upper[ax] for ax in group_axes)) + eqn = f"{dy1},{uA1},{uB1},{dy2},{uA2},{uB2}->{out}" + gram = jnp.einsum(eqn, dy_m, uA_m, uB_m, dy_m, uA_m, uB_m) + dims = (dy_m.shape[1], dy_m.shape[2], uA_m.shape[2], uB_m.shape[2]) + dim = _prod_int((dims[ax] for ax in group_axes)) + return gram.reshape(dim, dim) + + +def _merge_gradient_trace(dy_m, uA_m, uB_m, divisor): + squared_norm_sum = jnp.einsum( + "nij,nik,nil->", jnp.square(dy_m), jnp.square(uA_m), jnp.square(uB_m) + ) + return squared_norm_sum / divisor + + +def _trace_normalize_two_kron_marginals( + left_factor, right_factor, trace_mass, *, repeat_mass=1.0 +): + trace_mass = jnp.asarray(trace_mass, dtype=left_factor.dtype) + repeat_mass = jnp.asarray(repeat_mass, dtype=left_factor.dtype) + finite = jnp.isfinite(trace_mass) & jnp.isfinite(repeat_mass) + no_mass = finite & ((trace_mass <= 0) | (repeat_mass <= 0)) + safe_trace = jnp.where(no_mass, jnp.ones_like(trace_mass), trace_mass) + safe_repeat = jnp.where(no_mass, jnp.zeros_like(repeat_mass), repeat_mass) + factor_scale = jnp.sqrt(safe_repeat / safe_trace) + normalized_left = jnp.where( + no_mass, jnp.zeros_like(left_factor), factor_scale * left_factor + ) + normalized_right = jnp.where( + no_mass, jnp.zeros_like(right_factor), factor_scale * right_factor + ) + normalized_left = jnp.where( + finite, normalized_left, jnp.full_like(left_factor, jnp.nan) + ) + normalized_right = jnp.where( + finite, normalized_right, jnp.full_like(right_factor, jnp.nan) + ) + return (normalized_left, normalized_right) + + +def _identity_wma(dim, dtype): + return kfac_utils.WeightedMovingAverage( + value=jnp.eye(dim, dtype=dtype), weight=jnp.asarray(1.0, dtype=dtype) + ) + + +def _scalar_wma(value, dtype): + return kfac_utils.WeightedMovingAverage( + value=jnp.asarray(value, dtype=dtype), weight=jnp.asarray(1.0, dtype=dtype) + ) + + +def _poison_cached_inverse_on_failure(state, factor_key, certified): + if "-1" in state.cache: + cached = state.cache["-1"][factor_key] + state.cache["-1"][factor_key] = jnp.where( + certified, cached, jnp.full_like(cached, jnp.nan) + ) + return state + + +STRUCTURAL_QUADRILINEAR_MERGE_TAG_VARIANT = "structural_quadrilinear_merge" + + +def _structural_name_kw(name: str | None) -> dict[str, str]: + return {} if name is None else {"name": name} + + +def register_structural_quadrilinear_merge( + y, + x_l, + x_r, + T, + structural_mask, + *, + scan_shared: bool, + repeat_ndim: int, + name: str | None = None, +): + if tuple(x_l.shape) != tuple(x_r.shape): + raise ValueError( + f"quadrilinear input shapes differ: {x_l.shape} vs {x_r.shape}" + ) + if tuple(structural_mask.shape) != tuple(x_l.shape[:-1]): + raise ValueError( + f"quadrilinear structural mask must match local leading shape: mask={structural_mask.shape}, input={x_l.shape}" + ) + return layer_tag.bind( + y, + x_l, + x_r, + structural_mask, + T, + meta=LayerMetaData( + variant=STRUCTURAL_QUADRILINEAR_MERGE_TAG_VARIANT, + outputs_index=(0,), + inputs_index=(1, 2, 3), + params_index=(4,), + ), + scan_shared=bool(scan_shared), + repeat_ndim=int(repeat_ndim), + **_structural_name_kw(name), + ) + + +@kfac_utils.register_state_class +class _QuadrilinearMergeState(kfac_jax.CurvatureBlock.State): + sigma_left: kfac_utils.WeightedMovingAverage + sigma_right: kfac_utils.WeightedMovingAverage + + +class _QuadrilinearMergeBlock(kfac_jax.CurvatureBlock): + State = _QuadrilinearMergeState + + @property + def parameters_canonical_order(self) -> tuple[int, ...]: + return (0,) + + def _init( + self, rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues + ): + del rng + shape = tuple(self.parameters_shapes[0]) + left_axes, right_axes = _balanced_axis_partition(shape) + left_dim = _prod_int((shape[i] for i in left_axes)) + right_dim = _prod_int((shape[i] for i in right_axes)) + + def _eye_wma(d): + return kfac_utils.WeightedMovingAverage( + value=jnp.eye(d, dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ) + + return _QuadrilinearMergeState( + cache=_init_two_kron_cache( + left_dim, + right_dim, + self.dtype, + exact_powers_to_cache, + approx_powers_to_cache, + cache_eigenvalues, + ), + sigma_left=_eye_wma(left_dim), + sigma_right=_eye_wma(right_dim), + ) + + def sync(self, state, pmap_axis_name): + state = state.copy() + for f in (state.sigma_left, state.sigma_right): + f.sync(pmap_axis_name) + return state + + def update_curvature_matrix_estimate( + self, state, estimation_data, ema_old, ema_new, identity_weight, batch_size + ): + del identity_weight, batch_size + state = state.copy() + u_a, u_b = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + [T_param] = estimation_data.primals.params + G, d_j, d_k, d_l = T_param.shape + d_m_eff = G * d_l + + def _find_last_feature_axis(arr, size): + for i in range(arr.ndim - 1, -1, -1): + if arr.shape[i] == size: + return i + return arr.ndim - 1 + + ax_dy = _find_last_feature_axis(dy, d_m_eff) + ax_uA = _find_last_feature_axis(u_a, d_m_eff) + ax_uB = _find_last_feature_axis(u_b, d_m_eff) + dy_f = jnp.moveaxis(dy, ax_dy, -1).reshape(-1, G, d_j) + uA_f = jnp.moveaxis(u_a, ax_uA, -1).reshape(-1, G, d_k) + uB_f = jnp.moveaxis(u_b, ax_uB, -1).reshape(-1, G, d_l) + is_active = 1.0 - jnp.all(dy_f == 0.0, axis=(-2, -1), keepdims=True).astype( + dy_f.dtype + ) + n_active = jnp.sum(is_active) + normalizer = jnp.maximum(n_active, 1.0).astype(dy_f.dtype) + inv_n = jnp.reciprocal(normalizer) + dy_m = dy_f * is_active + uA_m = uA_f * is_active + uB_m = uB_f * is_active + shape = tuple(T_param.shape) + left_axes, right_axes = _balanced_axis_partition(shape) + sigma_left_new = ( + _two_kron_marginal_from_merge(dy_m, uA_m, uB_m, left_axes) * inv_n + ) + sigma_right_new = ( + _two_kron_marginal_from_merge(dy_m, uA_m, uB_m, right_axes) * inv_n + ) + trace_mass = _merge_gradient_trace(dy_m, uA_m, uB_m, normalizer) + sigma_left_new, sigma_right_new = _trace_normalize_two_kron_marginals( + sigma_left_new, sigma_right_new, trace_mass + ) + sigma_left_new = 0.5 * (sigma_left_new + sigma_left_new.T) + sigma_right_new = 0.5 * (sigma_right_new + sigma_right_new.T) + state.sigma_left.update(sigma_left_new, ema_old, ema_new) + state.sigma_right.update(sigma_right_new, ema_old, ema_new) + return state + + _MATPOWER_EPSILON_FLOOR: float = 1e-06 + + def _multiply_matpower_unscaled( + self, state, vector, identity_weight, power, exact_power, use_cached + ): + if exact_power and power != 1: + raise NotImplementedError( + "QuadrilinearMergeBlock implements approximate inverse powers only." + ) + [grad_T] = vector + shape = tuple(self.parameters_shapes[0]) + left_axes, right_axes = _balanced_axis_partition(shape) + grad_mat = _matricize(grad_T, left_axes, right_axes) + if power == -1: + if use_cached: + inv_left = state.cache["-1"]["left_factor"] + inv_right = state.cache["-1"]["right_factor"] + else: + eps = self._MATPOWER_EPSILON_FLOOR + inv_left, inv_right = kfac_utils.pi_adjusted_kronecker_inverse( + _floor_matrix_avg_diag(state.sigma_left.value, eps), + _floor_matrix_avg_diag(state.sigma_right.value, eps), + damping=identity_weight, + ) + new_mat = jnp.einsum("pP,qQ,PQ->pq", inv_left, inv_right, grad_mat) + elif power == 1: + curvature_product = jnp.einsum( + "pP,qQ,PQ->pq", + state.sigma_left.value, + state.sigma_right.value, + grad_mat, + ) + if use_cached: + curvature_product = ( + self.state_dependent_scale(state) * curvature_product + ) + new_mat = curvature_product + identity_weight * grad_mat + else: + raise NotImplementedError( + f"QuadrilinearMergeBlock: power={power} not implemented (only ±1 supported)." + ) + new_T = _unmatricize(new_mat, shape, left_axes, right_axes) + return (new_T,) + + def _eigenvalues_unscaled(self, state, use_cached): + if use_cached: + return state.cache["eigenvalues"] + s_left, _ = kfac_utils.safe_psd_eigh(state.sigma_left.value) + s_right, _ = kfac_utils.safe_psd_eigh(state.sigma_right.value) + return jnp.einsum("p,q->pq", s_left, s_right).reshape(-1) + + def _update_cache( + self, state, identity_weight, exact_powers, approx_powers, eigenvalues + ): + eps = self._MATPOWER_EPSILON_FLOOR + return _update_two_kron_cache( + state, + state.sigma_left.value, + state.sigma_right.value, + identity_weight, + exact_powers, + approx_powers, + eigenvalues, + inverse_epsilon=eps, + ) + + def _to_dense_unscaled(self, state): + return jnp.kron(state.sigma_left.value, state.sigma_right.value) + + def _norm_unscaled(self, state, norm_type): + n_left = kfac_utils.psd_matrix_norm(state.sigma_left.value, norm_type=norm_type) + n_right = kfac_utils.psd_matrix_norm( + state.sigma_right.value, norm_type=norm_type + ) + return n_left * n_right + + +@kfac_utils.register_state_class +class _StructuralQuadrilinearMergeState(_QuadrilinearMergeState): + average_repeats: kfac_utils.WeightedMovingAverage + + +class StructuralQuadrilinearMergeBlock(_QuadrilinearMergeBlock): + State = _StructuralQuadrilinearMergeState + + def _init( + self, rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues + ): + base = super()._init( + rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues + ) + return self.State( + **base.__dict__, + average_repeats=kfac_utils.WeightedMovingAverage( + value=jnp.ones((), dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + ) + + def sync(self, state, pmap_axis_name): + state = super().sync(state, pmap_axis_name) + state.average_repeats.sync(pmap_axis_name) + return state + + def state_dependent_scale(self, state): + return 1.0 / jnp.where( + state.average_repeats.value > 0, state.average_repeats.value, 1.0 + ) + + def _update_cache( + self, state, identity_weight, exact_powers, approx_powers, eigenvalues + ): + state = super()._update_cache( + state, identity_weight, exact_powers, approx_powers, eigenvalues + ) + scale = self.state_dependent_scale(state) + if eigenvalues: + state.cache["eigenvalues"] = scale * state.cache["eigenvalues"] + if -1 in approx_powers: + factor_scale = jnp.sqrt(scale) + state.cache["-1"]["left_factor"] /= factor_scale + state.cache["-1"]["right_factor"] /= factor_scale + return state + + def update_curvature_matrix_estimate( + self, state, estimation_data, ema_old, ema_new, identity_weight, batch_size + ): + del identity_weight + from hamiltonzero.optim.blocks import ( + align_structural_mask_to_leading, + structural_group_repeats, + ) + + state = state.copy() + u_a, u_b, structural_mask = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + [T_param] = estimation_data.primals.params + scan_shared = bool(self._layer_tag_eq.params["scan_shared"]) + repeat_ndim = int(self._layer_tag_eq.params["repeat_ndim"]) + structural_mask = align_structural_mask_to_leading( + structural_mask, dy.shape[:-1], repeat_ndim=repeat_ndim + ) + ua_g, mask_g, logical_batch, _ = structural_group_repeats( + u_a, + structural_mask, + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + ub_g, _, ub_batch, _ = structural_group_repeats( + u_b, + structural_mask, + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + dy_g, _, dy_batch, _ = structural_group_repeats( + dy, + structural_mask, + scan_shared=scan_shared, + repeat_ndim=repeat_ndim, + feature_ndim=1, + ) + if logical_batch != ub_batch or logical_batch != dy_batch: + raise ValueError("quadrilinear structural logical batches differ") + G, d_j, d_k, d_l = T_param.shape + row_mask = mask_g.astype(dy_g.dtype)[..., None, None] + dy_f = dy_g.reshape(-1, G, d_j) * row_mask.reshape(-1, 1, 1) + uA_f = ua_g.reshape(-1, G, d_k) * row_mask.reshape(-1, 1, 1) + uB_f = ub_g.reshape(-1, G, d_l) * row_mask.reshape(-1, 1, 1) + sample_divisor = jnp.maximum( + jnp.asarray(batch_size, dtype=dy_f.dtype), + jnp.asarray(1.0, dtype=dy_f.dtype), + ) + logical_divisor = jnp.asarray(logical_batch, dtype=dy_f.dtype) + shape = tuple(T_param.shape) + left_axes, right_axes = _balanced_axis_partition(shape) + sigma_left = ( + _two_kron_marginal_from_merge(dy_f, uA_f, uB_f, left_axes) / sample_divisor + ) + sigma_right = ( + _two_kron_marginal_from_merge(dy_f, uA_f, uB_f, right_axes) / sample_divisor + ) + repeats = jnp.sum(mask_g) / logical_divisor + trace_mass = _merge_gradient_trace(dy_f, uA_f, uB_f, sample_divisor) + sigma_left, sigma_right = _trace_normalize_two_kron_marginals( + sigma_left, sigma_right, trace_mass, repeat_mass=repeats + ) + sigma_left = 0.5 * (sigma_left + sigma_left.T) + sigma_right = 0.5 * (sigma_right + sigma_right.T) + state.sigma_left.update(sigma_left, ema_old, ema_new) + state.sigma_right.update(sigma_right, ema_old, ema_new) + state.average_repeats.update(repeats, ema_old, ema_new) + return state + + +kfac_jax.set_default_tag_to_block_ctor( + STRUCTURAL_QUADRILINEAR_MERGE_TAG_VARIANT, StructuralQuadrilinearMergeBlock +) +SMALL_FULL_TAG_VARIANT = "small_full" +_SMALL_FULL_MAX_SIZE = 4096 + + +@kfac_utils.register_state_class +class _SmallFullBlockState(kfac_jax.CurvatureBlock.State): + matrix: kfac_utils.WeightedMovingAverage + + +class SmallFullBlock(kfac_jax.CurvatureBlock): + State = _SmallFullBlockState + + @property + def parameters_canonical_order(self) -> tuple[int, ...]: + return (0,) + + def _param_size(self) -> int: + shape = self.parameters_shapes[0] + n = 1 + for s in shape: + n *= int(s) + return n + + @staticmethod + def _safe_eigh(matrix): + matrix = 0.5 * (matrix + matrix.T) + diagonal_scale = jnp.max(jnp.abs(jnp.diagonal(matrix))) + floor = jnp.maximum( + jnp.asarray(1e-06, dtype=matrix.dtype), + jnp.asarray(0.0001, dtype=matrix.dtype) * diagonal_scale, + ) + matrix = matrix + floor * jnp.eye(matrix.shape[0], dtype=matrix.dtype) + scale = jnp.maximum( + jnp.max(jnp.abs(matrix)), jnp.asarray(1.0, dtype=matrix.dtype) + ) + eigenvalues, eigenvectors = kfac_utils.safe_psd_eigh(matrix / scale) + return (eigenvalues * scale, eigenvectors) + + def _init( + self, rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues + ): + del rng + n = self._param_size() + powers_to_cache = set(exact_powers_to_cache) | set(approx_powers_to_cache) + unsupported = powers_to_cache - {-1} + if unsupported: + raise NotImplementedError( + f"SmallFullBlock does not cache powers {sorted(unsupported)}." + ) + cache = {} + if -1 in powers_to_cache: + cache["-1"] = jnp.eye(n, dtype=self.dtype) + if cache_eigenvalues: + cache["eigenvalues"] = jnp.zeros((n,), dtype=self.dtype) + return self.State( + cache=cache, + matrix=kfac_utils.WeightedMovingAverage( + value=jnp.eye(n, dtype=self.dtype), + weight=jnp.asarray(1.0, dtype=self.dtype), + ), + ) + + def sync(self, state, pmap_axis_name): + state = state.copy() + state.matrix.sync(pmap_axis_name) + return state + + def update_curvature_matrix_estimate( + self, state, estimation_data, ema_old, ema_new, identity_weight, batch_size + ): + del identity_weight + state = state.copy() + [dy] = estimation_data.tangents.outputs + n = self._param_size() + d2 = dy.reshape(-1, n) + divisor = jnp.maximum( + jnp.asarray(batch_size, dtype=d2.dtype), jnp.asarray(1.0, dtype=d2.dtype) + ) + fisher = d2.T @ d2 / divisor + fisher = 0.5 * (fisher + fisher.T) + state.matrix.update(fisher, ema_old, ema_new) + return state + + def _multiply_matpower_unscaled( + self, state, vector, identity_weight, power, exact_power, use_cached + ): + del exact_power + [v] = vector + n = self._param_size() + vf = v.reshape(n) + if power == -1: + if use_cached: + out = state.cache["-1"] @ vf + else: + m = 0.5 * (state.matrix.value + state.matrix.value.T) + w_eig, q_eig = self._safe_eigh(m) + w_eig = w_eig + identity_weight + out = q_eig @ (q_eig.T @ vf / w_eig) + elif power == 1: + m = 0.5 * (state.matrix.value + state.matrix.value.T) + out = m @ vf + identity_weight * vf + else: + raise NotImplementedError( + f"SmallFullBlock: power={power} not implemented (only ±1)." + ) + return (out.reshape(v.shape),) + + def _eigenvalues_unscaled(self, state, use_cached): + if use_cached: + return state.cache["eigenvalues"] + matrix = 0.5 * (state.matrix.value + state.matrix.value.T) + eigenvalues, _ = self._safe_eigh(matrix) + return eigenvalues + + def _update_cache( + self, state, identity_weight, exact_powers, approx_powers, eigenvalues + ): + powers = set(exact_powers) | set(approx_powers) + unsupported = powers - {-1} + if unsupported: + raise NotImplementedError( + f"SmallFullBlock does not cache powers {sorted(unsupported)}." + ) + state = state.copy() + if eigenvalues or -1 in powers: + m = 0.5 * (state.matrix.value + state.matrix.value.T) + w_eig, q_eig = self._safe_eigh(m) + eig_ok = jnp.all(jnp.isfinite(w_eig)) & jnp.all(jnp.isfinite(q_eig)) + if eigenvalues: + state.cache["eigenvalues"] = jnp.where( + eig_ok, w_eig, state.cache["eigenvalues"] + ) + if -1 in powers: + inv_eig = 1.0 / (w_eig + identity_weight) + candidate_inverse = q_eig * inv_eig[None, :] @ q_eig.T + inverse_ok = eig_ok & jnp.all(jnp.isfinite(candidate_inverse)) + state.cache["-1"] = jnp.where( + inverse_ok, candidate_inverse, state.cache["-1"] + ) + return state + + def _to_dense_unscaled(self, state): + return state.matrix.value + + def _norm_unscaled(self, state, norm_type): + del norm_type + n = self._param_size() + return jnp.trace(state.matrix.value) / n + + +kfac_jax.set_default_tag_to_block_ctor(SMALL_FULL_TAG_VARIANT, SmallFullBlock) + + +def register_small_full(param, *, tag_id: str = ""): + if param.size > _SMALL_FULL_MAX_SIZE: + raise ValueError( + f"register_small_full: param size {param.size} exceeds {_SMALL_FULL_MAX_SIZE}; use a structured block instead." + ) + return layer_tag.bind( + param, + meta=LayerMetaData( + variant=SMALL_FULL_TAG_VARIANT, + outputs_index=(0,), + inputs_index=(), + params_index=(0,), + ), + ) diff --git a/src/hamiltonzero/optim/targets.py b/src/hamiltonzero/optim/targets.py new file mode 100644 index 0000000000000000000000000000000000000000..798b2723b46fc00ec72bef6b7352d89a4061317f --- /dev/null +++ b/src/hamiltonzero/optim/targets.py @@ -0,0 +1,117 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import jax +import jax.numpy as jnp + + +_ROUTE_SAMPLES = 8 + + +def _clip_one_channel(x, width: float): + median = jnp.nanmedian(x, axis=1, keepdims=True) + mean_ad = jnp.nanmean(jnp.abs(x - median), axis=1, keepdims=True) + delta = jnp.asarray(width, dtype=x.dtype) * mean_ad + return jnp.clip(x, median - delta, median + delta) + + +def _mad_clip_per_system(local_energies, width: float): + re = jnp.real(local_energies) + im = jnp.imag(local_energies) + re_clipped = _clip_one_channel(re, width) + im_clipped = _clip_one_channel(im, width) + if jnp.iscomplexobj(local_energies): + clipped = (re_clipped + 1j * im_clipped).astype(local_energies.dtype) + else: + clipped = re_clipped.astype(local_energies.dtype) + return clipped + + +def process_route_targets( + sampled_energy, + baseline_energy, + sigma1, + baseline_weights, + *, + mad_width: float = 5.0, +): + systems = int(sampled_energy.shape[0]) + if systems % _ROUTE_SAMPLES: + raise ValueError("router targets require a system axis divisible by K=8") + sigma = sigma1.astype(sampled_energy.real.dtype) + sampled_normalized = sampled_energy / sigma[:, None] + baseline_normalized = baseline_energy / sigma[:, None] + sampled_clipped = _mad_clip_per_system( + sampled_normalized, + mad_width, + ) + centered = sampled_clipped - jnp.mean( + sampled_clipped, + axis=1, + keepdims=True, + ) + variance = jnp.mean( + centered.real**2 + centered.imag**2, + axis=1, + keepdims=True, + ) + group_variance = jnp.mean( + variance.reshape(systems // _ROUTE_SAMPLES, _ROUTE_SAMPLES), + axis=1, + keepdims=True, + ) + group_std = jnp.sqrt( + jnp.broadcast_to( + group_variance, + (systems // _ROUTE_SAMPLES, _ROUTE_SAMPLES), + ) + ).reshape(systems, 1) + scale = jnp.maximum(group_std, jnp.asarray(1.0, dtype=group_std.dtype)) + sampled_target = sampled_clipped / scale + baseline_target = baseline_normalized / scale + sampled_rewards = jnp.mean(sampled_target.real, axis=1).astype(jnp.float32) + baseline_rewards = jnp.sum( + baseline_weights * baseline_target.real, + axis=1, + ).astype(jnp.float32) + grouped_baseline = baseline_rewards.reshape((-1, _ROUTE_SAMPLES)) + group_is_finite = jnp.all( + jnp.isfinite(grouped_baseline), + axis=1, + keepdims=True, + ) + baseline_rewards = jnp.where( + group_is_finite, + grouped_baseline, + jnp.zeros_like(grouped_baseline), + ).reshape(baseline_rewards.shape) + reward_delta = sampled_rewards.reshape( + (-1, _ROUTE_SAMPLES) + ) - baseline_rewards.reshape((-1, _ROUTE_SAMPLES)) + advantage = ( + float(_ROUTE_SAMPLES) + / float(_ROUTE_SAMPLES - 1) + * (reward_delta - jnp.mean(reward_delta, axis=1, keepdims=True)) + ) + advantage = jax.lax.stop_gradient( + advantage.reshape(sampled_rewards.shape).astype(jnp.float32) + ) + return sampled_target, advantage + + +def process_finetune_targets( + energy, + sigma1, + *, + mad_width: float = 5.0, +): + normalized = energy / sigma1.astype(energy.real.dtype)[:, None] + clipped = _mad_clip_per_system(normalized, mad_width) + centered = clipped - jnp.mean(clipped, axis=1, keepdims=True) + std = jnp.sqrt(jnp.mean(centered.real**2 + centered.imag**2, axis=1)) + return clipped / jnp.maximum(std, 1.0)[:, None] + + +__all__ = ["process_finetune_targets", "process_route_targets"] diff --git a/src/hamiltonzero/renyi2.py b/src/hamiltonzero/renyi2.py new file mode 100644 index 0000000000000000000000000000000000000000..c30c7aacde04f2f2d905793e0933fd52e1cd9d5f --- /dev/null +++ b/src/hamiltonzero/renyi2.py @@ -0,0 +1,678 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import math +from typing import NamedTuple, Sequence + +import equinox as eqx +import jax +import jax.numpy as jnp +import numpy as np + +from hamiltonzero.inference import PreparedInference + + +class BasisSamplerState(NamedTuple): + bits: jax.Array + log_abs: jax.Array + key: jax.Array + accepted: jax.Array + proposed: jax.Array + + +class Renyi2Result(NamedTuple): + purity: float | None + imaginary_mean: float | None + standard_error: float | None + renyi2_nats: float | None + renyi2_bits: float | None + resolved: bool + failure_reasons: tuple[str, ...] + naive_standard_error: float | None + imaginary_standard_error: float | None + imaginary_naive_standard_error: float | None + integrated_autocorrelation_time_blocks: float | None + integrated_autocorrelation_time_imaginary_blocks: float | None + effective_blocks: float + effective_imaginary_blocks: float + largest_absolute_block_fraction: float | None + renyi2_standard_error_nats: float | None + renyi2_lower_3sigma_nats: float | None + n_blocks: int + mean_log_abs: float + mean_phase: float + block_log_abs: np.ndarray + block_phase: np.ndarray + swap_log_abs: np.ndarray + swap_phase: np.ndarray + valid_denominator: np.ndarray + + +def _geometry(prepared: PreparedInference) -> tuple[jax.Array, int]: + if not isinstance(prepared, PreparedInference): + raise TypeError("prepared must be a PreparedInference") + route_host = np.asarray(jax.device_get(prepared.route), dtype=np.int32) + if route_host.ndim != 1: + raise ValueError("prepared route must have shape [N]") + if not np.array_equal(np.sort(route_host), np.arange(route_host.size)): + raise ValueError("prepared route must be a permutation") + mask = np.asarray(jax.device_get(prepared._initial_context.mask), dtype=np.bool_) + n_spins = int(np.sum(mask)) + if mask.shape != route_host.shape or not np.array_equal( + mask, np.arange(mask.size) < n_spins + ): + raise ValueError("prepared physical sites must be a contiguous prefix") + return jnp.asarray(prepared.route, dtype=jnp.int32), n_spins + + +def _as_bits(bits, *, n_spins: int) -> jax.Array: + value = jnp.asarray(bits) + if ( + value.ndim not in (2, 3) + or any(size < 1 for size in value.shape[:-1]) + or value.shape[-1] != n_spins + ): + raise ValueError( + f"bits must have shape [pairs, {n_spins}] or [blocks, pairs, {n_spins}]" + ) + if value.dtype != jnp.bool_: + invalid = jnp.any((value != 0) & (value != 1)) + if bool(jax.device_get(invalid)): + raise ValueError("bits must contain only zero and one") + return value.astype(jnp.bool_) + + +def _routed_corners(bits, route, *, n_spins: int): + width = int(route.shape[0]) + padding = width - int(n_spins) + if padding < 0: + raise ValueError("physical spin count exceeds compiled width") + full_bits = jnp.pad(bits, ((0, 0), (0, padding)), constant_values=False) + routed = jnp.take(full_bits, route, axis=-1) + up = jnp.logical_not(routed).astype(jnp.float32) + down = routed.astype(jnp.float32) + zeros = jnp.zeros_like(up) + return jnp.stack((up, zeros, down, zeros), axis=-1) + + +def _basis_log_wavefunction(wavefunction, bits, route, *, n_spins: int): + q = _routed_corners(bits, route, n_spins=n_spins) + return wavefunction(q, None, 0.0) + + +def _basis_log_abs(wavefunction, bits, route, *, n_spins: int): + log_abs, _phase = _basis_log_wavefunction( + wavefunction, bits, route, n_spins=n_spins + ) + return log_abs + + +def _metropolis_log_acceptance(current_log_abs, proposed_log_abs): + current_finite = jnp.isfinite(current_log_abs) + proposed_finite = jnp.isfinite(proposed_log_abs) + log_ratio = 2.0 * (proposed_log_abs - current_log_abs) + log_accept = jnp.minimum(jnp.zeros_like(log_ratio), log_ratio) + both_zero = jnp.isneginf(current_log_abs) & jnp.isneginf(proposed_log_abs) + recover_to_finite = jnp.logical_not(current_finite) & proposed_finite + invalid_proposal = jnp.logical_not(proposed_finite) & jnp.logical_not(both_zero) + log_accept = jnp.where(both_zero | recover_to_finite, 0.0, log_accept) + return jnp.where(invalid_proposal, -jnp.inf, log_accept) + + +def _state_from_bits(key, wavefunction, bits, route, *, n_spins: int): + log_abs = _basis_log_abs(wavefunction, bits, route, n_spins=n_spins) + return BasisSamplerState( + bits=bits, + log_abs=log_abs, + key=key, + accepted=jnp.asarray(0, dtype=jnp.int32), + proposed=jnp.asarray(0, dtype=jnp.int32), + ) + + +def _basis_step(state, wavefunction, route, *, n_spins: int): + site_key, accept_key, next_key = jax.random.split(state.key, 3) + batch_size = state.bits.shape[0] + sites = jax.random.randint(site_key, (batch_size,), 0, n_spins, dtype=jnp.int32) + rows = jnp.arange(batch_size) + proposed_bits = state.bits.at[rows, sites].set( + jnp.logical_not(state.bits[rows, sites]) + ) + proposed_log_abs = _basis_log_abs( + wavefunction, proposed_bits, route, n_spins=n_spins + ) + log_accept = _metropolis_log_acceptance(state.log_abs, proposed_log_abs) + log_uniform = jnp.log( + jax.random.uniform(accept_key, state.log_abs.shape, dtype=state.log_abs.dtype) + ) + accept = log_uniform < log_accept + return BasisSamplerState( + bits=jnp.where(accept[:, None], proposed_bits, state.bits), + log_abs=jnp.where(accept, proposed_log_abs, state.log_abs), + key=next_key, + accepted=state.accepted + jnp.sum(accept, dtype=jnp.int32), + proposed=state.proposed + jnp.asarray(state.bits.shape[0], dtype=jnp.int32), + ) + + +@eqx.filter_jit +def _run_basis_steps(state, wavefunction, route, *, n_spins: int, n_steps: int): + def one_step(carry, _): + return _basis_step(carry, wavefunction, route, n_spins=n_spins), None + + state, _ = jax.lax.scan(one_step, state, xs=None, length=n_steps) + return state + + +def _require_finite_state(state: BasisSamplerState) -> None: + finite = np.asarray(jax.device_get(jnp.isfinite(state.log_abs))) + if not np.all(finite): + raise RuntimeError( + "basis burn-in ended with a zero or non-finite wavefunction coefficient" + ) + + +def burn_in_basis( + prepared: PreparedInference, + key, + *, + batch_size: int = 256, + burn_in: int = 1024, +): + if int(batch_size) < 1: + raise ValueError("batch_size must be positive") + if int(burn_in) < 0: + raise ValueError("burn_in must be non-negative") + route, n_spins = _geometry(prepared) + bits_key, state_key = jax.random.split(key) + bits = jax.random.bernoulli(bits_key, shape=(int(batch_size), n_spins)) + state = _state_from_bits( + state_key, + prepared.wavefunction, + bits, + route, + n_spins=n_spins, + ) + state = _run_basis_steps( + state, + prepared.wavefunction, + route, + n_spins=n_spins, + n_steps=int(burn_in), + ) + jax.block_until_ready(state.log_abs) + _require_finite_state(state) + return state, state.bits + + +def step_basis( + prepared: PreparedInference, + state: BasisSamplerState, + *, + steps: int = 24, +): + if not isinstance(state, BasisSamplerState): + raise TypeError("state must be a BasisSamplerState") + if int(steps) < 1: + raise ValueError("steps must be positive") + route, n_spins = _geometry(prepared) + if state.bits.shape[-1] != n_spins: + raise ValueError("basis state width does not match the prepared system") + state = _run_basis_steps( + state, + prepared.wavefunction, + route, + n_spins=n_spins, + n_steps=int(steps), + ) + jax.block_until_ready(state.log_abs) + _require_finite_state(state) + return state, state.bits + + +def _subsystem_mask( + subsystem: Sequence[int] | Sequence[bool] | np.ndarray, + *, + n_spins: int, +) -> jax.Array: + value = np.asarray(subsystem) + if value.ndim == 1 and value.size == 0: + mask = np.zeros((n_spins,), dtype=np.bool_) + elif value.dtype == np.bool_: + if value.shape != (n_spins,): + raise ValueError(f"boolean subsystem mask must have shape [{n_spins}]") + mask = value + else: + if value.ndim != 1: + raise ValueError("subsystem site indices must be one-dimensional") + if not np.issubdtype(value.dtype, np.integer): + raise TypeError("subsystem must contain integer sites or booleans") + sites = value.astype(np.int64) + if len(np.unique(sites)) != sites.size: + raise ValueError("subsystem site indices must be unique") + if np.any(sites < 0) or np.any(sites >= n_spins): + raise ValueError("subsystem site index is out of range") + mask = np.zeros((n_spins,), dtype=np.bool_) + mask[sites] = True + return jnp.asarray(mask) + + +@eqx.filter_jit +def _swap_log_ratios( + wavefunction, + replica_x, + replica_y, + route, + region_mask, + *, + n_spins: int, +): + swapped_x = jnp.where(region_mask[None, :], replica_y, replica_x) + swapped_y = jnp.where(region_mask[None, :], replica_x, replica_y) + denominator_x, phase_x = _basis_log_wavefunction( + wavefunction, replica_x, route, n_spins=n_spins + ) + denominator_y, phase_y = _basis_log_wavefunction( + wavefunction, replica_y, route, n_spins=n_spins + ) + numerator_x, numerator_phase_x = _basis_log_wavefunction( + wavefunction, swapped_x, route, n_spins=n_spins + ) + numerator_y, numerator_phase_y = _basis_log_wavefunction( + wavefunction, swapped_y, route, n_spins=n_spins + ) + valid_denominator = jnp.isfinite(denominator_x) & jnp.isfinite(denominator_y) + log_abs = numerator_x + numerator_y - denominator_x - denominator_y + phase = numerator_phase_x + numerator_phase_y - phase_x - phase_y + phase = jnp.arctan2(jnp.sin(phase), jnp.cos(phase)) + log_abs = jnp.where(valid_denominator, log_abs, -jnp.inf) + phase = jnp.where(valid_denominator, phase, 0.0) + exact_identity = jnp.logical_or( + jnp.all(jnp.logical_not(region_mask)), jnp.all(region_mask) + ) + log_abs = jnp.where(exact_identity & valid_denominator, 0.0, log_abs) + phase = jnp.where(exact_identity & valid_denominator, 0.0, phase) + return log_abs, phase, valid_denominator + + +def _complex_mean_log_polar(log_abs, phase) -> tuple[float, float]: + logs = np.asarray(log_abs, dtype=np.float64).reshape(-1) + phases = np.asarray(phase, dtype=np.float64).reshape(-1) + if logs.shape != phases.shape or logs.size == 0: + raise ValueError("log_abs and phase must have matching nonempty shapes") + if ( + np.any(np.isnan(logs)) + or np.any(np.isposinf(logs)) + or np.any(~np.isfinite(phases)) + ): + raise ValueError("non-finite SWAP log-polar sample") + finite = np.isfinite(logs) + if not np.any(finite): + return -math.inf, 0.0 + pivot = float(np.max(logs[finite])) + scaled = np.zeros(logs.shape, dtype=np.complex128) + scaled[finite] = np.exp(logs[finite] - pivot + 1j * phases[finite]) + scaled_mean = np.mean(scaled) + magnitude = float(abs(scaled_mean)) + if magnitude == 0.0: + return -math.inf, 0.0 + return pivot + math.log(magnitude), float(np.angle(scaled_mean)) + + +def _log_polar_to_complex(log_abs: float, phase: float) -> complex: + if log_abs == -math.inf: + return 0.0j + if not math.isfinite(log_abs) or not math.isfinite(phase): + raise ValueError("log-polar scalar must be finite or exact zero") + if log_abs > math.log(np.finfo(np.float64).max): + raise OverflowError("complex SWAP block mean exceeds float64 range") + return complex(math.exp(log_abs) * np.exp(1j * phase)) + + +def _integrated_autocorrelation_time(values) -> float: + x = np.asarray(values, dtype=np.float64).reshape(-1) + if x.size < 2: + return 0.5 + x = x - np.mean(x) + variance = float(np.dot(x, x) / x.size) + if not math.isfinite(variance) or variance <= 0.0: + return 0.5 + correlations = [] + for lag in range(1, x.size): + covariance = float(np.dot(x[:-lag], x[lag:]) / x.size) + correlations.append(covariance / variance) + tau = 0.5 + previous_pair = math.inf + for offset in range(0, len(correlations) - 1, 2): + pair = correlations[offset] + correlations[offset + 1] + if not math.isfinite(pair) or pair <= 0.0: + break + pair = min(pair, previous_pair) + tau += pair + previous_pair = pair + return max(0.5, float(tau)) + + +def _summarize_blocks(block_log_abs, block_phase): + logs = np.asarray(block_log_abs, dtype=np.float64).reshape(-1) + phases = np.asarray(block_phase, dtype=np.float64).reshape(-1) + if logs.shape != phases.shape or logs.size == 0: + raise ValueError("block log magnitudes and phases must match") + try: + blocks = np.asarray( + [ + _log_polar_to_complex(float(value), float(angle)) + for value, angle in zip(logs, phases, strict=True) + ], + dtype=np.complex128, + ) + except (OverflowError, ValueError): + return { + "n_blocks": int(logs.size), + "purity": None, + "naive_standard_error": None, + "standard_error": None, + "imaginary_mean": None, + "imaginary_naive_standard_error": None, + "imaginary_standard_error": None, + "integrated_autocorrelation_time_blocks": None, + "integrated_autocorrelation_time_imaginary_blocks": None, + "effective_blocks": 0.0, + "effective_imaginary_blocks": 0.0, + "largest_absolute_block_fraction": None, + "resolved": False, + "failure_reasons": ("block_mean_float64_overflow_or_nonfinite",), + "renyi2_nats": None, + "renyi2_bits": None, + "renyi2_standard_error_nats": None, + "renyi2_lower_3sigma_nats": None, + } + n_blocks = int(blocks.size) + real = blocks.real + imaginary = blocks.imag + purity = float(np.mean(real)) + imaginary_mean = float(np.mean(imaginary)) + naive_standard_error = ( + float(np.std(real, ddof=1) / math.sqrt(n_blocks)) if n_blocks > 1 else math.inf + ) + imaginary_naive_standard_error = ( + float(np.std(imaginary, ddof=1) / math.sqrt(n_blocks)) + if n_blocks > 1 + else math.inf + ) + tau = _integrated_autocorrelation_time(real) + effective_blocks = float(n_blocks / (2.0 * tau)) + imaginary_tau = _integrated_autocorrelation_time(imaginary) + effective_imaginary_blocks = float(n_blocks / (2.0 * imaginary_tau)) + standard_error = ( + float(np.std(real, ddof=1) / math.sqrt(effective_blocks)) + if n_blocks > 1 + else math.inf + ) + imaginary_standard_error = ( + float(np.std(imaginary, ddof=1) / math.sqrt(effective_imaginary_blocks)) + if n_blocks > 1 + else math.inf + ) + absolute_sum = float(np.sum(np.abs(blocks))) + tail_fraction = ( + float(np.max(np.abs(blocks)) / absolute_sum) if absolute_sum > 0.0 else 1.0 + ) + failures = [] + if n_blocks < 16: + failures.append("too_few_blocks") + if not math.isfinite(purity) or not math.isfinite(standard_error): + failures.append("nonfinite_real_estimate") + elif purity <= 3.0 * standard_error: + failures.append("purity_not_resolved_above_zero") + if math.isfinite(imaginary_standard_error): + if abs(imaginary_mean) > 3.0 * imaginary_standard_error: + failures.append("imaginary_null_test_failed") + else: + failures.append("nonfinite_imaginary_uncertainty") + if purity > 1.0: + failures.append("purity_point_above_physical_upper_bound") + if effective_blocks < 8.0: + failures.append("insufficient_effective_blocks") + if tail_fraction > 0.25: + failures.append("single_block_tail_dominance") + resolved = not failures + entropy = -math.log(purity) if resolved else None + entropy_standard_error = standard_error / purity if resolved else None + purity_upper_3sigma = min(1.0, purity + 3.0 * standard_error) if resolved else None + entropy_lower_3sigma = ( + max(0.0, -math.log(purity_upper_3sigma)) + if purity_upper_3sigma is not None + else None + ) + return { + "n_blocks": n_blocks, + "purity": purity, + "naive_standard_error": naive_standard_error, + "standard_error": standard_error, + "imaginary_mean": imaginary_mean, + "imaginary_naive_standard_error": imaginary_naive_standard_error, + "imaginary_standard_error": imaginary_standard_error, + "integrated_autocorrelation_time_blocks": tau, + "integrated_autocorrelation_time_imaginary_blocks": imaginary_tau, + "effective_blocks": effective_blocks, + "effective_imaginary_blocks": effective_imaginary_blocks, + "largest_absolute_block_fraction": tail_fraction, + "resolved": resolved, + "failure_reasons": tuple(failures), + "renyi2_nats": entropy, + "renyi2_bits": entropy / math.log(2.0) if entropy is not None else None, + "renyi2_standard_error_nats": entropy_standard_error, + "renyi2_lower_3sigma_nats": entropy_lower_3sigma, + } + + +def _evaluate_swap( + prepared, + x, + y, + route, + mask, + *, + n_spins: int, + chunk_size: int, +): + logs = [] + phases = [] + valid = [] + for start in range(0, x.shape[0], chunk_size): + stop = min(start + chunk_size, x.shape[0]) + values = _swap_log_ratios( + prepared.wavefunction, + x[start:stop], + y[start:stop], + route, + mask, + n_spins=n_spins, + ) + values = jax.device_get(values) + logs.append(np.asarray(values[0])) + phases.append(np.asarray(values[1])) + valid.append(np.asarray(values[2])) + return ( + np.concatenate(logs), + np.concatenate(phases), + np.concatenate(valid), + ) + + +def _result(log_abs, phase, valid) -> Renyi2Result: + swap_log_abs = np.asarray(log_abs) + swap_phase = np.asarray(phase) + valid_denominator = np.asarray(valid, dtype=np.bool_) + if ( + swap_log_abs.ndim != 2 + or swap_log_abs.shape != swap_phase.shape + or swap_log_abs.shape != valid_denominator.shape + ): + raise ValueError("SWAP blocks must have aligned shape [blocks, pairs]") + if not np.all(valid_denominator): + raise RuntimeError("SWAP denominator contains a zero wavefunction coefficient") + block_values = [ + _complex_mean_log_polar(logs, phases) + for logs, phases in zip(swap_log_abs, swap_phase, strict=True) + ] + block_log_abs = np.asarray([value[0] for value in block_values], dtype=np.float64) + block_phase = np.asarray([value[1] for value in block_values], dtype=np.float64) + summary = _summarize_blocks(block_log_abs, block_phase) + mean_log_abs, mean_phase = _complex_mean_log_polar(block_log_abs, block_phase) + return Renyi2Result( + purity=summary["purity"], + imaginary_mean=summary["imaginary_mean"], + standard_error=summary["standard_error"], + renyi2_nats=summary["renyi2_nats"], + renyi2_bits=summary["renyi2_bits"], + resolved=summary["resolved"], + failure_reasons=summary["failure_reasons"], + naive_standard_error=summary["naive_standard_error"], + imaginary_standard_error=summary["imaginary_standard_error"], + imaginary_naive_standard_error=summary["imaginary_naive_standard_error"], + integrated_autocorrelation_time_blocks=summary[ + "integrated_autocorrelation_time_blocks" + ], + integrated_autocorrelation_time_imaginary_blocks=summary[ + "integrated_autocorrelation_time_imaginary_blocks" + ], + effective_blocks=summary["effective_blocks"], + effective_imaginary_blocks=summary["effective_imaginary_blocks"], + largest_absolute_block_fraction=summary["largest_absolute_block_fraction"], + renyi2_standard_error_nats=summary["renyi2_standard_error_nats"], + renyi2_lower_3sigma_nats=summary["renyi2_lower_3sigma_nats"], + n_blocks=summary["n_blocks"], + mean_log_abs=mean_log_abs, + mean_phase=mean_phase, + block_log_abs=block_log_abs, + block_phase=block_phase, + swap_log_abs=swap_log_abs, + swap_phase=swap_phase, + valid_denominator=valid_denominator, + ) + + +def renyi2_purity( + prepared: PreparedInference, + replica_x, + replica_y, + subsystem: Sequence[int] | Sequence[bool] | np.ndarray, + *, + chunk_size: int = 256, +) -> Renyi2Result: + if int(chunk_size) < 1: + raise ValueError("chunk_size must be positive") + route, n_spins = _geometry(prepared) + x = _as_bits(replica_x, n_spins=n_spins) + y = _as_bits(replica_y, n_spins=n_spins) + if x.shape != y.shape: + raise ValueError("replica batches must have the same shape") + if x.ndim == 2: + x = x[None, ...] + y = y[None, ...] + n_blocks, pairs_per_block = x.shape[:2] + mask = _subsystem_mask(subsystem, n_spins=n_spins) + values = _evaluate_swap( + prepared, + x.reshape((n_blocks * pairs_per_block, n_spins)), + y.reshape((n_blocks * pairs_per_block, n_spins)), + route, + mask, + n_spins=n_spins, + chunk_size=int(chunk_size), + ) + return _result( + values[0].reshape((n_blocks, pairs_per_block)), + values[1].reshape((n_blocks, pairs_per_block)), + values[2].reshape((n_blocks, pairs_per_block)), + ) + + +def measure_renyi2( + prepared: PreparedInference, + replica_x: BasisSamplerState, + replica_y: BasisSamplerState, + subsystem: Sequence[int] | Sequence[bool] | np.ndarray, + *, + blocks: int = 16, + samples_per_block: int = 1, + steps_between: int = 24, + chunk_size: int = 256, +): + if not isinstance(replica_x, BasisSamplerState) or not isinstance( + replica_y, BasisSamplerState + ): + raise TypeError("replicas must be BasisSamplerState values") + if int(blocks) < 1: + raise ValueError("blocks must be positive") + if int(samples_per_block) < 1: + raise ValueError("samples_per_block must be positive") + if int(steps_between) < 1: + raise ValueError("steps_between must be positive") + if int(chunk_size) < 1: + raise ValueError("chunk_size must be positive") + route, n_spins = _geometry(prepared) + if replica_x.bits.shape != replica_y.bits.shape: + raise ValueError("replica states must have the same walker shape") + if replica_x.bits.ndim != 2 or replica_x.bits.shape[-1] != n_spins: + raise ValueError("basis state width does not match the prepared system") + _require_finite_state(replica_x) + _require_finite_state(replica_y) + mask = _subsystem_mask(subsystem, n_spins=n_spins) + block_logs = [] + block_phases = [] + block_valid = [] + for _ in range(int(blocks)): + logs = [] + phases = [] + valid = [] + for _ in range(int(samples_per_block)): + replica_x = _run_basis_steps( + replica_x, + prepared.wavefunction, + route, + n_spins=n_spins, + n_steps=int(steps_between), + ) + replica_y = _run_basis_steps( + replica_y, + prepared.wavefunction, + route, + n_spins=n_spins, + n_steps=int(steps_between), + ) + values = _evaluate_swap( + prepared, + replica_x.bits, + replica_y.bits, + route, + mask, + n_spins=n_spins, + chunk_size=int(chunk_size), + ) + logs.append(values[0]) + phases.append(values[1]) + valid.append(values[2]) + block_logs.append(np.concatenate(logs)) + block_phases.append(np.concatenate(phases)) + block_valid.append(np.concatenate(valid)) + result = _result( + np.stack(block_logs), + np.stack(block_phases), + np.stack(block_valid), + ) + return replica_x, replica_y, result + + +__all__ = [ + "BasisSamplerState", + "Renyi2Result", + "burn_in_basis", + "measure_renyi2", + "renyi2_purity", + "step_basis", +] diff --git a/src/hamiltonzero/router/__init__.py b/src/hamiltonzero/router/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..b05be1cf444b70416cd5fd7559d483359e30489d --- /dev/null +++ b/src/hamiltonzero/router/__init__.py @@ -0,0 +1,40 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from .api import ( + batch_context, + route_context, + route_state, + select_frozen_route, + strip_router, +) +from .baseline import snis_mode_baseline +from .compiled import bind_router_kernel, compile_router_static +from .decode import ( + GLOBAL_BEAM_WIDTH, + ROUTE_SAMPLES, + build_beam16, + build_route_sampler, +) +from .permutation import permute_ctx_prefix, permute_multi_ctx_prefix, permute_q_prefix +from .state import rebase_cold_samples, reframe_state_context + +__all__ = [ + "batch_context", + "GLOBAL_BEAM_WIDTH", + "ROUTE_SAMPLES", + "bind_router_kernel", + "build_beam16", + "build_route_sampler", + "compile_router_static", + "permute_ctx_prefix", + "permute_multi_ctx_prefix", + "permute_q_prefix", + "rebase_cold_samples", + "reframe_state_context", + "route_context", + "route_state", + "select_frozen_route", + "snis_mode_baseline", + "strip_router", +] diff --git a/src/hamiltonzero/router/api.py b/src/hamiltonzero/router/api.py new file mode 100644 index 0000000000000000000000000000000000000000..780cc8397b90a6878a108bcc1d0f00631ffc4e6c --- /dev/null +++ b/src/hamiltonzero/router/api.py @@ -0,0 +1,81 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +import jax.numpy as jnp + +from hamiltonzero.model.context import MultiSystemContext, SpinContext +from .permutation import permute_ctx_prefix, permute_multi_ctx_prefix, permute_q_prefix + + +def batch_context(context: SpinContext) -> MultiSystemContext: + return MultiSystemContext.from_single(context) + + +def select_frozen_route( + model, + context: SpinContext, + *, + tau: float, +): + node, edge, global_features = model.route_features(context) + permutations, _log_probabilities = model.route_decoder.beam_search( + node, + edge, + context.bmask, + global_feat=global_features, + tau=tau, + beam_width=8, + real_mask=context.mask, + first_orbit_ids=( + context.route_quotient_node_key, + context.route_quotient_edge_key, + ), + ) + return permutations[0].astype(jnp.int32) + + +def route_context(context, perm): + if isinstance(context, MultiSystemContext): + perms = perm if perm.ndim == 2 else perm[None, :] + routed = permute_multi_ctx_prefix(context, perms) + return eqx.tree_at(lambda c: c.route_perm, routed, perms) + routed = permute_ctx_prefix(context, perm) + return eqx.tree_at(lambda c: c.route_perm, routed, perm) + + +def route_state(state, perm): + if state.q.ndim == 5: + perms = perm if perm.ndim == 2 else perm[None, :] + q = permute_q_prefix(state.q, perms) + grad = permute_q_prefix(state.grad_log_p, perms) + mask = jnp.take_along_axis(state.mask, perms, axis=-1) + else: + p = perm if perm.ndim == 1 else perm[0] + q = permute_q_prefix(state.q[None, ...], p[None, :])[0] + grad = permute_q_prefix(state.grad_log_p[None, ...], p[None, :])[0] + mask = jnp.take_along_axis(state.mask, p, axis=-1) + return eqx.tree_at( + lambda s: (s.q, s.grad_log_p, s.mask), + state, + (q, grad, mask), + ) + + +def strip_router(model): + return eqx.tree_at( + lambda m: (m.route_decoder, m.route_contextualizer), + model, + (None, None), + ) + + +__all__ = [ + "batch_context", + "route_context", + "route_state", + "select_frozen_route", + "strip_router", +] diff --git a/src/hamiltonzero/router/baseline.py b/src/hamiltonzero/router/baseline.py new file mode 100644 index 0000000000000000000000000000000000000000..6e978bc7b2288cf0acd99da4c1145f8064ca33d7 --- /dev/null +++ b/src/hamiltonzero/router/baseline.py @@ -0,0 +1,26 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import jax + + +def snis_mode_baseline( + total, + candidate_log_p, + sampled_log_p, +): + if total.shape != candidate_log_p.shape: + raise ValueError("baseline energy and log-density shapes differ") + if sampled_log_p.shape != total.shape: + raise ValueError("sampled log-density must match baseline energy") + weights = jax.nn.softmax( + candidate_log_p.astype(total.real.dtype) + - sampled_log_p.astype(total.real.dtype), + axis=-1, + ) + return weights + + +__all__ = ["snis_mode_baseline"] diff --git a/src/hamiltonzero/router/compiled.py b/src/hamiltonzero/router/compiled.py new file mode 100644 index 0000000000000000000000000000000000000000..68780f8e2e511491e19d5cd062220f80421c272d --- /dev/null +++ b/src/hamiltonzero/router/compiled.py @@ -0,0 +1,83 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import jax +import jax.numpy as jnp + +from hamiltonzero.compiled.types import SharedTrunk +from hamiltonzero.model.route_pointer import TreePrefixPointerMHSEA, lca_gaussian_decay + +from .types import RouterKernel, RouterStatic + + +def bind_router_kernel(model) -> RouterKernel: + decoder = model.route_decoder + if not isinstance(decoder, TreePrefixPointerMHSEA): + raise ValueError("learned-router train requires TreePrefixPointerMHSEA") + return RouterKernel( + contextualizer=model.route_contextualizer, + global_fork=model.gladder_fork_route, + decoder=decoder, + ) + + +def compile_router_static( + kernel: RouterKernel, + trunk: SharedTrunk, + quotient_node_key, + quotient_edge_key, + needs_fwl2, +) -> RouterStatic: + decoder = kernel.decoder + node_input, raw_edge, global_input = kernel.contextualizer.with_edge( + trunk.node_raw, + trunk.edge_raw, + trunk.real_mask, + trunk.balanced_mask, + g=trunk.global_stream, + ) + global_input = kernel.global_fork(global_input, raw_edge, trunk.balanced_mask) + (node_raw, node_projected), _ = decoder._prepare_nodes( + node_input, trunk.balanced_mask + ) + global_raw, global_projected = decoder._project_global( + global_input, node_input.dtype + ) + ((_, _, initial_suffix, _, _), edge_messages) = ( + decoder._initial_summaries_and_edge_messages( + raw_edge, trunk.balanced_mask, node_input.dtype + ) + ) + prefix_edge_messages, suffix_edge_messages = edge_messages + n = node_input.shape[0] + idx = jnp.arange(n, dtype=jnp.int32) + return RouterStatic( + node_input=node_raw, + node_projected=node_projected, + global_input=global_raw, + global_projected=global_projected, + raw_edge=raw_edge, + initial_suffix=initial_suffix, + prefix_edge_messages=prefix_edge_messages, + suffix_edge_messages=suffix_edge_messages, + order_decay=lca_gaussian_decay( + idx, idx, decoder.order_decay_w[0], decoder.order_decay_b[0] + ), + virtual_decay=lca_gaussian_decay( + idx, idx, decoder.virt_decay_w[0], decoder.virt_decay_b[0] + ), + tree_pair_messages=decoder._tree_pair_messages(raw_edge, trunk.balanced_mask), + static_bias_tables=decoder._pack_heavy_static_bias_tables(raw_edge), + quadratic_base_static=(), + forked_global_static=(), + quotient_node_key=quotient_node_key, + quotient_edge_key=quotient_edge_key, + real_mask=trunk.real_mask, + routable_mask=trunk.balanced_mask, + needs_fwl2=jnp.asarray(needs_fwl2, dtype=jnp.bool_), + ) + + +__all__ = ["bind_router_kernel", "compile_router_static"] diff --git a/src/hamiltonzero/router/decode.py b/src/hamiltonzero/router/decode.py new file mode 100644 index 0000000000000000000000000000000000000000..8c3c6a7ab3f5c9c939b0c2e2d1ba02fa572ccffb --- /dev/null +++ b/src/hamiltonzero/router/decode.py @@ -0,0 +1,127 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import functools + +import jax +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +ROUTE_SAMPLES = 8 +GLOBAL_BEAM_WIDTH = 16 + + +def _decode_one(decoder, static, key, tau): + return decoder._decode( + static.node_input, + static.raw_edge, + static.routable_mask, + tau=tau, + key=key, + real_mask=static.real_mask, + first_orbit_ids=( + static.quotient_node_key, + static.quotient_edge_key, + static.needs_fwl2, + ), + router_static=static, + ) + + +def _beam_local(decoder, static, tau, *, lanes): + permutations, _log_probabilities = decoder.beam_search( + static.node_input, + static.raw_edge, + static.routable_mask, + global_feat=static.global_input, + tau=tau, + beam_width=GLOBAL_BEAM_WIDTH, + real_mask=static.real_mask, + first_orbit_ids=( + static.quotient_node_key, + static.quotient_edge_key, + static.needs_fwl2, + ), + router_static=static, + distributed_axis_name="systems", + distributed_lanes=lanes, + ) + return permutations[0] + + +def build_beam16(mesh: Mesh, decoder, static): + lanes = int(mesh.shape["systems"]) + if tuple(mesh.axis_names) != ("systems",) or GLOBAL_BEAM_WIDTH % lanes: + raise ValueError("beam16 requires a one-dimensional divisible systems mesh") + mapped = jax.shard_map( + functools.partial(_beam_local, lanes=lanes), + mesh=mesh, + in_specs=( + jax.tree_util.tree_map(lambda _: P(), decoder), + jax.tree_util.tree_map(lambda _: P(), static), + P(), + ), + out_specs=P(), + check_vma=False, + ) + replicated = NamedSharding(mesh, P()) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _: replicated, decoder), + jax.tree_util.tree_map(lambda _: replicated, static), + replicated, + ), + out_shardings=replicated, + ) + + +def build_route_sampler(mesh: Mesh, decoder, static): + if tuple(mesh.axis_names) != ("systems",) or mesh.shape["systems"] != ROUTE_SAMPLES: + raise ValueError("learned-router train requires an eight-lane systems mesh") + replicated = NamedSharding(mesh, P()) + route_vector = NamedSharding(mesh, P("systems", None)) + local_specs = ( + jax.tree_util.tree_map(lambda _: P(), decoder), + jax.tree_util.tree_map(lambda _: P(), static), + P(), + P(), + ) + + def local(decoder_value, static_value, key, tau): + lane_key = jax.random.fold_in(key, jax.lax.axis_index("systems")) + sample_key = jax.random.split(lane_key, 1)[0] + permutation = _decode_one( + decoder_value, + static_value, + sample_key, + tau, + ) + return permutation[None] + + mapped = jax.shard_map( + local, + mesh=mesh, + in_specs=local_specs, + out_specs=P("systems", None), + check_vma=False, + ) + return jax.jit( + mapped, + in_shardings=( + jax.tree_util.tree_map(lambda _: replicated, decoder), + jax.tree_util.tree_map(lambda _: replicated, static), + replicated, + replicated, + ), + out_shardings=route_vector, + ) + + +__all__ = [ + "GLOBAL_BEAM_WIDTH", + "ROUTE_SAMPLES", + "build_beam16", + "build_route_sampler", +] diff --git a/src/hamiltonzero/router/permutation.py b/src/hamiltonzero/router/permutation.py new file mode 100644 index 0000000000000000000000000000000000000000..eeeda7d8eefc7ba954bab1046285f7514fa38061 --- /dev/null +++ b/src/hamiltonzero/router/permutation.py @@ -0,0 +1,76 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +import jax +import jax.numpy as jnp +from jaxtyping import Array, Float, Int + +from hamiltonzero.model.context import MultiSystemContext, SpinContext + + +def permute_ctx_prefix(ctx: SpinContext, perm: Int[Array, "n"]) -> SpinContext: + return eqx.tree_at( + lambda c: ( + c.h_prime, + c.J_double_prime, + c.mask, + c.route_quotient_node_key, + c.route_quotient_edge_key, + ), + ctx, + ( + ctx.h_prime[perm], + ctx.J_double_prime[perm][:, perm, :], + ctx.mask[perm], + ctx.route_quotient_node_key[perm], + ( + ctx.route_quotient_edge_key + if ctx.route_quotient_edge_key.shape[-1] == 0 + else ctx.route_quotient_edge_key[perm][:, perm] + ), + ), + ) + + +def permute_multi_ctx_prefix( + ctx: MultiSystemContext, + perms: Int[Array, "s n"], +) -> MultiSystemContext: + return eqx.tree_at( + lambda c: ( + c.h_prime, + c.J_double_prime, + c.mask, + c.route_quotient_node_key, + c.route_quotient_edge_key, + ), + ctx, + ( + jax.vmap(lambda x, p: x[p])(ctx.h_prime, perms), + jax.vmap(lambda x, p: x[p][:, p, :])(ctx.J_double_prime, perms), + jax.vmap(lambda x, p: x[p])(ctx.mask, perms), + jax.vmap(lambda x, p: x[p])(ctx.route_quotient_node_key, perms), + ( + ctx.route_quotient_edge_key + if ctx.route_quotient_edge_key.shape[-1] == 0 + else jax.vmap(lambda x, p: x[p][:, p])( + ctx.route_quotient_edge_key, perms + ) + ), + ), + ) + + +def permute_q_prefix(q: Float[Array, "s b r n d"], perms: Int[Array, "s n"]): + idx = jnp.broadcast_to(perms[:, None, None, :, None], q.shape) + return jnp.take_along_axis(q, idx, axis=3) + + +__all__ = [ + "permute_ctx_prefix", + "permute_multi_ctx_prefix", + "permute_q_prefix", +] diff --git a/src/hamiltonzero/router/state.py b/src/hamiltonzero/router/state.py new file mode 100644 index 0000000000000000000000000000000000000000..ff0bf2ae3fda5366cb8a3b035c2bfde061ad4020 --- /dev/null +++ b/src/hamiltonzero/router/state.py @@ -0,0 +1,41 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +import jax +import jax.numpy as jnp + +from .permutation import permute_multi_ctx_prefix, permute_q_prefix + + +def reframe_state_context(state, context, old_perms, new_perms): + if old_perms.shape != new_perms.shape: + raise ValueError("old and new route permutations must have equal shapes") + old_inverse = jnp.argsort(old_perms, axis=-1) + composed = jnp.take_along_axis(old_inverse, new_perms, axis=-1) + q = permute_q_prefix(state.q, composed) + grad = permute_q_prefix(state.grad_log_p, composed) + mask = jnp.take_along_axis(state.mask, composed, axis=-1) + state = eqx.tree_at( + lambda value: (value.q, value.grad_log_p, value.mask), + state, + (q, grad, mask), + ) + context = permute_multi_ctx_prefix(context, composed) + context = eqx.tree_at( + lambda value: value.route_perm, + context, + new_perms, + ) + return state, context + + +def rebase_cold_samples(q_routed, perms): + inverse = jnp.argsort(perms, axis=-1) + index = jnp.broadcast_to(inverse[:, None, :, None], q_routed.shape) + return jnp.take_along_axis(q_routed, index, axis=-2) + + +__all__ = ["rebase_cold_samples", "reframe_state_context"] diff --git a/src/hamiltonzero/router/types.py b/src/hamiltonzero/router/types.py new file mode 100644 index 0000000000000000000000000000000000000000..b7e623c4656d18fda6882d5837997c55adcc3032 --- /dev/null +++ b/src/hamiltonzero/router/types.py @@ -0,0 +1,38 @@ +# Copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import equinox as eqx +from jaxtyping import Array + + +class RouterKernel(eqx.Module): + contextualizer: object + global_fork: object + decoder: object + + +class RouterStatic(eqx.Module): + node_input: Array + node_projected: Array + global_input: Array + global_projected: Array + raw_edge: Array + initial_suffix: Array + prefix_edge_messages: Array + suffix_edge_messages: Array + order_decay: Array + virtual_decay: Array + tree_pair_messages: Array + static_bias_tables: tuple[Array, ...] + quadratic_base_static: tuple[Array, ...] + forked_global_static: tuple[Array, ...] + quotient_node_key: Array + quotient_edge_key: Array + real_mask: Array + routable_mask: Array + needs_fwl2: Array + + +__all__ = ["RouterKernel", "RouterStatic"] diff --git a/src/kfac_jax/__init__.py b/src/kfac_jax/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..be18f6d7464a5f4fd09864fadf5731dfb84c5b81 --- /dev/null +++ b/src/kfac_jax/__init__.py @@ -0,0 +1,216 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""kfac-jax public APIs.""" + +from kfac_jax._src import curvature_blocks +from kfac_jax._src import curvature_estimator +from kfac_jax._src import layers_and_loss_tags +from kfac_jax._src import loss_functions +from kfac_jax._src import optimizer +from kfac_jax._src import patches_second_moment +from kfac_jax._src import tag_graph_matcher +from kfac_jax._src import tracer +from kfac_jax._src import utils + + +__version__ = "0.0.8" + +# Patches Second Moments +patches_moments = patches_second_moment.patches_moments +patches_moments_explicit = patches_second_moment.patches_moments_explicit + +# Layers and loss tags +LayerData = layers_and_loss_tags.LayerData +LayerMetaData = layers_and_loss_tags.LayerMetaData +LossTag = layers_and_loss_tags.LossTag +LayerTag = layers_and_loss_tags.LayerTag +register_generic = layers_and_loss_tags.register_generic +register_dense = layers_and_loss_tags.register_dense +register_conv2d = layers_and_loss_tags.register_conv2d +register_scale_and_shift = layers_and_loss_tags.register_scale_and_shift + +# Tag graph matcher +auto_register_tags = tag_graph_matcher.auto_register_tags + +# Tracer +ProcessedJaxpr = tracer.ProcessedJaxpr +LayerVjpData = tracer.LayerVjpData +loss_tags_vjp = tracer.loss_tags_vjp +loss_tags_jvp = tracer.loss_tags_jvp +loss_tags_hvp = tracer.loss_tags_hvp +layer_tags_vjp = tracer.layer_tags_vjp + +# Loss functions +LossFunction = loss_functions.LossFunction +NegativeLogProbLoss = loss_functions.NegativeLogProbLoss +DistributionNegativeLogProbLoss = loss_functions.DistributionNegativeLogProbLoss +NormalMeanNegativeLogProbLoss = loss_functions.NormalMeanNegativeLogProbLoss +NormalMeanVarianceNegativeLogProbLoss = ( + loss_functions.NormalMeanVarianceNegativeLogProbLoss) +MultiBernoulliNegativeLogProbLoss = ( + loss_functions.MultiBernoulliNegativeLogProbLoss) +CategoricalLogitsNegativeLogProbLoss = ( + loss_functions.CategoricalLogitsNegativeLogProbLoss) +OneHotCategoricalLogitsNegativeLogProbLoss = ( + loss_functions.OneHotCategoricalLogitsNegativeLogProbLoss) +register_sigmoid_cross_entropy_loss = ( + loss_functions.register_sigmoid_cross_entropy_loss) +register_multi_bernoulli_predictive_distribution = ( + loss_functions.register_multi_bernoulli_predictive_distribution) +register_softmax_cross_entropy_loss = ( + loss_functions.register_softmax_cross_entropy_loss) +register_categorical_predictive_distribution = ( + loss_functions.register_categorical_predictive_distribution) +register_squared_error_loss = loss_functions.register_squared_error_loss +register_normal_predictive_distribution = ( + loss_functions.register_normal_predictive_distribution) + +# Curvature blocks +CurvatureBlock = curvature_blocks.CurvatureBlock +ScaledIdentity = curvature_blocks.ScaledIdentity +Diagonal = curvature_blocks.Diagonal +Full = curvature_blocks.Full +KroneckerFactored = curvature_blocks.KroneckerFactored +NaiveDiagonal = curvature_blocks.NaiveDiagonal +NaiveFull = curvature_blocks.NaiveFull +NaiveTNT = curvature_blocks.NaiveTNT +DenseDiagonal = curvature_blocks.DenseDiagonal +DenseFull = curvature_blocks.DenseFull +DenseTwoKroneckerFactored = curvature_blocks.DenseTwoKroneckerFactored +RepeatedDenseKroneckerFactored = curvature_blocks.RepeatedDenseKroneckerFactored +DenseTNT = curvature_blocks.DenseTNT +Conv2DDiagonal = curvature_blocks.Conv2DDiagonal +Conv2DFull = curvature_blocks.Conv2DFull +Conv2DTwoKroneckerFactored = curvature_blocks.Conv2DTwoKroneckerFactored +Conv2DTNT = curvature_blocks.Conv2DTNT +ScaleAndShiftDiagonal = curvature_blocks.ScaleAndShiftDiagonal +ScaleAndShiftFull = curvature_blocks.ScaleAndShiftFull +set_max_parallel_elements = curvature_blocks.set_max_parallel_elements +get_max_parallel_elements = curvature_blocks.get_max_parallel_elements +set_default_eigen_decomposition_threshold = ( + curvature_blocks.set_default_eigen_decomposition_threshold) +get_default_eigen_decomposition_threshold = ( + curvature_blocks.get_default_eigen_decomposition_threshold) + +# Curvature estimators +CurvatureEstimator = curvature_estimator.CurvatureEstimator +BlockDiagonalCurvature = curvature_estimator.BlockDiagonalCurvature +ExplicitExactCurvature = curvature_estimator.ExplicitExactCurvature +ImplicitExactCurvature = curvature_estimator.ImplicitExactCurvature +set_default_tag_to_block_ctor = ( + curvature_estimator.set_default_tag_to_block_ctor) +get_default_tag_to_block_ctor = ( + curvature_estimator.get_default_tag_to_block_ctor) +OptaxPreconditioner = curvature_estimator.OptaxPreconditioner +OptaxPreconditionState = curvature_estimator.OptaxPreconditionState + +# Optimizers +Optimizer = optimizer.Optimizer + +HAIKU_BIASES = optimizer.HAIKU_BIASES +HAIKU_BIASES_AND_NORMS = optimizer.HAIKU_BIASES_AND_NORMS + +__all__ = ( + # Modules + "utils", + "patches_second_moment", + "layers_and_loss_tags", + "loss_functions", + "tag_graph_matcher", + "tracer", + "curvature_blocks", + "curvature_estimator", + "optimizer", + # Patches second moments + "patches_moments", + "patches_moments_explicit", + # Layer and loss tags + "LossTag", + "LayerTag", + "register_generic", + "register_dense", + "register_conv2d", + "register_scale_and_shift", + # Tag graph matcher + "auto_register_tags", + # Tracer + "ProcessedJaxpr", + "loss_tags_vjp", + "loss_tags_jvp", + "loss_tags_hvp", + "layer_tags_vjp", + # Loss functions + "LossFunction", + "NegativeLogProbLoss", + "DistributionNegativeLogProbLoss", + "NormalMeanNegativeLogProbLoss", + "NormalMeanVarianceNegativeLogProbLoss", + "MultiBernoulliNegativeLogProbLoss", + "CategoricalLogitsNegativeLogProbLoss", + "OneHotCategoricalLogitsNegativeLogProbLoss", + "register_sigmoid_cross_entropy_loss", + "register_multi_bernoulli_predictive_distribution", + "register_softmax_cross_entropy_loss", + "register_categorical_predictive_distribution", + "register_squared_error_loss", + "register_normal_predictive_distribution", + # Curvature blocks + "CurvatureBlock", + "ScaledIdentity", + "Diagonal", + "Full", + "KroneckerFactored", + "NaiveDiagonal", + "NaiveFull", + "NaiveTNT", + "DenseDiagonal", + "DenseFull", + "DenseTwoKroneckerFactored", + "RepeatedDenseKroneckerFactored", + "DenseTNT", + "Conv2DDiagonal", + "Conv2DFull", + "Conv2DTwoKroneckerFactored", + "Conv2DTNT", + "ScaleAndShiftDiagonal", + "ScaleAndShiftFull", + "set_max_parallel_elements", + "get_max_parallel_elements", + "set_default_eigen_decomposition_threshold", + "get_default_eigen_decomposition_threshold", + # Estimators + "CurvatureEstimator", + "BlockDiagonalCurvature", + "ExplicitExactCurvature", + "ImplicitExactCurvature", + "set_default_tag_to_block_ctor", + "get_default_tag_to_block_ctor", + # Optimizers + "Optimizer", +) + +# _________________________________________ +# / Please don't use symbols in `_src` they \ +# \ are not part of the KFAC Jax public API./ +# ----------------------------------------- +# \ ^__^ +# \ (oo)\_______ +# (__)\ )\/\ +# ||----w | +# || || +# +try: + del _src # pylint: disable=undefined-variable +except NameError: + pass diff --git a/src/kfac_jax/_src/__init__.py b/src/kfac_jax/_src/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..352a587601d262f97c92a6176c89720fe3217654 --- /dev/null +++ b/src/kfac_jax/_src/__init__.py @@ -0,0 +1,13 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. diff --git a/src/kfac_jax/_src/curvature_blocks/__init__.py b/src/kfac_jax/_src/curvature_blocks/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..94fcf871efcd27cf218862f42f0db95ada49be33 --- /dev/null +++ b/src/kfac_jax/_src/curvature_blocks/__init__.py @@ -0,0 +1,51 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC curvature approximation to single layer blocks.""" +from kfac_jax._src.curvature_blocks import curvature_block +from kfac_jax._src.curvature_blocks import diagonal +from kfac_jax._src.curvature_blocks import full +from kfac_jax._src.curvature_blocks import kronecker_factored +from kfac_jax._src.curvature_blocks import tnt +from kfac_jax._src.curvature_blocks import utils + +CurvatureBlock = curvature_block.CurvatureBlock +ScaledIdentity = curvature_block.ScaledIdentity +ScalarOrSequence = curvature_block.ScalarOrSequence + +Diagonal = diagonal.Diagonal +Full = full.Full +KroneckerFactored = kronecker_factored.KroneckerFactored +NaiveDiagonal = diagonal.NaiveDiagonal +NaiveFull = full.NaiveFull +NaiveTNT = tnt.NaiveTNT +DenseDiagonal = diagonal.DenseDiagonal +DenseFull = full.DenseFull +DenseTwoKroneckerFactored = kronecker_factored.DenseTwoKroneckerFactored +RepeatedDenseKroneckerFactored = ( + kronecker_factored.RepeatedDenseKroneckerFactored) +DenseTNT = tnt.DenseTNT +Conv2DDiagonal = diagonal.Conv2DDiagonal +Conv2DFull = full.Conv2DFull +Conv2DTwoKroneckerFactored = kronecker_factored.Conv2DTwoKroneckerFactored +Conv2DTNT = tnt.Conv2DTNT +ScaleAndShiftDiagonal = diagonal.ScaleAndShiftDiagonal +ScaleAndShiftFull = full.ScaleAndShiftFull + +set_max_parallel_elements = utils.set_max_parallel_elements +get_max_parallel_elements = utils.get_max_parallel_elements +set_default_eigen_decomposition_threshold = ( + utils.set_default_eigen_decomposition_threshold) +get_default_eigen_decomposition_threshold = ( + utils.get_default_eigen_decomposition_threshold) +to_real_set = utils.to_real_set diff --git a/src/kfac_jax/_src/curvature_blocks/curvature_block.py b/src/kfac_jax/_src/curvature_blocks/curvature_block.py new file mode 100644 index 0000000000000000000000000000000000000000..5e962aba758bcbb6665c0b1c5233331276be0ce0 --- /dev/null +++ b/src/kfac_jax/_src/curvature_blocks/curvature_block.py @@ -0,0 +1,621 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the abstract base class for curvature blocks.""" + +import abc +from typing import Any, Sequence + +import jax +import jax.extend as jex +import jax.numpy as jnp +import jax.scipy +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import tag_graph_matcher as tgm +from kfac_jax._src import tracer +from kfac_jax._src import utils +from kfac_jax._src.curvature_blocks import utils as cb_utils +import numpy as np + + +# Types for annotation +Array = utils.Array +Scalar = utils.Scalar +Numeric = utils.Numeric +PRNGKey = utils.PRNGKey +Shape = utils.Shape +DType = utils.DType +ScalarOrSequence = Scalar | Sequence[Scalar] + + +class CurvatureBlock(utils.Finalizable): + """Abstract class for curvature approximation blocks. + + A CurvatureBlock defines a curvature matrix to be estimated, and gives methods + to multiply powers of this with a vector. Powers can be computed exactly or + with a class-determined approximation. Cached versions of the powers can be + pre-computed to make repeated multiplications cheaper. During initialization, + you would have to explicitly specify all powers that you will need to cache. + """ + + @utils.register_state_class + class State(utils.State): + """Persistent state of the block. + + Any subclasses of :class:`~CurvatureBlock` should also internally extend + this class, with any attributes needed for the curvature estimation. + + Attributes: + cache: A dictionary, containing any state data that is updated on + irregular intervals, such as inverses, eigenvalues, etc. Elements of + this are updated via calls to :func:`~CurvatureBlock.update_cache`, and + do not necessarily correspond to the most up-to-date curvature estimate. + """ + cache: dict[str, Array | dict[str, Array]] | None + + def __init__(self, layer_tag_eq: tags.LayerTagEqn): + """Initializes the block. + + Args: + layer_tag_eq: The Jax equation corresponding to the layer tag that this + block will approximate the curvature to. + """ + super().__init__() + + self._layer_tag_eq = layer_tag_eq + + self.finalize() + + @property + def name(self) -> str: + return tags.layer_eqn_name(self._layer_tag_eq) + + @property + def layer_tag_primitive(self) -> tags.LayerTag: + """The :class:`jex.core.Primitive` corresponding to the block's tag equation.""" + + primitive = self._layer_tag_eq.primitive + assert isinstance(primitive, tgm.tags.LayerTag) + + return primitive + + @property + def parameter_variables(self) -> tuple[jex.core.Var, ...]: + """The parameter variables of the underlying Jax equation.""" + + param_vars = [] + + for p in tags.layer_eqn_data(self._layer_tag_eq).params: + + assert isinstance(p, jex.core.Var) + param_vars.append(p) + + return tuple(param_vars) + + @property + def outputs_shapes(self) -> tuple[Shape, ...]: + """The shapes of the output variables of the block's tag equation.""" + + output_vars = tags.layer_eqn_data(self._layer_tag_eq).outputs + + return jax.tree.map(lambda x: x.aval.shape, output_vars) + + @property + def inputs_shapes(self) -> tuple[Shape, ...]: + """The shapes of the input variables of the block's tag equation.""" + + input_vars = tags.layer_eqn_data(self._layer_tag_eq).inputs + + return jax.tree.map(lambda x: x.aval.shape, input_vars) + + @property + def parameters_shapes(self) -> tuple[Shape, ...]: + """The shapes of the parameter variables of the block's tag equation.""" + return tuple(jax.tree.map( + lambda x: tuple(x.aval.shape), self.parameter_variables)) + + @property + def dtype(self) -> DType: + dtypes = set(p.aval.dtype for p in self.parameter_variables) # pytype: disable=attribute-error + if len(dtypes) > 1: + raise ValueError("Not all parameters are the same dtype.") + return dtypes.pop() + + @property + def parameters_canonical_order(self) -> tuple[int, ...]: + """The canonical order of the parameter variables.""" + + meta = self._layer_tag_eq.params.get("meta") + assert meta is not None and isinstance(meta, tags.LayerMetaData) + if meta.params_canonical_order is not None: + return meta.params_canonical_order + + # Before JAX 0.11, Var.count encoded the global Jaxpr variable order. + # Retain this fallback for blocks constructed outside ProcessedJaxpr on + # older JAX versions. ProcessedJaxpr records the semantic parameter order + # above, because JAX 0.11 removed Var.count entirely. + counts = [getattr(p, "count", None) for p in self.parameter_variables] + if all(count is not None for count in counts): + return tuple(np.argsort(counts)) + raise ValueError( + "Canonical parameter order is unavailable. Construct curvature blocks " + "from a ProcessedJaxpr so global parameter indices are recorded." + ) + + @property + def layer_tag_extra_params(self) -> dict[str, Any]: + """Any extra parameters of passed into the Jax primitive of this block.""" + + return self._layer_tag_eq.params + + @property + def number_of_parameters(self) -> int: + """Number of parameter variables of this block.""" + + return len(self.parameters_shapes) + + @property + def dim(self) -> int: + """The number of elements of all parameter variables together.""" + + return sum(utils.product(shape) for shape in self.parameters_shapes) + + def scale(self, state: State, use_cache: bool) -> Numeric: + """A scalar pre-factor of the curvature approximation. + + Importantly, all methods assume that whenever a user requests cached values, + any state dependant scale is taken into account by the cache (e.g. either + stored explicitly and used or mathematically added to values). + + Args: + state: The state for this block. + use_cache: Whether the method requesting this is using cached values or + not. + + Returns: + A scalar value to be multiplied with any unscaled block representation. + """ + + # TODO(jamesmartens,botev): This way of handling state dependent scale is + # a bit hacky and leads to complexity in other parts of the code that must + # be aware of how this part works. Should try to replace this with something + # better. + + if use_cache: + return self.fixed_scale() + + return self.fixed_scale() * self.state_dependent_scale(state) + + def fixed_scale(self) -> Numeric: + """A fixed scalar pre-factor of the curvature (e.g. constant).""" + return 1.0 + + def state_dependent_scale(self, state: State) -> Numeric: + """A scalar pre-factor of the curvature, computed from the most fresh curvature estimate.""" + del state # Unused + return 1.0 + + def __str__(self): + return (f"{self.__class__.__name__}, tag name: {self.name}, " + f"params shapes: {self.parameters_shapes!r}") + + @utils.auto_scope_method + def init( + self, + rng: PRNGKey, + exact_powers_to_cache: ScalarOrSequence | None, + approx_powers_to_cache: ScalarOrSequence | None, + cache_eigenvalues: bool, + ) -> State: + """Initializes the state for this block. + + Args: + rng: The PRNGKey which to be used for any randomness of the initialization + exact_powers_to_cache: A single value, or multiple values in a list, which + specify which exact matrix powers the block should be caching. Matrix + powers, which are expected to be used in + :func:`~CurvatureBlock.multiply_matpower`, + :func:`~CurvatureBlock.multiply_inverse` or + :func:`~CurvatureBlock.multiply` with ``exact_power=True`` and + ``use_cached=True`` must be provided here. + approx_powers_to_cache: A single value, or multiple values in a list, + which specify approximate matrix powers the block should be caching. + Matrix powers, which are expected to be used in + :func:`~CurvatureBlock.multiply_matrix_power`, + :func:`~CurvatureBlock.multiply_inverse` or + :func:`~CurvatureBlock.multiply` with ``exact_power=False`` and + ``use_cached=True`` must be provided here. + cache_eigenvalues: Specifies whether the block should be caching the + eigenvalues of its approximate curvature. + Returns: + A dictionary with the initialized state. + """ + return self._init( + rng=rng, + exact_powers_to_cache=cb_utils.to_real_set(exact_powers_to_cache), + approx_powers_to_cache=cb_utils.to_real_set(approx_powers_to_cache), + cache_eigenvalues=cache_eigenvalues) + + @abc.abstractmethod + def _init( + self, + rng: PRNGKey, + exact_powers_to_cache: set[Scalar], + approx_powers_to_cache: set[Scalar], + cache_eigenvalues: bool, + ) -> State: + """The non-public interface of ``init``.""" + + @abc.abstractmethod + def sync( + self, + state: State, + pmap_axis_name: str, + ) -> State: + """Syncs the state across different devices (does not sync the cache).""" + + @utils.auto_scope_method + def multiply_matpower( + self, + state: State, + vector: Sequence[Array], + identity_weight: Numeric, + power: Scalar, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + """Computes ``(BlockMatrix + identity_weight I)**power`` times ``vector``. + + Args: + state: The state for this block. + vector: A tuple of arrays that should have the same shapes as the block's + parameters_shapes, which represent the vector you want to multiply. + identity_weight: A scalar specifying the weight on the identity matrix + that is added to the block matrix before raising it to a power. If + ``use_cached=False`` it is guaranteed that this argument will be used in + the computation. When returning cached values, this argument *may* be + ignored in favor whatever value was last passed to + :func:`~CurvatureBlock.update_cache`. The precise semantics of this + depend on the concrete subclass and its particular behavior in regard to + caching. + power: The power to which to raise the matrix. + exact_power: Specifies whether to compute the exact matrix power of + ``BlockMatrix + identity_weight I``. When this argument is ``False`` + the exact behaviour will depend on the concrete subclass and the + result will *in general* be an approximation to + ``(BlockMatrix + identity_weight I)^power``, although some subclasses + may still compute the exact matrix power. + use_cached: Whether to use a cached version for computing the product or + to use the most recent curvature estimates. The cached version is + going to be *at least* as fresh as the value provided to the last call + to :func:`~CurvatureBlock.update_cache` with the same value of ``power`` + + Returns: + A tuple of arrays, representing the result of the matrix-vector product. + """ + + scale = self.scale(state, use_cached) + + result = self._multiply_matpower_unscaled( + state=state, + vector=vector, + identity_weight=identity_weight / scale, + power=power, + exact_power=exact_power, + use_cached=use_cached, + ) + + return utils.scalar_mul(result, jnp.power(scale, power)) + + @abc.abstractmethod + def _multiply_matpower_unscaled( + self, + state: State, + vector: Sequence[Array], + identity_weight: Numeric, + power: Scalar, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + """Performs matrix-vector multiplication, ignoring ``self.scale``.""" + + def multiply( + self, + state: State, + vector: Sequence[Array], + identity_weight: Numeric, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + """Computes ``(BlockMatrix + identity_weight I)`` times ``vector``.""" + + return self.multiply_matpower( + state=state, + vector=vector, + identity_weight=identity_weight, + power=1, + exact_power=exact_power, + use_cached=use_cached, + ) + + def multiply_inverse( + self, + state: State, + vector: Sequence[Array], + identity_weight: Numeric, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + """Computes ``(BlockMatrix + identity_weight I)^-1`` times ``vector``.""" + + return self.multiply_matpower( + state=state, + vector=vector, + identity_weight=identity_weight, + power=-1, + exact_power=exact_power, + use_cached=use_cached, + ) + + @utils.auto_scope_method + def eigenvalues( + self, + state: State, + use_cached: bool, + ) -> Array: + """Computes the eigenvalues for this block approximation. + + Args: + state: The state dict for this block. + use_cached: Whether to use a cached versions of the eigenvalues or to use + the most recent curvature estimates to compute them. The cached version + are going to be *at least* as fresh as the last time you called + :func:`~CurvatureBlock.update_cache` with ``eigenvalues=True``. + + Returns: + An array containing the eigenvalues of the block. + """ + eigenvalues = self._eigenvalues_unscaled(state, use_cached) + + assert eigenvalues.size == self.dim + + return self.scale(state, use_cached) * eigenvalues + + @abc.abstractmethod + def _eigenvalues_unscaled( + self, + state: State, + use_cached: bool, + ) -> Array: + """Computes the eigenvalues for this block, ignoring `self.scale`.""" + + @abc.abstractmethod + def update_curvature_matrix_estimate( + self, + state: State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> State: + """Updates the block's curvature estimates using the ``info`` provided. + + Each block *in general* estimates a moving average of its associated + curvature matrix. If you don't want a moving average you can set + ``ema_old=0`` and ``ema_new=1``. + + Args: + state: The state dict for this block to update. + estimation_data: A map containing data used for updating the curvature + matrix estimate for this block. This can be computed by calling the + function returned from :func:`~layer_tags_vjp`. Please see its + implementation for more details on the name of the fields and how they + are constructed. + ema_old: Specifies the weight of the old value when computing the updated + estimate in the moving average. + ema_new: Specifies the weight of the new value when computing the updated + estimate in the moving average. + identity_weight: The weight of the identity added to the block's curvature + matrix before computing the cached matrix power. + batch_size: The batch size used in computing the values in ``info``. + """ + + @utils.auto_scope_method + def update_cache( + self, + state: State, + identity_weight: Numeric, + exact_powers: ScalarOrSequence | None, + approx_powers: ScalarOrSequence | None, + eigenvalues: bool, + ) -> State: + """Updates the cached estimates of the different powers specified. + + Args: + state: The state dict for this block to update. + identity_weight: The weight of the identity added to the block's curvature + matrix before computing the cached matrix power. + exact_powers: Specifies any cached exact matrix powers to be updated. + approx_powers: Specifies any cached approximate matrix powers to be + updated. + eigenvalues: Specifies whether to update the cached eigenvalues + of the block. If they have not been cached before, this will create an + entry with them in the block's cache. + + Returns: + The updated state. + """ + return self._update_cache( + state=state, + identity_weight=identity_weight / self.scale(state, False), + exact_powers=cb_utils.to_real_set(exact_powers), + approx_powers=cb_utils.to_real_set(approx_powers), + eigenvalues=eigenvalues, + ) + + @abc.abstractmethod + def _update_cache( + self, + state: State, + identity_weight: Numeric, + exact_powers: set[Scalar], + approx_powers: set[Scalar], + eigenvalues: bool, + ) -> State: + """The cache updating function, ignoring ``self.scale``.""" + + @utils.auto_scope_method + def to_dense_matrix(self, state: State) -> Array: + """Returns a dense representation of the curvature matrix.""" + return self.scale(state, False) * self._to_dense_unscaled(state) + + @abc.abstractmethod + def _to_dense_unscaled(self, state: State) -> Array: + """A dense representation of the curvature, ignoring ``self.scale``.""" + + def undamped_diagonal(self, state: State) -> tuple[Array, ...]: + """Returns the diagonal of the undamped curvature.""" + return utils.scalar_mul(self._undamped_diagonal_unscaled(state), + self.scale(state, False)) + + def _undamped_diagonal_unscaled(self, state: State) -> tuple[Array, ...]: + """Returns the diagonal of the undamped curvature, ignoring ``self.scale``.""" + raise NotImplementedError() + + def norm(self, state: State, norm_type: str) -> Numeric: + """Computes the norm of the curvature block, according to ``norm_type``.""" + + return self.scale(state, False) * self._norm_unscaled(state, norm_type) + + @abc.abstractmethod + def _norm_unscaled( + self, + state: State, + norm_type: str + ) -> Numeric: + """Like ``norm`` but with ``self.scale`` not included.""" + + +class ScaledIdentity(CurvatureBlock): + """A block that assumes that the curvature is a scaled identity matrix.""" + + def __init__( + self, + layer_tag_eq: tags.LayerTagEqn, + scale: Numeric = 1.0, + ): + """Initializes the block. + + Args: + layer_tag_eq: The Jax equation corresponding to the layer tag, that this + block will approximate the curvature to. + scale: The scale of the identity matrix. + """ + self._scale = scale + super().__init__(layer_tag_eq) + + def fixed_scale(self) -> Numeric: + return self._scale + + def _init( + self, + rng: PRNGKey, + exact_powers_to_cache: set[Scalar], + approx_powers_to_cache: set[Scalar], + cache_eigenvalues: bool, + ) -> CurvatureBlock.State: + + del rng, exact_powers_to_cache, approx_powers_to_cache # Unused + + return CurvatureBlock.State( + cache=None, + ) + + def sync( + self, + state: CurvatureBlock.State, + pmap_axis_name: str, + ) -> CurvatureBlock.State: + return state + + def _multiply_matpower_unscaled( + self, + state: CurvatureBlock.State, + vector: Sequence[Array], + identity_weight: Numeric, + power: Scalar, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + + del exact_power # Unused + + # state_dependent_scale needs to be included because it won't be by the + # caller of this function (multiply_matpower) when use_cached=True + scale = self.state_dependent_scale(state) if use_cached else 1.0 + + identity_weight = identity_weight + scale + + if power == 1: + return jax.tree.map(lambda x: identity_weight * x, vector) + + elif power == -1: + return jax.tree.map(lambda x: x / identity_weight, vector) + + else: + identity_weight = jnp.power(identity_weight, power) + return jax.tree.map(lambda x: identity_weight * x, vector) + + def _eigenvalues_unscaled( + self, + state: CurvatureBlock.State, + use_cached: bool, + ) -> Array: + return jnp.ones([self.dim]) + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: CurvatureBlock.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> CurvatureBlock.State: + + return state.copy() + + def _update_cache( + self, + state: CurvatureBlock.State, + identity_weight: Numeric, + exact_powers: set[Scalar], + approx_powers: set[Scalar], + eigenvalues: bool, + ) -> CurvatureBlock.State: + + return state.copy() + + def _to_dense_unscaled(self, state: CurvatureBlock.State) -> Array: + del state # not used + return jnp.eye(self.dim) + + def _norm_unscaled( + self, + state: CurvatureBlock.State, + norm_type: str + ) -> Numeric: + + return utils.psd_matrix_norm(jnp.ones([self.dim]), norm_type=norm_type) diff --git a/src/kfac_jax/_src/curvature_blocks/diagonal.py b/src/kfac_jax/_src/curvature_blocks/diagonal.py new file mode 100644 index 0000000000000000000000000000000000000000..58ae74ef500fe0db06d158ee357b587f6896186a --- /dev/null +++ b/src/kfac_jax/_src/curvature_blocks/diagonal.py @@ -0,0 +1,376 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the diagonal curvature blocks.""" +import abc +import functools +from typing import Sequence + +import jax +import jax.numpy as jnp +import jax.scipy +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import tracer +from kfac_jax._src import utils +from kfac_jax._src.curvature_blocks import curvature_block +from kfac_jax._src.curvature_blocks import utils as cb_utils + +# Types for annotation +Array = utils.Array +Scalar = utils.Scalar +Numeric = utils.Numeric +PRNGKey = utils.PRNGKey +CurvatureBlock = curvature_block.CurvatureBlock + + +class Diagonal(CurvatureBlock, abc.ABC): + """An abstract class for approximating only the diagonal of curvature.""" + + @utils.register_state_class + class State(CurvatureBlock.State): + """Persistent state of the block. + + Attributes: + diagonal_factors: A tuple of the moving averages of the estimated + diagonals of the curvature for each parameter that is part of the + associated layer. + """ + diagonal_factors: tuple[utils.WeightedMovingAverage, ...] + + def _init( + self, + rng: PRNGKey, + exact_powers_to_cache: set[Scalar], + approx_powers_to_cache: set[Scalar], + cache_eigenvalues: bool, + ) -> State: + + del rng + + return Diagonal.State( + cache=None, + diagonal_factors=tuple( + utils.WeightedMovingAverage.zeros_array(shape, self.dtype) + for shape in self.parameters_shapes + ), + ) + + def sync( + self, + state: State, + pmap_axis_name: str, + ) -> State: + + # Copy this first since we mutate it later in this function. + state = state.copy() + + for factor in state.diagonal_factors: + factor.sync(pmap_axis_name) + + return state + + def _multiply_matpower_unscaled( + self, + state: State, + vector: Sequence[Array], + identity_weight: Numeric, + power: Scalar, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + + # state_dependent_scale needs to be included because it won't be by the + # caller of this function (multiply_matpower) when use_cached=True + scale = self.state_dependent_scale(state) if use_cached else 1.0 + + factors = tuple(scale * f.value + identity_weight + for f in state.diagonal_factors) + + assert len(factors) == len(vector) + + if power == 1: + return tuple(f * v for f, v in zip(factors, vector)) + elif power == -1: + return tuple(v / f for f, v in zip(factors, vector)) + else: + return tuple(jnp.power(f, power) * v for f, v in zip(factors, vector)) + + def _eigenvalues_unscaled( + self, + state: State, + use_cached: bool, + ) -> Array: + return jnp.concatenate([f.value.flatten() for f in state.diagonal_factors], + axis=0) + + def _update_cache( + self, + state: State, + identity_weight: Numeric, + exact_powers: set[Scalar], + approx_powers: set[Scalar], + eigenvalues: bool, + ) -> State: + + return state.copy() + + def _to_dense_unscaled(self, state: State) -> Array: + + # Extract factors in canonical order + factors = [state.diagonal_factors[i].value.flatten() + for i in self.parameters_canonical_order] + + # Construct diagonal matrix + return jnp.diag(jnp.concatenate(factors, axis=0)) + + def _norm_unscaled( + self, + state: CurvatureBlock.State, + norm_type: str + ) -> Numeric: + + return utils.product( + utils.psd_matrix_norm(f.value.flatten(), norm_type=norm_type) + for f in state.diagonal_factors) + + +class NaiveDiagonal(Diagonal): + """Approximates the diagonal of the curvature with in the most obvious way. + + The update to the curvature estimate is computed by ``(sum_i g_i) ** 2 / N``. + where `g_i` is the gradient of each individual data point, and ``N`` is the + batch size. + """ + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Diagonal.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Diagonal.State: + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + for factor, dw in zip( + state.diagonal_factors, estimation_data.tangents.params + ): + factor.update(dw * dw / batch_size, ema_old, ema_new) + + return state + + +class DenseDiagonal(Diagonal): + """A `Diagonal` block specifically for dense layers.""" + + @property + def has_bias(self) -> bool: + """Whether the layer has a bias parameter.""" + return len(self.parameter_variables) == 2 + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Diagonal.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Diagonal.State: + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + diagonals = (jnp.matmul((x * x).T, dy * dy) / batch_size,) + if self.has_bias: + diagonals += (jnp.mean(dy * dy, axis=0),) + + assert len(diagonals) == self.number_of_parameters + + for diagonal_factor, diagonal in zip(state.diagonal_factors, diagonals): + diagonal_factor.update(diagonal, ema_old, ema_new) + + return state + + +class Conv2DDiagonal(Diagonal): + """A :class:`~Diagonal` block specifically for 2D convolution layers.""" + + def __init__( + self, + layer_tag_eq: tags.LayerTagEqn, + max_elements_for_vmap: int | None = None, + ): + """Initializes the block. + + Since there is no 'nice' formula for computing the average of the + tangents for a 2D convolution, what we do is that we have a function - + ``self.conv2d_tangent_squared`` - that computes for a single feature map the + square of the tangents for the kernel of the convolution. To average over + the batch we have two choices - vmap or loop over the batch sequentially + using scan. This utility function provides a trade-off by being able to + specify the maximum number of batch size that we can vmap over. This means + that the maximum memory usage will be ``max_batch_size_for_vmap`` times the + memory needed when calling ``self.conv2d_tangent_squared``. And the actual + ``vmap`` will be called ``ceil(total_batch_size / max_batch_size_for_vmap)`` + number of times in a loop to find the final average. + + Args: + layer_tag_eq: The Jax equation corresponding to the layer tag, that this + block will approximate the curvature to. + max_elements_for_vmap: The threshold used for determining how much + computation to the in parallel and how much in serial manner. If + ``None`` will use the value returned by + :func:`~get_max_parallel_elements`. + """ + self._averaged_kernel_squared_tangents = utils.loop_and_parallelize_average( + func=self.conv2d_tangent_squared, + max_parallel_size=max_elements_for_vmap or + cb_utils.get_max_parallel_elements(), + ) + super().__init__(layer_tag_eq) + + @property + def has_bias(self) -> bool: + return len(self.parameter_variables) == 2 + + def conv2d_tangent_squared( + self, + image_features_map: Array, + output_tangent: Array, + ) -> Array: + """Computes the elementwise square of a tangent for a single feature map.""" + + extra_params = {k: v for k, v in self.layer_tag_extra_params.items() + if k not in ("lhs_shape", "rhs_shape", "meta")} + + _, vjp = jax.vjp( + functools.partial( + jax.lax.conv_general_dilated, + **extra_params + ), + image_features_map[None], jnp.zeros(self.parameters_shapes[0]) + ) + + return jnp.square(vjp(output_tangent[None])[1]) + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Diagonal.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Diagonal.State: + + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + diagonals = (self._averaged_kernel_squared_tangents(x, dy),) + + if self.has_bias: + sum_axis = tuple(range(1, dy.ndim - len(self.parameters_shapes[1]))) + bias_dy = jnp.sum(dy, axis=sum_axis) + diagonals += (jnp.mean(bias_dy * bias_dy, axis=0),) + + assert len(diagonals) == self.number_of_parameters + + for diagonal_factor, diagonal in zip(state.diagonal_factors, diagonals): + diagonal_factor.update(diagonal, ema_old, ema_new) + + return state + + +class ScaleAndShiftDiagonal(Diagonal): + """A diagonal approximation specifically for a scale and shift layers.""" + + @property + def has_scale(self) -> bool: + """Whether this layer's equation has a scale.""" + return self._layer_tag_eq.params["has_scale"] + + @property + def has_shift(self) -> bool: + """Whether this layer's equation has a shift.""" + return self._layer_tag_eq.params["has_shift"] + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Diagonal.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Diagonal.State: + + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + if self.has_scale: + + assert state.diagonal_factors[0].shape == self.parameters_shapes[0] + + scale_shape = estimation_data.primals.params[0].shape + + d_scale = cb_utils.compatible_sum(x * dy, scale_shape, skip_axes=[0]) + + scale_diag_update = jnp.sum( + d_scale * d_scale, + axis=0, keepdims=d_scale.ndim == len(scale_shape) + ) / batch_size + + state.diagonal_factors[0].update(scale_diag_update, ema_old, ema_new) + + if self.has_shift: + + shift_shape = estimation_data.primals.params[-1].shape + d_shift = cb_utils.compatible_sum(dy, shift_shape, skip_axes=[0]) + + shift_diag_update = jnp.sum( + d_shift * d_shift, + axis=0, keepdims=d_shift.ndim == len(shift_shape) + ) / batch_size + + state.diagonal_factors[-1].update(shift_diag_update, ema_old, ema_new) + + return state diff --git a/src/kfac_jax/_src/curvature_blocks/full.py b/src/kfac_jax/_src/curvature_blocks/full.py new file mode 100644 index 0000000000000000000000000000000000000000..047e2d61db2f7c03e9a5d07fa0dd60de5291bc0e --- /dev/null +++ b/src/kfac_jax/_src/curvature_blocks/full.py @@ -0,0 +1,543 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the full matrix curvature blocks.""" +import abc +import functools +from typing import Sequence + +import jax +import jax.numpy as jnp +import jax.scipy +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import tracer +from kfac_jax._src import utils +from kfac_jax._src.curvature_blocks import curvature_block +from kfac_jax._src.curvature_blocks import utils as cb_utils + +# Types for annotation +Array = utils.Array +Scalar = utils.Scalar +Numeric = utils.Numeric +PRNGKey = utils.PRNGKey +CurvatureBlock = curvature_block.CurvatureBlock + + +class Full(CurvatureBlock, abc.ABC): + """An abstract class for approximating the block matrix with a full matrix.""" + + @utils.register_state_class + class State(CurvatureBlock.State): + """Persistent state of the block. + + Attributes: + matrix: A moving average of the estimated curvature matrix for all + parameters that are part of the associated layer. + """ + matrix: utils.WeightedMovingAverage + + def __init__( + self, + layer_tag_eq: tags.LayerTagEqn, + eigen_decomposition_threshold: int | None = None, + ): + """Initializes the block. + + Args: + layer_tag_eq: The Jax equation corresponding to the layer tag that this + block will approximate the curvature to. + eigen_decomposition_threshold: During calls to ``init`` and + ``update_cache`` if higher number of matrix powers than this threshold + are requested, instead of computing individual approximate powers, will + directly compute the eigen-decomposition instead (which provide access to + any matrix power). If this is ``None`` will use the value returned from + :func:`~get_default_eigen_decomposition_threshold()`. + """ + + if eigen_decomposition_threshold is None: + threshold = cb_utils.get_default_eigen_decomposition_threshold() + self._eigen_decomposition_threshold = threshold + + else: + self._eigen_decomposition_threshold = eigen_decomposition_threshold + + super().__init__(layer_tag_eq) + + def parameters_list_to_single_vector( + self, + parameters_shaped_list: Sequence[Array], + ) -> Array: + """Converts values corresponding to parameters of the block to vector.""" + + if len(parameters_shaped_list) != self.number_of_parameters: + + raise ValueError(f"Expected a list of {self.number_of_parameters} values," + f" but got {len(parameters_shaped_list)} instead.") + + for array, shape in zip(parameters_shaped_list, self.parameters_shapes): + + if array.shape != shape: + raise ValueError(f"Expected a value of shape {shape}, but got " + f"{array.shape} instead.") + + return jnp.concatenate([v.flatten() for v in parameters_shaped_list]) + + def single_vector_to_parameters_list( + self, + vector: Array, + ) -> tuple[Array, ...]: + """Reverses the transformation ``self.parameters_list_to_single_vector``.""" + + if vector.ndim != 1: + raise ValueError(f"Expecting a vector, got {vector.ndim}-tensor.") + + if vector.size != self.dim: + raise ValueError(f"Expected a vector of size {self.dim}, but got " + f"{vector.size} instead.") + + parameters_shaped_list = [] + index = 0 + + for shape in self.parameters_shapes: + + size = utils.product(shape) + parameters_shaped_list.append(vector[index: index + size].reshape(shape)) + index += size + + assert index == self.dim + + return tuple(parameters_shaped_list) + + def _init( + self, + rng: PRNGKey, + exact_powers_to_cache: set[Scalar], + approx_powers_to_cache: set[Scalar], + cache_eigenvalues: bool, + ) -> State: + + del rng + + # This block does not have any notion of "approximate" powers + exact_powers_to_cache = exact_powers_to_cache | approx_powers_to_cache + cache = {} + + if len(exact_powers_to_cache) > self._eigen_decomposition_threshold: + cache["eigenvalues"] = jnp.zeros([self.dim], self.dtype) + cache["eigen_vectors"] = jnp.zeros([self.dim, self.dim], self.dtype) + + elif cache_eigenvalues: + cache["eigenvalues"] = jnp.zeros([self.dim], self.dtype) + + if len(exact_powers_to_cache) <= self._eigen_decomposition_threshold: + for power in exact_powers_to_cache: + cache[str(power)] = jnp.zeros([self.dim, self.dim], self.dtype) + + return Full.State( + cache=cache, + matrix=utils.WeightedMovingAverage.zeros_array( + [self.dim, self.dim], self.dtype), + ) + + def sync( + self, + state: State, + pmap_axis_name: str, + ) -> State: + + # Copy this first since we mutate it later in this function. + state = state.copy() + + state.matrix.sync(pmap_axis_name) + + return state + + def _multiply_matpower_unscaled( + self, + state: State, + vector: Sequence[Array], + identity_weight: Numeric, + power: Scalar, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + + vector = self.parameters_list_to_single_vector(vector) + + if power == 1: + + result = jnp.matmul(state.matrix.value, vector) + + if use_cached: + # state_dependent_scale needs to be included here because it won't be by + # the caller of this function (multiply_matpower) when use_cached=True. + # This is not an issue for other powers because they bake in + # state_dependent_scale. + result *= self.state_dependent_scale(state) + + result += identity_weight * vector + + elif not use_cached: + + matrix = state.matrix.value + identity_weight * jnp.eye(self.dim) + + if power == -1: + result = utils.psd_solve(matrix, vector) + else: + if power == -0.5: + matrix = utils.inverse_sqrt_psd_matrices(matrix) + elif power == 0.5: + matrix = jnp.dot(matrix, utils.inverse_sqrt_psd_matrices(matrix)) + else: + raise ValueError(f"Unsupported power: {power}") + # TODO(jamesmartens,botev): investigate this for determinism on GPUs + # NOTE: this function only works for integer powers + result = jnp.matmul(matrix, vector) + + else: + + if str(power) in state.cache: + result = jnp.matmul(state.cache[str(power)], vector) + + else: + s = state.cache["eigenvalues"] + q = state.cache["eigen_vectors"] + + result = jnp.matmul(jnp.transpose(q), vector) + result = jnp.power(s + identity_weight, power) * result + result = jnp.matmul(q, result) + + return self.single_vector_to_parameters_list(result) + + def _eigenvalues_unscaled( + self, + state: State, + use_cached: bool, + ) -> Array: + + if not use_cached: + return utils.safe_psd_eigh(state.matrix.value)[0] + + else: + return state.cache["eigenvalues"] + + def _update_cache( + self, + state: State, + identity_weight: Numeric, + exact_powers: set[Scalar], + approx_powers: set[Scalar], + eigenvalues: bool, + ) -> State: + + # Copy this first since we mutate it later in this function. + state = state.copy() + + scale = self.state_dependent_scale(state) + + # This block does not have any notion of "approximate" powers + exact_powers = exact_powers | approx_powers + + if len(exact_powers) > self._eigen_decomposition_threshold: + + s, q = utils.safe_psd_eigh(state.matrix.value) + state.cache = dict(eigenvalues=scale * s, eigen_vectors=q) + + else: + + if eigenvalues: + state.cache["eigenvalues"] = scale * utils.safe_psd_eigh( + state.matrix.value)[0] + + for power in exact_powers: + + if power == -1: + state.cache[str(power)] = utils.psd_inv( + state.matrix.value + identity_weight * jnp.eye(self.dim)) / scale + else: + matrix = state.matrix.value + identity_weight * jnp.eye(self.dim) + state.cache[str(power)] = ( + (scale ** power) * jnp.linalg.matrix_power(matrix, power)) + + return state + + def _to_dense_unscaled(self, state: State) -> Array: + + # Permute the matrix according to the parameters canonical order + return utils.block_permuted( + state.matrix.value, + block_sizes=[utils.product(shape) for shape in self.parameters_shapes], + block_order=self.parameters_canonical_order + ) + + def _norm_unscaled( + self, + state: CurvatureBlock.State, + norm_type: str + ) -> Numeric: + + return utils.psd_matrix_norm(state.matrix.value, norm_type=norm_type) + + def _undamped_diagonal_unscaled(self, state: State) -> tuple[Array, ...]: + diag_vec = jnp.diag(state.matrix.value) + return self.single_vector_to_parameters_list(diag_vec) + + +class NaiveFull(Full): + """Approximates the full curvature with in the most obvious way. + + The update to the curvature estimate is computed by + ``(sum_i g_i) (sum_i g_i)^T / N``, where ``g_i`` is the gradient of each + individual data point, and ``N`` is the batch size. + """ + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Full.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Full.State: + del identity_weight + + # This method supports the case where the param tangents have an extra + # leading dimension that should be summed over (after the outer products). + # TODO(jamesmartens): add support for this to NaiveDiagonal + + # Copy this first since we mutate it later in this function. + state = state.copy() + + params_tangents = jax.tree_util.tree_leaves( + estimation_data.tangents.params) + + params_tangents_flattened = [] + + assert len(params_tangents) == self.number_of_parameters + + for p_shape, pt in zip(self.parameters_shapes, params_tangents): + + if p_shape: + assert ( + pt.shape[-len(p_shape) :] == p_shape + ), f"{pt.shape=} and {p_shape=}" + + p_size = utils.product(p_shape) + + params_tangents_flattened.append(pt.reshape([-1, p_size])) + + tangents = jnp.concatenate(params_tangents_flattened, axis=1) + + if jnp.iscomplexobj(tangents): + stats = ( + jnp.einsum("ay,az->yz", tangents.real, tangents.real) + - jnp.einsum("ay,az->yz", tangents.imag, tangents.imag)) / batch_size + else: + stats = jnp.einsum("ay,az->yz", tangents, tangents) / batch_size + + state.matrix.update(stats, ema_old, ema_new) + + return state + + +class DenseFull(Full): + """A `Full` block specifically for dense layers.""" + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Full.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Full.State: + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + params_tangents = x[:, :, None] * dy[:, None, :] + + if self.number_of_parameters == 2: + params_tangents = jnp.concatenate([params_tangents, dy[:, None]], axis=1) + + params_tangents = jnp.reshape(params_tangents, [batch_size, -1]) + + matrix_update = jnp.matmul(params_tangents.T, params_tangents) / batch_size + state.matrix.update(matrix_update, ema_old, ema_new) + + return state + + +class Conv2DFull(Full): + """A :class:`~Full` block specifically for 2D convolution layers.""" + + def __init__( + self, + layer_tag_eq: tags.LayerTagEqn, + max_elements_for_vmap: int | None = None, + ): + """Initializes the block. + + Since there is no 'nice' formula for computing the average of the + tangents for a 2D convolution, what we do is that we have a function - + ``self.conv2d_tangent_squared`` - that computes for a single feature map the + square of the tangents for the kernel of the convolution. To average over + the batch we have two choices - vmap or loop over the batch sequentially + using scan. This utility function provides a trade-off by being able to + specify the maximum batch that that will be handled in a single iteration + of the loop. This means that the maximum memory usage will be + ``max_batch_size_for_vmap`` times the memory needed when calling + ``self.conv2d_tangent_squared``. And the actual ``vmap`` will be + called ``ceil(total_batch_size / max_batch_size_for_vmap)`` number of times + in a loop to find the final average. + + Args: + layer_tag_eq: The Jax equation corresponding to the layer tag, that this + block will approximate the curvature to. + max_elements_for_vmap: The threshold used for determining how much + computation to the in parallel and how much in serial manner. If + ``None`` will use the value returned by + :func:`~get_max_parallel_elements`. + """ + + self._averaged_tangents_outer_product = utils.loop_and_parallelize_average( + func=self.conv2d_tangent_outer_product, + max_parallel_size=max_elements_for_vmap or + cb_utils.get_max_parallel_elements(), + ) + + super().__init__(layer_tag_eq) + + def conv2d_tangent_outer_product( + self, + inputs: Array, + tangent_of_outputs: Array, + ) -> Array: + """Computes the outer product of a tangent for a single feature map.""" + + extra_params = {k: v for k, v in self.layer_tag_extra_params.items() + if k not in ("lhs_shape", "rhs_shape", "meta")} + + _, vjp = jax.vjp( + functools.partial( + jax.lax.conv_general_dilated, + **extra_params + ), + inputs[None], jnp.zeros(self.parameters_shapes[0]) + ) + + tangents = (vjp(tangent_of_outputs[None])[1],) + + if self.number_of_parameters == 2: + num_axis = tangent_of_outputs.ndim - len(self.parameters_shapes[1]) + sum_axis = tuple(range(num_axis)) + tangents += (jnp.sum(tangent_of_outputs, axis=sum_axis),) + + flat_tangents = self.parameters_list_to_single_vector(tangents) + + return jnp.outer(flat_tangents, flat_tangents) + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Full.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Full.State: + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + assert utils.first_dim_is_size(batch_size, x, dy) + + matrix_update = self._averaged_tangents_outer_product(x, dy) + state.matrix.update(matrix_update, ema_old, ema_new) + + return state + + +class ScaleAndShiftFull(Full): + """A full dense approximation specifically for a scale and shift layers.""" + + @property + def _has_scale(self) -> bool: + """Whether this layer's equation has a scale.""" + return self._layer_tag_eq.params["has_scale"] + + @property + def _has_shift(self) -> bool: + """Whether this layer's equation has a shift.""" + return self._layer_tag_eq.params["has_shift"] + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: Full.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> Full.State: + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + assert utils.first_dim_is_size(batch_size, x, dy) + + tangents = [] + + if self._has_scale: + # Scale tangent + scale_shape = estimation_data.primals.params[0].shape + + d_scale = cb_utils.compatible_sum(x * dy, scale_shape, skip_axes=[0]) + d_scale = d_scale.reshape([batch_size, -1]) + + tangents.append(d_scale) + + if self._has_shift: + # Shift tangent + + shift_shape = estimation_data.primals.params[-1].shape + + d_shift = cb_utils.compatible_sum(dy, shift_shape, skip_axes=[0]) + d_shift = d_shift.reshape([batch_size, -1]) + + tangents.append(d_shift) + + tangents = jnp.concatenate(tangents, axis=1) + matrix_update = jnp.matmul(tangents.T, tangents) / batch_size + + state.matrix.update(matrix_update, ema_old, ema_new) + + return state diff --git a/src/kfac_jax/_src/curvature_blocks/kronecker_factored.py b/src/kfac_jax/_src/curvature_blocks/kronecker_factored.py new file mode 100644 index 0000000000000000000000000000000000000000..61d39e066ca9d375d211c5b4aa0ebfaebb9108e8 --- /dev/null +++ b/src/kfac_jax/_src/curvature_blocks/kronecker_factored.py @@ -0,0 +1,721 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the Kronecker factored curvature blocks.""" +import abc +import math +from typing import Any, Sequence + +import jax +import jax.numpy as jnp +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import patches_second_moment as psm +from kfac_jax._src import tracer +from kfac_jax._src import utils +from kfac_jax._src.curvature_blocks import curvature_block +from kfac_jax._src.curvature_blocks import utils as cb_utils +from typing_extensions import Self + + +# Types for annotation +Array = utils.Array +Scalar = utils.Scalar +Numeric = utils.Numeric +PRNGKey = utils.PRNGKey +Shape = utils.Shape +CurvatureBlock = curvature_block.CurvatureBlock + + +class KroneckerFactored(CurvatureBlock, abc.ABC): + """An abstract class for approximating the block with a Kronecker product. + + The constructor takes two special arguments: + - parameters_specs: A list, where each element specifies for each + parameter a "rearrange string". This is in the format `abc->b(ca)` + similar to `einops.rearrange`. + - parameters_concat_axis: The axis along which the parameters will be + concatenated to form a single array after each parameter has been + rearranged according to its "rearrange string". + + The above implies that: + - All parameters must have the same rank after they have been rearranged. + - All parameters must have the same size along all axes except the + concatenation axis after they have been rearranged. + + By default, each parameter is rearanged to a matrix, by merging all dimensions + except the last one. If a parameter is a vector (rank 1), it is rearranged to + a matrix with the first dimension being 1. Then concatenation is done along + axis=0. + """ + + @utils.register_state_class + class State(CurvatureBlock.State): + """Persistent state of the block. + + Attributes: + factors: A tuple of the moving averages of the estimated factors of the + curvature for each axis group. + """ + + factors: tuple[utils.WeightedMovingAverage, ...] + + @classmethod + def from_dict(cls, dict_rep: dict[str, Any]) -> Self: + class_name = dict_rep.pop("__class__", cls.__name__) + assert class_name == cls.__name__ + return cls( + factors=tuple( + utils.WeightedMovingAverage.from_dict(rep) + for rep in dict_rep["factor"] + ) + ) + + def __init__( + self, + layer_tag_eq: tags.LayerTagEqn, + parameters_specs: Sequence[str] | None = None, + parameters_concat_axis: int = 0, + ): + + # Even though the superclass constructor will set this later, we need to do + # it now since it's used below. + self._layer_tag_eq = layer_tag_eq + + if parameters_specs is None: + parameters_specs = [] + + for shape in self.parameters_shapes: + + if len(shape) == 1: + parameters_specs.append("a -> 1a") + + else: + in_str = cb_utils.ALPHABET[:len(shape)] + out_str = f"({in_str[:-1]}){in_str[-1]}" + parameters_specs.append(f"{in_str} -> {out_str}") + + else: + assert len(parameters_specs) == self.number_of_parameters + + self.parameters_specs = parameters_specs + self.parameters_concat_axis = parameters_concat_axis + + super().__init__(layer_tag_eq) + + def __str__(self): + return ( + f"{self.__class__.__name__}(parameter_specs={self.parameters_specs}, " + f"parameters_concat_axis={self.parameters_concat_axis}), " + f"tag name: {self.name}, params shapes: {self.parameters_shapes!r}" + ) + + def parameters_shaped_list_to_array( + self, + parameters_shaped_list: Sequence[Array], + ) -> Array: + """Combines all parameters to a single array.""" + values = [] + for p, spec in zip( + parameters_shaped_list, + self.parameters_specs, + strict=True, + ): + values.append(utils.rearrange(p, spec)) + + return jnp.concatenate(values, axis=self.parameters_concat_axis) + + def array_to_parameters_shaped_list(self, array: Array) -> tuple[Array, ...]: + """An inverse transformation of ``self.parameters_shaped_list_to_array``.""" + parameters_list = [] + n = 0 + index = [slice(None)] * array.ndim + + for shape, spec in zip( + self.parameters_shapes, + self.parameters_specs, + strict=True, + ): + zero = utils.rearrange(jnp.zeros(shape), spec) + d = zero.shape[self.parameters_concat_axis] + index[self.parameters_concat_axis] = slice(n, n + d) + p = array[tuple(index)] + parameters_list.append(p.reshape(shape)) + n += d + + return tuple(parameters_list) + + @property + def array_shape(self) -> Shape: + """The shape of the single non axis grouped array.""" + avals = [jnp.zeros(shape) for shape in self.parameters_shapes] + return self.parameters_shaped_list_to_array(avals).shape + + @property + def array_ndim(self) -> int: + """The number of dimensions of the single non axis grouped array.""" + return len(self.array_shape) + + def _init( + self, + rng: PRNGKey, + exact_powers_to_cache: set[Scalar], + approx_powers_to_cache: set[Scalar], + cache_eigenvalues: bool, + ) -> State: + + cache = {} + factors = [] + + for i, d in enumerate(self.array_shape): + + factors.append( + utils.WeightedMovingAverage.zeros_array((d, d), self.dtype) + ) + + if cache_eigenvalues or exact_powers_to_cache: + cache[f"{i}_factor_eigenvalues"] = jnp.zeros((d,), dtype=self.dtype) + + if exact_powers_to_cache: + cache[f"{i}_factor_eigen_vectors"] = jnp.zeros((d, d), dtype=self.dtype) + + for power in approx_powers_to_cache: + + if power != -1: + raise NotImplementedError( + f"Approximations for power {power} is not yet implemented." + ) + + if str(power) not in cache: + cache[str(power)] = {} + + cache[str(power)][f"{i}_factor"] = jnp.zeros((d, d), dtype=self.dtype) + + return KroneckerFactored.State( + cache=cache, + factors=tuple(factors), + ) + + def sync( + self, + state: State, + pmap_axis_name: str, + ) -> State: + + # Copy this first since we mutate it later in this function. + state = state.copy() + + for factor in state.factors: + factor.sync(pmap_axis_name) + + return state + + def _multiply_matpower_unscaled( + self, + state: State, + vector: Sequence[Array], + identity_weight: Numeric, + power: Scalar, + exact_power: bool, + use_cached: bool, + ) -> tuple[Array, ...]: + + assert len(state.factors) == self.array_ndim + + vector = self.parameters_shaped_list_to_array(vector) + + if power == 1: + + factors = [f.value for f in state.factors] + + # state_dependent_scale needs to be included here because it won't be by + # the caller of this function (multiply_matpower) when use_cached=True. + # This is not an issue for other powers because they bake in + # state_dependent_scale. + scale = self.state_dependent_scale(state) if use_cached else 1.0 + + if exact_power: + result = scale * utils.kronecker_product_axis_mul_v(factors, vector) + result = result + identity_weight * vector + + else: + # If compute pi_adjusted_kronecker_factors used a more expensive matrix + # norm in its computation, it might make sense to cache it. But we + # currently don't do that. + + result = scale * utils.kronecker_product_axis_mul_v( + utils.pi_adjusted_kronecker_factors( + *factors, damping=identity_weight / scale), + vector) + + elif exact_power: + + if use_cached: + s = [ + state.cache[f"{i}_factor_eigenvalues"] + for i in range(len(state.factors)) + ] + q = [ + state.cache[f"{i}_factor_eigen_vectors"] + for i in range(len(state.factors)) + ] + + else: + s, q = zip( + *[utils.safe_psd_eigh(factor.value) for factor in state.factors] + ) + + eigenvalues = utils.outer_product(*s) + identity_weight + eigenvalues = jnp.power(eigenvalues, power) + + result = utils.kronecker_eigen_basis_axis_mul_v(q, eigenvalues, vector) + + else: + + if power not in [-1, -0.5, 0.5]: + raise NotImplementedError( + f"Approximations for power {power} is not yet implemented." + ) + + if use_cached: + + assert power != -0.5 + + factors = [ + state.cache[str(power)][f"{i}_factor"] + for i in range(len(state.factors)) + ] + + else: + + factors = [factor.value for factor in state.factors] + + factors = utils.pi_adjusted_kronecker_factors( + *factors, damping=identity_weight) + + if power == -1: + factors = utils.invert_psd_matrices(factors) + elif power == -0.5: + factors = utils.inverse_sqrt_psd_matrices(factors) + # TODO(timothycnguyen): Hacky psd square root. Will find a better way. + elif power == 0.5: + inverse_sqrt_factors = utils.inverse_sqrt_psd_matrices(factors) + + def matmul(x, y): + if x.ndim == y.ndim == 2: + return jnp.dot(x, y) + assert x.ndim == y.ndim == 1 + return x * y + + factors = jax.tree_util.tree_map( + matmul, factors, inverse_sqrt_factors + ) + else: + raise NotImplementedError() + + result = utils.kronecker_product_axis_mul_v(factors, vector) + + return self.array_to_parameters_shaped_list(result) + + def _eigenvalues_unscaled( + self, + state: State, + use_cached: bool, + ) -> Array: + + assert len(state.factors) == self.array_ndim + + if use_cached: + s = [ + state.cache[f"{i}_factor_eigenvalues"] + for i in range(len(state.factors)) + ] + else: + s_q = [utils.safe_psd_eigh(factor.value) for factor in state.factors] + s, _ = zip(*s_q) + + return utils.outer_product(*s) + + def _update_cache( + self, + state: State, + identity_weight: Numeric, + exact_powers: set[Scalar], + approx_powers: set[Scalar], + eigenvalues: bool, + ) -> State: + + assert len(state.factors) == self.array_ndim + + # Copy this first since we mutate it later in this function. + state = state.copy() + + scale = self.state_dependent_scale(state) + factor_scale = jnp.power(scale, 1.0 / self.array_ndim) + + if eigenvalues or exact_powers: + + s_q = [utils.safe_psd_eigh(factor.value) for factor in state.factors] + + s, q = zip(*s_q) + + for i in range(len(state.factors)): + state.cache[f"{i}_factor_eigenvalues"] = factor_scale * s[i] + + if exact_powers: + state.cache[f"{i}_factor_eigen_vectors"] = q[i] + + for power in approx_powers: + + if power != -1: + raise NotImplementedError( + f"Approximations for power {power} is not yet implemented." + ) + + cache = state.cache[str(power)] + + # This computes the approximate inverse factors using the generalization + # of the pi-adjusted inversion from the original KFAC paper. + inv_factors = utils.pi_adjusted_kronecker_inverse( + *[factor.value for factor in state.factors], + damping=identity_weight, + ) + + for i in range(len(state.factors)): + cache[f"{i}_factor"] = inv_factors[i] / factor_scale + + return state + + def _norm_unscaled( + self, + state: CurvatureBlock.State, + norm_type: str + ) -> Numeric: + + return utils.product( + utils.psd_matrix_norm(f.value, norm_type=norm_type) + for f in state.factors) + + def _to_dense_unscaled(self, state: "KroneckerFactored.State") -> Array: + + # We currently support this only for 2 parameters + assert 0 < self.number_of_parameters <= 2 + inputs_factor = state.factors[0].value + + if (self.number_of_parameters == 2 and + self.parameters_canonical_order[0] != 0): + + # Permute the matrix according to the parameters canonical order + inputs_factor = utils.block_permuted( + state.factors[0].value, + block_sizes=[state.factors[0].shape[0] - 1, 1], + block_order=(1, 0), + ) + + return jnp.kron(inputs_factor, state.factors[1].value) + + +class DenseTwoKroneckerFactored(KroneckerFactored): + """A :class:`~TwoKroneckerFactored` block specifically for dense layers.""" + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: KroneckerFactored.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> KroneckerFactored.State: + del identity_weight + assert 1 <= self.number_of_parameters <= 2 + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + if self.number_of_parameters == 2: + x_one = jnp.ones_like(x[:, :1]) + x = jnp.concatenate([x, x_one], axis=1) + + input_stats = jnp.einsum("ay,az->yz", x, x) / batch_size + output_stats = jnp.einsum("ay,az->yz", dy, dy) / batch_size + + state.factors[0].update(input_stats, ema_old, ema_new) + state.factors[1].update(output_stats, ema_old, ema_new) + + return state + + +class RepeatedDenseKroneckerFactored(DenseTwoKroneckerFactored): + """Block for dense layers applied to tensors with extra time/loc dims.""" + + @utils.register_state_class + class State(KroneckerFactored.State): + """Persistent state of the block. + + Attributes: + average_repeats: A decayed average of the per-case number of non-masked + repeats in the data used to compute the block's statistics. We use the + same decayed averaging for this quantity that we do for the statistics, + so that they "match". + """ + + average_repeats: utils.WeightedMovingAverage + + def __init__( + self, + layer_tag_eq: tags.LayerTagEqn, + use_masking: bool = True, + parameters_specs: Sequence[str] | None = None, + parameters_concat_axis: int = 0, + ): + self._use_masking = use_masking + super().__init__( + layer_tag_eq=layer_tag_eq, + parameters_specs=parameters_specs, + parameters_concat_axis=parameters_concat_axis, + ) + + def _init( + self, + rng: PRNGKey, + exact_powers_to_cache: set[Scalar], + approx_powers_to_cache: set[Scalar], + cache_eigenvalues: bool, + ) -> "RepeatedDenseKroneckerFactored.State": + + super_state = super()._init( + rng, exact_powers_to_cache, approx_powers_to_cache, cache_eigenvalues + ) + + return RepeatedDenseKroneckerFactored.State( + average_repeats=utils.WeightedMovingAverage.zeros_array((), self.dtype), + **super_state.__dict__, + ) + + def state_dependent_scale( + self, + state: "RepeatedDenseKroneckerFactored.State", + ) -> Numeric: + return 1.0 / state.average_repeats.value + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: KroneckerFactored.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> KroneckerFactored.State: + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + if self._use_masking: + + # hack: we identify masked repeats by checking if all corresponding + # entries of dy are zero + mask = 1.0 - jnp.all(dy == 0.0, axis=-1, keepdims=True) + + # zero out corresponding elts of x + x = x * mask + + # compute total non-masked + total = jnp.sum(mask) + + else: + total = math.prod(dy.shape[:-1]) + + x = x.reshape([-1, x.shape[-1]]) + dy = dy.reshape([-1, dy.shape[-1]]) + + if self.number_of_parameters == 2: + x_one = jnp.ones_like(x[:, :1]) + x = jnp.concatenate([x, x_one], axis=1) + + input_stats = jnp.einsum("ay,az->yz", x, x) / batch_size + output_stats = jnp.einsum("ay,az->yz", dy, dy) / batch_size + + state.factors[0].update(input_stats, ema_old, ema_new) + state.factors[1].update(output_stats, ema_old, ema_new) + state.average_repeats.update(total / batch_size, ema_old, ema_new) + + return state + + +class Conv2DTwoKroneckerFactored(KroneckerFactored): + """A :class:`~TwoKroneckerFactored` block specifically for 2D convolution layers.""" + + def fixed_scale(self) -> Numeric: + return float(self.num_locations) + + @property + def kernel_output_axis(self) -> int: + return self._layer_tag_eq.params["dimension_numbers"].rhs_spec[0] + + @property + def outputs_channel_index(self) -> int: + """The ``channels`` index in the outputs of the layer.""" + return self._layer_tag_eq.params["dimension_numbers"].out_spec[1] + + @property + def inputs_channel_index(self) -> int: + """The ``channels`` index in the inputs of the layer.""" + return self._layer_tag_eq.params["dimension_numbers"].lhs_spec[1] + + @property + def weights_output_channel_index(self) -> int: + """The ``channels`` index in weights of the layer.""" + return self._layer_tag_eq.params["dimension_numbers"].rhs_spec[0] + + @property + def weights_spatial_shape(self) -> Shape: + spatial_index = self._layer_tag_eq.params["dimension_numbers"].rhs_spec[2:] + return tuple(self.parameters_shapes[0][i] for i in spatial_index) + + @property + def weights_spatial_size(self) -> int: + """The spatial filter size of the weights.""" + return utils.product(dim for dim in self.weights_spatial_shape) + + @property + def inputs_spatial_shape(self) -> Shape: + spatial_index = self._layer_tag_eq.params["dimension_numbers"].lhs_spec[2:] + return tuple(self.inputs_shapes[0][i] for i in spatial_index) + + @property + def num_locations(self) -> int: + """The number of spatial locations that each filter is applied to.""" + return psm.num_conv_locations( + self.inputs_spatial_shape, + self.weights_spatial_shape, + self._layer_tag_eq.params["window_strides"], + self._layer_tag_eq.params["padding"]) + + def input_size(self) -> int: + if self.has_bias: + return self.num_inputs_channels * self.weights_spatial_size + 1 + else: + return self.num_inputs_channels * self.weights_spatial_size + + def output_size(self) -> int: + return self.num_outputs_channels + + @property + def num_inputs_channels(self) -> int: + """The number of channels in the inputs to the layer.""" + return self._layer_tag_eq.invars[0].aval.shape[ # pytype: disable=attribute-error + self.inputs_channel_index] + + @property + def num_outputs_channels(self) -> int: + """The number of channels in the outputs to the layer.""" + return self._layer_tag_eq.invars[1].aval.shape[ # pytype: disable=attribute-error + self.weights_output_channel_index] + + def compute_inputs_stats( + self, + inputs: Array, + weighting_array: Array | None = None, + ) -> Array: + """Computes the statistics for the inputs factor.""" + batch_size = inputs.shape[0] + + input_cov_m, input_cov_v = psm.patches_moments( + inputs, + kernel_spatial_shape=self.weights_spatial_shape, + strides=self._layer_tag_eq.params["window_strides"], + padding=self._layer_tag_eq.params["padding"], + data_format=None, + dim_numbers=self._layer_tag_eq.params["dimension_numbers"], + precision=self._layer_tag_eq.params.get("precision"), + weighting_array=weighting_array, + ) + + # Flatten the kernel and channels dimensions + k, h, c = input_cov_v.shape + input_cov_v = jnp.reshape(input_cov_v, (k * h * c,)) + input_cov_m = jnp.reshape(input_cov_m, (k * h * c, k * h * c)) + + # Normalize by the `batch size` * `num_locations` + normalizer = batch_size * self.num_locations + input_cov_m = input_cov_m / normalizer + input_cov_v = input_cov_v / normalizer + + if self.number_of_parameters == 1: + return input_cov_m + + if weighting_array is None: + corner = jnp.ones([1], dtype=input_cov_m.dtype) + else: + corner = jnp.mean(weighting_array).reshape([1]) + + input_cov = jnp.concatenate([input_cov_m, input_cov_v[None]], axis=0) + input_cov_v = jnp.concatenate([input_cov_v, corner], axis=0) + + return jnp.concatenate([input_cov, input_cov_v[:, None]], axis=1) + + def compute_outputs_stats(self, tangent_of_output: Array) -> Array: + """Computes the statistics for the outputs factor.""" + lhs_str = utils.replace_char( + cb_utils.ALPHABET[:4], "y", self.outputs_channel_index) + rhs_str = utils.replace_char( + cb_utils.ALPHABET[:4], "z", self.outputs_channel_index) + ein_str = f"{lhs_str},{rhs_str}->yz" + stats = jnp.einsum(ein_str, tangent_of_output, tangent_of_output) + + # Normalize by the `batch size` * `num_locations` + normalizer = tangent_of_output.shape[0] * self.num_locations + return stats / normalizer + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: KroneckerFactored.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> KroneckerFactored.State: + del identity_weight + assert 1 <= self.number_of_parameters <= 2 + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + input_stats = self.compute_inputs_stats(x) + output_stats = self.compute_outputs_stats(dy) + + state.factors[0].update(input_stats, ema_old, ema_new) + state.factors[1].update(output_stats, ema_old, ema_new) + + return state diff --git a/src/kfac_jax/_src/curvature_blocks/tnt.py b/src/kfac_jax/_src/curvature_blocks/tnt.py new file mode 100644 index 0000000000000000000000000000000000000000..dcdc006ce1e2a99be723f7ff50553061bade4eab --- /dev/null +++ b/src/kfac_jax/_src/curvature_blocks/tnt.py @@ -0,0 +1,250 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing TNT curvature blocks.""" +from typing import Sequence + +import jax +import jax.numpy as jnp +import jax.scipy +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import tracer +from kfac_jax._src import utils +from kfac_jax._src.curvature_blocks import kronecker_factored +from kfac_jax._src.curvature_blocks import utils as cb_utils + + +# Types for annotation +Array = utils.Array +Numeric = utils.Numeric +KroneckerFactored = kronecker_factored.KroneckerFactored + + +class NaiveTNT(KroneckerFactored): + """A standard TNT block for a single parameter, or weights + bias. + + Each factor of the standard TNT curvature approximation estimates the expected + value of the contraction of the gradients with themselves along all but a + single axis `i`: + ``F_i ~~ E[contract_all_but_one(g, g, i)].`` + where `contrat_all_but_one` is defined as the contraction over all axes except + the i-th of its first two inputs, e.g.: + ``contract_all_but_one(A, B, 1)[a,b] = sum_{i,j} A[i, a, j] B[i, b, j]`` + + The estimation is performed in a naive way by contracting the sum of each + examples' gradients and then dividing by the batch size: + ``F_i = contract_all_but_one(sum_n g_n, sum_n g_n, i) / N`` + where `g_n` is the model gradient of a single example and `N` is the + batch size. Since the expectations of the gradients over the model + distribution is zero and they are independent across cases, this is still an + unbiased estimator. + """ + + def state_dependent_scale( + self, + state: "NaiveTNT.State", + ) -> Numeric: + return utils.tnt_scale([factor.value for factor in state.factors]) + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: KroneckerFactored.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> KroneckerFactored.State: + + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + dw = self.parameters_shaped_list_to_array(estimation_data.tangents.params) + + assert dw.ndim == len(state.factors) + + in_str = cb_utils.ALPHABET[: dw.ndim] + + for i, factor in enumerate(state.factors): + # For factor i we contract the gradient with itself along all axes, + # except the i-th. + + lhs_str = utils.replace_char(in_str, "y", i) + rhs_str = utils.replace_char(in_str, "z", i) + + # This is a rank-1 mod since it's like we flattened all but dim i together + # and then did an outer product + factor_update = ( + jnp.einsum(f"{lhs_str},{rhs_str}->yz", dw, dw) / batch_size + ) + + factor.update(factor_update, ema_old, ema_new) + + return state + + +class DenseTNT(kronecker_factored.DenseTwoKroneckerFactored): + """A TNT block for dense layers. + + This TNT block modifies :class:`~NaiveTNTBlock` by the way it estimates each + factor specifically for a dense layer. Instead of using the contraction over + the summed gradient, it performs the contraction over each individual batch + elements and then averages over the batch: + ``F_i = sum_n contract_all_but_one(g_n, g_n, i) / N`` + The estimator is unbiased, and will have lower variance then the naive one. + """ + + def state_dependent_scale(self, state: "DenseTNT.State") -> Numeric: + return utils.tnt_scale([factor.value for factor in state.factors]) + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: KroneckerFactored.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> KroneckerFactored.State: + + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + if self.number_of_parameters == 2: + x_one = jnp.ones_like(x[:, :1]) + x = jnp.concatenate([x, x_one], axis=1) + + # We multiply each x by the norm_y, and each dy by the norm of x + dy_norms = jnp.linalg.norm(dy, axis=-1, keepdims=True) + x_norms = jnp.linalg.norm(x, axis=-1, keepdims=True) + x = x * dy_norms + dy = dy * x_norms + + input_stats = jnp.einsum("ay,az->yz", x, x) / batch_size + output_stats = jnp.einsum("ay,az->yz", dy, dy) / batch_size + + state.factors[0].update(input_stats, ema_old, ema_new) + state.factors[1].update(output_stats, ema_old, ema_new) + + return state + + +class Conv2DTNT(kronecker_factored.Conv2DTwoKroneckerFactored): + """A TNT block for Conv2D layers. + + This TNT block modifies :class:`~NaiveTNTBlock` by the way it estimates each + factor specifically for a conv2D layer. Importantly, it assumes "location + independence" similar to :class:~`Conv2DTwoKroneckerFactored`. Given this + assumption, instead of using the contraction over the summed gradient, it + performs the contraction for each individual example in the batch, and each + individual spatial location, and then averages over these: + ``F_i = sum_n sum_t contract_all_but_one(g_{n,t}, g_{n,t}, i) / (N * T)`` + where T here is the number of spatial locations. The estimator is unbiased + (under the "location independence" approximation), and will have lower + variance then the naive one. + + If the argument `weighting_per_location` is set to `False`, then the block + uses a mixture between location-independence and not, in the sense that it + computes the contractions per example, while the matrix factor statistics + still assume location independence. + """ + + def __init__( + self, + layer_tag_eq: tags.LayerTagEqn, + weighting_per_location: bool = True, + parameters_specs: Sequence[str] | None = None, + parameters_concat_axis: int = 0, + ): + self.weighting_per_location = weighting_per_location + super().__init__( + layer_tag_eq=layer_tag_eq, + parameters_specs=parameters_specs, + parameters_concat_axis=parameters_concat_axis, + ) + + def state_dependent_scale( + self, state: "Conv2DTNT.State" + ) -> Numeric: + return utils.tnt_scale([factor.value for factor in state.factors]) + + def x_squared_spatial_norms(self, x: Array) -> Array: + + kernel_shape = list(self.parameters_shapes[0]) + kernel_shape[self.kernel_output_axis] = 1 + + return jax.lax.conv_general_dilated( + lhs=x * x, + rhs=jnp.ones(kernel_shape), + window_strides=self.layer_tag_extra_params["window_strides"], + padding=self.layer_tag_extra_params["padding"], + lhs_dilation=self.layer_tag_extra_params["lhs_dilation"], + rhs_dilation=self.layer_tag_extra_params["rhs_dilation"], + dimension_numbers=self.layer_tag_extra_params["dimension_numbers"], + feature_group_count=self.layer_tag_extra_params["feature_group_count"], + precision=self.layer_tag_extra_params["precision"], + preferred_element_type= + self.layer_tag_extra_params["preferred_element_type"], + ) + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: KroneckerFactored.State, + estimation_data: tracer.LayerVjpData[Array], + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + ) -> KroneckerFactored.State: + + del identity_weight + + # Copy this first since we mutate it later in this function. + state = state.copy() + + [x] = estimation_data.primals.inputs + [dy] = estimation_data.tangents.outputs + + assert utils.first_dim_is_size(batch_size, x, dy) + + # We multiply each x by the norm_y, and each dy by the norm of x + dy_sq_norms = jnp.sum(dy * dy, axis=self.outputs_channel_index) + x_sq_norms = self.x_squared_spatial_norms(x) + + if self.number_of_parameters == 2: + # When we have a bias we need to add 1 coming from it to the squared norm + x_sq_norms = x_sq_norms + 1 + + if not self.weighting_per_location: + dy_sq_norms = jnp.sum(dy_sq_norms, axis=[1, 2]) + x_sq_norms = jnp.sum(x_sq_norms, axis=[1, 2, 3], keepdims=True) + + input_cov = self.compute_inputs_stats(x, weighting_array=dy_sq_norms) + output_cov = self.compute_outputs_stats(dy * jnp.sqrt(x_sq_norms)) + + state.factors[0].update(input_cov, ema_old, ema_new) + state.factors[1].update(output_cov, ema_old, ema_new) + + return state diff --git a/src/kfac_jax/_src/curvature_blocks/utils.py b/src/kfac_jax/_src/curvature_blocks/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..0dbad99bc989c3fcb9f440d3bc626a608e9d2bb9 --- /dev/null +++ b/src/kfac_jax/_src/curvature_blocks/utils.py @@ -0,0 +1,137 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing utility functions for curvature blocks.""" +import collections +import string +from typing import Sequence + +import jax.numpy as jnp +from kfac_jax._src import utils + + +# Types for annotation +Scalar = utils.Scalar +ScalarOrSequence = Scalar | Sequence[Scalar] + +# Special global variables +# This is used for einsum strings +ALPHABET = string.ascii_lowercase +# The default value that would be used for the argument +# ``max_elements_for_vmap``, when it is set to ``None`` in the +# ``Conv2DDiagonal`` and ``Conv2DFull` curvature blocks. +_MAX_PARALLEL_ELEMENTS: int = 2 ** 23 +# The default value that would be used for the argument +# ``eigen_decomposition_threshold``, when it is set to ``None`` in any of the +# curvature blocks that inherit from ``Full`. +_DEFAULT_EIGEN_DECOMPOSITION_THRESHOLD = 5 + + +def set_max_parallel_elements(value: int): + """Sets the default value of maximum parallel elements in the module. + + This value is used to determine the parallel-to-memory tradeoff in the + curvature estimation procedure of :class:`~Conv2DDiagonal` and + :class:`~Conv2DFull`. See their corresponding docs for further details. + + Args: + value: The default value for maximum number of parallel elements. + """ + global _MAX_PARALLEL_ELEMENTS + _MAX_PARALLEL_ELEMENTS = value + + +def get_max_parallel_elements() -> int: + """Returns the default value of maximum parallel elements in the module. + + This value is used to determine the parallel-to-memory tradeoff in the + curvature estimation procedure of :class:`~Conv2DDiagonal` and + :class:`~Conv2DFull`. See their corresponding docs for further details. + + Returns: + The default value for maximum number of parallel elements. + """ + return _MAX_PARALLEL_ELEMENTS + + +def set_default_eigen_decomposition_threshold(value: int): + """Sets the default value of the eigen decomposition threshold. + + This value is used in :class:`~Full` to determine when updating the cache, + at what number of different powers to switch the implementation from a simple + matrix power to an eigenvector decomposition. + + Args: + value: The default value for eigen decomposition threshold. + """ + global _DEFAULT_EIGEN_DECOMPOSITION_THRESHOLD + _DEFAULT_EIGEN_DECOMPOSITION_THRESHOLD = value + + +def get_default_eigen_decomposition_threshold() -> int: + """Returns the default value of the eigen decomposition threshold. + + This value is used in :class:`~Full` to determine when updating the cache, + at what number of different powers to switch the implementation from a simple + matrix power to an eigenvector decomposition. + + Returns: + The default value of the eigen decomposition threshold. + """ + return _DEFAULT_EIGEN_DECOMPOSITION_THRESHOLD + + +def to_real_set( + number_or_sequence: ScalarOrSequence | None +) -> set[Scalar]: + """Converts the optional number or sequence to a set.""" + if number_or_sequence is None: + return set() + elif isinstance(number_or_sequence, set): + return number_or_sequence + elif isinstance(number_or_sequence, (float, int)): + return {number_or_sequence} + elif (isinstance(number_or_sequence, collections.abc.Sequence) and + all(isinstance(x, (int, float)) for x in number_or_sequence)): + return set(number_or_sequence) + else: + raise ValueError(f"Expecting a real-number or a sequence of reals, but got " + f"{type(number_or_sequence)}.") + + +def compatible_shapes(ref_shape, target_shape): + + if len(target_shape) > len(ref_shape): + raise ValueError("Target shape should be smaller.") + + for ref_d, target_d in zip(reversed(ref_shape), reversed(target_shape)): + if ref_d != target_d and target_d != 1: + raise ValueError(f"{target_shape} is incompatible with {ref_shape}.") + + +def compatible_sum(tensor, target_shape, skip_axes): + """Compute sum over ``tensor`` to achieve shape given by ``target_shape``.""" + + compatible_shapes(tensor.shape, target_shape) + + n = tensor.ndim - len(target_shape) + + axis = [i + n for i, t in enumerate(target_shape) + if t == 1 and i + n not in skip_axes] + + tensor = jnp.sum(tensor, axis=axis, keepdims=True) + + axis = [i for i in range(tensor.ndim - len(target_shape)) + if i not in skip_axes] + + return jnp.sum(tensor, axis=axis) diff --git a/src/kfac_jax/_src/curvature_estimator/__init__.py b/src/kfac_jax/_src/curvature_estimator/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..8189108431d9a1b49ef6290ff860db5c03748b1f --- /dev/null +++ b/src/kfac_jax/_src/curvature_estimator/__init__.py @@ -0,0 +1,81 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC curvature explicit and implicit estimators. + +Curvature matrices are always defined in terms of some single differentiable +function of the parameters and inputs. In all cases in this module this quantity +is not the output from the model function (usually provided as argument to the +constructor of each curvature matrix), but is the sum of all losses +(weighted accordingly) which have been registered with a loss tag in the +computation graph of the model function. This quantity is referred to as the +``total_loss``. + +In this module there are three curvature matrices considered: + ``H`` - the Hessian matrix + ``F`` - the Fisher matrix + ``G`` - The Generalized Gauss-Newton(GGN) matrix +Vectors that are multiplied by a curvature matrix (or any of its matrix powers) +are always represented as a PyTree structure, equivalent to the parameters of +the model function. In all functions such vector is named +``parameter_structured_vector`` in the argument list. + +Factors of a matrix ``M`` are defined as matrices ``B`` such that ``BB^T = M``. +If we have to left-multiply ``B`` with a vector ``v``, than ``v`` has the same +format as if we have to multiply the whole curvature matrix ``M``. However the +second size of ``B`` is not clearly defined (and can be different for the +different curvature matrices). In all methods working with factors, e.g. if we +need to right multiply ``B`` with a vector ``v`` or the result of left +multiplying ``B`` by a parameter structured vector, then the provided vector +``v`` should be a list of lists of arrays. Each element of ``v`` corresponds to +a single loss registered in the model function, and its elements should have the +shapes as the corresponding ``loss.XXX_inner_shapes`` (XXX=Hessian, Fisher or +GGN). In all function such vector is named ``loss_vectors`` in the argument +list. + +See for example: www.cs.utoronto.ca/~jmartens/docs/HF_book_chapter.pdf and +https://arxiv.org/abs/1412.1193 for more information about the Hessian, Fisher +and GGN matrices and how to compute matrix-vector products. +""" + +from kfac_jax._src.curvature_estimator import block_diagonal +from kfac_jax._src.curvature_estimator import curvature_estimator +from kfac_jax._src.curvature_estimator import explicit_exact +from kfac_jax._src.curvature_estimator import implicit_exact +from kfac_jax._src.curvature_estimator import optax_interface + + +BlockDiagonalCurvature = block_diagonal.BlockDiagonalCurvature +set_default_tag_to_block_ctor = ( + block_diagonal.set_default_tag_to_block_ctor) +get_default_tag_to_block_ctor = ( + block_diagonal.get_default_tag_to_block_ctor) +set_multi_default_tag_to_block_ctor = ( + block_diagonal.set_multi_default_tag_to_block_ctor) + +StateType = curvature_estimator.StateType +CurvatureBlockCtor = curvature_estimator.CurvatureBlockCtor +CurvatureEstimator = curvature_estimator.CurvatureEstimator + +ExplicitExactCurvature = explicit_exact.ExplicitExactCurvature + +ImplicitExactCurvature = implicit_exact.ImplicitExactCurvature +LossFunction = implicit_exact.LossFunction +LossFunctionsTuple = implicit_exact.LossFunctionsTuple +LossFunctionsSequence = implicit_exact.LossFunctionsSequence +LossFunctionInputs = implicit_exact.LossFunctionInputs +LossFunctionInputsSequence = implicit_exact.LossFunctionInputsSequence +LossFunctionInputsTuple = implicit_exact.LossFunctionInputsTuple + +OptaxPreconditioner = optax_interface.OptaxPreconditioner +OptaxPreconditionState = optax_interface.OptaxPreconditionState diff --git a/src/kfac_jax/_src/curvature_estimator/block_diagonal.py b/src/kfac_jax/_src/curvature_estimator/block_diagonal.py new file mode 100644 index 0000000000000000000000000000000000000000..32affbb686fca5f8c2e08c7a88b8f214edef0611 --- /dev/null +++ b/src/kfac_jax/_src/curvature_estimator/block_diagonal.py @@ -0,0 +1,1126 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the BlockDiagonalCurvature class.""" + +import functools +from typing import Any, Callable, Sequence, Mapping +from absl import logging +import jax +from jax import scipy +import jax.numpy as jnp + +from kfac_jax._src import curvature_blocks +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import loss_functions +from kfac_jax._src import tracer +from kfac_jax._src import utils +from kfac_jax._src.curvature_estimator import curvature_estimator + +import numpy as np + + +# Types for annotation +Array = utils.Array +PRNGKey = utils.PRNGKey +Numeric = utils.Numeric +Scalar = utils.Scalar +Shape = utils.Shape +CurvatureBlockCtor = Callable[ + [tags.LayerTagEqn], + curvature_blocks.CurvatureBlock +] + +_DEFAULT_TAG_TO_BLOCK_CTOR: dict[str, CurvatureBlockCtor] = dict( + dense=curvature_blocks.DenseTwoKroneckerFactored, + conv2d=curvature_blocks.Conv2DTwoKroneckerFactored, + generic=curvature_blocks.NaiveDiagonal, + scale_and_shift=curvature_blocks.ScaleAndShiftDiagonal, + repeated_dense=curvature_blocks.RepeatedDenseKroneckerFactored, +) + + +def get_default_tag_to_block_ctor( + tag_name: str +) -> CurvatureBlockCtor | None: + """Returns the default curvature block constructor for the give tag name.""" + if tag_name.endswith("_tag"): + raise ValueError( + "You are using the old style of tag names. Remove the '_tag' suffix." + ) + return _DEFAULT_TAG_TO_BLOCK_CTOR.get(tag_name) + + +def set_default_tag_to_block_ctor( + tag_name: str, + block_ctor: CurvatureBlockCtor +) -> None: + """Sets the default curvature block constructor for the given tag.""" + if tag_name.endswith("_tag"): + raise ValueError( + "You are using the old style of tag names. Remove the '_tag' suffix." + ) + _DEFAULT_TAG_TO_BLOCK_CTOR[tag_name] = block_ctor + + +def set_multi_default_tag_to_block_ctor( + tags_to_block_ctor: Mapping[str, CurvatureBlockCtor] +): + _DEFAULT_TAG_TO_BLOCK_CTOR.update(tags_to_block_ctor) + + +class BlockDiagonalCurvature( + curvature_estimator.CurvatureEstimator["BlockDiagonalCurvature.State"]): + """Block diagonal curvature estimator class. + + NOTE: all of these options, except for fisher_empirical_direct_synced, will + perform their computations per-device (using the per-device mini-batch and + batch size), and then average the resulting 2nd-order statistics (used to + define the curvature approximation or preconditioner matrix) across the + devices. + + Supports for the following estimation modes: + + * fisher_gradients - the basic estimation approach from the original + K-FAC paper. + + * fisher_curvature_prop - method which estimates the Fisher using + self-products of random 1/-1 vectors times "half-factors" of the + Fisher, as described `here `__. + + * fisher_exact - is the obvious generalization of Curvature + Propagation to compute the exact Fisher (modulo any additional + diagonal or Kronecker approximations) by looping over one-hot vectors + for each coordinate of the output instead of using 1/-1 vectors. It is + more expensive to compute than the other three options by a factor + equal to the output dimension, roughly speaking. Because this assumes + independence of elements in the batch, it won't work properly for models + that violate this (e.g. ones that use batch normalization). This mode + also assumes that all registered losses have an input tensor whose first + dimension is the batch dimension (whereas the other modes don't care + about this). + + * fisher_empirical - computes the 'empirical' Fisher information + matrix (which uses the data's distribution for the targets, as + opposed to the true Fisher which uses the model's distribution) and + requires that each registered loss have specified targets. This mode + will introduce an additional normalization of the 2nd-order statistics + by the batch size (on top of the normalization by batch size done in the + curvature block classes). + + * fisher_empirical_direct - similar to fisher_empirical, but bypasses the + loss registration machinery, and just computes the gradients for the + statistics by reading the gradient from ``func_and_grad`` if provided + (or applying ``jax.grad`` to ``func`` instead). This will typically be + faster than fisher_empirical (as the gradient computation will be merged + by XLA with the one done for the optimizer). Currently, this mode only + supports curvature block approximations that use parameter gradient + information. + + * fisher_empirical_direct_synced - the same as fisher_empirical_direct, + but syncs the gradients across the devices before computing the + 2nd-order statistics from these. This option is useful to exactly + reproduce the preconditioner used in RMSProp/Adam, for example. + + * ggn_curvature_prop - Analogous to fisher_curvature_prop, but + estimates the Generalized Gauss-Newton matrix (GGN). + + * ggn_exact - Analogous to fisher_exact, but estimates the Generalized + Gauss-Newton matrix (GGN). + """ + + @utils.register_state_class + class State(utils.State): + """Persistent state of the estimator. + + Attributes: + synced: A Jax boolean, specifying if the state has been synced across + devices (this does not include the cache, which is never explicitly + synced). + blocks_states: A tuple of the state of the estimator corresponding to each + block. + """ + synced: Array + blocks_states: tuple[curvature_blocks.CurvatureBlock.State, ...] + + def __init__( + self, + func: utils.Func | None = None, + func_and_grad: utils.ValueAndGradFunc | None = None, + default_estimation_mode: str | None = None, + layer_tag_to_block_ctor: + Mapping[str, CurvatureBlockCtor] | None = None, + index_to_block_ctor: + Mapping[tuple[int, ...], CurvatureBlockCtor] | None = None, + auto_register_tags: bool = True, + distributed_multiplies: bool = True, + distributed_cache_updates: bool = True, + num_samples: int = 1, + should_vmap_samples: bool = False, + auto_register_kwargs: dict[str, Any] | None = None, + shared_forward_value_func: utils.Func | None = None, + **kwargs: Any, + ): + """Initializes the BlockDiagonalCurvature instance. + + Args: + func: The loss function, which should have at least one registered loss. + Should return only the loss value and not any auxiliary data. For the + purposes of the ``'fisher_empirical_direct[_synced]'`` estimation modes, + the loss value should be normalized by the batch size. (For other + estimation modes it won't matter either way.) Only one of ``func`` or + ``func_and_grad`` should be provided. + func_and_grad: A function returning the loss value and the gradient as a + tuple. For the purposes of the ``'fisher_empirical_direct[_synced]'`` + estimation modes, the loss value should be normalized by the batch size. + (For other estimation modes it won't matter either way.) Only one of + ``func`` or ``func_and_grad`` should be provided. + default_estimation_mode: The estimation mode which to use by default when + calling ``self.update_curvature_matrix_estimate``. If ``None`` this will + be ``'ggn_curvature_prop'``. + layer_tag_to_block_ctor: An optional dict mapping tags to specific classes + of block approximations, which to override the default ones. + index_to_block_ctor: An optional dict mapping a specific block parameter + indices to specific classes of block approximation, which to override + the default ones. To get the correct indices check + ``estimator.indices_to_block_map``. + auto_register_tags: Whether to automatically register layer tags for + parameters that have not been manually registered. For further details + see ``tag_graph_matcher.auto_register_tags``. + distributed_multiplies: Whether to distribute the curvature matrix + multiplication operations across the different devices in a block-wise + fashion. If False, each device will (redundantly) perform the operations + for all of the blocks. + distributed_cache_updates: Whether to distribute the cache + update multiplication operations across the different devices in a + block-wise fashion. If False, each device will (redundantly) perform + the operations for all of the blocks. + num_samples: Number of samples (per case) to use when computing stochastic + curvature matrix estimates. This option is only used when + ``estimation_mode == 'fisher_gradients'`` or ``estimation_mode == + '[fisher,ggn]_curvature_prop'``. + should_vmap_samples: Whether to use ``jax.vmap`` to compute samples + when ``num_samples > 1``. + shared_forward_value_func: Optional tagged function returning + ``(loss, gradient_surrogate)``. The surrogate's parameter gradient + must be the training gradient associated with ``loss``. When provided, + exact-Fisher curvature and the training gradient can share one model + primal evaluation. + auto_register_kwargs: Keyword arguments to pass to into the + layer auto-registration function. + **kwargs: Addiional keyword arguments passed to the superclass + ``CurvatureEstimator``. + """ + + super().__init__( + func=func, + func_and_grad=func_and_grad, + default_estimation_mode=default_estimation_mode or "ggn_curvature_prop", + **kwargs, + ) + + self._index_to_block_ctor = index_to_block_ctor or dict() + self._layer_tag_to_block_ctor = layer_tag_to_block_ctor or dict() + self._auto_register_tags = auto_register_tags + self._auto_register_kwargs = auto_register_kwargs or {} + self._vjp, self._jaxpr_extractor = tracer.layer_tags_vjp( + func=self.func, + params_index=self.params_index, + auto_register_tags=auto_register_tags, + **self._auto_register_kwargs + ) + if shared_forward_value_func is None: + self._vjp_and_value_and_grad = None + self._value_and_grad_jaxpr_extractor = None + else: + ( + self._vjp_and_value_and_grad, + self._value_and_grad_jaxpr_extractor, + ) = tracer.layer_tags_vjp_and_value_and_grad( + value_func=shared_forward_value_func, + params_index=self.params_index, + auto_register_tags=auto_register_tags, + **self._auto_register_kwargs, + ) + + # Initialized during finalization + self._jaxpr: tracer.ProcessedJaxpr | None = None + self._blocks: tuple[curvature_blocks.CurvatureBlock, ...] | None = None + + self._distributed_multiplies = distributed_multiplies + self._distributed_cache_updates = distributed_cache_updates + + self._num_samples = num_samples + self._should_vmap_samples = should_vmap_samples + + @property + def valid_estimation_modes(self) -> tuple[str, ...]: + """The valid estimation modes for this estimator.""" + return ("fisher_gradients", "fisher_empirical", "fisher_exact", + "fisher_curvature_prop", "ggn_exact", "ggn_curvature_prop", + "fisher_empirical_direct", "fisher_empirical_direct_synced") + + def _check_finalized(self): + if not self.finalized: + raise ValueError("The estimator has not been finalized. Call `init` or " + "`finalize` first.") + + def _create_blocks(self): + """Creates all the curvature blocks instances in ``self._blocks``.""" + + assert self._jaxpr is not None + + blocks_list = [] + + for tag_eqn, idx in zip(self._jaxpr.layer_tags, self._jaxpr.layer_indices): + meta = tag_eqn.params.get("meta") + assert meta is not None and isinstance(meta, tags.LayerMetaData) + assert not meta.nesting + + # Correctly get the block class + if idx in self._index_to_block_ctor: + cls = self._index_to_block_ctor[idx] + + elif meta.variant in self._layer_tag_to_block_ctor: + cls = self._layer_tag_to_block_ctor[meta.variant] + + else: + cls = get_default_tag_to_block_ctor(meta.variant) + if cls is None: + raise ValueError( + "Did not find anywhere a block class for layer tag variant " + f"{meta.variant}." + ) + + blocks_list.append(cls(tag_eqn)) + + self._blocks = tuple(blocks_list) + + @property + def blocks(self) -> tuple[curvature_blocks.CurvatureBlock, ...] | None: + """The tuple of :class:`~CurvatureBlock` instances used for each layer.""" + self._check_finalized() + return self._blocks + + @property + def num_blocks(self) -> int: + """The number of separate blocks that this estimator has.""" + return len(self.blocks) + + @property + def block_dims(self) -> Shape: + """The number of elements of all parameter variables for each block.""" + return tuple(block.dim for block in self.blocks) + + @property + def dim(self) -> int: + """The number of elements of all parameter variables together.""" + return sum(self.block_dims) + + @property + def jaxpr(self) -> tracer.ProcessedJaxpr: + self._check_finalized() + assert self._jaxpr is not None + return self._jaxpr + + @property + def params_structure_vector_of_indices(self) -> utils.Params: + """A tree structure with parameters replaced by their indices.""" + return jax.tree_util.tree_unflatten( + self.jaxpr.params_tree, range(len(self.jaxpr.params_vars_flat)) + ) + + @property + def indices_to_block_map( + self + ) -> Mapping[tuple[int, ...], curvature_blocks.CurvatureBlock]: + """A mapping of parameter indices to their associated blocks.""" + return dict(zip(self.jaxpr.layer_indices, self.blocks)) + + @property + def params_block_index(self) -> utils.Params: + """A structure, which shows each parameter to which block it corresponds. + + Returns: + A parameter-like structure, where each parameter is replaced by an integer + index. This index specifies the block (found by ``self.blocks[index]``) + which approximates the part of the curvature matrix associated with the + parameter. + """ + params_block_index: list[int | None] = [None] * self.num_params_variables + + for i, block_indices in enumerate(self.jaxpr.layer_indices): + for index in block_indices: + params_block_index[index] = i + + assert all(x is not None for x in params_block_index) + + return jax.tree_util.tree_unflatten( + self.jaxpr.params_tree, params_block_index) + + @property + def num_params_variables(self) -> int: + """The number of separate parameter variables of the model.""" + return len(self.jaxpr.params_vars_flat) + + @utils.auto_scope_method + def _compute_losses_vjp(self, func_args: utils.FuncArgs): + """Computes all model statistics needed for estimating the curvature.""" + return self._vjp(func_args) + + @property + def param_order(self): + block_order = tuple( + index + for block_indices in self.jaxpr.layer_indices + for index in block_indices + ) + return np.argsort(block_order) + + def log_registrations(self): + if self._blocks is None: + raise ValueError( + "You must initialize the estimator before calling this method." + ) + + logging.info("BlockDiagonalCurvature blocks:") + for block in self._blocks: + logging.info(str(block)) + logging.info("=" * 50) + + def params_vector_to_blocks_vectors( + self, + parameter_structured_vector: utils.Params, + ) -> tuple[tuple[Array, ...], ...]: + """Splits the parameters to values for each corresponding block.""" + + params_values_flat = jax.tree_util.tree_leaves(parameter_structured_vector) + blocks_vectors: list[tuple[Array, ...]] = [] + + for indices in self.jaxpr.layer_indices: + blocks_vectors.append(tuple(params_values_flat[i] for i in indices)) + + return tuple(blocks_vectors) + + def blocks_vectors_to_params_vector( + self, + blocks_vectors: Sequence[Sequence[Array]], + ) -> utils.Params: + """Reverses the effect of ``self.vectors_to_blocks``.""" + + if len(blocks_vectors) != self.num_blocks: + raise ValueError("Incorrect number of block vectors. Expected " + f"{self.num_blocks}, but got {len(blocks_vectors)}.") + + values_flat: list[Array | None] = [None] * self.num_params_variables + + for idx, (indices, vectors) in enumerate( + zip(self.jaxpr.layer_indices, blocks_vectors)): + + if len(indices) != len(vectors): + raise ValueError(f"Expected len(block_vectors[{idx}])=={len(indices)}, " + f"not {len(vectors)}.") + + for i, v in zip(indices, vectors): + assert values_flat[i] is None + values_flat[i] = v + + assert not any(v is None for v in values_flat) + + return jax.tree_util.tree_unflatten(self.jaxpr.params_tree, values_flat) + + def _finalize(self, func_args: utils.FuncArgs): + self._jaxpr = self._jaxpr_extractor(func_args) + if self._value_and_grad_jaxpr_extractor is not None: + shared_jaxpr = self._value_and_grad_jaxpr_extractor(func_args) + + def aval_signature(var): + return (tuple(var.aval.shape), var.aval.dtype) + + def layer_signature(tag): + meta = tag.params.get("meta") + data = tag.primitive.layer_data(tag.invars, tag.params) + return ( + getattr(meta, "variant", None), + getattr(meta, "inputs_index", None), + getattr(meta, "outputs_index", None), + getattr(meta, "params_index", None), + getattr(meta, "params_canonical_order", None), + tuple(aval_signature(var) for var in data.inputs), + tuple(aval_signature(var) for var in data.outputs), + tuple(aval_signature(var) for var in data.params), + ) + + def loss_signature(tag): + meta = tag.params.get("meta") + if not isinstance(meta, tags.LossMetaData): + return (type(meta),) + return ( + meta.loss_class, + meta.parameter_dependants, + meta.parameter_independants, + meta.argument_names, + tuple(aval_signature(var) for var in tag.invars), + ) + + shared_layers = tuple( + (tag.primitive, layer_signature(tag)) + for tag in shared_jaxpr.layer_tags + ) + reference_layers = tuple( + (tag.primitive, layer_signature(tag)) + for tag in self._jaxpr.layer_tags + ) + if ( + shared_jaxpr.layer_indices != self._jaxpr.layer_indices + or shared_layers != reference_layers + ): + raise ValueError( + "The shared-forward graph changed layer registrations." + ) + shared_losses = tuple( + loss_signature(tag) for tag in shared_jaxpr.loss_tags + ) + reference_losses = tuple( + loss_signature(tag) for tag in self._jaxpr.loss_tags + ) + if shared_losses != reference_losses: + raise ValueError( + "The shared-forward graph changed loss registrations." + ) + self._create_blocks() + self.log_registrations() + + @utils.auto_scope_method + def init( + self, + rng: PRNGKey, + func_args: utils.FuncArgs, + exact_powers_to_cache: curvature_blocks.ScalarOrSequence | None, + approx_powers_to_cache: curvature_blocks.ScalarOrSequence | None, + cache_eigenvalues: bool = False, + ) -> State: + + if not self.finalized: + self.finalize(func_args) + + blocks_init = [] + blocks_rng = jax.random.split(rng, self.num_blocks) + + for block, block_rng in zip(self.blocks, blocks_rng): + + block_init = block.init( + rng=block_rng, + exact_powers_to_cache=exact_powers_to_cache, + approx_powers_to_cache=approx_powers_to_cache, + cache_eigenvalues=cache_eigenvalues) + + blocks_init.append(block_init) + + return BlockDiagonalCurvature.State( + synced=jnp.asarray(True), + blocks_states=tuple(blocks_init), + ) + + def _sync_state( + self, + state: State, + pmap_axis_name: str | None, + ) -> State: + + block_states = [] + + for block, block_state in zip(self.blocks, state.blocks_states): + block_states.append(block.sync(block_state.copy(), pmap_axis_name)) + + return BlockDiagonalCurvature.State( + synced=jnp.asarray(True), + blocks_states=tuple(block_states), + ) + + @utils.auto_scope_method + def sync( + self, + state: State, + pmap_axis_name: str | None, + ) -> State: + + return jax.lax.cond( + state.synced, + lambda s: s, + functools.partial(self._sync_state, pmap_axis_name=pmap_axis_name), + state, + ) + + @utils.auto_scope_method + def multiply_matpower( + self, + state: State, + parameter_structured_vector: utils.Params, + identity_weight: Numeric | Sequence[Numeric], + power: Scalar, + exact_power: bool, + use_cached: bool, + pmap_axis_name: str | None, + norm_to_scale_identity_weight_per_block: str | None = None, + ) -> utils.Params: + + blocks_vectors = self.params_vector_to_blocks_vectors( + parameter_structured_vector) + + identity_weight = utils.to_tuple_or_repeat(identity_weight, self.num_blocks) + + def make_thunk(block, block_state, block_vector, block_identity_weight): + + def thunk(): + + weight = block_identity_weight + + if (norm_to_scale_identity_weight_per_block is not None + and norm_to_scale_identity_weight_per_block != "none"): + + weight *= block.norm( + block_state, norm_to_scale_identity_weight_per_block) + + return block.multiply_matpower( + state=block_state, + vector=block_vector, + identity_weight=weight, + power=power, + exact_power=exact_power, + use_cached=use_cached, + ) + + return thunk + + thunks = [] + for block, block_state, block_vector, block_identity_weight in zip( + self.blocks, state.blocks_states, blocks_vectors, identity_weight): + + thunks.append( + make_thunk(block, block_state, block_vector, block_identity_weight)) + + if self._distributed_multiplies and pmap_axis_name is not None: + result = utils.distribute_thunks(thunks, pmap_axis_name) + else: + result = tuple(thunk() for thunk in thunks) + + parameter_structured_result = self.blocks_vectors_to_params_vector(result) + + assert utils.abstract_objects_equal( + parameter_structured_vector, parameter_structured_result) + + return parameter_structured_result + + @utils.auto_scope_method + def block_eigenvalues( + self, + state: State, + use_cached: bool, + ) -> tuple[Array, ...]: + """Computes the eigenvalues for each block of the curvature estimator. + + Args: + state: The state of the estimator. + use_cached: Whether to use a cached versions of the eigenvalues or to use + the most recent curvature estimates to compute them. The cached version + are going to be *at least* as fresh as the last time you called + :func:`~CurvatureEstimator.update_cache` with ``eigenvalues=True``. + + Returns: + A tuple of arrays containing the eigenvalues for each block. The + order of this tuple corresponds to the ordering of ``self.blocks``. + To understand which parameters correspond to which block you can call + ``self.parameters_block_index``. + """ + return tuple(block.eigenvalues(b_state, use_cached=use_cached) + for block, b_state in zip(self.blocks, state.blocks_states)) + + @utils.auto_scope_method + def eigenvalues( + self, + state: State, + use_cached: bool, + ) -> Array: + + blocks_eigenvalues = self.block_eigenvalues(state, use_cached) + return jnp.concatenate(blocks_eigenvalues, axis=0) + + def _package_params_curvature_into_blocks_info( + self, + params: utils.Params, + params_curvature: utils.Params, + ) -> list[tracer.LayerVjpData]: + + params_flat = jax.tree_util.tree_leaves(params) + params_curvature_flat = jax.tree_util.tree_leaves(params_curvature) + blocks_info: list[tracer.LayerVjpData] = [] + + for indices in self.jaxpr.layer_indices: + + primals = tags.LayerData( + inputs=(), + outputs=(), + params=tuple(params_flat[i] for i in indices) + ) + tangents = tags.LayerData( + inputs=(), + outputs=(), + params=tuple(params_curvature_flat[i] for i in indices) + ) + + blocks_info.append(tracer.LayerVjpData( + primals=primals, + tangents=tangents, + )) + + assert ( + len(blocks_info) == self.num_blocks + ), f"{len(blocks_info)=}, {self.num_blocks=}" + + return blocks_info + + # Helper function that updates the blocks given a vjp vector + def _update_blocks( + self, + blocks_info, + state, + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: int, + ): + + assert len(blocks_info) == self.num_blocks + + new_state = [] + for block, block_state, block_info in zip( + self.blocks, state.blocks_states, blocks_info): + + new_state.append( + block.update_curvature_matrix_estimate( + block_state, + block_info, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + ) + ) + + return BlockDiagonalCurvature.State( + synced=jnp.asarray(False), + blocks_states=tuple(new_state), + ) + + def _maybe_do_multiple_updates(self, update_func, state, rng, ema_old): + + if self._num_samples > 1 and self._should_vmap_samples: + + def f(rng_i): + return update_func(state, rng_i, ema_old) + + states = jax.vmap(f)(jax.random.split(rng, self._num_samples)) + + # This implementation is quick and hacky and might break in the future. + # It works by averaging the states only for their floating point leaves, + # which are assumed to be statistics tensors. + return jax.tree_util.tree_map( + lambda x: ( # pylint: disable=g-long-lambda + jnp.mean(x, axis=0) if jnp.issubdtype(x.dtype, jnp.floating) + else x[0]), + states) + + elif self._num_samples > 1: + + def f(carry, rng_i): + + state_i, ema_old_i = carry + new_state_i = update_func(state_i, rng_i, ema_old_i) + + return (new_state_i, jnp.ones_like(ema_old_i)), None + + (new_state, _), _ = jax.lax.scan( + f, + init=(state, jnp.asarray(ema_old)), + xs=jax.random.split(rng, self._num_samples) + ) + return new_state + + elif self._num_samples == 1: + return update_func(state, rng, ema_old) + + else: + # Don't update the preconditioner at all. + return state + + def _update_exact_curvature_matrix_estimate( + self, + *, + losses, + losses_vjp, + state: State, + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + estimation_mode: str, + ) -> State: + """Updates exact Fisher/GGN factors from an already-built layer VJP.""" + + zero_tangents = jax.tree_util.tree_map( + jnp.zeros_like, + list(loss.parameter_dependants for loss in losses), + ) + if estimation_mode == "fisher_exact": + inner_shapes = [loss.fisher_factor_inner_shape[1:] for loss in losses] + else: + inner_shapes = [loss.ggn_factor_inner_shape[1:] for loss in losses] + + total_num_indices = sum(utils.product(shape) for shape in inner_shapes) + ema_new = ema_new / total_num_indices + + for i, (loss, shape) in enumerate(zip(losses, inner_shapes)): + for index in np.ndindex(shape): + vjp_vec = zero_tangents.copy() + if estimation_mode == "fisher_exact": + vjp_vec[i] = loss.multiply_fisher_factor_replicated_one_hot(index) + else: + vjp_vec[i] = loss.multiply_ggn_factor_replicated_one_hot(index) + + if isinstance(vjp_vec[i], Array): + vjp_vec[i] = (vjp_vec[i],) + vjp_vec[i] = jax.tree_util.tree_map( + lambda x: x * jnp.sqrt(total_num_indices), + vjp_vec[i], + ) + state = self._update_blocks( + losses_vjp(tuple(vjp_vec)), + state=state, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + ) + ema_old = 1.0 + + return state + + @utils.auto_scope_method + def update_curvature_matrix_estimate_and_value_and_grad( + self, + state: State, + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + rng: PRNGKey, + func_args: utils.FuncArgs, + estimation_mode: str | None = None, + pmap_axis_name: str | None = None, + ) -> tuple[State, Array, utils.Params]: + """Updates exact Fisher factors and returns loss/grads from one forward.""" + + del rng, pmap_axis_name + if not self.finalized: + self.finalize(func_args) + + estimation_mode = estimation_mode or self.default_estimation_mode + if estimation_mode != "fisher_exact": + raise ValueError( + "Shared curvature/value-and-grad execution supports only " + f"`fisher_exact`; got {estimation_mode!r}." + ) + if self._vjp_and_value_and_grad is None: + raise ValueError( + "This estimator has no `shared_forward_value_func`." + ) + + losses, losses_vjp, loss, grads = self._vjp_and_value_and_grad(func_args) + if any( + not isinstance(loss_obj, loss_functions.NegativeLogProbLoss) + for loss_obj in losses + ): + raise ValueError( + "One of the losses is incompatible with `fisher_exact`." + ) + state = self._update_exact_curvature_matrix_estimate( + losses=losses, + losses_vjp=losses_vjp, + state=state, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + estimation_mode=estimation_mode, + ) + return state, loss, grads + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: State, + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + rng: PRNGKey, + func_args: utils.FuncArgs, + estimation_mode: str | None = None, + pmap_axis_name: str | None = None, # only used for fisher_empirical_direct_synced # pylint: disable=line-too-long + ) -> State: + + if not self.finalized: + self.finalize(func_args) + + estimation_mode = estimation_mode or self.default_estimation_mode + + # Compute the losses and the VJP function from the function inputs + losses, losses_vjp = self._compute_losses_vjp(func_args) + + if "fisher" in estimation_mode: + if any(not isinstance(l, loss_functions.NegativeLogProbLoss) + for l in losses): + raise ValueError( + f"One of the losses in the function is not an instance of " + f"`loss_functions.NegativeLogProbLoss`, which is incompatible " + f"with the estimation mode provided - {estimation_mode}.") + + if estimation_mode == "fisher_gradients": + + def update_func(state_i, rng_i, ema_old_i): + + keys = jax.random.split( + rng_i, len(losses)) if len(losses) > 1 else [rng_i] + + vjp_vec = tuple( + loss.grad_of_evaluate_on_sample(key, coefficient_mode="sqrt") + for loss, key in zip(losses, keys)) + + return self._update_blocks( + losses_vjp(vjp_vec), + state=state_i, + ema_old=ema_old_i, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + ) + + return self._maybe_do_multiple_updates(update_func, state, rng, ema_old) + + elif estimation_mode == "fisher_empirical": + + vjp_vec = tuple( + loss.grad_of_evaluate(None, coefficient_mode="regular") + for loss in losses) + + # losses_vjp(vjp_vec) will return a gradient which isn't normalized by + # batch_size. Meanwhile, the curvature block classes will normalize their + # statistics (which are 2nd-order stats of the gradients) by just + # batch_size, instead of the batch_size**2. Thus, we need to divide by + # sqrt(batch_size) here to correct for that (so that the 2nd-order stats + # are effectively normalized by batch_size an additional time). Note that + # we *don't* perform this extra normalization for the regular Fisher or + # GGN computations, because the proper thing to do for their statistics is + # to normalize by just batch_size, and not batch_size**2. (This is due to + # how the vectors used to form the statistics for the Fisher/GGN are, by + # construction, statistically uncorrelated across the mini-batch.) + vjp_vec = utils.scalar_div(vjp_vec, jnp.sqrt(batch_size)) + + return self._update_blocks( + losses_vjp(vjp_vec), + state=state, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + ) + + elif estimation_mode in {"fisher_empirical_direct", + "fisher_empirical_direct_synced"}: + + # This gradient computation is redundant with the one computed inside of + # the optimizer. Fortunately, XLA should optimize it away. One way that + # this can fail to happen is if self.func is defined using a slightly + # different function than was used to construct func_and_grad (e.g., if + # self.func was defined as value part of an external jax.value_and_grad + # call). + if self.func_and_grad is not None: + params_grad = self.func_and_grad(*func_args)[1] + else: + params_grad = jax.grad(self.func, self.params_index)(*func_args) + + if estimation_mode == "fisher_empirical_direct_synced": + params_grad = utils.pmean_if_pmap(params_grad, pmap_axis_name) + + # Since self.func should be normalized by batch_size, and the curvature + # block classes will normalize their statistics (which are 2nd-order stats + # of the gradients) by batch_size an additional time, the overall + # normalization will be batch_size**3. Thus, we need to multiply the + # gradients by sqrt(batch_size) here to correct for that. + params_grad = utils.scalar_mul(params_grad, jnp.sqrt(batch_size)) + + block_info = self._package_params_curvature_into_blocks_info( + func_args[self.params_index], params_grad) + + return self._update_blocks( + block_info, + state=state, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + ) + + elif estimation_mode in ("fisher_curvature_prop", "ggn_curvature_prop"): + + def update_func(state_i, rng_i, ema_old_i): + + keys = jax.random.split( + rng_i, len(losses)) if len(losses) > 1 else [rng_i] + + vjp_vec = [] + + for loss, key in zip(losses, keys): + + if estimation_mode == "fisher_curvature_prop": + shape = loss.fisher_factor_inner_shape + random_sign = jax.random.rademacher(key, shape=shape) + vjp_vec.append(loss.multiply_fisher_factor(random_sign)) + + else: + shape = loss.ggn_factor_inner_shape + random_sign = jax.random.rademacher(key, shape=shape) + vjp_vec.append(loss.multiply_ggn_factor(random_sign)) + + return self._update_blocks( + losses_vjp(tuple(vjp_vec)), + state=state_i, + ema_old=ema_old_i, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + ) + + return self._maybe_do_multiple_updates(update_func, state, rng, ema_old) + + elif estimation_mode in ("fisher_exact", "ggn_exact"): + return self._update_exact_curvature_matrix_estimate( + losses=losses, + losses_vjp=losses_vjp, + state=state, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=identity_weight, + batch_size=batch_size, + estimation_mode=estimation_mode, + ) + + else: + raise ValueError(f"Unrecognised estimation_mode {estimation_mode}.") + + @utils.auto_scope_method + def update_cache( + self, + state: State, + identity_weight: Numeric | Sequence[Numeric], + exact_powers: curvature_blocks.ScalarOrSequence | None, + approx_powers: curvature_blocks.ScalarOrSequence | None, + eigenvalues: bool, + pmap_axis_name: str | None, + norm_to_scale_identity_weight_per_block: str | None = None, + ) -> State: + + identity_weight = utils.to_tuple_or_repeat(identity_weight, self.num_blocks) + + def make_thunk(block, block_state, block_identity_weight): + + def thunk(): + + weight = block_identity_weight + + if (norm_to_scale_identity_weight_per_block is not None + and norm_to_scale_identity_weight_per_block != "none"): + + weight *= block.norm( + block_state, norm_to_scale_identity_weight_per_block) + + return block.update_cache( + state=block_state, + identity_weight=block_identity_weight, + exact_powers=exact_powers, + approx_powers=approx_powers, + eigenvalues=eigenvalues, + ) + + return thunk + + thunks = [] + for block, block_state, block_identity_weight in zip(self.blocks, + state.blocks_states, + identity_weight): + + thunks.append(make_thunk(block, block_state, block_identity_weight)) + + if self._distributed_cache_updates and pmap_axis_name is not None: + + def filter_outputs(thunk, vals): + + # We must precompute the matches outside of the thunk itself, as the + # thunk will be traced separately from the current compiled context + # (since it's called within a lax.switch statement). + matches = jax.tree_util.tree_map(lambda o, v: o is v, thunk(), vals) + + def new_thunk(): + return jax.tree_util.tree_map( + lambda o, m: None if m else o, thunk(), matches + ) + return new_thunk + + # Create new thunks that only return the state arrays that they actually + # modify. This should reduce the communication costs associated with the + # syncs performed by utils.distribute_thunks. + filtered_thunks = tuple( + filter_outputs(thunk, block_state) + for thunk, block_state in zip(thunks, state.blocks_states)) + + new_states = utils.distribute_thunks(filtered_thunks, pmap_axis_name) + + # Restore all of the unmodified state arrays. + new_states = jax.tree_util.tree_map(lambda s, n: s if n is None else n, + state.blocks_states, new_states) + + else: + new_states = tuple(thunk() for thunk in thunks) + + return BlockDiagonalCurvature.State( + synced=state.synced, + blocks_states=new_states, + ) + + def undamped_diagonal(self, state: State) -> utils.Params: + result = tuple( + block.undamped_diagonal(block_state) + for block, block_state in zip(self.blocks, state.blocks_states)) + + return self.blocks_vectors_to_params_vector(result) + + @utils.auto_scope_method + def to_diagonal_block_dense_matrix(self, state: State) -> tuple[Array, ...]: + """Returns a tuple of arrays with explicit dense matrices of each block.""" + return tuple(block.to_dense_matrix(block_state) for block, block_state in + zip(self.blocks, state.blocks_states)) + + @utils.auto_scope_method + def to_dense_matrix(self, state: State) -> Array: + return scipy.linalg.block_diag(*self.to_diagonal_block_dense_matrix(state)) diff --git a/src/kfac_jax/_src/curvature_estimator/curvature_estimator.py b/src/kfac_jax/_src/curvature_estimator/curvature_estimator.py new file mode 100644 index 0000000000000000000000000000000000000000..a9eea1bcc3003184b8549337d8007b5fa116e0a9 --- /dev/null +++ b/src/kfac_jax/_src/curvature_estimator/curvature_estimator.py @@ -0,0 +1,365 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the abstract class for curvature estimators.""" +import abc +from typing import Callable, Sequence, Generic, TypeVar +import jax.numpy as jnp +from kfac_jax._src import curvature_blocks +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import loss_functions +from kfac_jax._src import tracer +from kfac_jax._src import utils + +# Types for annotation +Array = utils.Array +PRNGKey = utils.PRNGKey +Numeric = utils.Numeric +Scalar = utils.Scalar +Shape = utils.Shape +LossFunction = loss_functions.LossFunction +LossFunctionsTuple = tuple[loss_functions.LossFunction, ...] +LossFunctionsSequence = Sequence[loss_functions.LossFunction] +LossFunctionInputs = loss_functions.LossFunctionInputs +LossFunctionInputsSequence = Sequence[loss_functions.LossFunctionInputs] +LossFunctionInputsTuple = tuple[loss_functions.LossFunctionInputs, ...] +CurvatureBlockCtor = Callable[ + [tags.LayerTagEqn], + curvature_blocks.CurvatureBlock +] +StateType = TypeVar("StateType") + + +class CurvatureEstimator(Generic[StateType], utils.Finalizable): + """An abstract curvature estimator class. + + This is a class that abstracts away the process of estimating a curvature + matrix and provides many useful functionalities for interacting with it. + The state of the estimator contains two parts: the estimated curvature + internal representation, as well as potential cached values of different + expression involving the curvature matrix (for example matrix powers). + The cached values are only updated once you call the method + :func:`~CurvatureEstimator.update_cache`. Multiple methods contain the keyword + argument ``use_cached`` which specify whether you want to compute the + corresponding expression using the current curvature estimate or using a + cached version. + + Attributes: + func: The model evaluation function. + params_index: The index of the parameters argument in arguments list of + ``func``. + batch_index: The index of the batch data argument in arguments list of + ``func``. + default_estimation_mode: The estimation mode which to use by default when + calling :func:`~CurvatureEstimator.update_curvature_matrix_estimate`. + """ + + def __init__( + self, + func: utils.Func | None = None, + func_and_grad: utils.ValueAndGradFunc | None = None, + params_index: int = 0, + batch_index: int = 1, + default_estimation_mode: str = "ggn_curvature_prop", + ): + """Initializes the CurvatureEstimator instance. + + Args: + func: The loss function, which should have at least one registered loss. + Should return only the loss value and not any auxiliary data. Only one + of ``func`` or ``func_and_grad`` should be provided. + func_and_grad: A function returning the loss value and the gradient as a + tuple. Only one of ``func`` or ``func_and_grad`` should be provided. + params_index: The index of the parameters argument in arguments list of + ``func``. + batch_index: The index of the batch data argument in arguments list of + ``func``. + default_estimation_mode: The estimation mode which to use by default when + calling :func:`~CurvatureEstimator.update_curvature_matrix_estimate`. + """ + + if (func is None) == (func_and_grad is None): + raise ValueError("Exactly one of `func` and `func_and_grad` must be " + "provided.") + + if default_estimation_mode not in self.valid_estimation_modes: + raise ValueError( + f"Unsupported estimation mode: {default_estimation_mode}. This class " + f"currently only supports ones in {self.valid_estimation_modes}.") + + super().__init__() + + if func_and_grad is not None: + self.func = lambda *args, **kwargs: func_and_grad(*args, **kwargs)[0] + else: + self.func = func + + self.func_and_grad = func_and_grad + + self.params_index = params_index + self.batch_index = batch_index + self.default_estimation_mode = default_estimation_mode + self.compute_losses, _ = tracer.compute_all_losses( + func=self.func, params_index=params_index + ) + + @property + def default_mat_type(self) -> str: + """The type of matrix that this estimator is approximating.""" + idx = self.default_estimation_mode.index("_") + return self.default_estimation_mode[:idx] + + @property + @abc.abstractmethod + def valid_estimation_modes(self) -> tuple[str, ...]: + """The valid estimation modes for this estimator.""" + + @property + @abc.abstractmethod + def dim(self) -> int: + """The number of elements of all parameter variables together.""" + + @abc.abstractmethod + def init( + self, + rng: PRNGKey, + func_args: utils.FuncArgs, + exact_powers_to_cache: curvature_blocks.ScalarOrSequence | None, + approx_powers_to_cache: curvature_blocks.ScalarOrSequence | None, + cache_eigenvalues: bool = False, + ) -> StateType: + """Initializes the state for the estimator. + + Args: + rng: The PRNGKey which to be used for any randomness of the + initialization. + func_args: Example function arguments, which to be used to trace the model + function and initialize the state. + exact_powers_to_cache: A single value, or multiple values in a list, which + specify which exact matrix powers that each block should be caching. + Matrix powers for which you intend to call + ``self.multiply_matrix_power``, ``self.multiply_inverse`` or + ``self.multiply`` with ``exact_power=True`` and ``use_cached=True`` must + be provided here. + approx_powers_to_cache: A single value, or multiple values in a list, + which specify approximate matrix powers that each block should be + caching. Matrix powers for which you intend to call + ``self.multiply_matrix_power``, ``self.multiply_inverse`` or + ``self.multiply`` with ``exact_power=False`` and ``use_cached=True`` + must be provided here. + cache_eigenvalues: Specifies whether each block should be caching the + eigenvalues of its approximate curvature. + Returns: + The initialized state of the estimator. + """ + + @abc.abstractmethod + def sync( + self, + state: StateType, + pmap_axis_name: str | None, + ) -> StateType: + """Synchronizes across devices the state of the estimator.""" + + @abc.abstractmethod + def multiply_matpower( + self, + state: StateType, + parameter_structured_vector: utils.Params, + identity_weight: Numeric, + power: Scalar, + exact_power: bool, + use_cached: bool, + pmap_axis_name: str | None, + norm_to_scale_identity_weight_per_block: str | None = None, + ) -> utils.Params: + """Computes ``(CurvatureMatrix + identity_weight I)**power`` times ``vector``. + + Args: + state: The state of the estimator. + parameter_structured_vector: A vector in the same structure as the + parameters of the model. + identity_weight: Specifies the weight of the identity element that is + added to the curvature matrix. This can be either a scalar value or a + list/tuple of scalar in which case each value specifies the weight + individually for each block. + power: The power to which you want to raise the matrix + ``(EstimateCurvature + identity_weight I)``. + exact_power: When set to ``True`` the matrix power of + ``EstimateCurvature + identity_weight I`` is computed exactly. + Otherwise this method might use a cheaper approximation, which *may* + vary across different blocks. + use_cached: Whether to use a cached (and possibly stale) version of the + curvature matrix estimate. + pmap_axis_name: The name of any pmap axis, which will be used for + aggregating any computed values over multiple devices, as well as + parallelizing the computation over devices in a block-wise fashion. + norm_to_scale_identity_weight_per_block: The name of a norm to use to + compute extra per-block scaling for identity_weight. See + psd_matrix_norm() in utils/math.py for the definition of these. + + Returns: + A parameter structured vector containing the product. + """ + + def multiply( + self, + state: StateType, + parameter_structured_vector: utils.Params, + identity_weight: Numeric, + exact_power: bool, + use_cached: bool, + pmap_axis_name: str | None, + norm_to_scale_identity_weight_per_block: str | None = None, + ) -> utils.Params: + """Computes ``(CurvatureMatrix + identity_weight I)`` times ``vector``.""" + + return self.multiply_matpower( + state=state, + parameter_structured_vector=parameter_structured_vector, + identity_weight=identity_weight, + power=1, + exact_power=exact_power, + use_cached=use_cached, + pmap_axis_name=pmap_axis_name, + norm_to_scale_identity_weight_per_block=norm_to_scale_identity_weight_per_block, + ) + + def multiply_inverse( + self, + state: StateType, + parameter_structured_vector: utils.Params, + identity_weight: Numeric, + exact_power: bool, + use_cached: bool, + pmap_axis_name: str | None, + norm_to_scale_identity_weight_per_block: str | None = None, + ) -> utils.Params: + """Computes ``(CurvatureMatrix + identity_weight I)^-1`` times ``vector``.""" + + return self.multiply_matpower( + state=state, + parameter_structured_vector=parameter_structured_vector, + identity_weight=identity_weight, + power=-1, + exact_power=exact_power, + use_cached=use_cached, + pmap_axis_name=pmap_axis_name, + norm_to_scale_identity_weight_per_block=norm_to_scale_identity_weight_per_block, + ) + + @abc.abstractmethod + def eigenvalues( + self, + state: StateType, + use_cached: bool, + ) -> Array: + """Computes the eigenvalues of the curvature matrix. + + Args: + state: The state of the estimator. + use_cached: Whether to use a cached versions of the eigenvalues or to use + the most recent curvature estimates to compute them. The cached version + are going to be *at least* as fresh as the last time you called + :func:`~CurvatureEstimator.update_cache` with ``eigenvalues=True``. + + Returns: + A single array containing the eigenvalues of the curvature matrix. + """ + + @abc.abstractmethod + def update_curvature_matrix_estimate( + self, + state: StateType, + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + rng: PRNGKey, + func_args: utils.FuncArgs, + estimation_mode: str | None = None, + pmap_axis_name: str | None = None, + ) -> StateType: + """Updates the estimator's curvature estimates. + + Args: + state: The state of the estimator to update. + ema_old: Specifies the weight of the old value when computing the updated + estimate in the moving average. + ema_new: Specifies the weight of the new value when computing the updated + estimate in the moving average. + identity_weight: The weight of the identity added to the block's curvature + matrix before computing the cached matrix power. + batch_size: The batch size. + rng: A PRNGKey to be used for any potential sampling in the estimation + process. + func_args: A structure with the values of the inputs to the traced + function (the ``tagged_func`` passed into the constructor) which to be + used for the estimation process. Should have the same structure as the + argument ``func_args`` passed in the constructor. + estimation_mode: The type of curvature estimator to use. By default + (e.g. if ``None``) will use ``self.default_estimation_mode``. Must be + one of ``self.valid_estimation_modes``. + pmap_axis_name: The name of a pmap axis, which will be used for + aggregating values over multiple devices if needed. + + Returns: + The updated state. + """ + + @abc.abstractmethod + def update_cache( + self, + state: StateType, + identity_weight: Numeric, + exact_powers: curvature_blocks.ScalarOrSequence | None, + approx_powers: curvature_blocks.ScalarOrSequence | None, + eigenvalues: bool, + pmap_axis_name: str | None, + ) -> StateType: + """Updates the estimator cached values. + + Args: + state: The state of the estimator to update. + identity_weight: Specified the weight of the identity element that is + added to the curvature matrix. This can be either a scalar value or a + list/tuple of scalar in which case each value specifies the weight + individually for each block. + exact_powers: Specifies which exact matrix powers in the cache should be + updated. + approx_powers: Specifies which approximate matrix powers in the cache + should be updated. + eigenvalues: Specifies whether to update the cached eigenvalues + of each block. If they have not been cached before, this will create + an entry with them in the block's cache. + pmap_axis_name: The name of any pmap axis, which will be used for + aggregating any computed values over multiple devices, as well as + parallelizing the computation over devices in a block-wise fashion. + + Returns: + The updated state. + """ + + @abc.abstractmethod + def to_dense_matrix(self, state: StateType) -> Array: + """Returns an explicit dense array representing the curvature matrix.""" + + def compute_func_from_registered(self, func_args, batch_size) -> Array: + + losses = self.compute_losses(func_args) + + loss_values = tuple( + jnp.sum(loss.evaluate(None, coefficient_mode="regular")) + for loss in losses) + + return sum(loss_values) / batch_size diff --git a/src/kfac_jax/_src/curvature_estimator/explicit_exact.py b/src/kfac_jax/_src/curvature_estimator/explicit_exact.py new file mode 100644 index 0000000000000000000000000000000000000000..6181c3a5e18c78a754e88024fc56491dcbc3f793 --- /dev/null +++ b/src/kfac_jax/_src/curvature_estimator/explicit_exact.py @@ -0,0 +1,161 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the ExplicitExactCurvature class.""" +from typing import Any, Callable, Mapping +import jax +from kfac_jax._src import curvature_blocks +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import utils +from kfac_jax._src.curvature_estimator import block_diagonal + +# Types for annotation +PRNGKey = utils.PRNGKey +Numeric = utils.Numeric +CurvatureBlockCtor = Callable[ + [tags.LayerTagEqn], + curvature_blocks.CurvatureBlock +] +BlockDiagonalState = block_diagonal.BlockDiagonalCurvature.State + + +class ExplicitExactCurvature(block_diagonal.BlockDiagonalCurvature): + """Explicit exact full curvature estimator class. + + This class estimates the full curvature matrix by looping over the batch + dimension of the input data and for each single example computes an estimate + of the curvature matrix and then averages over all examples in the input data. + This implies that the computation scales linearly (without parallelism) with + the batch size. The class stores the estimated curvature as a dense matrix, + hence its memory requirement is (number of parameters)^2. If + ``estimation_mode`` is ``fisher_exact`` or ``ggn_exact`` then this would + compute the exact curvature, but other modes are also supported. As a result + of looping over the input data this class needs to know the index of the batch + in the arguments to the model function and additionally, since the loop is + achieved through indexing, each array leaf of that argument must have the same + first dimension size, which will be interpreted as the batch size. + """ + + def __init__( + self, + func: utils.Func | None = None, + func_and_grad: utils.ValueAndGradFunc | None = None, + default_estimation_mode: str | None = None, + layer_tag_to_block_ctor: + Mapping[str, CurvatureBlockCtor] | None = None, + auto_register_tags: bool = False, + param_order: tuple[int, ...] | None = None, + **kwargs: Any, + ): + """Initializes the curvature instance. + + Args: + func: The loss function, which should have at least one registered loss. + Should return only the loss value and not any auxiliary data. Only one + of ``func`` or ``func_and_grad`` should be provided. + func_and_grad: A function returning the loss value and the gradient as a + tuple. Only one of ``func`` or ``func_and_grad`` should be provided. + default_estimation_mode: The estimation mode which to use by default when + calling ``self.update_curvature_matrix_estimate``. If ``None`` this will + be ``'ggn_curvature_prop'``. + layer_tag_to_block_ctor: An optional dict mapping tags to specific classes + of block approximations, which to override the default ones. + auto_register_tags: This argument will be ignored since this subclass + doesn't use automatic registration. + param_order: An optional tuple of ints specifying the order of parameters + (with the reference order being the one used by ``func``). If not + specified, the reference order is used. The parameter order will + determine the order of blocks returned by + ``to_diagonal_block_dense_matrix``, and the order of the rows and + columns of ``to_dense_matrix``. + **kwargs: Addiional keyword arguments passed to the superclass + ``BlockDiagonalCurvature``. + """ + + if (func is None) == (func_and_grad is None): + raise ValueError("Exactly one of `func` and `func_and_grad` must be " + "provided.") + + if layer_tag_to_block_ctor is None: + layer_tag_to_block_ctor = dict(generic=curvature_blocks.NaiveFull) + + if func_and_grad is not None: + func = lambda *args, **kwargs: func_and_grad(*args, **kwargs)[0] + + def retagged_func(params, *args): + + params_flat, params_treedef = jax.tree_util.tree_flatten(params) + + if param_order is not None: + params_flat_canonical_order = [params_flat[i] for i in param_order] + params_flat[param_order[0]] = tags.register_generic( + *params_flat_canonical_order) + + else: + params_flat[0] = tags.register_generic(*params_flat) + + params = jax.tree_util.tree_unflatten(params_treedef, params_flat) + + return func(params, *args) + + super().__init__( + func=retagged_func, + default_estimation_mode=default_estimation_mode or "ggn_curvature_prop", + layer_tag_to_block_ctor=layer_tag_to_block_ctor, + auto_register_tags=False, + **kwargs, + ) + + @utils.auto_scope_method + def update_curvature_matrix_estimate( + self, + state: BlockDiagonalState, + ema_old: Numeric, + ema_new: Numeric, + identity_weight: Numeric, + batch_size: Numeric, + rng: PRNGKey, + func_args: utils.FuncArgs, + estimation_mode: str | None = None, + pmap_axis_name: str | None = None, + ) -> BlockDiagonalState: + + rng = jax.random.split(rng, batch_size) + + super_ = super() + + def single_state_update( + index: Numeric, + state_: BlockDiagonalState + ) -> BlockDiagonalState: + + is_first = index == 0 + args = list(func_args) + + # Index the batch for the `index` arguments. + args[self.batch_index] = jax.tree_util.tree_map( + lambda x: x[index][None], args[self.batch_index]) + + return super_.update_curvature_matrix_estimate( + state=state_, + ema_old=is_first * ema_old + (1 - is_first) * 1.0, + ema_new=ema_new / batch_size, + identity_weight=identity_weight, + batch_size=1, + rng=rng[index], + func_args=args, + estimation_mode=estimation_mode, + pmap_axis_name=pmap_axis_name, + ) + + return jax.lax.fori_loop(0, batch_size, single_state_update, state) diff --git a/src/kfac_jax/_src/curvature_estimator/implicit_exact.py b/src/kfac_jax/_src/curvature_estimator/implicit_exact.py new file mode 100644 index 0000000000000000000000000000000000000000..08068afed2ef8220d495066656a5e4ddb8c2003d --- /dev/null +++ b/src/kfac_jax/_src/curvature_estimator/implicit_exact.py @@ -0,0 +1,505 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Module containing the ImplicitExactCurvature class.""" +from typing import Callable, Sequence +import jax +import jax.numpy as jnp +from kfac_jax._src import loss_functions +from kfac_jax._src import tracer +from kfac_jax._src import utils + +# Types for annotation +Array = utils.Array +Shape = utils.Shape +LossFunction = loss_functions.LossFunction +LossFunctionsTuple = tuple[loss_functions.LossFunction, ...] +LossFunctionsSequence = Sequence[loss_functions.LossFunction] +LossFunctionInputs = loss_functions.LossFunctionInputs +LossFunctionInputsSequence = Sequence[loss_functions.LossFunctionInputs] +LossFunctionInputsTuple = tuple[loss_functions.LossFunctionInputs, ...] + + +class ImplicitExactCurvature: + """Represents all exact curvature matrices never constructed explicitly.""" + + def __init__( + self, + func: utils.Func, + params_index: int = 0, + batch_size_extractor: Callable[[utils.Batch], int] = + utils.default_batch_size_extractor, + ): + """Initializes the ImplicitExactCurvature instance. + + Args: + func: The model function, which should have at least one registered loss. + params_index: The index of the parameters argument in arguments list of + ``func``. + batch_size_extractor: A function that takes as input the function + arguments and returns the batch size for a single device. + (Default: ``kfac.utils.default_batch_size_extractor``) + """ + self.compute_losses = tracer.compute_all_losses( + func=func, + params_index=params_index + ) + self._loss_tags_vjp, _ = tracer.loss_tags_vjp( + func=func, + params_index=params_index + ) + self._loss_tags_jvp, _ = tracer.loss_tags_jvp( + func=func, + params_index=params_index, + ) + self._loss_tags_hvp, _ = tracer.loss_tags_hvp( + func=func, + params_index=params_index, + ) + self._batch_size_extractor = batch_size_extractor + + def batch_size(self, func_args: utils.FuncArgs) -> int: + """The expected batch size given a list of loss instances.""" + return self._batch_size_extractor(func_args[-1]) + + def multiply_loss_fisher( + self, + losses: Sequence[loss_functions.NegativeLogProbLoss], + loss_vectors: LossFunctionInputsSequence, + ) -> LossFunctionInputsTuple: + """Multiplies ``loss_vectors`` by the Fisher of the total loss.""" + assert len(losses) == len(loss_vectors) + return tuple(loss.multiply_fisher(vec) + for loss, vec in zip(losses, loss_vectors)) + + def multiply_loss_ggn( + self, + losses: LossFunctionsSequence, + loss_vectors: LossFunctionInputsSequence, + ) -> LossFunctionInputsTuple: + """Multiplies ``loss_vectors`` by the GGN of the total loss.""" + return tuple(loss.multiply_ggn(vec) + for loss, vec in zip(losses, loss_vectors)) + + def multiply_loss_fisher_factor( + self, + losses: Sequence[loss_functions.NegativeLogProbLoss], + loss_inner_vectors: Sequence[Array], + ) -> LossFunctionInputsTuple: + """Multiplies the vectors with the Fisher factors of each loss. + + Args: + losses: A sequence of loss instances. + loss_inner_vectors: A sequence of vectors, each corresponding to one + instance of a loss in losses. + + Returns: + The product of all vectors with the factors of the Fisher of each the + losses. + """ + assert len(losses) == len(loss_inner_vectors) + return tuple(loss.multiply_fisher_factor(vec) + for loss, vec in zip(losses, loss_inner_vectors)) + + def multiply_loss_ggn_factor( + self, + losses: Sequence[loss_functions.LossFunction], + loss_inner_vectors: Sequence[Array], + ) -> LossFunctionInputsTuple: + """Multiplies the vectors with the GGN factors of each loss. + + Args: + losses: A sequence of loss instances. + loss_inner_vectors: A sequence of vectors, each corresponding to one + instance of a loss in losses. + + Returns: + The product of all vectors with the factors of the GGN of each the + losses. + """ + return tuple(loss.multiply_ggn_factor(vec) + for loss, vec in zip(losses, loss_inner_vectors)) + + def multiply_loss_fisher_factor_transpose( + self, + losses: Sequence[loss_functions.NegativeLogProbLoss], + loss_vectors: LossFunctionInputsSequence, + ) -> tuple[Array, ...]: + """Multiplies the vectors with the transposed Fisher factors of each loss. + + Args: + losses: A sequence of loss instances. + loss_vectors: A sequence of vectors, each corresponding to one instance of + a loss in losses. + + Returns: + The product of all vectors with the factors of the Fisher of each the + losses. + """ + assert len(losses) == len(loss_vectors) + return tuple(loss.multiply_fisher_factor_transpose(vec) + for loss, vec in zip(losses, loss_vectors)) + + def multiply_loss_ggn_factor_transpose( + self, + losses: LossFunctionsSequence, + loss_vectors: LossFunctionInputsSequence, + ) -> tuple[Array, ...]: + """Multiplies the vectors with the transposed GGN factors of each loss. + + Args: + losses: A sequence of loss instances. + loss_vectors: A sequence of vectors, each corresponding to one instance of + a loss in losses. + + Returns: + The product of all vectors with the factors of the GGN of each the + losses. + """ + return tuple(loss.multiply_ggn_factor_transpose(vec) + for loss, vec in zip(losses, loss_vectors)) + + @classmethod + def _assert_losses_same( + cls, + losses1: Sequence[loss_functions.LossFunction], + losses2: Sequence[loss_functions.LossFunction], + ) -> None: + """Asserts that the two losses sequence are equivalent.""" + assert len(losses1) == len(losses2) + for loss1, loss2 in zip(losses1, losses2): + assert isinstance(loss1, type(loss2)) + inputs1 = jax.tree_util.tree_leaves(loss1.parameter_dependants) + inputs2 = jax.tree_util.tree_leaves(loss2.parameter_dependants) + for in1, in2 in zip(inputs1, inputs2): + assert in1.shape == in2.shape + assert in1.dtype == in2.dtype + + @utils.auto_scope_method + def multiply_hessian( + self, + func_args: utils.FuncArgs, + parameter_structured_vector: utils.Params, + ) -> utils.Params: + """Multiplies the vector with the Hessian matrix of the total loss. + + Args: + func_args: The inputs to the model function, on which to evaluate the + Hessian matrix. + parameter_structured_vector: The vector which to multiply with the Hessian + matrix. + + Returns: + The product ``Hv``. + """ + vector, _ = self._loss_tags_hvp(func_args, parameter_structured_vector) + batch_size = self.batch_size(func_args) + + assert utils.abstract_objects_equal(parameter_structured_vector, vector) + + return utils.scalar_div(vector, batch_size) + + @utils.auto_scope_method + def multiply_jacobian( + self, + func_args: utils.FuncArgs, + parameter_structured_vector: utils.Params, + return_loss_objects: bool = False, + ) -> ( + LossFunctionInputsTuple | + tuple[LossFunctionInputsTuple, LossFunctionsTuple] + ): + """Multiplies a vector by the model's Jacobian. + + Args: + func_args: The inputs to the model function. + parameter_structured_vector: A vector in the same structure as the + parameters of the model. + return_loss_objects: If set to `True` will return as an additional output + the loss objects evaluated at the provided function arguments. + + Returns: + The product ``J v``, where ``J`` is the model's Jacobian and ``v`` is + given by ``parameter_structured_vector``. + """ + losses, jacobian_vectors = self._loss_tags_jvp( + func_args, parameter_structured_vector) + + if return_loss_objects: + return jacobian_vectors, losses + + return jacobian_vectors + + @utils.auto_scope_method + def multiply_jacobian_transpose( + self, + func_args: utils.FuncArgs, + loss_input_vectors: LossFunctionInputsSequence, + return_loss_objects: bool = False, + ) -> utils.Params | tuple[utils.Params, LossFunctionsTuple]: + """Multiplies a vector by the model's transposed Jacobian. + + Args: + func_args: The inputs to the model function. + loss_input_vectors: A sequence over losses of sequences of arrays that + are the size of the loss's inputs. This represents the vector to be + multiplied. + return_loss_objects: If set to `True` will return as an additional output + the loss objects evaluated at the provided function arguments. + + Returns: + The product ``J^T v``, where ``J`` is the model's Jacobian and ``v`` is + given by ``loss_inner_vectors``. + """ + losses, vjp = self._loss_tags_vjp(func_args) + + vector = vjp(loss_input_vectors) + + if return_loss_objects: + return vector, losses + + return vector + + @utils.auto_scope_method + def multiply_fisher( + self, + func_args: utils.FuncArgs, + parameter_structured_vector: utils.Params, + ) -> utils.Params: + """Multiplies the vector with the Fisher matrix of the total loss. + + Args: + func_args: The inputs to the model function, on which to evaluate the + Fisher matrix. + parameter_structured_vector: The vector which to multiply with the Fisher + matrix. + + Returns: + The product ``Fv``. + """ + jacobian_vectors, losses = self.multiply_jacobian( + func_args, parameter_structured_vector, True) + + losses: Sequence[loss_functions.NegativeLogProbLoss] + + if any(not isinstance(l, loss_functions.NegativeLogProbLoss) + for l in losses): + raise ValueError("To use `multiply_fisher` all registered losses must " + "be a subclass of `NegativeLogProbLoss`.") + + loss_fisher_jacobian_vectors = self.multiply_loss_fisher( + losses, jacobian_vectors) + + vector = self.multiply_jacobian_transpose( + func_args, loss_fisher_jacobian_vectors) + + assert utils.abstract_objects_equal(parameter_structured_vector, vector) + + return utils.scalar_div(vector, self.batch_size(func_args)) + + @utils.auto_scope_method + def multiply_ggn( + self, + func_args: utils.FuncArgs, + parameter_structured_vector: utils.Params, + ) -> utils.Params: + """Multiplies the vector with the GGN matrix of the total loss. + + Args: + func_args: The inputs to the model function, on which to evaluate the GGN + matrix. + parameter_structured_vector: The vector which to multiply with the GGN + matrix. + + Returns: + The product ``Gv``. + """ + jacobian_vectors, losses = self.multiply_jacobian( + func_args, parameter_structured_vector, True) + + loss_ggn_jacobian_vectors = self.multiply_loss_ggn( + losses, jacobian_vectors) + + vector = self.multiply_jacobian_transpose( + func_args, loss_ggn_jacobian_vectors) + + assert utils.abstract_objects_equal(parameter_structured_vector, vector) + + return utils.scalar_div(vector, self.batch_size(func_args)) + + @utils.auto_scope_method + def multiply_fisher_factor_transpose( + self, + func_args: utils.FuncArgs, + parameter_structured_vector: utils.Params, + ) -> tuple[Array, ...]: + """Multiplies the vector with the transposed factor of the Fisher matrix. + + Args: + func_args: The inputs to the model function, on which to evaluate the + Fisher matrix. + parameter_structured_vector: The vector which to multiply with the Fisher + matrix. + + Returns: + The product ``B^T v``, where ``F = BB^T``. + """ + jacobian_vectors, losses = self.multiply_jacobian( + func_args, parameter_structured_vector, True) + + losses: Sequence[loss_functions.NegativeLogProbLoss] + + if any(not isinstance(l, loss_functions.NegativeLogProbLoss) + for l in losses): + raise ValueError("To use `multiply_fisher` all registered losses must " + "be a subclass of `NegativeLogProbLoss`.") + + loss_vectors = self.multiply_loss_fisher_factor_transpose( + losses, jacobian_vectors) + + return utils.scalar_div(loss_vectors, jnp.sqrt(self.batch_size(func_args))) + + @utils.auto_scope_method + def multiply_ggn_factor_transpose( + self, + func_args: utils.FuncArgs, + parameter_structured_vector: utils.Params, + ) -> tuple[Array, ...]: + """Multiplies the vector with the transposed factor of the GGN matrix. + + Args: + func_args: The inputs to the model function, on which to evaluate the GGN + matrix. + parameter_structured_vector: The vector which to multiply with the GGN + matrix. + + Returns: + The product ``B^T v``, where ``G = BB^T``. + """ + jacobian_vectors, losses = self.multiply_jacobian( + func_args, parameter_structured_vector, True) + + vectors = self.multiply_loss_ggn_factor_transpose( + losses, jacobian_vectors) + + return utils.scalar_div(vectors, jnp.sqrt(self.batch_size(func_args))) + + @utils.auto_scope_method + def multiply_fisher_factor( + self, + func_args: utils.FuncArgs, + loss_inner_vectors: Sequence[Array], + ) -> utils.Params: + """Multiplies the vector with the factor of the Fisher matrix. + + Args: + func_args: The inputs to the model function, on which to evaluate the + Fisher matrix. + loss_inner_vectors: The vector which to multiply with the Fisher factor + matrix. + + Returns: + The product ``Bv``, where ``F = BB^T``. + """ + losses, vjp = self._loss_tags_vjp(func_args) + losses: Sequence[loss_functions.NegativeLogProbLoss] + + if any(not isinstance(l, loss_functions.NegativeLogProbLoss) + for l in losses): + raise ValueError("To use `multiply_fisher` all registered losses must " + "be a subclass of `NegativeLogProbLoss`.") + + fisher_factor_vectors = self.multiply_loss_fisher_factor( + losses, loss_inner_vectors) + + vectors = vjp(fisher_factor_vectors) + + return utils.scalar_div(vectors, jnp.sqrt(self.batch_size(func_args))) + + @utils.auto_scope_method + def multiply_ggn_factor( + self, + func_args: utils.FuncArgs, + loss_inner_vectors: Sequence[Array], + ) -> utils.Params: + """Multiplies the vector with the factor of the GGN matrix. + + Args: + func_args: The inputs to the model function, on which to evaluate the GGN + matrix. + loss_inner_vectors: The vector which to multiply with the GGN factor + matrix. + + Returns: + The product ``Bv``, where ``G = BB^T``. + """ + losses, vjp = self._loss_tags_vjp(func_args) + + ggn_factor_vectors = self.multiply_loss_ggn_factor( + losses, loss_inner_vectors) + + vectors = vjp(ggn_factor_vectors) + + return utils.scalar_div(vectors, jnp.sqrt(self.batch_size(func_args))) + + def get_loss_inner_vector_shapes_and_batch_size( + self, + func_args: utils.FuncArgs, + mode: str + ) -> tuple[tuple[Shape, ...], int]: + """Get shapes of loss inner vectors, and the batch size. + + Args: + func_args: The inputs to the model function. + mode: A string representing the type of curvature matrix for the loss + inner vectors. Can be "fisher" or "ggn". + + Returns: + Shapes of loss inner vectors in a tuple, and the batch size as an int. + """ + losses, _ = self._loss_tags_vjp(func_args) + shapes = [] + + for loss in losses: + if mode == "fisher": + if not isinstance(loss, loss_functions.NegativeLogProbLoss): + raise ValueError(f"To use {mode=}, each loss must be a subclass of " + "`NegativeLogProbLoss`.") + shapes.append(loss.fisher_factor_inner_shape) + elif mode == "ggn": + shapes.append(loss.ggn_factor_inner_shape) + + else: + raise ValueError(f"Unrecognized mode: {mode}") + + return tuple(shapes), self.batch_size(func_args) + + def get_loss_input_shapes_and_batch_size( + self, + func_args: utils.FuncArgs, + ) -> tuple[tuple[tuple[Shape, ...], ...], int]: + """Get shapes of loss input vectors, and the batch size. + + Args: + func_args: The inputs to the model function. + + Returns: + A tuple over losses of tuples containing the shapes of their different + inputs, and the batch size (as an int). + """ + losses, _ = self._loss_tags_vjp(func_args) # pytype: disable=attribute-error # always-use-return-annotations + batch_size = self.batch_size(func_args) + + return (tuple(tuple(x.shape for x in loss.parameter_dependants) + for loss in losses), + batch_size) diff --git a/src/kfac_jax/_src/curvature_estimator/optax_interface.py b/src/kfac_jax/_src/curvature_estimator/optax_interface.py new file mode 100644 index 0000000000000000000000000000000000000000..a2e52c8c715d0e5a21a1d7ab62bde9a3bc81b9b0 --- /dev/null +++ b/src/kfac_jax/_src/curvature_estimator/optax_interface.py @@ -0,0 +1,506 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Utilities for setting up different optimizers.""" + +import functools +from typing import Any, Callable, NamedTuple, Sequence + +import jax +from jax import lax +import jax.numpy as jnp +from kfac_jax._src import utils +from kfac_jax._src.curvature_estimator import block_diagonal +from kfac_jax._src.curvature_estimator import curvature_estimator +import optax + + +Array = utils.Array +Numeric = utils.Numeric +PRNGKey = utils.PRNGKey +Params = utils.Params +Batch = utils.Batch +ValueFunc = utils.ValueFunc +FuncArgs = utils.FuncArgs +ScheduleType = utils.ScheduleType +EstimatorState = block_diagonal.BlockDiagonalCurvature.State + + +class OptaxPreconditionState(NamedTuple): + count: Array + estimator_state: EstimatorState + + +class OptaxPreconditioner: + """An Optax-compatible K-FAC preconditioner.""" + + def __init__( + self, + value_func: ValueFunc, + l2_reg: Numeric = 0.0, + damping: float | None = None, + damping_schedule: ScheduleType | None = None, + norm_constraint: Numeric | None = None, + estimation_mode: str = "fisher_gradients", + curvature_ema: Numeric = 0.95, + curvature_update_period: int = 1, + inverse_update_period: int = 5, + use_exact_inverses: bool = False, + use_sqrt_inv: bool = False, + register_only_generic: bool = False, + patterns_to_skip: Sequence[str] = (), + auto_register_kwargs: dict[str, Any] | None = None, + layer_tag_to_block_ctor: ( + dict[str, curvature_estimator.CurvatureBlockCtor] | None + ) = None, + pmap_axis_name: str = "batch_axis", + batch_size_extractor: Callable[ + [Batch], Numeric + ] = utils.default_batch_size_extractor, + distributed_inverses: bool = True, + distributed_precon_apply: bool = True, + num_samples: int = 1, + should_vmap_samples: bool = False, + norm_to_scale_identity_weight_per_block: str | None = None, + ): + """Initializes the curvature estimator and preconditioner. + + Args: + value_func: Callable. The function should return the value of the loss to + be optimized. + l2_reg: Scalar. Set this value to tell the optimizer what L2 + regularization coefficient you are using (if any). Note the coefficient + appears in the regularizer as ``coeff / 2 * sum(param**2)``. This adds + an additional diagonal term to the curvature and hence will affect the + quadratic model when using adaptive damping. Note that the user is still + responsible for adding regularization to the loss. (Default: ``0.``) + damping: Scalar. The fixed damping that will be used throughput the + lifespan of Preconditioner. (Default: ``None``) + damping_schedule: Callable. A schedule for the damping. This should take + as input the current step number and return a single array that + represents the learning rate. (Default: ``None``) + norm_constraint: Scalar. If specified, the update is scaled down so that + its approximate squared Fisher norm ``v^T F v`` is at most the specified + value. (Note that here ``F`` is the approximate curvature matrix, not + the exact.) (Default: ``None``) + estimation_mode: String. The type of estimator to use for the curvature + matrix. See the documentation for :class:`~CurvatureEstimator` for a + detailed description of the possible options. (Default: + ``fisher_gradients``). + curvature_ema: The decay factor used when calculating the covariance + estimate moving averages. (Default: ``0.95``) + curvature_update_period: Int. The number of steps in between updating the + the curvature estimates. (Default: ``1``) + inverse_update_period: Int. The number of steps in between updating the + the computation of the inverse curvature approximation. (Default: ``5``) + use_exact_inverses: Bool. If ``True``, preconditioner inverses are + computed "exactly" without the pi-adjusted factored damping approach. + Note that this involves the use of eigendecompositions, which can + sometimes be much more expensive. (Default: ``False``) + use_sqrt_inv: Bool. If ``True``, we use inverse square roots for + preconditioner instead of inverse. (Default: ``False``) + register_only_generic: Boolean. Whether when running the auto-tagger to + register only generic parameters, or allow it to use the graph matcher + to automatically pick up any kind of layer tags. (Default: ``False``) + patterns_to_skip: Tuple. A list of any patterns that should be skipped by + the graph matcher when auto-tagging. (Default: ``()``) + auto_register_kwargs: Any additional kwargs to be passed down to + :func:`~auto_register_tags`, which is called by the curvature estimator. + (Default: ``None``) + layer_tag_to_block_ctor: Dictionary. A mapping from layer tags to block + classes which to override the default choices of block approximation for + that specific tag. See the documentation for + :class:`~CurvatureEstimator` for a more detailed description. (Default: + ``None``) + pmap_axis_name: String. The name of the pmap axis to use when + ``multi_device`` is set to True. (Default: ``batch_axis``) + batch_size_extractor: A function that takes as input the function + arguments and returns the batch size for a single device. (Default: + ``kfac.utils.default_batch_size_extractor``) + distributed_inverses: Boolean. Whether to distribute the inverse + computations (required to compute the preconditioner) across the + different devices in a layer-wise fashion. If False, each device will + (redundantly) perform the required computations for all of the layers. + (Default: True) + distributed_precon_apply: Boolean. Whether to distribute the application + of the preconditioner across the different devices in a layer-wise + fashion. If False, each device will (redundantly) perform the required + operations for all of the layers. (Default: True) + num_samples: Number of samples (per case) to use when computing stochastic + curvature matrix estimates. This option is only used when + ``estimation_mode == 'fisher_gradients'`` or ``estimation_mode == + '[fisher,ggn]_curvature_prop'``. (Default: 1) + should_vmap_samples: Whether to use ``jax.vmap`` to compute samples + when ``num_samples > 1``. (Default: False) + norm_to_scale_identity_weight_per_block: The name of a norm to use to + compute extra per-block scaling for the damping. See psd_matrix_norm() + in utils/math.py for the definition of these. (Default: None) + """ + + if curvature_update_period != 1: + raise ValueError( + "`curvature_update_period` must be 1 until XLA can support sharing " + "of computation across cond barriers." + ) + + self._l2_reg = l2_reg + self._damping = damping + self._damping_schedule = damping_schedule + + if (self._damping_schedule is None) == (self._damping is None): + raise ValueError( + "Only one of `damping_schedule` or `damping` has to be specified." + ) + self._norm_constraint = norm_constraint + self._curvature_ema = curvature_ema + self._curvature_update_period = curvature_update_period + self._inverse_update_period = inverse_update_period + self._pmap_axis_name = pmap_axis_name + self._batch_size_extractor = batch_size_extractor + + self._use_cached_inverses = self._inverse_update_period != 1 + self._use_exact_inverses = use_exact_inverses + + self._use_sqrt_inv = use_sqrt_inv + + self._norm_to_scale_identity_weight_per_block = ( + norm_to_scale_identity_weight_per_block + ) + + auto_register_kwargs = auto_register_kwargs or {} + auto_register_kwargs.update(dict( + register_only_generic=register_only_generic, + patterns_to_skip=patterns_to_skip, + )) + # Curvature estimator + self._estimator = block_diagonal.BlockDiagonalCurvature( + func=value_func, + default_estimation_mode=estimation_mode, + params_index=0, + layer_tag_to_block_ctor=layer_tag_to_block_ctor, + distributed_multiplies=distributed_precon_apply, + distributed_cache_updates=distributed_inverses, + num_samples=num_samples, + should_vmap_samples=should_vmap_samples, + auto_register_kwargs=auto_register_kwargs, + ) + + def init( + self, + func_args: FuncArgs, + rng: PRNGKey, + ) -> OptaxPreconditionState: + """Initializes the preconditioner and returns the state.""" + + return OptaxPreconditionState( + count=jnp.array(0, dtype=jnp.int32), + estimator_state=self.estimator.init( + rng=rng, + func_args=func_args, + exact_powers_to_cache=self._exact_powers_to_cache, + approx_powers_to_cache=self._approx_powers_to_cache, + cache_eigenvalues=False, + ), + ) + + @property + def _exact_powers_to_cache(self) -> int | None: + if self._use_exact_inverses and self._use_cached_inverses: + return -1 + return None + + @property + def _approx_powers_to_cache(self) -> int | None: + if not self._use_exact_inverses and self._use_cached_inverses: + return -1 + return None + + @property + def estimator(self) -> block_diagonal.BlockDiagonalCurvature: + """The underlying curvature estimator used by the preconditioner.""" + return self._estimator + + @property + def pmap_axis_name(self): + return self._pmap_axis_name + + def get_identity_weight( + self, state: OptaxPreconditionState + ) -> Array | float: + + damping = self._damping + + if damping is None: + damping = self._damping_schedule(state.count) + + return damping + self._l2_reg + + def sync_estimator_state( + self, + state: OptaxPreconditionState, + ) -> OptaxPreconditionState: + """Syncs the estimator state.""" + + return OptaxPreconditionState( + count=state.count, + estimator_state=self.estimator.sync( + state.estimator_state, pmap_axis_name=self.pmap_axis_name), + ) + + def should_update_estimator_curvature( + self, state: OptaxPreconditionState + ) -> Array | bool: + """Whether at the current step the preconditioner should update the curvature estimates.""" + + if self._curvature_update_period == 1: + return True + + return state.count % self._curvature_update_period == 0 + + def should_sync_estimate_curvature( + self, state: OptaxPreconditionState + ) -> Array | bool: + """Whether at the current step the preconditioner should synchronize (pmean) the curvature estimates.""" + + # sync only before inverses are calculated (either for updating the + # cache or for preconditioning). + if not self._use_cached_inverses: + return True + + return self.should_update_inverse_cache(state) + + def should_update_inverse_cache( + self, state: OptaxPreconditionState + ) -> Array | bool: + """Whether at the current step the preconditioner should update the inverse cache.""" + + if not self._use_cached_inverses: + return False + + return state.count % self._inverse_update_period == 0 + + def maybe_update( + self, + state: OptaxPreconditionState, + func_args: FuncArgs, + rng: PRNGKey, + ) -> OptaxPreconditionState: + """Updates the estimates if it is the right iteration.""" + + # NOTE: This maybe update curvatures and inverses at an iteration. But + # if curvatures should be accumulated for multiple iterations + # before updating inverses (for micro-batching), call + # `maybe_update_estimator_curvature` and `maybe_update_inverse_cache` + # separately, instead of calling this method. + state = self.maybe_update_estimator_curvature( + state=state, + func_args=func_args, + rng=rng, + sync=self.should_sync_estimate_curvature(state), + ) + + state = self.maybe_update_inverse_cache(state) + + return OptaxPreconditionState(state.count, state.estimator_state) + + def _update_estimator_curvature( + self, + estimator_state: EstimatorState, + func_args: FuncArgs, + rng: PRNGKey, + ema_old: Numeric, + ema_new: Numeric, + sync: Array | bool = True + ) -> EstimatorState: + """Updates the curvature estimator state.""" + + state = self.estimator.update_curvature_matrix_estimate( + state=estimator_state, + ema_old=ema_old, + ema_new=ema_new, + # Note that the batch is always the last entry of FuncArgsVariantsdef + batch_size=self._batch_size_extractor(func_args[-1]), + identity_weight=self.get_identity_weight(estimator_state), + rng=rng, + func_args=func_args, + pmap_axis_name=self.pmap_axis_name, + ) + + return jax.lax.cond( + sync, + functools.partial(self.estimator.sync, + pmap_axis_name=self.pmap_axis_name), + lambda state_: state_, + state, + ) + + def maybe_update_estimator_curvature( + self, + state: OptaxPreconditionState, + func_args: FuncArgs, + rng: PRNGKey, + decay_old_ema: Array | bool = True, + sync: Array | bool = True, + ) -> OptaxPreconditionState: + """Updates the curvature estimates if it is the right iteration.""" + + ema_old = decay_old_ema * self._curvature_ema + (1.0 - decay_old_ema) * 1.0 + + return self._maybe_update_estimator_state( + state, + self.should_update_estimator_curvature(state), + self._update_estimator_curvature, + func_args=func_args, + rng=rng, + ema_old=ema_old, + ema_new=1.0, + sync=sync, + ) + + def maybe_update_inverse_cache( + self, + state: OptaxPreconditionState, + ) -> OptaxPreconditionState: + """Updates the estimator state cache if it is the right iteration.""" + + if state.count is None: + raise ValueError( + "PreconditionState is not initialized. Call" + " `maybe_update_estimator_curvature` first." + ) + + return self._maybe_update_estimator_state( + state, + self.should_update_inverse_cache(state), + self.estimator.update_cache, + identity_weight=self.get_identity_weight(state), + exact_powers=self._exact_powers_to_cache, + approx_powers=self._approx_powers_to_cache, + eigenvalues=False, + pmap_axis_name=self.pmap_axis_name, + norm_to_scale_identity_weight_per_block=self._norm_to_scale_identity_weight_per_block, + ) + + def _maybe_update_estimator_state( + self, + state: OptaxPreconditionState, + should_update: Array | bool, + update_func: Callable[..., EstimatorState], + **update_func_kwargs, + ) -> OptaxPreconditionState: + """Updates the estimator state if it should update.""" + + # Because XLA currently doesn't support sharing of computation across cond + # barriers, using a cond here isn't a good idea, for any update_func that + # does expensive computations that share intermediate results with things + # outside the cond. If should_update can be evaluated statically (e.g. if + # it is like i % 1 == 0, for type(i) == int), then the cond will be pruned + # away by XLA and things will be efficient. This is the case when + # curvature_update_period == 1, for example. + # TODO(jamesmartens): Redesign this or wait until XLA is + # updated. + + estimator_state = lax.cond( + should_update, + functools.partial(update_func, **update_func_kwargs), + lambda s: s, + state.estimator_state, + ) + + return OptaxPreconditionState(state.count, estimator_state) + + def apply( + self, + updates: optax.Updates, + state: OptaxPreconditionState, + ) -> optax.Updates: + """Preconditions (= multiplies the inverse curvature estimation matrix to) updates.""" + + new_updates = self.estimator.multiply_matpower( + state=state.estimator_state, + parameter_structured_vector=updates, + identity_weight=self.get_identity_weight(state), + power=-1 if not self._use_sqrt_inv else -0.5, + exact_power=self._use_exact_inverses, + use_cached=self._use_cached_inverses, + pmap_axis_name=self.pmap_axis_name, + norm_to_scale_identity_weight_per_block=self._norm_to_scale_identity_weight_per_block, + ) + + if self._norm_constraint is not None: + + sq_norm_grads = utils.inner_product(new_updates, updates) + del updates + + max_coefficient = jnp.sqrt(self._norm_constraint / sq_norm_grads) + coeff = jnp.minimum(max_coefficient, 1) + + new_updates = utils.scalar_mul(new_updates, coeff) + + else: + del updates + + return new_updates + + def multiply_curvature( + self, + updates: optax.Updates, + state: OptaxPreconditionState, + ) -> optax.Updates: + """Multiplies the (non-inverse) curvature estimation matrix to updates.""" + + # NOTE: Currently, `exact_power` and `use_cached` arguments are not used + # in `self.estimator.multiply()`, and the exact power (of 1) is always used. + # Therefore, the way `identity_weight` (damping) is used with + # `estimator.multiply()` is different from how it's used in + # `estimator.multiply_inverse()` (in `Preconditioner.apply()`) when + # `use_exact_inverses == False` (default). In particular, the former uses + # non-factored damping while the latter uses factored one, and the two are + # NOT the exact inverses of each other. + return self.estimator.multiply( + state=state.estimator_state, + parameter_structured_vector=updates, + identity_weight=self.get_identity_weight(state), + exact_power=self._use_exact_inverses, + use_cached=self._use_cached_inverses, + pmap_axis_name=self.pmap_axis_name, + norm_to_scale_identity_weight_per_block=self._norm_to_scale_identity_weight_per_block, + ) + + def as_gradient_transform( + self, use_inverse: bool = True + ) -> optax.GradientTransformationExtraArgs: + """Multiplies the inverse or non-inverse curvature estimation matrix to updates.""" + + def init_fn(params): + del params + return optax.EmptyState() + + multiply_fn = self.apply if use_inverse else self.multiply_curvature + + def update_fn( + updates, + state, + params=None, + *, + precond_state: OptaxPreconditionState, + **extra_args, + ): + del params, extra_args + return multiply_fn(updates, precond_state), state + + return optax.GradientTransformationExtraArgs(init_fn, update_fn) + + def increment_count(self, state: OptaxPreconditionState): + count_inc = optax.safe_int32_increment(state.count) + return OptaxPreconditionState(count_inc, state.estimator_state) diff --git a/src/kfac_jax/_src/layers_and_loss_tags.py b/src/kfac_jax/_src/layers_and_loss_tags.py new file mode 100644 index 0000000000000000000000000000000000000000..f4b7af1432cc694e2e8e350d9dcde0675dcf7c6c --- /dev/null +++ b/src/kfac_jax/_src/layers_and_loss_tags.py @@ -0,0 +1,485 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC losses and layers tagging Jax primitives.""" +import dataclasses +import functools +from typing import Any, Generic, Sequence, TypeVar + +import jax +import jax.extend as jex + +try: + # JAX v0.10.0 and newer + Effects: type[Any] = jex.core.Effects + no_effects: Effects = jex.core.no_effects +except AttributeError: + # JAX v0.9.2 and older + Effects = jax.core.Effects + no_effects = jax.core.no_effects + + +# Types for annotation +T = TypeVar("T") +Array = jax.Array +Arrays = tuple[Array, ...] + + +@dataclasses.dataclass(frozen=True, kw_only=True, unsafe_hash=True) +@jax.tree_util.register_pytree_node_class +class LayerData(Generic[T]): + """A compact class for all data related to a single layer.""" + + inputs: tuple[T, ...] + outputs: tuple[T, ...] + params: tuple[T, ...] + + def tree_flatten(self) -> tuple[ + tuple[tuple[T, ...], tuple[T, ...], tuple[T, ...]], + None, + ]: + return (self.inputs, self.outputs, self.params), None + + @classmethod + def tree_unflatten(cls, aux_data, children): + assert aux_data is None + inputs, outputs, params = children + return cls(inputs=inputs, outputs=outputs, params=params) + + +@dataclasses.dataclass(kw_only=True, unsafe_hash=True) +class LayerMetaData: + """A compact class for all metadata related to a single layer.""" + variant: str + outputs_index: tuple[int, ...] + inputs_index: tuple[int, ...] + params_index: tuple[int, ...] + name: str | None = None + nesting: tuple[str, ...] = () + params_canonical_order: tuple[int, ...] | None = None + + +@dataclasses.dataclass(kw_only=True, unsafe_hash=True) +class LossMetaData(Generic[T]): + """A compact class for all metadata related to a single layer.""" + + loss_class: type[T] + parameter_dependants: tuple[str, ...] + parameter_independants: tuple[str, ...] + argument_names: tuple[str, ...] + + +def get_and_verify_loss_meta( + args: Sequence[Any], + params: Any, + err_suffix: str = "", +) -> LossMetaData: + """Verifies that the number of arguments matches expectations.""" + meta = params.get("meta") + if meta is None or not isinstance(meta, LossMetaData): + raise ValueError(f"Meta must be LossMetaData, but found {meta=}.") + if len(args) != len(meta.argument_names): + raise ValueError(f"Number of arguments {len(args)} must match the " + f"number of argument names {len(meta.argument_names)} for" + f" {err_suffix}.") + return meta + + +def get_loss_outputs( + args: Sequence[T], + params: dict[str, Any], + err_suffix: str = "", +) -> tuple[T, ...]: + meta = get_and_verify_loss_meta(args, params, err_suffix) + kwargs = dict(zip(meta.argument_names, args)) + return tuple(kwargs[name] for name in meta.parameter_dependants) + + +class LossTag(jex.core.Primitive): + """A Jax primitive for tagging K-FAC losses. + + The primitive is no-op at runtime, however its goal is to tag (annotate) the + Jax computation graph what expression exactly is the loss and what type of + loss it represents. This is the only way for K-FAC to know how to compute the + curvature matrix. + """ + + # Whether the primitive returns multiple outputs (from jex.core.Primitive) + multiple_results = True + + def __init__(self): + """Initializes a loss tag primitive for the given :class:`~LossFunction` class. + + When the primitive is created, the constructor automatically registers it + with the standard Jax machinery for differentiation, :func:`jax.vmap` and + XLA lowering. For further details see please take a look at the JAX + documentation on `primitives + `__. + """ + super().__init__("loss_tag") + + jax.interpreters.mlir.register_lowering(self, self._mlir_lowering) + jax.interpreters.ad.primitive_jvps[self] = self._jvp + + # This line defines how does the tag behave under vmap. It is required for + # any primitive that can be used inside a vmap. The reason why we want to + # allow this is two fold - one to not break user code when the tags are not + # used at all, and two - to be able to define a network with code for a + # single example which is the vmap-ed for a batch. + jax.interpreters.batching.primitive_batchers[self] = self._batching + + def impl(self, *args: Array, **params: Any) -> tuple[Array, ...]: + return get_loss_outputs(args, params) + + def abstract_eval( + self, + *args: Array, + **params: Any, + ) -> tuple[Arrays, Effects]: + + return get_loss_outputs(args, params), no_effects + + def _mlir_lowering( + self, + _: jax.interpreters.mlir.LoweringRuleContext, + *args, + **params: Any, + ) -> tuple[Any, ...]: + """The XLA translation rule for this primitive (creates a no-op tuple).""" + return get_loss_outputs(args, params) + + def _jvp( + self, + arg_values: Sequence[Array], + arg_tangents: Sequence[Array], + **params: Any, + ) -> tuple[Arrays, Arrays]: + """Computes the Jacobian-vector product for the primitive.""" + + if len(arg_values) != len(arg_tangents): + raise ValueError("Values and tangents are not the same length.") + + primal_output = self.bind(*arg_values, **params) + tangent_output = get_loss_outputs(arg_tangents, params) + + return primal_output, tangent_output + + def _batching( + self, + batched_args: Sequence[Array], + batched_dims: int | tuple[int, ...], + **params: Any, + ) -> tuple[Array, int | tuple[int, ...]]: + """Defines how the primitive behaves under :func:`jax.vmap`.""" + + return self.bind(*batched_args, **params), batched_dims[:1] + + +def loss_eqn_parameter_dependants( + eqn: jex.core.JaxprEqn, + raise_an_error: bool = True, +) -> list[jex.core.Var]: + """Returns the parameter dependants variables from the give loss equation.""" + if not isinstance(eqn.primitive, LossTag): + if raise_an_error: + raise ValueError("Primitive must be a LossTag.") + return [] + + meta = eqn.params.get("meta") + assert meta is not None and isinstance(meta, LossMetaData) + assert len(eqn.invars) == len(meta.argument_names) + kwargs = dict(zip(meta.argument_names, eqn.invars)) + return [kwargs[name] for name in meta.parameter_dependants] + + +def loss_eqn_construct_loss( + eqn: jex.core.JaxprEqn, + *args: Array, +) -> Any: + """Constructs an instance of the corresponding :class:`~LossFunction` class.""" + if not isinstance(eqn.primitive, LossTag): + raise ValueError("Primitive must be a LossTag.") + + meta: LossMetaData[T] = eqn.params.get("meta") # pytype: disable=invalid-annotation + assert meta is not None and isinstance(meta, LossMetaData) + assert len(eqn.invars) == len(meta.argument_names) + kwargs = dict(zip(meta.argument_names, args)) + return meta.loss_class(**kwargs) + + +def loss_eqn_class_name(eqn: jex.core.JaxprEqn) -> str: + """The name of the underlying `~LossFunction` class.""" + + if not isinstance(eqn.primitive, LossTag): + raise ValueError("Primitive must be a LossTag.") + + meta: LossMetaData[T] = eqn.params.get("meta") # pytype: disable=invalid-annotation + assert meta is not None and isinstance(meta, LossMetaData) + + return meta.loss_class.__name__ + + +def get_and_verify_layer_meta( + args: Sequence[Any], + params: dict[str, Any], + err_suffix: str = "", +) -> LayerMetaData: + """Verifies that the number of arguments matches expectations.""" + + meta = params.get("meta") + + if meta is None or not isinstance(meta, LayerMetaData): + raise ValueError(f"Meta must be LayerMetaData, but found {meta=}.") + + n = len(args) + + for i in meta.inputs_index: + if i >= n: + raise ValueError( + f"Meta data has {meta.input_index=}, but only {n} " + f"arguments passed for {err_suffix}.") + + for i in meta.outputs_index: + if i >= n: + raise ValueError( + f"Meta data has {meta.output_index=}, but only {n} " + f"arguments passed for {err_suffix}.") + + for i in meta.params_index: + if i >= n: + raise ValueError( + f"Meta data has {meta.params_index=}, but only {n} " + f"arguments passed for {err_suffix}.") + + return meta + + +class LayerTag(jex.core.Primitive): + """A Jax primitive for tagging K-FAC layers. + + The primitive is no-op at runtime, however its goal is to tag (annotate) the + Jax computation graph what expressions represents a single unique layer type. + This is the only way for K-FAC to know how to compute the curvature matrix. + """ + + def __init__(self): + """Initializes a layer tag primitive with the given name. + + Any layer tag primitive must have the following interface `layer_tag( + *outputs, *inputs, *parameters, **kwargs)`. We refer collectively to + ``inputs`` , ``outputs`` and ``parameters`` as operands. All operands must + be Jax arrays, while any of the values in ``kwargs`` must be hashable fixed + constants. + + When the primitive is created, the constructor automatically registers it + with the standard Jax machinery for differentiation, :func:`jax.vmap` and + XLA lowering. For further details see please take a look at the JAX + documentation on `primitives + `__. + + """ + super().__init__(name="layer_tag") + jax.interpreters.mlir.register_lowering(self, self._mlir_lowering) + jax.interpreters.ad.deflinear(self, self._transpose) + jax.interpreters.ad.primitive_transposes[self] = self._transpose + # This line defines how does the tag behave under vmap. It is required for + # any primitive that can be used inside a vmap. The reason why we want to + # allow this is two fold - one to not break user code when the tags are not + # used at all, and two - to be able to define a network with code for a + # single example which is the vmap-ed for a batch. + jax.interpreters.batching.primitive_batchers[self] = self._batching + + def layer_data( # pytype: disable=invalid-annotation + self, + args: Sequence[T], + params: dict[str, Any], + err_suffix: str = "", + exclude_inputs: bool = False, + ) -> LayerData[T]: + """Splits the operands of the primitive into a ``LayerData`` object.""" + + meta = get_and_verify_layer_meta(args, params, err_suffix) + + return LayerData( + inputs=(() if exclude_inputs + else tuple(args[i] for i in meta.inputs_index)), + outputs=tuple(args[i] for i in meta.outputs_index), + params=tuple(args[i] for i in meta.params_index), + ) + + def _mlir_lowering( + self, + _: jax.interpreters.mlir.LoweringRuleContext, + *args: T, + **params: Any, + ) -> tuple[T, ...]: + """The XLA translation rule for this primitive - returns the ``outputs`` .""" + return self.layer_data(args, params).outputs + + @classmethod + def _transpose( + cls, + cotangent: Array, + *args: Array, + **_: Any, + ) -> tuple[Array | None, ...]: + """Computes the cotangents of the operands from those of the primitive.""" + del cls # not used + return (cotangent,) + (None,) * (len(args) - 1) + + def impl(self, *args: Array, **params: Any) -> Array: + # For now we support only single output + [output] = self.layer_data(args, params).outputs + return output + + def abstract_eval( + self, + *args: Array, + **params: Any, + ) -> tuple[Array, Effects]: + # For now we support only single output + [output] = self.layer_data(args, params).outputs + return output, no_effects + + def _batching( + self, + batched_args: Sequence[Array], + batched_dims: int | tuple[int, ...], + **params: Any, + ) -> tuple[Array, int]: + """Defines how the primitive behaves under :func:`jax.vmap`.""" + return self.bind(*batched_args, **params), batched_dims[0] + + +def layer_eqn_data( # pytype: disable=invalid-annotation + eqn: jex.core.JaxprEqn, + raise_an_error: bool = True, +) -> LayerData[jex.core.Var]: + + if isinstance(eqn.primitive, LayerTag): + return eqn.primitive.layer_data(eqn.invars, eqn.params, str(eqn)) + + if raise_an_error: + raise ValueError("Primitive must be a LayerTag.") + else: + return LayerData(inputs=(), outputs=(), params=()) + + +def layer_eqn_name(eqn: jex.core.JaxprEqn) -> str: + meta = get_and_verify_layer_meta(eqn.invars, eqn.params) + if meta.name is None: + raise ValueError("Layer name must be provided at this stage.") + return meta.name + + +loss_tag = LossTag() +layer_tag = LayerTag() + + +def register_generic(*args: Array) -> Array: + """Registers a generic tag around the provided parameters array.""" + return layer_tag.bind( + *args, + meta=LayerMetaData( + variant="generic", + inputs_index=(), + outputs_index=(0,), + params_index=tuple(range(len(args))), + ), + ) + + +def register_dense( + y: Array, + x: Array, + w: Array, + b: Array | None = None, + variant: str = "dense", + **kwargs, +) -> Array: + """Registers a dense layer: ``y = matmul(x, w) + b``.""" + args = (y, x, w) if b is None else (y, x, w, b) + return layer_tag.bind( + *args, + meta=LayerMetaData( + variant=variant, + outputs_index=(0,), + inputs_index=(1,), + params_index=tuple(i + 2 for i in range(len(args) - 2)), + ), + **kwargs, + ) + + +def register_conv2d( + y: Array, + x: Array, + w: Array, + b: Array | None = None, + variant: str = "conv2d", + **kwargs: Any +) -> Array: + """Registers a 2d convolution layer: ``y = conv2d(x, w) + b``.""" + args = (y, x, w) if b is None else (y, x, w, b) + return layer_tag.bind( + *args, + meta=LayerMetaData( + variant=variant, + outputs_index=(0,), + inputs_index=(1,), + params_index=tuple(i + 2 for i in range(len(args) - 2)), + ), + **kwargs, + ) + + +def register_scale_and_shift( + y: Array, + x: Array, + scale: Array | None = None, + shift: Array | None = None, + variant: str = "scale_and_shift", + **kwargs: Any, +) -> Array: + """Registers a scale and shift layer: ``y = x * scale + shift``.""" + args = tuple(a for a in (y, x, scale, shift) if a is not None) + if len(args) < 3: + raise ValueError("At least one of `scale` and `shift` must be provided.") + return layer_tag.bind( + *args, + meta=LayerMetaData( + variant=variant, + outputs_index=(0,), + inputs_index=(1,), + params_index=tuple(i + 2 for i in range(len(args) - 2)), + ), + has_scale=scale is not None, + has_shift=shift is not None, + **kwargs, + ) + + +register_repeated_dense = functools.partial( + register_dense, + variant="repeated_dense", +) + + +class LossTagEqn(jex.core.JaxprEqn): + """A class used only for annotation purposes.""" + primitive: LossTag + + +class LayerTagEqn(jex.core.JaxprEqn): + """A class used only for annotation purposes.""" + primitive: LayerTag diff --git a/src/kfac_jax/_src/loss_functions.py b/src/kfac_jax/_src/loss_functions.py new file mode 100644 index 0000000000000000000000000000000000000000..aeb444774ec6e9e4cd2af15c4fd1010de47687f4 --- /dev/null +++ b/src/kfac_jax/_src/loss_functions.py @@ -0,0 +1,1492 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""""K-FAC loss functions objects, tags and registration functions.""" +import abc +from typing import Sequence, Any + +import distrax +import jax +import jax.numpy as jnp + +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import utils +from typing_extensions import Self + + +Array = utils.Array +Numeric = utils.Numeric +PRNGKey = utils.PRNGKey +Shape = utils.Shape +DType = utils.DType +LossFunctionInputs = tuple[Array, ...] + + +def filter_none( + **kwargs: Numeric | None, +) -> tuple[tuple[Numeric, ...], tuple[str, ...]]: + args, args_names = zip( + *[(value, name) for name, value in kwargs.items() if value is not None] + ) + return args, args_names + + +# pylint: disable=g-one-element-tuple + + +class LossFunction(utils.Finalizable): + """Abstract base class for loss functions. + + Note that unlike typical loss functions used in neural networks these are + neither summed nor averaged over the batch and the output of evaluate() will + not be a scalar. It is up to the user to then to correctly manipulate them as + needed. + + Note that all of the GGN and Fisher (factor) multiplication methods defined in + this class refer to the GGN and Fisher of just the loss function itself, which + does *not* include the parameterized function (e.g. neural network) that + generates the inputs (e.g. predictions or logits) feeding into the loss + function. + """ + + def __init__(self, weight: Numeric): + """Initializes the loss instance. + + Args: + weight: The weight attributed to the loss. + """ + if not isinstance(weight, (int, float)) and type(weight) is not object: # pylint: disable=unidiomatic-typecheck + if not isinstance(weight, Array) or weight.size > 1: + raise ValueError("`weight` must be a scalar value.") + super().__init__() + self._weight = weight + self.finalize() + + @property + def dtype(self) -> DType: + return self.parameter_dependants[0].dtype + + @property + def weight(self) -> Numeric: + """The weight of the loss.""" + return self._weight + + @property + @abc.abstractmethod + def targets(self) -> Array | None: + """The targets (if present) used for evaluating the loss.""" + + @property + @abc.abstractmethod + def parameter_dependants(self) -> tuple[Array, ...]: + """All the parameter dependent arrays of the loss.""" + + @property + def num_parameter_dependants(self) -> int: + """Number of parameter dependent arrays of the loss.""" + return len(self.parameter_dependants) + + @property + @abc.abstractmethod + def parameter_independants(self) -> tuple[Numeric | None, ...]: + """All the parameter independent arrays of the loss.""" + + @property + def num_parameter_independants(self) -> int: + """Number of parameter independent arrays of the loss.""" + return len(self.parameter_independants) + + def copy_with_different_inputs( + self, + parameter_dependants: Sequence[Array], + ) -> Self: + """Creates a copy of the loss function object, but with different inputs.""" + + array_args, aux = self.tree_flatten() + + assert len(array_args) == ( + self.num_parameter_dependants + self.num_parameter_independants + ) + + array_args = (tuple(parameter_dependants) + + tuple(array_args[self.num_parameter_dependants:])) + + return self.tree_unflatten(aux, array_args) + + def tree_flatten( + self, + ) -> tuple[tuple[Numeric | None, ...], dict[str, Any] | None]: + return self.parameter_dependants + self.parameter_independants, None + + @classmethod + def tree_unflatten( + cls, + aux: dict[str, Any] | None, + children: tuple[Numeric | None, ...], + ) -> Self: + return cls(*children, **(aux or {})) # pytype: disable=not-instantiable + + def evaluate( + self, + targets: Array | None = None, + coefficient_mode: str = "regular", + ) -> Array: + """Evaluates the loss function on the targets. + + Args: + targets: The targets, on which to evaluate the loss. If this is set to + ``None`` will use ``self.targets`` instead. + coefficient_mode: Specifies how to use the weight of the loss in + the returned value. There are three options: + + 1. 'regular' - returns ``self.weight * loss(targets)`` + + 2. 'sqrt' - returns ``sqrt(self.weight) * loss(targets)`` + + 3. 'off' - returns ``loss(targets)`` + + Returns: + The value of the loss scaled appropriately by ``self.weight`` according to + the coefficient mode. + Raises: + ValueError if both ``targets`` and ``self.targets`` are ``None``. + """ + if targets is None and self.targets is None: + raise ValueError("Cannot evaluate losses with unspecified targets.") + elif targets is None: + targets = self.targets + if coefficient_mode == "regular": + multiplier = self.weight + elif coefficient_mode == "sqrt": + multiplier = jnp.sqrt(self.weight) + elif coefficient_mode == "off": + multiplier = 1.0 + else: + raise ValueError(f"Unrecognized coefficient_mode={coefficient_mode}.") + return self._evaluate(targets) * multiplier + + @abc.abstractmethod + def _evaluate(self, targets: Array) -> Array: + """Evaluates the value of the loss, disregarding the weight.""" + + def grad_of_evaluate( + self, + targets: Array | None, + coefficient_mode: str, + ) -> tuple[Array, ...]: + """Evaluates the gradient of the loss function w.r.t. its inputs. + + Args: + targets: The targets at which to evaluate the loss. If this is ``None`` + will use ``self.targets`` instead. + coefficient_mode: The coefficient mode to use for evaluation. See + ``self.evaluate`` for more details. + + Returns: + The gradient of the loss function w.r.t. its inputs, at the provided + targets. + """ + def evaluate_sum(inputs: Sequence[Array]) -> Array: + """Evaluates the loss summed over all axis, including batch etc.""" + instance = self.copy_with_different_inputs(inputs) + return jnp.sum(instance.evaluate(targets, coefficient_mode)) + + return jax.grad(evaluate_sum)(self.parameter_dependants) + + def multiply_ggn( + self, + vector: Sequence[Array], + ) -> tuple[Array, ...]: + """Right-multiplies a vector by the GGN of the loss function. + + Args: + vector: The vector to multiply. Must have the same shape(s) as + ``self.inputs``. + + Returns: + The vector right-multiplied by the GGN. Will have the same shape(s) as + ``self.inputs``. + """ + return utils.scalar_mul(self.multiply_ggn_unweighted(vector), self.weight) + + @abc.abstractmethod + def multiply_ggn_unweighted( + self, + vector: Sequence[Array], + ) -> tuple[Array, ...]: + """Unweighted version of :func:`~LossFunction.multiply_ggn`.""" + + def multiply_ggn_factor( + self, + vector: Array, + ) -> tuple[Array, ...]: + """Right-multiplies a vector by a factor B of the GGN. + + Note that B can be any matrix satisfying ``B * B^T = G`` where ``G`` is the + GGN, but will agree with the one used in the other methods of this class. + + Args: + vector: The vector to multiply. Must be of the shape(s) given by + 'self.ggn_factor_inner_shape'. + + Returns: + The vector right-multiplied by B. Will be of the same shape(s) as + ``self.inputs``. + """ + return utils.scalar_mul( + self.multiply_ggn_factor_unweighted(vector), jnp.sqrt(self.weight)) + + @abc.abstractmethod + def multiply_ggn_factor_unweighted( + self, vector: Array + ) -> tuple[Array, ...]: + """Unweighted version of :func:`~LossFunction.multiply_ggn_factor`.""" + + def multiply_ggn_factor_transpose( + self, + vector: Sequence[Array], + ) -> Array: + """Right-multiplies a vector by the transpose of a factor B of the GGN. + + Note that B can be any matrix satisfying ``B * B^T = G`` where G is the GGN, + but will agree with the one used in the other methods of this class. + + Args: + vector: The vector to multiply. Must have the same shape(s) as + ``self.inputs``. + + Returns: + The vector right-multiplied by B^T. Will be of the shape(s) given by + ``self.ggn_factor_inner_shape``. + """ + return utils.scalar_mul( + self.multiply_ggn_factor_transpose_unweighted(vector), + jnp.sqrt(self.weight)) + + @abc.abstractmethod + def multiply_ggn_factor_transpose_unweighted( + self, + vector: Sequence[Array], + ) -> Array: + """Unweighted version of :func:`~LossFunction.multiply_ggn_factor_transpose`.""" + + def multiply_ggn_factor_replicated_one_hot( + self, + index: tuple[int, ...], + ) -> tuple[Array, ...]: + """Right-multiplies a replicated-one-hot vector by a factor B of the GGN. + + A replicated-one-hot vector means a tensor which, for each slice along the + batch dimension (assumed to be dimension 0), is 1.0 in the entry + corresponding to the given index and 0 elsewhere. + + Note that B can be any matrix satisfying ``B * B^T = G`` where G is the GGN, + but will agree with the one used in the other methods of this class. + + The reason that we have this special method and don't just use + multiply_ggn_factor is that we can be more efficient by using knowledge of + the zero-entries in the replicated-one-hot vector. + + Args: + index: A tuple representing in the index of the entry in each slice that + is 1.0 (excluding the batch dimension). Note that len(index) must be + equal to the number of elements of the ``ggn_factor_inner_shape`` tensor + minus one. + + Returns: + The vector right-multiplied by B^T. Will be of the same shape(s) as the + ``inputs`` property. + """ + return utils.scalar_mul( + self.multiply_ggn_factor_replicated_one_hot_unweighted(index), + jnp.sqrt(self.weight)) + + @abc.abstractmethod + def multiply_ggn_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array, ...]: + """Unweighted version of :func:`~LossFunction.multiply_ggn_factor_replicated_one_hot`.""" + + @property + @abc.abstractmethod + def ggn_factor_inner_shape(self) -> Shape: + """The shape of the array returned by `self.multiply_ggn_factor`.""" + + +class NegativeLogProbLoss(LossFunction): + """Base class for loss functions that represent negative log-probability.""" + + @property + def parameter_dependants(self) -> tuple[Array, ...]: + return self.params + + @property + @abc.abstractmethod + def params(self) -> tuple[Array, ...]: + """Parameters to the underlying distribution.""" + + def multiply_fisher( + self, + vector: Sequence[Array], + ) -> tuple[Array, ...]: + """Right-multiplies a vector by the Fisher. + + Args: + vector: The vector to multiply. Must have the same shape(s) as + ``self.inputs``. + + Returns: + The vector right-multiplied by the Fisher. Will have of the same shape(s) + as ``self.inputs``. + """ + return utils.scalar_mul( + self.multiply_fisher_unweighted(vector), self.weight) + + @abc.abstractmethod + def multiply_fisher_unweighted( + self, + vector: Sequence[Array], + ) -> tuple[Array, ...]: + """Unweighted version of :func:`~LossFunction.multiply_fisher`.""" + + def multiply_fisher_factor( + self, + vector: Array, + ) -> tuple[Array, ...]: + """Right-multiplies a vector by a factor B of the Fisher. + + Note that B can be any matrix satisfying ``B * B^T = F`` where F is the + Fisher, but will agree with the one used in the other methods of this class. + + Args: + vector: The vector to multiply. Must have the same shape(s) as + ``self.fisher_factor_inner_shape``. + + Returns: + The vector right-multiplied by B. Will have the same shape(s) as + ``self.inputs``. + """ + return utils.scalar_mul( + self.multiply_fisher_factor_unweighted(vector), jnp.sqrt(self.weight)) + + @abc.abstractmethod + def multiply_fisher_factor_unweighted( + self, + vector: Array, + ) -> tuple[Array, ...]: + """Unweighted version of :func:`~LossFunction.multiply_fisher_factor`.""" + + def multiply_fisher_factor_transpose( + self, + vector: Sequence[Array], + ) -> Array: + """Right-multiplies a vector by the transpose of a factor B of the Fisher. + + Note that B can be any matrix satisfying ``B * B^T = F`` where F is the + Fisher, but will agree with the one used in the other methods of this class. + + Args: + vector: The vector to multiply. Must have the same shape(s) as + ``self.inputs``. + + Returns: + The vector right-multiplied by B^T. Will have the shape given by + ``self.fisher_factor_inner_shape``. + """ + return utils.scalar_mul( + self.multiply_fisher_factor_transpose_unweighted(vector), + jnp.sqrt(self.weight)) + + @abc.abstractmethod + def multiply_fisher_factor_transpose_unweighted( + self, + vector: Sequence[Array], + ) -> Array: + """Unweighted version of :func:`~LossFunction.multiply_fisher_factor_transpose`.""" + + def multiply_fisher_factor_replicated_one_hot( + self, + index: tuple[int, ...], + ) -> tuple[Array, ...]: + """Right-multiplies a replicated-one-hot vector by a factor B of the Fisher. + + A replicated-one-hot vector means a tensor which, for each slice along the + batch dimension (assumed to be dimension 0), is 1.0 in the entry + corresponding to the given index and 0 elsewhere. + + Note that B can be any matrix satisfying ``B * B^T = F`` where F is the + Fisher, but will agree with the one used in the other methods of this class. + + The reason that we have this special method and don't just use + multiply_fisher_factor is that we can be more efficient by using knowledge + of the zero-entries in the replicated-one-hot vector. + + Args: + index: A tuple representing in the index of the entry in each slice that + is 1.0 (excluding the batch dimension). Note that len(index) must be + equal to the number of elements of the ``fisher_factor_inner_shape`` + tensor minus one. + + Returns: + The vector right-multiplied by B. Will have the same shape(s) as + ``self.inputs``. + """ + return utils.scalar_mul( + self.multiply_fisher_factor_replicated_one_hot_unweighted(index), + jnp.sqrt(self.weight)) + + @abc.abstractmethod + def multiply_fisher_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array, ...]: + """Unweighted version of :func:`~LossFunction.multiply_fisher_factor_replicated_one_hot`.""" + + @property + @abc.abstractmethod + def fisher_factor_inner_shape(self) -> Shape: + """The shape of the array returned by :func:`~LossFunction.multiply_fisher_factor`.""" + + @abc.abstractmethod + def sample(self, rng: PRNGKey) -> Array: + """Sample ``targets`` from the underlying distribution.""" + + def grad_of_evaluate_on_sample( + self, + rng: Array, + coefficient_mode: str, + ) -> tuple[Array, ...]: + """Evaluates the gradient of the log probability on a random sample. + + Args: + rng: Jax PRNG key for sampling. + coefficient_mode: The coefficient mode to use for evaluation. + + Returns: + The gradient of the log probability of targets sampled from the + distribution. + """ + return self.grad_of_evaluate(self.sample(rng), coefficient_mode) + + +class NaturalParamsNegativeLogProbLoss(NegativeLogProbLoss, abc.ABC): + """Negative log-probability loss, whose inputs are natural parameters. + + We will take the GGN of the loss to be the Fisher associated with the + distribution, which also happens to be equal to the Hessian for this class + of loss functions. See https://arxiv.org/abs/1412.1193 for details. + + Natural parameters are defined for exponential-family models. See for + example `wikipedia `__. + """ + + def multiply_ggn_unweighted( + self, + vector: Sequence[Array], + ) -> tuple[Array, ...]: + return self.multiply_fisher_unweighted(vector) + + def multiply_ggn_factor_unweighted( + self, + vector: Array, + ) -> tuple[Array, ...]: + return self.multiply_fisher_factor_unweighted(vector) + + def multiply_ggn_factor_transpose_unweighted( + self, + vector: Sequence[Array], + ) -> Array: + return self.multiply_fisher_factor_transpose_unweighted(vector) + + def multiply_ggn_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array, ...]: + return self.multiply_fisher_factor_replicated_one_hot_unweighted(index) + + @property + def ggn_factor_inner_shape(self) -> Shape: + return self.fisher_factor_inner_shape + + +class DistributionNegativeLogProbLoss(NegativeLogProbLoss): + """Negative log-probability loss that uses a Distrax distribution.""" + + @property + @abc.abstractmethod + def dist(self) -> distrax.Distribution: + """The underlying Distrax distribution.""" + + def _evaluate(self, targets: Array) -> Array: + # keeps leading dims intact + return -self.dist.log_prob(targets) # pytype: disable=bad-return-type + + def sample(self, rng: PRNGKey) -> Array: + return self.dist.sample(seed=rng) # pytype: disable=bad-return-type + + +@jax.tree_util.register_pytree_node_class +class NormalMeanNegativeLogProbLoss(DistributionNegativeLogProbLoss, + NaturalParamsNegativeLogProbLoss): + """Loss log prob loss for a normal distribution parameterized by a mean vector. + + Note that the covariance is treated as the identity divided by 2. + Also note that the Fisher for such a normal distribution with respect the mean + parameter is given by: + + F = (1 / variance) * I + """ + + def __init__( + self, + mean: Array, + targets: Array | None = None, + variance: Numeric = 0.5, + weight: Numeric = 1.0, + normalize_log_prob: bool = True, + ): + """Initializes the loss instance. + + Args: + mean: The mean of the normal distribution. + targets: Optional targets to use for evaluation. + variance: The scalar variance of the normal distribution. + weight: The weight of the loss. + normalize_log_prob: Whether the log prob should include the standard + normalization constant for Gaussians (which is additive and depends + on the variance). + """ + + if not isinstance(variance, (int, float)) and type(variance) is not object: # pylint: disable=unidiomatic-typecheck + if not isinstance(variance, Array) or variance.size > 1: + raise ValueError("`variance` must be either a python scalar or a " + "scalar array.") + self._mean = mean + self._targets = targets + self._variance = variance + self._normalize_log_prob = normalize_log_prob + + super().__init__(weight=weight) + + @property + def mean(self) -> Array: + return self._mean + + @property + def variance(self) -> Numeric: + return self._variance + + @property + def targets(self) -> Array | None: + return self._targets + + @property + def normalize_log_prob(self) -> bool: + return self._normalize_log_prob + + @property + def parameter_independants(self) -> tuple[Numeric | None, ...]: + return self._targets, self._variance, self._weight + + @property + def dist(self) -> distrax.MultivariateNormalDiag: + scale_diag = jnp.full_like(self.mean, jnp.sqrt(self.variance)) + return distrax.MultivariateNormalDiag(loc=self.mean, scale_diag=scale_diag) + + @property + def params(self) -> tuple[Array]: + return (self.mean,) + + @property + def fisher_factor_inner_shape(self) -> Shape: + return self._mean.shape + + def _evaluate(self, targets: Array) -> Array: + + if self.normalize_log_prob: + return super()._evaluate(targets) + else: + # keeps leading dims intact + return 0.5 * jnp.sum(jnp.square( + self.mean - targets), axis=range(1, targets.ndim)) / self.variance + + def multiply_fisher_unweighted( + self, + vector: Sequence[Array] + ) -> tuple[Array]: + return (vector[0] / self.variance,) + + def multiply_fisher_factor_unweighted( + self, + vector: Array, + ) -> tuple[Array]: + return (vector / jnp.sqrt(self.variance),) + + def multiply_fisher_factor_transpose_unweighted( + self, + vector: Sequence[Array], + ) -> Array: + # it's symmetric + return self.multiply_fisher_factor_unweighted(vector[0])[0] + + def multiply_fisher_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array]: + + ones_slice = jnp.ones([self.mean.shape[0]] + [1] * (self.mean.ndim - 1)) + + output_slice = ones_slice / jnp.sqrt(self.variance) + + return (insert_slice_in_zeros(output_slice, self.mean.shape, (0,) + index),) + + +# TODO(jamesmartens): This class was copied from the TF K-FAC codebase and is +# untested with probable bugs. Test it? +@jax.tree_util.register_pytree_node_class +class NormalMeanVarianceNegativeLogProbLoss(DistributionNegativeLogProbLoss): + """Negative log prob loss for a normal distribution with mean and variance. + + This class parameterizes a multivariate normal distribution with n independent + dimensions. Unlike :class:`~NormalMeanNegativeLogProbLoss`, this class does + not assume the variance is held constant. The Fisher Information for n = 1 is + given by: + + F = [[1 / variance, 0], + [ 0, 0.5 / variance^2]] + + where the parameters of the distribution are concatenated into a single + vector as ``[mean, variance]``. For n > 1, the mean parameter vector is + concatenated with the variance parameter vector. For further details checkout + the Wikipedia `page + `__. + """ + + def __init__( + self, + mean: Array, + variance: Array, + targets: Array | None = None, + weight: Numeric = 1.0, + ): + """Initializes the loss instance. + + Args: + mean: The mean of the normal distribution. + variance: The variance of the normal distribution. + targets: Optional targets to use for evaluation. + weight: The weight of the loss. + """ + if mean.ndim != 2: + raise ValueError("Only 2D mean array is supported.") + if variance.ndim != 2: + raise ValueError("Only 2D variance array is supported.") + self._mean = mean + self._variance = variance + self._targets = targets + super().__init__(weight=weight) + + @property + def targets(self) -> Array | None: + return self._targets + + @property + def parameter_independants(self) -> tuple[Numeric | None, ...]: + return self._targets, self._weight + + @property + def dist(self) -> distrax.MultivariateNormalDiag: + return distrax.MultivariateNormalDiag( + loc=self._mean, scale_diag=jnp.sqrt(self._variance)) + + @property + def params(self) -> tuple[Array, Array]: + return self._mean, self._variance + + @property + def _fisher_mean(self) -> Array: + """The Fisher w.r.t. to the mean parameters.""" + return 1. / self._variance + + @property + def _fisher_mean_factor(self) -> Array: + """The Fisher factor w.r.t. to the mean parameters.""" + return jnp.sqrt(self._fisher_mean) + + @property + def _fisher_var(self) -> Array: + """The Fisher w.r.t. to the variance parameters.""" + return 1. / (2 * jnp.square(self._variance)) + + @property + def _fisher_var_factor(self) -> Array: + """The Fisher factor w.r.t. to the variance parameters.""" + return 1. / (jnp.sqrt(2.) * self._variance) + + def multiply_fisher_unweighted( + self, + vector: Sequence[Array], + ) -> tuple[Array, Array]: + mean_vec, var_vec = vector + return self._fisher_mean * mean_vec, self._fisher_var * var_vec + + def multiply_fisher_factor_unweighted( + self, + vector: Array, + ) -> tuple[Array, Array]: + mean_vec, var_vec = jnp.split(vector, 2, axis=-1) + result_mean_vec = self._fisher_mean_factor * mean_vec + result_var_vec = self._fisher_var_factor * var_vec + return result_mean_vec, result_var_vec + + def multiply_fisher_factor_transpose_unweighted( + self, + vector: Sequence[Array], + ) -> Array: + mean_vec, var_vec = vector + result_mean_vec = self._fisher_mean_factor * mean_vec + result_var_vec = self._fisher_var_factor * var_vec + return jnp.concatenate([result_mean_vec, result_var_vec], axis=-1) + + def multiply_fisher_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array, Array]: + [index] = index + + if index < int(self._mean.shape[-1]): + # Index corresponds to mean parameter. + mean_slice = self._fisher_mean_factor[:, index][..., None] + mean_output = insert_slice_in_zeros( + mean_slice, self._mean.shape, [0, index]) + var_output = jnp.zeros_like(mean_output) + + else: + index -= int(self._mean.shape[-1]) + # Index corresponds to variance parameter. + var_slice = self._fisher_var_factor[:, index][..., None] + var_output = insert_slice_in_zeros( + var_slice, self._variance.shape, [0, index]) + mean_output = jnp.zeros_like(var_output) + + return mean_output, var_output + + @property + def fisher_factor_inner_shape(self) -> Shape: + return self._mean.shape[:-1] + (self._mean.shape[-1] * 2,) + + def multiply_ggn_unweighted( + self, + vector: Sequence[Array], + ) -> tuple[Array, ...]: + raise NotImplementedError() + + def multiply_ggn_factor_unweighted( + self, vector: Array + ) -> tuple[Array, ...]: + raise NotImplementedError() + + def multiply_ggn_factor_transpose_unweighted( + self, + vector: Sequence[Array], + ) -> Array: + raise NotImplementedError() + + def multiply_ggn_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array, ...]: + raise NotImplementedError() + + @property + def ggn_factor_inner_shape(self) -> Shape: + raise NotImplementedError() + + +@jax.tree_util.register_pytree_node_class +class MultiBernoulliNegativeLogProbLoss(DistributionNegativeLogProbLoss, + NaturalParamsNegativeLogProbLoss): + """Negative log prob loss for multiple Bernoulli distributions parametrized by logits. + + Represents N independent Bernoulli distributions where N = len(logits). Its + Fisher Information matrix is given by ``F = diag(p * (1-p))``, where + ``p = sigmoid(logits)``. + + As F is diagonal with positive entries, its factor B is + ``B = diag(sqrt(p * (1-p)))``. + """ + + def __init__( + self, + logits: Array, + targets: Array | None = None, + mask: Array | None = None, + weight: Numeric = 1.0, + ): + """Initializes the loss instance. + + Args: + logits: The logits of the Bernoulli distribution. + targets: Optional targets to use for evaluation. + mask: Optional mask to apply to losses. Should be 0/1-valued and of + shape ``logits.shape``. The tensors returned by ``evaluate`` and + ``grad_of_evaluate``, as well as the various matrix vector products, + will be multiplied by mask. + weight: The weight of the loss. + """ + if (mask is not None and type(mask) is not object and # pylint: disable=unidiomatic-typecheck + mask.shape != logits.shape): + raise ValueError("If provided, mask.shape must be equal to " + "logits.shape.") + + self._logits = logits + self._targets = targets + self._mask = mask + + super().__init__(weight=weight) + + @property + def targets(self) -> Array | None: + return self._targets + + @property + def mask(self) -> Array | None: + return self._mask + + @property + def parameter_independants(self) -> tuple[Numeric | None, ...]: + return self._targets, self._mask, self._weight + + @property + def dist(self) -> distrax.Bernoulli: + return distrax.Bernoulli(logits=self._logits, dtype=jnp.int32) + + def _evaluate(self, targets: Array) -> Array: + + evl = super()._evaluate(targets) + + if self.mask is not None: + return evl * self.mask + else: + return evl + + @property + def _probs(self) -> Array: + """The probabilities of the underlying Bernoulli distribution.""" + if self.mask is not None: + return self.dist.probs * self.mask + else: + return self.dist.probs # pytype: disable=bad-return-type + + @property + def params(self) -> tuple[Array]: + return (self._logits,) + + @property + def fisher_factor_inner_shape(self) -> Shape: + return self._logits.shape + + def multiply_fisher_unweighted( + self, + vector: Sequence[Array] + ) -> tuple[Array]: + return (self._probs * (1 - self._probs) * vector[0],) + + def multiply_fisher_factor_unweighted( + self, + vector: Array + ) -> tuple[Array]: + return (utils.stable_sqrt(self._probs * (1 - self._probs)) * vector,) + + def multiply_fisher_factor_transpose_unweighted( + self, + vector: Sequence[Array] + ) -> Array: + # it's symmetric in this case + return self.multiply_fisher_factor_unweighted(vector[0])[0] + + def multiply_fisher_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array]: + + probs_slice = jnp.expand_dims(self._probs[(slice(None),) + index], + axis=range(1, len(self._probs.shape))) + + output_slice = utils.stable_sqrt(probs_slice * (1 - probs_slice)) + + return (insert_slice_in_zeros(output_slice, self._logits.shape, + (0,) + index),) + + +@jax.tree_util.register_pytree_node_class +class CategoricalLogitsNegativeLogProbLoss(DistributionNegativeLogProbLoss, + NaturalParamsNegativeLogProbLoss): + """Negative log prob loss for a categorical distribution parameterized by logits. + + + Note that the Fisher (for a single case) of a categorical distribution, with + respect to the natural parameters (i.e. the logits), is given by + ``F = diag(p) - p*p^T``, where ``p = softmax(logits)``. F can be factorized as + ``F = B * B^T``, where ``B = diag(q) - p*q^T`` and ``q`` is the entry-wise + square root of ``p``. This is easy to verify using the fact that ``q^T*q = 1`` + . + """ + + def __init__( + self, + logits: Array, + targets: Array | None = None, + mask: Array | None = None, + weight: Numeric = 1.0, + ): + """Initializes the loss instance. + + Args: + logits: The logits of the Categorical distribution. + targets: Optional targets to use for evaluation, which specify an integer + index of the correct class. Must be of shape ``logits.shape[:-1]``. + mask: Optional mask to apply to losses over the batch. Should be + 0/1-valued and of shape ``logits.shape[:-1]``. The tensors returned by + ``evaluate`` and ``grad_of_evaluate``, as well as the various matrix + vector products, will be multiplied by mask (with broadcasting to later + dimensions). + weight: The weight of the loss. + """ + if (mask is not None and type(mask) is not object and # pylint: disable=unidiomatic-typecheck + mask.shape != logits.shape[:-1]): + raise ValueError("If provided, mask.shape must be equal to " + "logits.shape[:-1].") + + self._logits = logits + self._targets = targets + self._mask = mask + + super().__init__(weight=weight) + + @property + def targets(self) -> Array | None: + return self._targets + + @property + def mask(self) -> Array | None: + return self._mask + + @property + def parameter_independants(self) -> tuple[Numeric | None, ...]: + return self._targets, self._mask, self._weight + + @property + def dist(self) -> distrax.Categorical: + return distrax.Categorical(logits=self._logits, dtype=jnp.int32) + + def _evaluate(self, targets: Array) -> Array: + + evl = super()._evaluate(targets) + + if self.mask is not None: + return evl * self.mask + else: + return evl + + @property + def _probs(self) -> Array: + """The probabilities of the underlying Categorical distribution.""" + + if self.mask is not None: + return self.dist.probs * self.mask[..., None] + else: + return self.dist.probs + + @property + def _sqrt_probs(self) -> Array: + """The square root of ``self.probs``.""" + + if self.mask is not None: + return utils.stable_sqrt(self.dist.probs) * self.mask[..., None] + else: + return utils.stable_sqrt(self.dist.probs) + + @property + def params(self) -> tuple[Array]: + return (self._logits,) + + @property + def fisher_factor_inner_shape(self) -> Shape: + return self._logits.shape + + def multiply_fisher_unweighted( + self, + vector: Sequence[Array] + ) -> tuple[Array]: + + assert len(vector) == 1 + + probs = self._probs + + fisher_product = vector[0] * probs - probs * jnp.sum( + vector[0] * probs, axis=-1, keepdims=True) + + return (fisher_product,) + + def multiply_fisher_factor_unweighted( + self, + vector: Array + ) -> tuple[Array]: + + probs = self._probs + + sqrt_probs = self._sqrt_probs + + return (sqrt_probs * vector - probs * jnp.sum( + sqrt_probs * vector, axis=-1, keepdims=True),) + + def multiply_fisher_factor_transpose_unweighted( + self, + vector: Sequence[Array] + ) -> Array: + + assert len(vector) == 1 + + probs = self._probs + + sqrt_probs = self._sqrt_probs + + return sqrt_probs * vector[0] - sqrt_probs * jnp.sum( + probs * vector[0], axis=-1, keepdims=True) + + def multiply_fisher_factor_replicated_one_hot_unweighted( + self, + index: tuple[int, ...], + ) -> tuple[Array]: + + probs = self._probs + + sqrt_probs_slice = jnp.expand_dims(self._sqrt_probs[(slice(None),) + index], + axis=range(1, len(probs.shape))) + + padded_slice = insert_slice_in_zeros( + sqrt_probs_slice, probs.shape, (0,) + index) + + return (padded_slice - probs * sqrt_probs_slice,) + + +@jax.tree_util.register_pytree_node_class +class OneHotCategoricalLogitsNegativeLogProbLoss( + CategoricalLogitsNegativeLogProbLoss): + """Neg log prob loss for a categorical distribution with one-hot targets. + + Identical to CategoricalLogitsNegativeLogProbLoss except that the underlying + distribution is OneHotCategorical as opposed to Categorical. ``targets`` is + vector-encoded instead of integer-encoded, and must have the same shape as + ``logits``. + """ + + @property + def dist(self) -> distrax.OneHotCategorical: + return distrax.OneHotCategorical(logits=self._logits, dtype=jnp.int32) + + +def insert_slice_in_zeros( + slice_to_insert: Array, + zeros_shape: Sequence[int], + position: Sequence[int], +) -> Array: + """Inserts slice into a larger array of zeros. + + Forms a new array of shape ``zeros_shape``, which is zeros everywhere except + for the slice given by the ``position`` argument. + + We assume that slice_to_insert.shape and zeros_shape are the same length, and + with ``slice_to_insert.shape[i] == zeros_shape[i]`` or ``1`` for all ``i``. + For ``i`` where slice_to_insert.shape[i] == 1, ``position[i]`` must be ``0``. + + Args: + slice_to_insert: The slice to insert. + zeros_shape: The shape of the new array. + position: The position of ``slice_to_insert`` in the new tensor. + + Returns: + The new array. + + Raises: + ValueError: If the slice's shape at the given dim is not 1. + """ + + assert slice_to_insert.ndim == len(zeros_shape) + assert slice_to_insert.ndim == len(position) + + pad_width = [] + + for i in range(slice_to_insert.ndim): + if slice_to_insert.shape[i] == 1: + pad_width.append((position[i], zeros_shape[i] - position[i] - 1)) + else: + assert slice_to_insert.shape[i] == zeros_shape[i] + assert position[i] == 0 + pad_width.append((0, 0)) + + return jnp.pad(slice_to_insert, pad_width) + + +def register_normal_predictive_distribution( + mean: Array, + targets: Array | None = None, + variance: float = 0.5, + weight: Numeric = 1.0, + normalize_log_prob: bool = True, +) -> None: + """Registers a normal predictive distribution. + + This corresponds to a squared error loss of the form + ``weight/(2*var) * jnp.sum((targets - mean)**2) / batch_size``. + + NOTE: this function assumes you are *not* averaging over non-batch dimensions + when computing the loss. e.g. if dimension 0 were the batch dimension, this + corresponds to + ``jnp.mean(jnp.sum((target - prediction)**2, + axis=range(1,target.ndims)), axis=0)`` + and not + ``jnp.mean((target - prediction)**2)``. + If your loss is of the latter form you can compensate for it by passing the + appropriate value to ``weight``. + + Args: + mean: An ND array defining the mean vector of the distribution. The first + dimension will usually be the batch size, but doesn't need to be (unless + using ``estimation_mode='fisher_exact'`` or + ``estimation_mode='ggn_exact'`` in the optimizer/estimator). + targets: (OPTIONAL) The targets for the loss function. Must have the same + shape as ``mean``. Only required if using + ``estimation_mode='fisher_empirical'`` in the optimizer/estimator. + (Default: None) + variance: The variance of the distribution. Must be a constant scalar, + independent of the network's parameters. Note that the default value of + 0.5 corresponds to a standard squared error loss + ``weight * jnp.sum((target - prediction)**2)``. If you want your squared + error loss to be of the form + ``0.5*coeff*jnp.sum((target - prediction)**2)`` you should use + variance=1.0. (Default: 0.5) + weight: A constant scalar coefficient that the log prob loss associated with + this distribution is multiplied by. In general this is NOT equivalent to + changing the temperature of the distribution, but in the case of normal + distributions it may be. Note that this must be constant and independent + of the network's parameters. (Default: 1.0) + normalize_log_prob: Whether the negative log prob loss associated to this + this distribution should include the additive normalization constant + (which is constant and depends on ``variance``) that makes it a true log + prob, and not just a squared error loss. Note that this has no effect on + the behavior of optimizer with the exception of in niche situations where + the loss value is computed from the registrations. e.g., when + ``include_registered_loss_in_stats=True`` is used. (Default: True) + """ + args, args_names = filter_none( + mean=mean, + targets=targets, + variance=variance, + weight=weight, + normalize_log_prob=normalize_log_prob, + ) + + tags.loss_tag.bind( + *args, + meta=tags.LossMetaData( + loss_class=NormalMeanNegativeLogProbLoss, + parameter_dependants=args_names[:1], + parameter_independants=args_names[1:], + argument_names=tuple(args_names), + ) + ) + + +def register_squared_error_loss( + prediction: Array, + targets: Array | None = None, + weight: Numeric = 1.0, +) -> None: + """Registers a squared error loss function. + + This assumes a squared error loss of the form + ``weight * jnp.sum((targets - prediction)**2) / batch_size``. + + If your loss uses a coefficient of 0.5 you need to set the ``weight`` argument + to reflect this. + + NOTE: this function assumes you are *not* averaging over non-batch dimensions + when computing the loss. e.g. if dimension 0 were the batch dimension, this + corresponds to + ``jnp.mean(jnp.sum((target - prediction)**2, + axis=range(1, target.ndims)), axis=0)`` + and not + ``jnp.mean((target - prediction)**2)`` + If your loss is of the latter form you can compensate for it by passing the + appropriate value to ``weight``. + + NOTE: even though ``prediction`` and ``targets`` are interchangeable in the + definition of the squared error loss, they are not interchangeable in this + function. ``prediction`` must be the output of your parameterized function + (e.g. neural network), and ``targets`` must not depend on the parameters. + Mixing the two up could lead to a silent failure of the curvature estimation. + + Args: + prediction: The prediction made by the network (i.e. its output) as an ND + array of floats. The first dimension will usually be the batch size, but + doesn't need to be (unless using ``estimation_mode='fisher_exact'`` or + ``estimation_mode='ggn_exact'`` in the optimizer/estimator). + targets: (OPTIONAL) The targets for the loss function. Must have the same + shape as ``prediction``. Only required if using + ``estimation_mode='fisher_empirical'`` in the optimizer/estimator. + (Default: None) + weight: The constant scalar coefficient which this loss is multiplied by. + Note that this must be constant and independent of the network's + parameters. (Default: 1.0) + """ + register_normal_predictive_distribution( + mean=prediction, + targets=targets, + variance=0.5, + weight=weight, + normalize_log_prob=False, + ) + + +def register_multi_bernoulli_predictive_distribution( + logits: Array, + targets: Array | None = None, + mask: Array | None = None, + weight: Numeric = 1.0, +) -> None: + """Registers a multi-Bernoulli predictive distribution. + + This corresponds to a sigmoid cross-entropy loss of the form + ``weight * jnp.sum(sigmoid_cross_entropy(logits, targets)) / batch_size``. + + NOTE: this function assumes you are *not* averaging over non-batch dimensions + when computing the loss. e.g. if dimension 0 were the batch dimension, this + corresponds to + ``jnp.mean(jnp.sum(sigmoid_cross_entropy(logits, targets), + axis=range(1, target.ndims)), axis=0)`` + and not + ``jnp.mean(sigmoid_cross_entropy(logits, targets))`` + If your loss is of the latter form you can compensate for it by passing the + appropriate value to ``weight``. + + NOTE: this is distinct from + :func:`~register_categorical_predictive_distribution` and should not be + confused with it. + + Args: + logits: The logits of the distribution (i.e. its parameters) as a ND array + of floats. The first dimension will usually be the batch size, but doesn't + need to be (unless using ``estimation_mode='fisher_exact'`` or + ``estimation_mode='ggn_exact'`` in the optimizer/estimator). + targets: (OPTIONAL) The targets for the loss function. Must be of the same + shape as ``logits``. Only required if using + ``estimation_mode='fisher_empirical'`` in the optimizer/estimator. + (Default: None) + mask: (OPTIONAL) Mask to apply to log probabilities generated by the + distribution. Should be 0/1-valued and of shape ``logits.shape``. + Log probabilities corresponding to mask values of 0 will be treated + as constant and equal to 0. (Default: None) + weight: The constant scalar coefficient that the log prob loss associated + with this distribution is multiplied by. This is NOT equivalent to + changing the temperature of the distribution since we don't renormalize + the log prob in the objective function. Note that this must be constant + and independent of the network's parameters. (Default: 1.0) + """ + args, args_names = filter_none( + logits=logits, + targets=targets, + mask=mask, + weight=weight, + ) + + tags.loss_tag.bind( + *args, + meta=tags.LossMetaData( + loss_class=MultiBernoulliNegativeLogProbLoss, + parameter_dependants=args_names[:1], + parameter_independants=args_names[1:], + argument_names=tuple(args_names), + ) + ) + + +def register_sigmoid_cross_entropy_loss( + logits: Array, + targets: Array | None = None, + mask: Array | None = None, + weight: Numeric = 1.0, +) -> None: + """Registers a sigmoid cross-entropy loss function. + + This assumes a sigmoid cross-entropy loss of the form + ``weight * jnp.sum(sigmoid_cross_entropy(logits, targets)) / batch_size``. + + NOTE: this function assumes you are *not* averaging over non-batch dimensions + when computing the loss. e.g. if dimension 0 were the batch dimension, this + corresponds to + ``jnp.mean(jnp.sum(sigmoid_cross_entropy(logits, targets), + axis=range(1, target.ndims)), axis=0)`` + and not + ``jnp.mean(sigmoid_cross_entropy(logits, targets))`` + If your loss is of the latter form you can compensate for this by passing the + appropriate value to ``weight``. + + NOTE: this function is distinct from + :func:`~register_softmax_cross_entropy_loss` and should not be confused with + it. It is similar to :func:`~register_multi_bernoulli_predictive_distribution` + but without the explicit probabilistic interpretation. It behaves identically + for now. + + Args: + logits: The input logits of the loss as a ND array of floats. The first + dimension will usually be the batch size, but doesn't need to be (unless + using ``estimation_mode='fisher_exact'`` or + ``estimation_mode='ggn_exact'`` in the optimizer/estimator). + targets: (OPTIONAL) The targets for the loss function. Must be of the same + shape as ``logits``. Only required if using + ``estimation_mode='fisher_empirical'`` in the optimizer/estimator. + (Default: None) + mask: (OPTIONAL) Mask to apply to losses. Should be 0/1-valued and of shape + ``logits.shape``. Losses corresponding to mask values of 0 will be + treated as constant and equal to 0. (Default: None) + weight: The constant scalar coefficient which this loss is multiplied by. + Note that this must be constant and independent of the network's + parameters. (Default: 1.0) + """ + register_multi_bernoulli_predictive_distribution( + logits=logits, + targets=targets, + mask=mask, + weight=weight, + ) + + +def register_categorical_predictive_distribution( + logits: Array, + targets: Array | None = None, + mask: Array | None = None, + weight: Numeric = 1.0, +) -> None: + """Registers a categorical predictive distribution. + + This corresponds to a softmax cross-entropy loss of the form + + ``weight * jnp.sum(softmax_cross_entropy(logits, targets)) / batch_size``, + + or in other words, the negative log probability of the distribution, + multiplied by ``weight``. + + NOTE: this is distinct from + :func:`~register_multi_bernoulli_predictive_distribution` and should not be + confused with it. + + Args: + logits: The logits of the distribution (i.e. its parameters) as an ND array + of floats. The first dimension will usually be the batch size, but doesn't + need to be (unless using ``estimation_mode='fisher_exact'`` or + ``estimation_mode='ggn_exact'`` in the optimizer/estimator). The final + dimension is the one over which the softmax is computed. + targets: (OPTIONAL) The values at which the log probability of this + distribution is evaluated (to give the loss). Must be a (N-1)D array of + integers with shape ``logits.shape[:-1]`` for integer-encoded targets, or + ``logits.shape`` for vector-encoded targets. Only required if using + ``estimation_mode='fisher_empirical'`` in the optimizer/estimator. + (Default: None) + mask: (OPTIONAL) Mask to apply to log probabilities generated by the + distribution. Should be 0/1-valued and of shape ``logits.shape[:-1]``. + Log probabilities corresponding to mask values of 0 will be treated + as constant and equal to 0. (Default: None) + weight: The constant scalar coefficient that the log prob loss associated + with this distribution is multiplied by. This is NOT equivalent to + changing the temperature of the distribution since we don't renormalize + the log prob in the objective function. Note that this must be constant + and independent of the network's parameters. (Default: 1.0) + """ + if targets is not None: + + if targets.ndim == logits.ndim: + loss_class = OneHotCategoricalLogitsNegativeLogProbLoss + + elif targets.ndim == logits.ndim - 1: + loss_class = CategoricalLogitsNegativeLogProbLoss + + else: + raise ValueError(f"The logits ndim is {logits.ndim} and the targets ndim " + f"must be either equal or one less than it, but is " + f"{targets.ndim}.") + + else: + loss_class = CategoricalLogitsNegativeLogProbLoss + + args, args_names = filter_none( + logits=logits, + targets=targets, + mask=mask, + weight=weight, + ) + + tags.loss_tag.bind( + *args, + meta=tags.LossMetaData( + loss_class=loss_class, + parameter_dependants=args_names[:1], + parameter_independants=args_names[1:], + argument_names=tuple(args_names), + ), + ) + + +def register_softmax_cross_entropy_loss( + logits: Array, + targets: Array | None = None, + mask: Array | None = None, + weight: Numeric = 1.0, +) -> None: + """Registers a softmax cross-entropy loss function. + + This assumes a softmax cross-entropy loss of the form + + ``weight * jnp.sum(softmax_cross_entropy(logits, targets)) / batch_size``. + + NOTE:this is distinct from :func:`~register_sigmoid_cross_entropy_loss` and + should not be confused with it. It is similar to + :func:`~register_categorical_predictive_distribution` but without the explicit + probabilistic interpretation. It behaves identically for now. + + Args: + logits: The input logits of the loss as an ND array of floats. The first + dimension will usually be the batch size, but doesn't need to be (unless + using ``estimation_mode='fisher_exact'`` or + ``estimation_mode='ggn_exact'`` in the optimizer/estimator). + The final dimension is the one over which the softmax is computed. + targets: (OPTIONAL) The targets for the loss function. Must be a (N-1)D + array of integers with shape ``logits.shape[:-1]`` for integer-encoded + targets, or ``logits.shape`` for vector-encoded targets. Only required if + using ``estimation_mode='fisher_empirical'`` in the optimizer/estimator. + (Default: None) + mask: (OPTIONAL) Mask to apply to losses. Should be 0/1-valued and of shape + ``logits.shape[:-1]``. Losses corresponding to mask values of 0 will be + treated as constant and equal to 0. (Default: None) + weight: The constant scalar coefficient which this loss is multiplied by. + Note that this must be constant and independent of the network's + parameters. (Default: 1.0) + """ + register_categorical_predictive_distribution( + logits=logits, + targets=targets, + mask=mask, + weight=weight, + ) diff --git a/src/kfac_jax/_src/optimizer.py b/src/kfac_jax/_src/optimizer.py new file mode 100644 index 0000000000000000000000000000000000000000..0c7f23a1ad3d11c70fda0b9503110c18ab9d68bb --- /dev/null +++ b/src/kfac_jax/_src/optimizer.py @@ -0,0 +1,2147 @@ +# Modifications copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""The kfac_jax optimizer (supporting K-FAC and other methods).""" + +import functools +from typing import Any, Callable, Generic, Iterator, Sequence +from absl import logging + +import jax +from jax import lax +import jax.numpy as jnp +from kfac_jax._src import curvature_estimator +from kfac_jax._src import utils +from typing_extensions import Self + + +# Types for annotation +Array = utils.Array +PRNGKey = utils.PRNGKey +Numeric = utils.Numeric +Params = utils.Params +Batch = utils.Batch +FuncState = Any +FuncAux = utils.FuncAux +Scalar = utils.Scalar +ScheduleType = utils.ScheduleType + +FuncArgsVariants = ( + tuple[Params, Batch] | + tuple[Params, FuncState, Batch] | + tuple[Params, PRNGKey, Batch] | + tuple[Params, FuncState, PRNGKey, Batch] +) +FuncOutputs = ( + Array | + tuple[Array, FuncState] | + tuple[Array, FuncAux] | + tuple[Array, tuple[FuncState, FuncAux]] +) +ValueFunc = Callable[..., FuncOutputs] +ValueAndGradFunc = Callable[..., tuple[FuncOutputs, Params]] +SharedForwardFunc = Callable[..., tuple[Array, Array]] +BlockDiagonalCurvature = curvature_estimator.BlockDiagonalCurvature + +ReturnEither = ( + tuple[Params, "Optimizer.State", FuncState, dict[str, Numeric]] | + tuple[Params, "Optimizer.State", dict[str, Numeric]] +) + +QuadModelParams = tuple[Array, Array, Array, Array] +# The quadratic model is given as +# Q(w) = w^T V^T (C + damping * I + reg * L) V w / 2.0 + w^T V^T g +# where (n - number of vectors, d - dimensions of each vector): +# damping - the damping value at the current iteration +# reg - the L2 regularization coefficient +# w (n,) - the vector of free weights (learning rate and momentum) +# V (d, n) - the matrix of proposed vectors for each weight +# C (d, d) - the curvature matrix (GGN/Fisher/Hessian) +# L (d, d) - the L2 regularization matrix. L is diagonal, with 1 on diagonal +# if the corresponding parameter is L2 regularised, and 0 +# otherwise. +# g (d,) - the gradient +# +# In QuadModelParams, we have the tuple (A, D, R, b) where: +# A = V^T C V +# D = V^T I V (for damping) +# R = V^T L V (for L2 regularization) +# b = V^T g +# +# See Optimizer._solve_quad_model for how these are used, and +# Optimizer._compute_exact_quad_model for how they are computed. + + +# Various lists of parameters that are biases and norms, to be +# used for registering parameters that are excluded from l2 regularization +# in Optimizer. +# "b" and "bias" for biases, "scale" for RMSNorm and LayerNorm, and +# "offset" for LayerNorm. + +HAIKU_BIASES = "b,bias" +HAIKU_BIASES_AND_NORMS = "b,bias,scale,offset" + + +class Optimizer(utils.WithStagedMethods): + """The kfac_jax optimizer (supporting K-FAC and other methods).""" + + @utils.register_state_class + class State(Generic[Params], utils.State): + r"""Persistent state of the optimizer. + + Attributes: + velocities: The update to the parameters from the previous step - + :math:`\theta_t - \theta_{t-1}`. + estimator_state: The persistent state for the curvature estimator. + damping: When using damping adaptation, this will contain the current + value. + data_seen: The number of training cases that the optimizer has processed. + step_counter: An integer giving the current step number :math:`t`. + """ + velocities: Params + estimator_state: BlockDiagonalCurvature.State + damping: Array + data_seen: Numeric + step_counter: Numeric + + @classmethod + def from_dict(cls, dict_representation: dict[str, Any]) -> Self: + dict_representation["estimator_state"] = ( + BlockDiagonalCurvature.State.from_dict( + dict_representation["estimator_state"] + ) + ) + return cls(**dict_representation) + + def __init__( + self, + value_and_grad_func: ValueAndGradFunc, + l2_reg: Numeric, + regularized_parameters_path_exclusions: str = "", + value_func_has_aux: bool = False, + value_func_has_state: bool = False, + value_func_has_rng: bool = False, + value_func_for_estimator: ValueFunc | None = None, + use_adaptive_learning_rate: bool = False, + learning_rate_schedule: ScheduleType | None = None, + use_adaptive_momentum: bool = False, + momentum_schedule: ScheduleType | None = None, + use_adaptive_damping: bool = False, + damping_schedule: ScheduleType | None = None, + initial_damping: Numeric | None = None, + use_initial_damping_calibration: bool = False, + min_damping: Numeric = 1e-8, + max_damping: Numeric = jnp.inf, + include_damping_in_quad_change: bool = False, + damping_adaptation_interval: int = 5, + damping_adaptation_decay: Numeric = 0.9, + damping_lower_threshold: Numeric = 0.25, + damping_upper_threshold: Numeric = 0.75, + always_use_exact_qmodel_for_damping_adjustment: bool = False, + precon_damping_mult: Numeric = 1.0, + precon_damping_schedule: ScheduleType | None = None, + use_step_rejection: bool = False, + reject_damping_increase_factor: float = 1.0, + norm_constraint: Numeric | None = None, + num_burnin_steps: int = 10, + estimation_mode: str | None = None, + custom_estimator_ctor: ( + Callable[..., BlockDiagonalCurvature] | None) = None, + curvature_ema: Numeric = 0.95, + curvature_update_period: int = 1, + inverse_update_period: int = 5, + use_exact_inverses: bool = False, + batch_process_func: Callable[[Batch], Batch] | None = None, + register_only_generic: bool = False, + patterns_to_skip: Sequence[str] = (), + use_automatic_registration: bool = True, + auto_register_kwargs: dict[str, Any] | None = None, + layer_tag_to_block_ctor: ( + dict[str, curvature_estimator.CurvatureBlockCtor] | None) = None, + multi_device: bool = False, + debug: bool = False, + invalid_metric_value: Numeric = jnp.nan, + batch_size_extractor: Callable[ + [Batch], Numeric + ] = utils.default_batch_size_extractor, + pmap_axis_name: str = "batch_axis", + forbid_setting_attributes_after_finalize: bool = True, + modifiable_attribute_exceptions: Sequence[str] = (), + include_norms_in_stats: bool = False, + include_per_param_norms_in_stats: bool = False, + include_registered_loss_in_stats: bool = False, + distributed_precon_apply: bool = True, + distributed_inverses: bool = True, + num_estimator_samples: int = 1, + should_vmap_estimator_samples: bool = False, + norm_to_scale_identity_weight_per_block: str | None = None, + step_stats_hook: Callable[..., dict[str, Array]] | None = None, + precon_power: Scalar = -1.0, + exact_quad_model_matrix_type: str | None = None, + value_func_for_shared_forward: SharedForwardFunc | None = None, + share_curvature_and_grad_forward: bool = False, + ): + """Initializes the kfac_jax optimizer with the provided settings. + + NOTE: Please read the docstring for this constructor carefully. Especially + the description of ``value_and_grad_func``. + + A note on the "damping" parameter: + + One of the main complications of using second-order optimizers like K-FAC is + the "damping" parameter. This parameter is multiplied by the identity matrix + and (approximately) added to the curvature matrix (i.e. the Fisher or GGN) + before it is inverted and multiplied by the gradient when computing the + update (before any learning rate scaling). The damping should follow the + scale of the objective, so that if you multiply your loss by some factor you + should do the same for the damping. Roughly speaking, larger damping values + constrain the update vector to a smaller region around zero, which is needed + in general since the second-order approximations that underlie second-order + methods can break down for large updates. (In gradient descent the learning + rate plays an analogous role.) The relationship between the damping + parameter and the radius of this region is complicated and depends on the + scale of the objective amongst other things. + + The optimizer provides a system for adjusting the damping automatically via + the ``use_adaptive_damping`` argument, although this system is not reliable, + especially for highly stochastic objectives. Using a fixed value or a + manually tuned schedule can work as good or better for some problems, while + it can be a very poor choice for others (like deep autoencoders). + Empirically we have found that using a fixed value works well enough for + common architectures like convnets and transformers. + + Args: + value_and_grad_func: Python callable. This function should return the + value of the loss to be optimized and its gradients, and optionally the + model state and auxiliary information in the form of a a dict mapping + strings to scalar arrays (usually statistics to log). Note that it + should *not* be jitted/pmapped or otherwise compiled by JAX, as this can + lead to errors. (Compilation is done internally by the optimizer.) The + interface of this function should be: ``out_args, loss_grads = + value_and_grad_func(*in_args)``. Here, ``in_args`` is ``(params, + func_state, rng, batch)``, with ``rng`` omitted if + ``value_func_has_rng`` is ``False``, and with ``func_state`` omitted if + ``value_func_has_state`` is ``False``. Meanwhile, ``out_args`` is + ``(loss, (func_state, aux))`` if ``value_func_has_state`` and + ``value_func_has_aux`` are both ``True``, ``(loss, func_state)`` if + ``value_func_has_state`` is ``True`` and ``value_func_has_aux`` is + ``False``, ``(loss, aux)`` if ``value_func_has_state`` is ``False`` and + ``value_func_has_aux`` is ``True``, and finally ``loss`` if + ``value_func_has_state`` and ``value_func_has_aux`` are both ``False``. + This should be consistent with how JAX's ``value_and_grad`` API function + is typically used. Note that the value (and its gradient) should be + normalized by the batch size, as is standard convention. Additional + normalization, such as by the sequence length, is up to the user, but + must by properly reported in the loss registration (by setting the + ``weight`` arguments in the loss registration functions.) + l2_reg: Scalar. Set this value to tell the optimizer what L2 + regularization coefficient you are using (if any). Note the coefficient + appears in the regularizer as ``coeff / 2 * sum(param**2)``. This adds + an additional diagonal term to the curvature and hence will affect the + quadratic model when using adaptive damping. Note that the user is still + responsible for adding regularization to the loss. + regularized_parameters_path_exclusions: str. A comma-separated list + specifying the names of parameters that should not be regularized. + A number of convenience examples are given in this module, e.g. + HAIKU_BIASES_AND_NORMS, which is ``"b,bias,scale,offset"``. + (Default: ``""``) + value_func_has_aux: Boolean. Specifies whether the provided callable + ``value_and_grad_func`` returns auxiliary data. (Default: ``False``) + value_func_has_state: Boolean. Specifies whether the provided callable + ``value_and_grad_func`` has a persistent state that is passed in and + out. (Default: ``False``) + value_func_has_rng: Boolean. Specifies whether the provided callable + ``value_and_grad_func`` additionally takes as input an rng key. + (Default: ``False``) + value_func_for_estimator: ValueFunc. If specified, this function will be + used by the preconditioner estimator instead of ``value_and_grad_func``. + This is useful for cases where the value function used for training is + expensive to add to the preconditioner, e.g. because it has costly + regularizers. (Default: ``None``) + value_func_for_shared_forward: Tagged function returning + ``(loss, gradient_surrogate)``. The surrogate's parameter gradient + must equal the training gradient. Required when + ``share_curvature_and_grad_forward=True``. + use_adaptive_learning_rate: Boolean. Specifies whether to use the special + rule from the original K-FAC paper for picking the learning rate at each + step. Note that this won't work well for stochastic objectives. If this + is ``False``, the user must use the ``learning_rate`` argument of the + step function, or the constructor argument ``learning_rate_schedule``. + (Default: ``False``) + learning_rate_schedule: Callable. A schedule for the learning rate. This + should take as input the current step number, and optionally the amount + of data seen so far as a keyword argument ``data_seen``, and return a + single array that represents the learning rate. (Default: ``None``) + use_adaptive_momentum: Boolean. Specifies whether to use the special rule + from the original K-FAC paper for picking the momentum "decay" parameter + at each step. Note that this won't work well for stochastic objectives. + If this is ``False``, the user must use the ``momentum`` argument of the + step function, or the constructor argument ``momentum_schedule``. + (Default: ``False``) + momentum_schedule: Callable. A schedule for the momentum parameter. This + should take as input the current step number, and optionally the amount + of data seen so far as a keyword argument ``data_seen``, and return a + single array that represents the momentum. (Default: ``None``) + use_adaptive_damping: Boolean. Specifies whether the optimizer will use + the Levenberg-Marquardt method to automatically adjust the damping every + ``damping_adaptation_interval`` iterations. If this is set to ``False`` + the user must provide a value to the damping argument of the step + function at each iteration, or use the ``damping_schedule`` constructor + argument. Note that the effectiveness of this technique seems to vary + between problems. (Default: ``False``) + damping_schedule: Callable. A schedule for the damping. This should take + as input the current step number, and optionally the amount of data seen + so far as a keyword argument ``data_seen``, and return a single array + that represents the learning rate. (Default: ``None``) + initial_damping: Scalar or None. This specifies the initial value of the + damping that the optimizer will use when using automatic damping + adaptation. (Default: ``None``) + use_initial_damping_calibration: Boolean. If ``True``, the initial damping + value, used to initialize the adaptive damping method, will be first + calibrated (after any burnin steps to estimate the preconditioner) so + that its value wouldn't be changed after the first step of optimization. + This calibration is done by essentially running the step function + multiple times without actually updating the parameters or sampling a + new mini-batch. ``num_burnin_steps`` must be greater than 0 to use this + option. (Default: ``False``) + min_damping: Scalar. Minimum value the damping parameter can take when + using automatic damping adaptation. Note that the default value of 1e-8 + is quite arbitrary, and you may have to adjust this up or down for your + particular problem. If you are using a non-zero value of l2_reg you + *may* be able to set this to zero. (Default: ``1e-8``) + max_damping: Scalar. Maximum value the damping parameter can take when + using automatic damping adaptation. (Default: ``Infinity``) + include_damping_in_quad_change: Boolean. Whether to include the + contribution of the damping in the quadratic model for the purposes + computing the reduction ration ("rho") in the Levenberg-Marquardt scheme + used for adapting the damping. Note that the contribution from the + ``l2_reg`` argument is always included. (Default: ``False``) + damping_adaptation_interval: Int. The number of steps in between adapting + the damping parameter. (Default: ``5``) + damping_adaptation_decay: Scalar. The damping parameter will be adjusted + up or down by ``damping_adaptation_decay ** + damping_adaptation_interval``, or remain unchanged, every + ``damping_adaptation_interval`` number of iterations. (Default: ``0.9``) + damping_lower_threshold: Scalar. The damping parameter is increased if the + reduction ratio is below this threshold. (Default: ``0.25``) + damping_upper_threshold: Scalar. The damping parameter is decreased if the + reduction ratio is below this threshold. (Default: ``0.75``) + always_use_exact_qmodel_for_damping_adjustment: Boolean. When using + learning rate and/or momentum adaptation, the quadratic model change + used for damping adaption is always computed using the exact curvature + matrix. Otherwise, there is an option to use either the exact or + approximate curvature matrix to compute the quadratic model change, + which is what this argument controls. When True, the exact curvature + matrix will be used, which is more expensive, but could possibly produce + a better damping schedule. (Default: ``False``) + precon_damping_mult: Scalar. When ``precon_damping_schedule`` is unset, + the regular damping is used for the preconditioner damping, multiplied + by this value. (Default: ``1.0``) + precon_damping_schedule: Similar to ``damping_schedule``, but for the + preconditioner only. If ``None``, the preconditioner will use the + regular damping, multiplied by ``precon_damping_mult``. + (Default: ``None``) + use_step_rejection: Whether or not to reject the step whenever the loss + on the current batch goes up after the update. This option offers + robustness at the cost of doing more work per step (unless adaptive + damping with Levenberg-Marquardt is used). (Default: ``False``) + reject_damping_increase_factor: The damping parameter is increased by this + factor if the step is rejected. (Default: ``1.0``) + norm_constraint: Scalar. If specified, the update is scaled down so that + its approximate squared Fisher norm ``v^T F v`` is at most the specified + value. (Note that here ``F`` is the approximate curvature matrix, not + the exact.) May only be used when ``use_adaptive_learning_rate`` is + ``False``. (Default: ``None``) + num_burnin_steps: Int. At the start of optimization, e.g. the first step, + before performing the actual step the optimizer will perform this many + times updates to the curvature approximation without updating the actual + parameters. (Default: ``10``) + estimation_mode: String. The type of estimator to use for the curvature + matrix. See the documentation for :class:`~BlockDiagonalCurvature` for a + detailed description of the possible options. If ``None`` will use + default estimation_mode mode of the used CurvatureEstimator subclass, + which is typically "ggn_curvature_prop". (Default: ``None``) + custom_estimator_ctor: Optional constructor for subclass of + :class:`~BlockDiagonalCurvature`. If specified, the optimizer will use + this conastructor instead of the default + :class:`~BlockDiagonalCurvature`. (Default: ``None``) + curvature_ema: The decay factor used when calculating the covariance + estimate moving averages. (Default: ``0.95``) + curvature_update_period: Int. The number of steps in between updating the + the curvature estimates. (Default: ``1``) + inverse_update_period: Int. The number of steps in between updating the + the computation of the inverse curvature approximation. (Default: ``5``) + use_exact_inverses: Bool. If ``True``, preconditioner inverses are + computed "exactly" without the pi-adjusted factored damping approach. + Note that this involves the use of eigendecompositions, which can + sometimes be much more expensive. (Default: ``False``) + batch_process_func: Callable. A function which to be called on each batch + before feeding to the KFAC on device. This could be useful for specific + device input optimizations. (Default: ``None``) + register_only_generic: Boolean. Whether when running the auto-tagger to + register only generic parameters, or allow it to use the graph matcher + to automatically pick up any kind of layer tags. (Default: ``False``) + patterns_to_skip: tuple. A list of any patterns that should be skipped by + the graph matcher when auto-tagging. (Default: ``()``) + use_automatic_registration: Bool. If ``True``, the optimizer will try to + automatically register the layers of your network. (Default: ``True``) + auto_register_kwargs: Any additional kwargs to be passed down to + :func:`~auto_register_tags`, which is called by the curvature estimator. + (Default: ``None``) + layer_tag_to_block_ctor: dictionary. A mapping from layer tags to block + classes which to override the default choices of block approximation for + that specific tag. See the documentation for + :class:`~CurvatureEstimator` for a more detailed description. (Default: + ``None``) + multi_device: Boolean. Whether to use pmap and run the optimizer on + multiple devices. (Default: ``False``) + debug: Boolean. If neither the step or init functions should be jitted. + Note that this also overrides ``multi_device`` and prevents using pmap, + instead using a "simulated pmap" that loops over the device index and + does everything on the default device. (Default: ``False``) + invalid_metric_value: Numeric. Certain metrics returned from the step + function are not always computed at each iteration, or may otherwise + be invalid. In such cases we need to return a value anyway. jnp.nan is + a natural choice, but can sometimes cause problems (e.g. false positives + JAX's automatic NaN checker). This argument allows the user to specify a + different value to return in such cases. (Default: ``jnp.nan``) + batch_size_extractor: A function that takes as input the function + arguments and returns the batch size for a single device. (Default: + ``kfac.utils.default_batch_size_extractor``) + pmap_axis_name: String. The name of the pmap axis to use when + ``multi_device`` is set to True. (Default: ``batch_axis``) + forbid_setting_attributes_after_finalize: Boolean. By default, after the + object is finalized, you can not set any of its properties. This is done + in order to protect the user from making changes to the object + attributes that would not be picked up by various internal methods after + they have been compiled. However, if you are extending this class, and + clearly understand the risks of modifying attributes, setting this to + ``False`` will remove the restriction. (Default: ``True``) + modifiable_attribute_exceptions: Sequence of strings. Gives a list of + names for attributes that can be modified after finalization even when + ``forbid_setting_attributes_after_finalize`` is ``True``. (Default: + ``()``) + include_norms_in_stats: Boolean. It True, the vector norms of the + gradient, preconditioned gradient, and parameter update are included in + the statistics returned by the step function. (Default: ``False``) + include_per_param_norms_in_stats: Boolean. It True, the per-parameter + vector norms of the gradient, preconditioned gradient, and parameter + update are included in the statistics returned by the step function. + (Default: ``False``) + include_registered_loss_in_stats: Boolean. If True, we include the loss, + as computed from the registered losses, in the stats. Also included is + the relative difference between this as the loss computed from + ``value_and_grad_func``. This is useful for debugging registration + errors. Note this for this option to work it's required that the targets + are passed for each loss function registration. (Default: ``False``) + distributed_precon_apply: Boolean. Whether to distribute the application + of the preconditioner across the different devices in a layer-wise + fashion. If False, each device will (redundantly) perform the required + operations for all the layers. (Default: True) + distributed_inverses: Boolean. Whether to distribute the inverse + computations (required to compute the preconditioner) across the + different devices in a layer-wise fashion. If False, each device will + (redundantly) perform the required computations for all the layers. + (Default: True) + num_estimator_samples: Number of samples (per case) to use when computing + stochastic curvature matrix estimates. This option is only used when + ``estimation_mode == 'fisher_gradients'`` or ``estimation_mode == + '[fisher,ggn]_curvature_prop'``. (Default: 1) + should_vmap_estimator_samples: Whether to use ``jax.vmap`` to compute + samples when ``num_estimator_samples > 1``. (Default: False) + share_curvature_and_grad_forward: Reuse the exact-Fisher tagged model + primal evaluation for the ordinary loss and training gradient on + curvature-update steps. (Default: ``False``) + norm_to_scale_identity_weight_per_block: The name of a norm to use to + compute extra per-block scaling for the damping. See psd_matrix_norm() + in utils/math.py for the definition of these. Note that this will not + affect the exact quadratic model that is used as part of the "adaptive" + learning rate, momentum, and damping methods. (Default: None) + step_stats_hook: Optional callable ``(estimator, grads, + preconditioned_gradient) -> dict`` invoked inside ``_step`` on the + PRE-norm-constraint preconditioned gradient; returned scalars are + merged into the step stats dict. Runs inside the step's jit — must + be trace-safe and cheap. (Default: None) + precon_power: The matrix power to use when computing the preconditioner. + K-FAC use -1 by default, but ``kfac_jax`` can simulate other optimizers + like RMSProp by using -0.5 (along with appropriate changes to + ``layer_tag_to_block_ctor`` and ``estimation_mode``). (Default: -1) + exact_quad_model_matrix_type: The type of matrix to use when computing the + exact quadratic model (used in the adaptive learning rate and momentum). + Can be ``'fisher'``, ``'ggn'``, or None. If None, will use the value + implied by ``estimation_mode``. (Default: None) + """ + + super().__init__( + multi_device=multi_device, + pmap_axis_name=pmap_axis_name if multi_device else None, + debug=debug, + forbid_setting_attributes_after_finalize= + forbid_setting_attributes_after_finalize, + excluded_attribute_names=modifiable_attribute_exceptions, + ) + + if use_adaptive_damping and initial_damping is None: + raise ValueError("When use_adaptive_damping is True you must provide a " + "value for initial_damping.") + if use_adaptive_learning_rate and learning_rate_schedule is not None: + raise ValueError("If you are using adaptive learning rate then " + "`learning_rate_schedule` should be None.") + if use_adaptive_momentum and momentum_schedule is not None: + raise ValueError("If you are using adaptive momentum then " + "`momentum_schedule` should be None.") + if use_adaptive_damping and damping_schedule is not None: + raise ValueError("If you are using adaptive damping then " + "`damping_schedule` should be None.") + + if num_burnin_steps <= 0 and use_initial_damping_calibration: + raise ValueError("num_burnin_steps must be > 0 if " + "use_initial_damping_calibration is True.") + + self._value_and_grad_func = value_and_grad_func + self._value_func_has_aux = value_func_has_aux + self._value_func_has_state = value_func_has_state + self._value_func_has_rng = value_func_has_rng + if share_curvature_and_grad_forward: + incompatible = [] + if estimation_mode != "fisher_exact": + incompatible.append("estimation_mode must be 'fisher_exact'") + if value_func_for_estimator is not None: + incompatible.append("value_func_for_estimator must be None") + if value_func_for_shared_forward is None: + incompatible.append("value_func_for_shared_forward must be provided") + if custom_estimator_ctor is not None: + incompatible.append("custom_estimator_ctor must be None") + if value_func_has_aux: + incompatible.append("value_func_has_aux must be False") + if value_func_has_state: + incompatible.append("value_func_has_state must be False") + if value_func_has_rng: + incompatible.append("value_func_has_rng must be False") + if include_registered_loss_in_stats: + incompatible.append( + "include_registered_loss_in_stats must be False" + ) + if incompatible: + raise ValueError( + "`share_curvature_and_grad_forward=True` is incompatible with: " + + "; ".join(incompatible) + + "." + ) + self._share_curvature_and_grad_forward = ( + share_curvature_and_grad_forward + ) + self._value_func: ValueFunc = convert_value_and_grad_to_value_func( + value_and_grad_func, + has_aux=value_func_has_aux or value_func_has_state, + ) + + self._l2_reg = l2_reg + self._regularized_parameters_path_exclusions = ( + regularized_parameters_path_exclusions.split(",")) + + self._use_adaptive_learning_rate = use_adaptive_learning_rate + self._learning_rate_schedule = learning_rate_schedule + self._use_adaptive_momentum = use_adaptive_momentum + self._momentum_schedule = momentum_schedule + + self._use_adaptive_damping = use_adaptive_damping + self._damping_schedule = damping_schedule + self._initial_damping = initial_damping + self._use_initial_damping_calibration = use_initial_damping_calibration + self._min_damping = min_damping + self._max_damping = max_damping + self._include_damping_in_quad_change = include_damping_in_quad_change + self._damping_adaptation_decay = damping_adaptation_decay + self._damping_adaptation_interval = damping_adaptation_interval + self._damping_lower_threshold = damping_lower_threshold + self._damping_upper_threshold = damping_upper_threshold + self._always_use_exact_qmodel_for_damping_adjustment = ( + always_use_exact_qmodel_for_damping_adjustment) + self._precon_damping_mult = precon_damping_mult + self._precon_damping_schedule = precon_damping_schedule + + self._use_step_rejection = use_step_rejection + self._reject_damping_increase_factor = reject_damping_increase_factor + + self._norm_constraint = norm_constraint + self._num_burnin_steps = num_burnin_steps + self._curvature_ema = curvature_ema + if curvature_update_period > inverse_update_period: + raise ValueError( + "curvature_update_period ({}) cannot be larger than" + " inverse_update_period ({}) as the identical matrix inversion would" + " be redundantly performed. Set inverse_update_period larger instead." + .format(curvature_update_period, inverse_update_period) + ) + self._curvature_update_period = curvature_update_period + self._inverse_update_period = inverse_update_period + self._layer_tag_to_block_cls = layer_tag_to_block_ctor + self._patterns_to_skip = patterns_to_skip + self._batch_process_func = batch_process_func or (lambda x: x) + self._include_norms_in_stats = include_norms_in_stats + self._include_per_param_norms_in_stats = include_per_param_norms_in_stats + self._include_registered_loss_in_stats = include_registered_loss_in_stats + self._batch_size_extractor = batch_size_extractor + + self.__invalid_metric_value = invalid_metric_value + + self._use_cached_inverses = (self._inverse_update_period != 1) + self._use_exact_inverses = use_exact_inverses + + self._norm_to_scale_identity_weight_per_block = ( + norm_to_scale_identity_weight_per_block + ) + + # Optional ``hook(estimator, grads, preconditioned_gradient) -> dict`` + # called inside ``_step`` on the PRE-norm-constraint preconditioned + # gradient; the returned scalars are merged into the step stats. + # Runs inside the step's jit — implementations must be trace-safe + # and cheap (reductions only). Lets clients log per-block/per-family + # update-allocation diagnostics without subclassing ``_step``. + self._step_stats_hook = step_stats_hook + + self._precon_power = precon_power + + self._exact_quad_model_matrix_type = exact_quad_model_matrix_type + + self._params_index = 0 + batch_index = int(value_func_has_state + value_func_has_rng + 1) + + if (norm_to_scale_identity_weight_per_block is not None + and norm_to_scale_identity_weight_per_block != "none"): + + assert (not use_adaptive_learning_rate and not use_adaptive_momentum + and not use_adaptive_damping) # not currently supported + + estimator_ctor = (custom_estimator_ctor or BlockDiagonalCurvature) + + auto_register_kwargs = auto_register_kwargs or {} + auto_register_kwargs.update(dict( + register_only_generic=register_only_generic, + patterns_to_skip=patterns_to_skip, + )) + + if value_func_for_estimator is None: + # The reason we pass value_and_grad_func to the estimator here and not + # value_func, is that the latter is usually produced using the primal + # computation which is part of jax.grad. For whatever reason, when JAX + # takes the gradient of this, it produces a slightly different graph than + # if we apply jax.grad directly to the original function. This then makes + # it impossible for XLA to merge the two computations, defeating the + # purpose of the estimation modes "fisher_empirical_direct[_synced]". + func_and_grad_for_estimator = convert_value_and_grad_to_clean_value_and_grad( # pylint: disable=line-too-long + value_and_grad_func, + has_aux=value_func_has_aux or value_func_has_state, + ) + + else: + func_and_grad_for_estimator = None + + estimator_extra_kwargs = {} + if share_curvature_and_grad_forward: + estimator_extra_kwargs["shared_forward_value_func"] = ( + value_func_for_shared_forward + ) + + # Curvature estimator + self._estimator = estimator_ctor( + func=value_func_for_estimator, + func_and_grad=func_and_grad_for_estimator, + default_estimation_mode=estimation_mode, + params_index=self._params_index, + batch_index=batch_index, + layer_tag_to_block_ctor=layer_tag_to_block_ctor, + distributed_multiplies=distributed_precon_apply, + distributed_cache_updates=distributed_inverses, + num_samples=num_estimator_samples, + should_vmap_samples=should_vmap_estimator_samples, + auto_register_tags=use_automatic_registration, + auto_register_kwargs=auto_register_kwargs, + **estimator_extra_kwargs, + ) + self._implicit = curvature_estimator.ImplicitExactCurvature( + self._value_func, + params_index=self._params_index, + batch_size_extractor=batch_size_extractor, + ) + + # Each subclass should call finalize on its own, so this gets called only + # for instances of exactly this class type. + if type(self) == Optimizer: # pylint: disable=unidiomatic-typecheck + self.finalize() + + @property + def _invalid_metric_value(self) -> Array: + return jnp.array(self.__invalid_metric_value, dtype=float) + + @property + def _damping_decay_factor(self) -> Numeric: + """How fast to decay the damping, when using damping adaptation.""" + return self._damping_adaptation_decay ** self._damping_adaptation_interval + + @property + def _exact_powers_to_cache(self) -> Numeric | Sequence[Numeric] | None: + if self._use_exact_inverses and self._use_cached_inverses: + return self._precon_power + else: + return None + + @property + def _approx_powers_to_cache(self) -> Numeric | Sequence[Numeric] | None: + if not self._use_exact_inverses and self._use_cached_inverses: + return self._precon_power + else: + return None + + @property + def _mat_type_for_exact_quad_model(self) -> str: + if self._exact_quad_model_matrix_type is None: + return self._estimator.default_mat_type + return self._exact_quad_model_matrix_type + + def _should_update_damping(self, step_counter: int) -> bool: + """Whether at the current step the optimizer should update the damping.""" + return ((step_counter + 1) % self._damping_adaptation_interval == 0) and ( + self._use_adaptive_damping + ) + + def _should_update_estimate_curvature(self, step_counter: int) -> bool: + """Whether at the current step the optimizer should update the curvature estimates.""" + return step_counter % self._curvature_update_period == 0 + + def _should_update_inverse_cache( + self, + state: State, + inverse_update_period: Numeric | None = None, + ) -> Array | bool: + """Whether at the current step the optimizer should update the inverse curvature approximation.""" + period = (self._inverse_update_period if inverse_update_period is None + else inverse_update_period) + return self._use_cached_inverses and ( + state.step_counter % period == 0) + + def _should_sync_estimator( + self, + state: State, + inverse_update_period: Numeric | None = None, + ) -> Array | bool: + """Whether at the current step the optimizer should update the inverse curvature approximation.""" + + if self._use_cached_inverses: + return self._should_update_inverse_cache(state, inverse_update_period) + + return True + + def set_live_hparams( + self, + *, + curvature_ema: Numeric | None = None, + curvature_update_period: int | None = None, + inverse_update_period: int | None = None, + ) -> None: + """Updates cadence/EMA hyperparameters on a live optimizer instance. + + All three take effect from the next call to ``step`` without triggering + recompilation: ``curvature_update_period`` only enters Python-side + executable selection, while ``curvature_ema`` and ``inverse_update_period`` + are threaded into the compiled step function as runtime scalars. + + ``inverse_update_period`` cannot be changed on an optimizer constructed + with ``inverse_update_period=1``, since that construction permanently + disables the inverse cache in the estimator state. + """ + new_curv = (self._curvature_update_period if curvature_update_period is None + else int(curvature_update_period)) + new_inv = (self._inverse_update_period if inverse_update_period is None + else int(inverse_update_period)) + if new_curv < 1 or new_inv < 1: + raise ValueError("Update periods must be positive integers.") + if new_inv != self._inverse_update_period and not self._use_cached_inverses: + raise ValueError( + "Cannot change inverse_update_period on an optimizer constructed " + "with inverse_update_period=1 (inverse cache disabled).") + if new_curv > new_inv: + raise ValueError( + "curvature_update_period ({}) cannot be larger than" + " inverse_update_period ({}).".format(new_curv, new_inv)) + if curvature_ema is not None and not 0.0 <= float(curvature_ema) <= 1.0: + raise ValueError("curvature_ema must be in [0, 1].") + self.unlock_attributes() + try: + self._curvature_update_period = new_curv + self._inverse_update_period = new_inv + if curvature_ema is not None: + self._curvature_ema = float(curvature_ema) + finally: + self.lock_attributes() + + def _live_step_scalars(self) -> tuple[Array, Array]: + """Current curvature_ema / inverse_update_period as traced step args. + + Plain rank-0 arrays, matching how callers pass learning_rate / + momentum / damping into ``step``: the staging layer broadcasts them, + so no per-device replication is required here (and + ``device_put_replicated`` no longer exists on modern JAX anyway). + """ + ema = jnp.asarray(self._curvature_ema, dtype=jnp.float32) + period = jnp.asarray(self._inverse_update_period, dtype=jnp.int32) + return ema, period + + @functools.partial(utils.staged, static_argnums=1) + def _rng_split(self, rng: PRNGKey, num: int) -> tuple[Array, ...]: + """Splits the ``rng`` key.""" + return tuple(jax.random.split(rng, num)) + + @utils.auto_scope_method + def _compute_loss_value(self, func_args: FuncArgsVariants) -> Array: + """Computes the value of the loss function being optimized.""" + return self._value_func(*func_args) + + def _verify_args_and_get_step_counter( + self, + step_counter: Array, + learning_rate: Array | None = None, + momentum: Array | None = None, + damping: Array | None = None, + global_step_int: int | None = None, + ) -> int: + """Verifies that the arguments passed to the step function are correct.""" + + # Verify correct arguments invocation + if self._use_adaptive_learning_rate and learning_rate is not None: + raise ValueError("When use_adaptive_learning_rate is set to True you " + "should not pass a value to the step function.") + + elif not self._use_adaptive_learning_rate and ( + self._learning_rate_schedule is None and learning_rate is None): + raise ValueError("When `use_adaptive_learning_rate` is set to False and " + "`learning_rate_schedule` is None you must provide a " + "value to the step function.") + + elif self._learning_rate_schedule is not None and learning_rate is not None: + raise ValueError("When you have passed a `learning_rate_schedule` you " + "should not pass a value to the step function.") + + if self._use_adaptive_momentum and momentum is not None: + raise ValueError("When `use_adaptive_momentum` is set to True you " + "should not pass a value to the step function.") + + elif not self._use_adaptive_momentum and ( + self._momentum_schedule is None and momentum is None): + raise ValueError("When `use_adaptive_momentum` is set to False and " + "`momentum_schedule` is None you must provide a value to" + " the step function.") + + elif self._momentum_schedule is not None and momentum is not None: + raise ValueError("When you have passed a `momentum_schedule` you should " + "not pass a value to the step function.") + + if self._use_adaptive_damping and damping is not None: + raise ValueError("When `use_adaptive_damping` is set to True you " + "should not pass a value to the step function.") + + elif not self._use_adaptive_damping and ( + self._damping_schedule is None and damping is None): + raise ValueError("When `use_adaptive_damping` is set to False and " + "`damping_schedule` is None you must provide a value to " + "the step function.") + + elif self._damping_schedule is not None and damping is not None: + raise ValueError("When you have passed a `damping_schedule` you should " + "not pass a value to the step function.") + + if global_step_int is None: + return int(self.get_first(step_counter)) + + return global_step_int + + @utils.staged + def _setup_state_and_schedules( + self, + learning_rate: Array | None, + momentum: Array | None, + damping: Array | None, + step_counter: Array, + data_seen: Array, + ) -> tuple[Numeric | None, Numeric | None, Numeric, Numeric]: + """Helper function for setting up learning rate, momentum and damping.""" + + # Compute schedules if applicable + if self._learning_rate_schedule is not None: + + assert learning_rate is None + learning_rate = utils.call_func_with_conditional_kwargs( + self._learning_rate_schedule, step_counter, data_seen=data_seen) + + if self._momentum_schedule is not None: + assert momentum is None + momentum = utils.call_func_with_conditional_kwargs( + self._momentum_schedule, step_counter, data_seen=data_seen) + + if self._damping_schedule is not None: + assert damping is None + damping = utils.call_func_with_conditional_kwargs( + self._damping_schedule, step_counter, data_seen=data_seen) + + else: + assert damping is not None + + if self._precon_damping_schedule is not None: + precon_damping = utils.call_func_with_conditional_kwargs( + self._precon_damping_schedule, step_counter, data_seen=data_seen) + else: + precon_damping = damping * self._precon_damping_mult + + return learning_rate, momentum, damping, precon_damping + + def _setup_func_args_and_rng( + self, + params: Params, + rng: PRNGKey, + batch: Batch, + func_state: FuncState | None, + ) -> tuple[FuncArgsVariants, Array]: + """Helper function for setting up the model function arguments correctly.""" + + # Preprocess the batch and construct correctly the function arguments + batch = self._batch_process_func(batch) + + # Correctly split rng + if self._value_func_has_rng: + rng, func_rng = jax.random.split(rng) + else: + func_rng = None + + # Make the function args + func_args = make_func_args( + params=params, + func_state=func_state, + rng=func_rng, + batch=batch, + has_state=self._value_func_has_state, + has_rng=self._value_func_has_rng, + ) + + return func_args, rng + + def _update_estimator_curvature( + self, + estimator_state: BlockDiagonalCurvature.State, + func_args: FuncArgsVariants, + rng: PRNGKey, + ema_old: Numeric, + ema_new: Numeric, + precon_damping: Numeric, + sync: Array | bool = True + ) -> BlockDiagonalCurvature.State: + """Updates the curvature estimator state.""" + + state = self._estimator.update_curvature_matrix_estimate( + state=estimator_state, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=self._l2_reg + precon_damping, + # Note that the batch is always the last entry of FuncArgsVariantsdef + batch_size=self._batch_size_extractor(func_args[-1]), + rng=rng, + func_args=func_args, + pmap_axis_name=self.pmap_axis_name, + ) + return jax.lax.cond( + sync, + functools.partial(self._estimator.sync, + pmap_axis_name=self.pmap_axis_name), + lambda state_: state_, + state, + ) + + def _update_estimator_curvature_and_value_and_grad( + self, + estimator_state: BlockDiagonalCurvature.State, + func_args: FuncArgsVariants, + rng: PRNGKey, + ema_old: Numeric, + ema_new: Numeric, + precon_damping: Numeric, + sync: Array | bool = True, + ) -> tuple[BlockDiagonalCurvature.State, Array, Params]: + """Updates exact-Fisher curvature and returns its shared loss/gradient.""" + + state, loss, grads = ( + self._estimator.update_curvature_matrix_estimate_and_value_and_grad( + state=estimator_state, + ema_old=ema_old, + ema_new=ema_new, + identity_weight=self._l2_reg + precon_damping, + batch_size=self._batch_size_extractor(func_args[-1]), + rng=rng, + func_args=func_args, + pmap_axis_name=self.pmap_axis_name, + ) + ) + state = jax.lax.cond( + sync, + functools.partial( + self._estimator.sync, + pmap_axis_name=self.pmap_axis_name, + ), + lambda state_: state_, + state, + ) + return state, loss, grads + + @utils.auto_scope_method + def _compute_loss_and_grads( + self, + func_args: FuncArgsVariants, + state: State | None = None, + ) -> tuple[Array, Params, FuncState | None, FuncAux | None]: + """Computes the model loss value and its gradients.""" + + del state + + out, grads = self._value_and_grad_func(*func_args) + + loss, func_state, aux = extract_func_outputs( + out, self._value_func_has_aux, self._value_func_has_state) + + if self._include_registered_loss_in_stats: + aux = aux or {} + aux["loss_registered"] = self._compute_loss_from_registrations(func_args) + + return loss, grads, func_state, aux + + @functools.partial(utils.staged, donate_argnums=0) + def _maybe_update_inverse_cache( + self, + state: State, + precon_damping: Array, + inverse_update_period: Array, + ) -> State: + """Updates the estimator state cache if it is the right iteration.""" + + # Copy this first since we mutate it later in this function. + state = state.copy() + + state.estimator_state = lax.cond( + self._should_update_inverse_cache(state, inverse_update_period), + functools.partial( + self._estimator.update_cache, + identity_weight=self._l2_reg + precon_damping, + exact_powers=self._exact_powers_to_cache, + approx_powers=self._approx_powers_to_cache, + eigenvalues=False, + pmap_axis_name=self.pmap_axis_name, + ), + lambda state_: state_, + state.estimator_state, + ) + + return state + + @functools.partial(utils.staged, static_argnums=3) + def _compute_preconditioned_gradient( + self, + state: State, + grads: Params, + precon_damping: Array, + can_distribute: bool = True, + ) -> Params: + """Computes the preconditioned gradient.""" + + return self._estimator.multiply_matpower( + state=state.estimator_state, + parameter_structured_vector=grads, + identity_weight=self._l2_reg + precon_damping, + power=self._precon_power, + exact_power=self._use_exact_inverses, + use_cached=self._use_cached_inverses, + pmap_axis_name=self.pmap_axis_name if can_distribute else None, + norm_to_scale_identity_weight_per_block=self._norm_to_scale_identity_weight_per_block, + ) + + @utils.staged + def _maybe_apply_norm_constraint( + self, grads: Params, preconditioned_grads: Params, coefficient: Array + ) -> tuple[Params, Params | None]: + """Scales precon grad to have curvature-weighted norm <= norm_constraint.""" + + if self._norm_constraint is None: + return preconditioned_grads, None + + assert not self._use_adaptive_learning_rate + + sq_norm_grads = utils.inner_product(preconditioned_grads, grads) + sq_norm_scaled_grads = sq_norm_grads * coefficient ** 2 + + max_coefficient = jnp.sqrt(self._norm_constraint / sq_norm_scaled_grads) + coefficient = jnp.minimum(max_coefficient, 1) + + precon_grad = utils.scalar_mul(preconditioned_grads, coefficient) + + return precon_grad, sq_norm_scaled_grads + + def _compute_quad_change_for_damping_adapt( + self, + state: State, + delta: Params, + grads: Params, + damping: Array, + func_args: FuncArgsVariants, + ) -> Array: + """The quadratic model change, when lr and momentum are non-adaptive.""" + + assert not (self._use_adaptive_learning_rate or self._use_adaptive_momentum) + + if self._always_use_exact_qmodel_for_damping_adjustment: + quad_model = self._compute_exact_quad_model_filtered( + [delta], grads, func_args, state=state) + else: + quad_model = self._compute_approx_quad_model(state, [delta], grads) + + w = jnp.ones([]) + return self._solve_quad_model(quad_model, damping, [w])[1] + + def _coefficients_and_quad_change( + self, + state: State, + vectors: Sequence[Params], + grads: Params, + learning_rate: Numeric | None, + momentum: Numeric | None, + damping: Numeric, + func_args: FuncArgsVariants, + should_update_damping: bool, + ) -> tuple[tuple[Numeric, Numeric], Numeric]: + """The correct update coefficients and corresponding quadratic change.""" + + # Compute the coefficients of the update vectors + # The learning rate is defined as the negative of the coefficient by which + # we multiply the gradients, while the momentum is the coefficient by + # which we multiply the velocities. + neg_learning_rate = -learning_rate if learning_rate is not None else None + fixed_coefficients = (neg_learning_rate, momentum) + + if self._use_adaptive_learning_rate or self._use_adaptive_momentum: + + assert fixed_coefficients[0] is None or fixed_coefficients[1] is None + + quad_model = self._compute_exact_quad_model_filtered( + vectors, grads, func_args, state=state, + fixed_coefficients=fixed_coefficients) + + return self._solve_quad_model(quad_model, damping, fixed_coefficients) + + else: + + assert all(c is not None for c in fixed_coefficients) + fixed_coefficients: tuple[Numeric, Numeric] + + if should_update_damping: + + delta = self._weighted_sum_of_objects(vectors, fixed_coefficients) + + quad_change = self._compute_quad_change_for_damping_adapt( + state, delta, grads, damping, func_args) + + else: + quad_change = self._invalid_metric_value + + return fixed_coefficients, quad_change + + @utils.staged + def _compute_loss_from_registrations( + self, + func_args: FuncArgsVariants + ) -> Array: + + loss = self._estimator.compute_func_from_registered( + func_args, self._batch_size_extractor(func_args[-1])) + + if self._l2_reg > 0.0: + + l2_reg_val = self._l2_reg / 2 * utils.squared_norm( + func_args[self._params_index]) + + loss += l2_reg_val + + return loss + + @utils.staged + def _init( + self, + params: Params, + rng: PRNGKey, + batch: Batch, + func_state: FuncState | None = None, + ) -> State: + """A staged function to initialize the optimizer state .""" + + # Note that we can reuse the rng in the func_args construction below, as + # these are just dummy values used to perform the tracing. + + return Optimizer.State( + velocities=jax.tree_util.tree_map(jnp.zeros_like, params), + estimator_state=self._estimator.init( + rng=rng, + func_args=make_func_args( + params=params, + func_state=func_state, + rng=rng, + batch=self._batch_process_func(batch), + has_state=self._value_func_has_state, + has_rng=self._value_func_has_rng, + ), + exact_powers_to_cache=self._exact_powers_to_cache, + approx_powers_to_cache=self._approx_powers_to_cache, + cache_eigenvalues=False + ), + damping=jnp.array( + (self._initial_damping if self._initial_damping is not None + else -1e10), dtype=float), + data_seen=jnp.array(0, dtype=int), + step_counter=jnp.array(0, dtype=int) + ) + + def init( + self, + params: Params, + rng: PRNGKey, + batch: Batch, + func_state: FuncState | None = None, + ) -> State: + """Initializes the optimizer and returns the appropriate optimizer state. + + NOTE: please do not jit/pmap or otherwise compile this function with JAX, + as this can lead to errors. Compilation is handled internally by the + optimizer. + + NOTE: when ``multi_device`` is ``True``, all of the JAX array arguments to + this function (including arrays inside of trees), should have an extra + leading axis the size of the number of local devices. + + Args: + params: Example models parameters (used for tracing and shape info). + rng: A Jax PRNG key. Unlike the ``rng`` in the step function, should be + the same for each host and for each slice in the leading axis (i.e. + corresponding to devices) when ``multi_device`` is ``True``. + batch: An example batch of the same size as the one passed to ``step`` + (or returned from the ``data_iterator``). Used for tracing and shape + info. + func_state: Example function state (used for tracing and shape info). + + Returns: + The initialized optimizer state. + """ + + if not self.finalized: + self.finalize(params, rng, batch, func_state) + + # Check that mask_out_unregularized_params works as intended. + _ = self._maybe_mask_out_unregularized_parameters(params, log_paths=True) + + return self._init(params, rng, batch, func_state) + + @functools.partial(utils.staged, donate_argnums=[1, 3, 5]) # pytype: disable=wrong-arg-types + def _burnin( + self, + params: Params, + state: State, + rng: Array, + batch: Batch, + func_state: FuncState | None, + damping: Array | None, + accumulator: utils.MultiChunkAccumulator, + sync: Array | bool, + ) -> tuple[State, utils.MultiChunkAccumulator]: + """A single burnin step, updating only the curvature estimate.""" + + _, _, _, precon_damping = self._setup_state_and_schedules( + None, None, + state.damping if self._use_adaptive_damping else damping, + state.step_counter, state.data_seen) + + # Copy this first since we mutate it later in this function. + accumulator = accumulator.copy() + + func_args, rng = self._setup_func_args_and_rng( + params, rng, batch, func_state) + + # Update curvature estimate + state.estimator_state = self._update_estimator_curvature( + state.estimator_state, + func_args, + rng, + ema_old=1.0, + ema_new=1.0, + precon_damping=precon_damping, + sync=sync, + ) + + # Optionally update func_state + if func_state is not None: + out, _ = self._value_and_grad_func(*func_args) + _, func_state, _ = extract_func_outputs( + out, self._value_func_has_aux, self._value_func_has_state) + + accumulator.add(func_state) + + return state, accumulator + + def _burnin_phase( + self, + num_steps: int, + params: Params, + state: State, + rng: PRNGKey, + data_iterator: Iterator[Batch], + func_state: FuncState | None = None, + damping: Array | None = None, + ) -> tuple[State, FuncState | None]: + """Runs all burnin steps required.""" + + if num_steps > 0: + + rng = self._rng_split(rng, num_steps) + + accumulator = utils.MultiChunkAccumulator.zeros_like( + func_state, self.multi_device) + + for i, rng_i in enumerate(rng): + batch = next(data_iterator) + + state, accumulator = self._burnin( + params, state, rng_i, batch, func_state, damping, accumulator, + i == num_steps - 1) + + func_state = accumulator.value_and_clear() + + return state, func_state + + @functools.partial( + utils.staged, donate_argnums=(0, 1, 4), static_argnums=(8, 9)) + @utils.auto_scope_method + def _step( + self, + params: Params, + state: State, + rng: Array, + batch: Batch, + func_state: FuncState | None, + learning_rate: Array | None, + momentum: Array | None, + damping: Array | None, + should_update_estimate_curvature: bool, + should_update_damping: bool, + curvature_ema: Numeric, + inverse_update_period: Numeric, + )-> ReturnEither: + """A single full step of the optimizer.""" + + # Copy this first since we mutate it later in this function. + state = state.copy() + + # Setup arguments + (learning_rate, momentum, damping, + precon_damping) = self._setup_state_and_schedules( + learning_rate, momentum, + state.damping if self._use_adaptive_damping else damping, + state.step_counter, state.data_seen) + + func_args, rng = self._setup_func_args_and_rng( + params, rng, batch, func_state) + + # Update curvature estimate + if should_update_estimate_curvature: + if self._share_curvature_and_grad_forward: + ( + state.estimator_state, + loss, + grads, + ) = self._update_estimator_curvature_and_value_and_grad( + state.estimator_state, + func_args, + rng, + ema_old=curvature_ema, + ema_new=1.0, + precon_damping=precon_damping, + sync=self._should_sync_estimator(state, inverse_update_period), + ) + else: + state.estimator_state = self._update_estimator_curvature( + state.estimator_state, + func_args, + rng, + ema_old=curvature_ema, + ema_new=1.0, + precon_damping=precon_damping, + sync=self._should_sync_estimator(state, inverse_update_period), + ) + + del rng # should not be used after this point! + + # Compute loss and gradients + if ( + should_update_estimate_curvature + and self._share_curvature_and_grad_forward + ): + func_state = None + aux = None + else: + loss, grads, func_state, aux = self._compute_loss_and_grads( + func_args, state=state) + + # Sync + loss, grads = utils.pmean_if_pmap((loss, grads), self.pmap_axis_name) + + # Update the inverse curvature + state = self._maybe_update_inverse_cache( + state, precon_damping, inverse_update_period) + + # Compute proposed directions + preconditioned_gradient = self._compute_preconditioned_gradient( + state, grads, precon_damping + ) + + # Client stats hook on the PRE-norm-constraint preconditioned + # gradient (the clip below is a scalar rescale). + if self._step_stats_hook is not None: + hook_stats = self._step_stats_hook( + self._estimator, grads, preconditioned_gradient) + else: + hook_stats = {} + + # constrain the norms + preconditioned_gradient, scaled_grad_norm_sq = ( + self._maybe_apply_norm_constraint( + grads, preconditioned_gradient, learning_rate, + ) + ) + + vectors = (preconditioned_gradient, state.velocities) + + # Compute the coefficients for the vectors + coefficients, quad_model_change = self._coefficients_and_quad_change( + state=state, + vectors=vectors, + grads=grads, + learning_rate=learning_rate, + momentum=momentum, + damping=damping, + func_args=func_args, + should_update_damping=should_update_damping, + ) + + # Compute the parameter update (delta) + delta = self._weighted_sum_of_objects(vectors, coefficients) + + # Update parameters + new_params = jax.tree_util.tree_map(jnp.add, params, delta) + + if should_update_damping or self._use_step_rejection: + + new_loss = self._compute_loss_value((new_params,) + func_args[1:]) + # Sync + new_loss = utils.pmean_if_pmap(new_loss, self.pmap_axis_name) + + else: + new_loss = self._invalid_metric_value + + # Optionally compute the reduction ratio and update the damping + if should_update_damping: + + state.damping, rho = self._compute_new_damping_and_rho( + loss, new_loss, quad_model_change, state.damping) + + else: + # If not adjusting the damping we don't compute these here and just set + # them to self._invalid_metric_value. + new_loss, rho = self._invalid_metric_value, self._invalid_metric_value + + if self._use_step_rejection: + + reject_step = jnp.logical_or(jnp.isnan(new_loss), new_loss > loss) + + params, state.velocities, state.damping = lax.cond( + reject_step, + lambda: (params, state.velocities, + self._reject_damping_increase_factor * state.damping), + lambda: (new_params, delta, state.damping)) + + else: + # stop the linter from complaining about uninitialized variable + reject_step = False + params, state.velocities = new_params, delta + + # Compute per-device and total batch size + batch_size = self._batch_size_extractor(func_args[-1]) + + if self.multi_device: + total_batch_size = batch_size * jax.device_count() + else: + total_batch_size = batch_size + + # Update data seen and step counter + state.data_seen = state.data_seen + total_batch_size + state.step_counter = state.step_counter + 1 + + # Statistics with useful information + # Unlike other norm stats, sq_norm_scaled_grads has to be computed if + # norm_constraint is not None, so log it by default even if the other + # norm stats are not logged. This reduces the overall computational cost if + # no other grad stats are desired. + stats = dict( + step=state.step_counter, + batch_size=jnp.asarray(total_batch_size, dtype=jnp.int32), + data_seen=state.data_seen, + loss=loss, + new_loss=new_loss, + learning_rate=-coefficients[0], + momentum=coefficients[1], + damping=damping, + precon_damping=precon_damping, + rho=rho, + quad_model_change=quad_model_change, + scaled_grad_norm_sq=scaled_grad_norm_sq, + ) + + if self._use_step_rejection: + stats["step_rejected"] = reject_step + + stats.update(hook_stats) + + if aux is not None: + aux = utils.pmean_if_pmap(aux, self.pmap_axis_name) + stats["aux"] = aux + + if self._include_norms_in_stats: + stats["param_norm"] = utils.norm(params) + stats["grad_norm"] = utils.norm(grads) + stats["precon_grad_norm"] = utils.norm(preconditioned_gradient) + stats["update_norm"] = utils.norm(delta) + + if self._include_per_param_norms_in_stats: + stats.update(utils.per_parameter_norm(params, "param_norm")) + stats.update(utils.per_parameter_norm(grads, "grad_norm")) + stats.update( + utils.per_parameter_norm(preconditioned_gradient, "precon_grad_norm") + ) + stats.update(utils.per_parameter_norm(delta, "update_norm")) + + if self._include_registered_loss_in_stats: + assert aux is not None + stats["loss_registered"] = aux.pop("loss_registered") + stats["loss_registered"] = utils.pmean_if_pmap(stats["loss_registered"], + self.pmap_axis_name) + stats["loss_registered_reldiff"] = ( + stats["loss_registered"] - loss) / loss + + if self._value_func_has_state: + return params, state, func_state, stats + + assert func_state is None + + return params, state, stats + + def step( + self, + params: Params, + state: State, + rng: PRNGKey, + data_iterator: Iterator[Batch] | None = None, + batch: Batch | None = None, + func_state: FuncState | None = None, + learning_rate: Array | None = None, + momentum: Array | None = None, + damping: Array | None = None, + global_step_int: int | None = None + )-> ReturnEither: + """Performs a single update step using the optimizer. + + NOTE: please do not jit/pmap or otherwise compile this function with JAX, + as this can lead to errors. Compilation is handled internally by the + optimizer. + + NOTE: when ``multi_device`` is ``True``, all of the JAX array arguments to + this function (including arrays inside of trees), should have an extra + leading axis the size of the number of local devices. Slices of ``batch`` + and ``rng`` should be different for each device, whereas the other arugments + should be identical for each slice. Passing the arguments any other way will + result in an exception, or possibly undefined behavior. + + Args: + params: The current parameters of the model. + state: The current state of the optimizer. + rng: A Jax PRNG key. Should be different for each iteration, each host, + and for each slice in the leading axis (i.e. corresponding to devices) + when ``multi_device`` is ``True``. + data_iterator: A data iterator to use (if not passing ``batch``). + batch: A single batch used to compute the update. Should only pass one + of ``data_iterator`` or ``batch``. + func_state: Any function state that gets passed in and returned. + learning_rate: Learning rate to use if the optimizer was created with + ``use_adaptive_learning_rate=False`` and + ``learning_rate_schedule=None``. Should be ``None`` otherwise. + momentum: Momentum to use if the optimizer was created with + ``use_adaptive_momentum=False`` and ``momentum_schedule=None``. Should + be ``None`` otherwise. + damping: Damping to use if the optimizer was created with + ``use_adaptive_damping=False`` and ``damping_schedule=None``. Should be + ``None`` otherwise. See discussion of constructor argument + ``initial_damping`` for more information about damping. + global_step_int: The global step as a python int. Note that this must + match the step internal to the optimizer that is part of its state. + + Returns: + (params, state, stats) if ``value_func_has_state=False`` and + (params, state, func_state, stats) otherwise, where + + * params is the updated model parameters. + + * state is the updated optimizer state. + + * func_state is the updated function state. + + * stats is a dictionary of useful statistics including the loss. + """ + + if (data_iterator is None) == (batch is None): + raise ValueError("Exactly one of the arguments ``data_iterator`` and " + "``batch`` must be provided.") + + step_counter_int = self._verify_args_and_get_step_counter( + step_counter=state.step_counter, + learning_rate=learning_rate, + momentum=momentum, + damping=damping, + global_step_int=global_step_int, + ) + + if step_counter_int == 0: + + if self._num_burnin_steps > 0: + + if data_iterator is None: + raise ValueError("If num_burnin_steps > 0, data_iterator must be " + "provided.") + + rng, burnin_rng = self._rng_split(rng, 2) + + state, func_state = self._burnin_phase( + num_steps=self._num_burnin_steps, + params=params, + state=state, + rng=burnin_rng, + data_iterator=data_iterator, + func_state=func_state, + damping=damping, + ) + + if data_iterator is not None: + batch = next(data_iterator) + + if (step_counter_int == 0 and self._use_adaptive_damping + and self._use_initial_damping_calibration): + + assert self._num_burnin_steps > 0 + + state = self._calibrate_initial_damping( + params, state, rng, batch, func_state, learning_rate, momentum) + + should_update_estimate_curvature = self._should_update_estimate_curvature( + step_counter_int + ) + should_update_damping = self._should_update_damping(step_counter_int) + + curvature_ema, inverse_update_period = self._live_step_scalars() + + return self._step( + params, state, rng, batch, func_state, learning_rate, momentum, damping, + should_update_estimate_curvature, should_update_damping, + curvature_ema, inverse_update_period) + + def _calibrate_initial_damping( + self, + params: Params, + state: State, + rng: PRNGKey, + batch: Batch, + func_state: FuncState | None = None, + learning_rate: Array | None = None, + momentum: Array | None = None, + ) -> State: + """Calibrates the initial damping parameter.""" + + # Instead of writing a custom compiled function to compute rho and update + # the damping, we're going to be lazy and just call the step function + # repeatedly, throwing out the new optimizer state, params, and stats, while + # keeping the rng and batch the same at each call. This is a bit hacky and + # somewhat wasteful, both in terms of a few extra (minor) computations done + # in step() that are pointless, as well as the extra memory required to + # store temporary copies of the optimizer state and model params. + + # TODO(jamesmartens): Improve the implementation if this feature is commonly + # used? + + while True: + + prev_damping = float(self.get_first(state.damping)) + + # Note that we need to copy params and func_state since _step() will + # donate them. A bette option might be to recompile _step() to not donate + # these arguments. + curvature_ema, inverse_update_period = self._live_step_scalars() + + ret = self._step( + self.copy_obj(params), self.copy_obj(state), rng, batch, + self.copy_obj(func_state), learning_rate, momentum, None, False, True, + curvature_ema, inverse_update_period) + + new_state = ret[1] + + new_damping = float(self.get_first(new_state.damping)) + state.damping = new_state.damping + + del new_state + + if prev_damping == new_damping: + return state + + @utils.auto_scope_method + def _compute_exact_quad_model_filtered( + self, + vectors: Sequence[Params], + grads: Params, + func_args: FuncArgsVariants, + state: State | None = None, + fixed_coefficients: Sequence[Numeric | None] | None = None, + **kwargs, + ) -> QuadModelParams: + """Computes the components of the exact quadratic model.""" + + # We check the fixed_coefficients for zeros to save computing the expensive + # matrix vector products for vectors that will eventually be multiplied by + # zero. If fixed_coefficients is None, we assume that all coefficients are + # free and compute the full model. + + if fixed_coefficients is None: # can we get rid of this? + return self._compute_exact_quad_model( + vectors, grads, func_args, state=state, **kwargs) + + assert len(vectors) == len(fixed_coefficients) + assert len(vectors) == 2 # only deal with the two vector case + + def if_momentum_coeff_zero(): + + # Only pass in the vectors that won't be multiplied by zero + quad_model = self._compute_exact_quad_model( + vectors[:1], grads, func_args, state=state, **kwargs) + + # Repad the quad model with zeroes for the removed entries + return tuple( + jnp.pad(arr, [(0, 1)] * arr.ndim, constant_values=0.0) + for arr in quad_model + ) + + # This saves compiling both branches in the static case + if (isinstance(fixed_coefficients[1], float) + and fixed_coefficients[1] == 0.0): + + return if_momentum_coeff_zero() + + # Due to how XLA cannot share computations across cond boundaries, such as + # network forward and backwards passes, we cannot use a cond here and remain + # efficient. If this behavior ever changes we can uncomment the block below. + + # return jax.lax.cond( + # fixed_coefficients[1] == 0.0, + # if_momentum_coeff_zero, + # lambda: self._compute_exact_quad_model( + # vectors, grads, func_args, state=state), + # ) + + return self._compute_exact_quad_model( + vectors, grads, func_args, state=state, **kwargs) + + def _maybe_mask_out_unregularized_parameters( + self, params: Params, log_paths: bool = False) -> Params: + """Mask out parameters that are not l2 regularized.""" + + if log_paths: + logging.info("Unregularized parameters masking info (for curvature " + "calculations and L2 regularization)") + + def maybe_mask_out_single_param( + path: tuple[Any, ...], + param: Array + ) -> Array: + """Zero out a single parameter.""" + str_path = [] + for p in path: + if isinstance(p, jax.tree_util.DictKey): + str_path.append(p.key) + elif isinstance(p, jax.tree_util.GetAttrKey): + str_path.append(p.name) + + should_mask = any( + p in str_path + for p in self._regularized_parameters_path_exclusions + ) + + if log_paths: + log_message = "Masking" if should_mask else "Not masking" + logging.info(" %s out %s", log_message, path) + + return jnp.zeros_like(param) if should_mask else param + + return jax.tree.map_with_path( + maybe_mask_out_single_param, params + ) + + @utils.auto_scope_method + def _compute_exact_quad_model( + self, + vectors: Sequence[Params], + grads: Params, + func_args: FuncArgsVariants, + state: State | None = None, + ) -> QuadModelParams: + """Computes the components of the exact quadratic model. + + See comments of QuadModelParams for a description of the returned tuple. + + Args: + vectors: sequence of update vectors `V`. + grads: The gradient `g` of the loss function. + func_args: The arguments to the model's value function. + state: The current optimizer state. + + Returns: + A `QuadModelParams` tuple (A, D, R, b). + """ + + del state + + if self._mat_type_for_exact_quad_model == "fisher": + c_factor_v = tuple(self._implicit.multiply_fisher_factor_transpose + (func_args, vi) for vi in vectors) + elif self._mat_type_for_exact_quad_model == "ggn": + c_factor_v = tuple(self._implicit.multiply_ggn_factor_transpose + (func_args, vi) for vi in vectors) + else: + raise ValueError(f"Unrecognized matrix type string for exact quad model:" + f"'{self._mat_type_for_exact_quad_model}'.") + + masked_vectors = tuple(self._maybe_mask_out_unregularized_parameters(vi) + for vi in vectors) + + # pylint: disable=invalid-name + A = utils.matrix_of_inner_products(c_factor_v) + D = utils.matrix_of_inner_products(vectors) + R = utils.matrix_of_inner_products(masked_vectors) + b = utils.vector_of_inner_products(grads, vectors) + # pylint: enable=invalid-name + + quad_model_params = (A, D, R, b) + + return utils.pmean_if_pmap(quad_model_params, self.pmap_axis_name) + + @functools.partial(utils.staged, donate_argnums=2) + @utils.auto_scope_method + def _compute_approx_quad_model( + self, + state: State, + vectors: Sequence[Params], + grads: Params, + ) -> QuadModelParams: + """Computes the components of the approximate quadratic model.""" + + # v_i^T C v_j + def c_times_v(v): + return self._estimator.multiply( + state=state.estimator_state, + parameter_structured_vector=v, + identity_weight=0.0, + exact_power=True, + use_cached=False, + pmap_axis_name=self.pmap_axis_name, + norm_to_scale_identity_weight_per_block=self._norm_to_scale_identity_weight_per_block, + ) + + c_vectors = [c_times_v(v_i) for v_i in vectors] + + return (utils.symmetric_matrix_inner_products(c_vectors, vectors), + utils.matrix_of_inner_products(vectors), + utils.matrix_of_inner_products(vectors), + utils.vector_of_inner_products(grads, vectors)) + + def _evaluate_quadratic_model( + self, + a: Array, + a_damped: Array, + b: Array, + w: Array, + ) -> Array: + """Computes the quadratic model value from the inputs provided.""" + + a_final = a_damped if self._include_damping_in_quad_change else a + + return jnp.dot(w, jnp.dot(a_final, w)) / 2 + jnp.dot(w, b) + + @utils.staged + def _solve_quad_model( + self, + quad_model_parameters: QuadModelParams, + damping: Array, + fixed_coefficients: Sequence[Numeric | None], + reg_coeff: Numeric | None = None, + ) -> tuple[tuple[Numeric, ...], Array]: + """Solves for the optimal learning rate and momentum of the quadratic model. + + Args: + quad_model_parameters: The computed matrices A, D, R, and vector b. + damping: The damping to use for evaluating the quadratic model. + fixed_coefficients: A list over the vectors of the fixed numerical values + to use for their coefficients. For each of these that is None, the + quadratic model is minimized to compute the 'optimal' coefficient value. + reg_coeff: The L2 regularization parameter to use. If None, the default + value from the optimizer is used. + + Returns: + A tuple of coefficients which are the solution (and include any values that + are not None from fixed_weights), and the value of the quadratic model + function for this solution (as a scalar). + + Raises: + The function currently supports only up to two vectors, hence if you + provide more, it will raise a ``NotImplementedError``. + """ + + if reg_coeff is None: + # use default l2 regularisation value. + reg_coeff = self._l2_reg + + # pylint: disable=invalid-name + A_no_diag, D, R, b = quad_model_parameters + A = A_no_diag + reg_coeff * R + A_damped = A + damping * D + + if all(c is None for c in fixed_coefficients): + # Adapt all coefficients + + if len(fixed_coefficients) == 1: + # This special case arises at the first iteration, because all + # velocities are zeros. + special_case = jnp.logical_and(A_damped[0, 0] == 0, b[0] == 0) + w = -lax.cond(special_case, lambda: b, lambda: b / A_damped[0]) + + elif len(fixed_coefficients) == 2: + w = -utils.psd_solve_maybe_zero_last_idx(A_damped, b) + + else: + raise NotImplementedError() + + elif all(c is not None for c in fixed_coefficients): + # No coefficients adapted + + w = jnp.asarray(fixed_coefficients) + + elif len(fixed_coefficients) == 2: + # Exactly one adapted coefficient + + w = [None, None] + index = fixed_coefficients.index(None) + w[1 - index] = fixed_coefficients[1 - index] + + b_extra = A_damped[1 - index, index] * w[1 - index] + # pylint: enable=invalid-name + + w[index] = -(b[index] + b_extra) / A_damped[index, index] + + else: + raise NotImplementedError() + + w = tuple(w) + w: tuple[Numeric, ...] + + quad_model_change = self._evaluate_quadratic_model( + A, A_damped, b, jnp.array(w)) + + return w, quad_model_change + + @utils.staged + def _compute_new_damping_and_rho( + self, + old_loss: Array, + new_loss: Array, + quad_change: Array, + current_damping: Array, + ) -> tuple[Array, Array]: + """Computes the reduction ratio and the updated value of the damping.""" + + # Reduction ratio + rho = (new_loss - old_loss) / quad_change + rho_not_nan = jnp.nan_to_num(rho, nan=-100.0) + + # Update damping + should_increase = rho_not_nan < self._damping_lower_threshold + increased_damping = current_damping / self._damping_decay_factor + should_decrease = rho_not_nan > self._damping_upper_threshold + decreased_damping = current_damping * self._damping_decay_factor + + damping = jnp.select([should_decrease, should_increase], + [decreased_damping, increased_damping], + default=current_damping) + + return jnp.clip(damping, self._min_damping, self._max_damping), rho + + @utils.staged + def _weighted_sum_of_objects( + self, + objects: Sequence[utils.PyTree], + coefficients: Sequence[Numeric], + ) -> utils.PyTree: + """Returns the weighted sum of the objects in the sequence.""" + return utils.weighted_sum_of_objects(objects, coefficients) + + +def convert_value_and_grad_to_value_func( + value_and_grad_func: ValueAndGradFunc, + has_aux: bool = False, +) -> ValueFunc: + """Converts a value_and_grad function to value_func only. + + Args: + value_and_grad_func: The function which computes the loss value and the + gradients w.r.t. parameters. + has_aux: Similar to the meaning in :func:`jax.grad`, whether the + ``value_and_grad_func`` returns with the loss value any auxiliary data. + + Returns: + A function that returns only the loss value. + """ + + def value_func(*args, **kwargs) -> Array: + out, _ = value_and_grad_func(*args, **kwargs) + return out[0] if has_aux else out + + return value_func + + +def convert_value_and_grad_to_clean_value_and_grad( + value_and_grad_func: ValueAndGradFunc, + has_aux: bool = False, +) -> utils.ValueAndGradFunc: + """Converts a value_and_grad function to return only (loss, grads). + + Args: + value_and_grad_func: The function which computes the loss value and the + gradients w.r.t. parameters. + has_aux: Similar to the meaning in :func:`jax.grad`, whether the + ``value_and_grad_func`` returns with the loss value any auxiliary data. + + Returns: + A function that returns `(loss, grads)`. + """ + + def clean_value_and_grad_func(*args, **kwargs) -> tuple[Array, Params]: + out, grads = value_and_grad_func(*args, **kwargs) + loss = out[0] if has_aux else out + return loss, grads + + return clean_value_and_grad_func + + +def make_func_args( + params: Params, + func_state: FuncState | None, + rng: PRNGKey | None, + batch: Batch, + has_state: bool, + has_rng: bool, +) -> FuncArgsVariants: + """Constructs the arguments to the model function in the pre-assumed order. + + The model function is assumed to take arguments in the following order: + params, func_state, rng, batch + If it has no function state or does not use an rng, those two arguments are + discarded. + + Args: + params: The model parameters. + func_state: The function state, if ``has_state`` is ``True``, ``None`` + otherwise. + rng: The PRNG, if ``has_rng`` is ``True``, ``None`` otherwise. + batch: The batch of data. + has_state: Whether the function has a function state. + has_rng: Whether the function uses an rng. + + Returns: + The arguments that need to be passed to the model function. + """ + if has_state and func_state is None: + raise ValueError("`func_state=None`, but argument `has_state=True`.") + + if has_rng and rng is None: + raise ValueError("`rng=None`, but argument `has_rng=True`.") + + if not has_state and not has_rng: + return params, batch + + elif not has_rng: + return params, func_state, batch + + elif not has_state: + return params, rng, batch + + else: + return params, func_state, rng, batch + + +def extract_func_outputs( + raw_outputs: FuncOutputs, + has_aux: bool, + has_state: bool, +) -> tuple[Array, FuncState | None, FuncAux | None]: + """Converts the raw output of the model function into loss,func_state and aux. + + Args: + raw_outputs: The direct output of the model function. + has_aux: Whether the model function returns also some auxiliary data. + has_state: Whether the model function has a function state. + + Returns: + A triple ``(loss, func_state, aux)``. If the model function does not return + any auxiliary data than ``aux`` will be ``None`` and if it does not have a + state ``func_state`` will be ``None``. + """ + + if not has_aux and not has_state: + assert isinstance(raw_outputs, Array) + return raw_outputs, None, None + + loss, other = raw_outputs + + if has_aux and has_state: + func_state, aux = other + elif has_aux: + func_state, aux = None, other + else: + func_state, aux = other, None + + return loss, func_state, aux diff --git a/src/kfac_jax/_src/patches_second_moment.py b/src/kfac_jax/_src/patches_second_moment.py new file mode 100644 index 0000000000000000000000000000000000000000..a6e5c088559bd48864171a080f8c64237e2a07d4 --- /dev/null +++ b/src/kfac_jax/_src/patches_second_moment.py @@ -0,0 +1,916 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC optimized functions for patches second moment(PSM) computation.""" +import functools +from typing import Sequence, TypeVar + +import jax +from jax import interpreters +from jax import lax +import jax.numpy as jnp +from kfac_jax._src import utils +from typing_extensions import Self + +# Types for annotation +T = TypeVar("T") +Array = utils.Array +Shape = utils.Shape +TracedType = interpreters.partial_eval.DynamicJaxprTracer +DimNumbers = tuple[Shape, Shape, Shape] +PaddingVariants = str | int | Sequence[int] | Sequence[tuple[int, int]] + +# Special global variables +_USE_4D_CONVOLUTION: bool = True + + +def set_use_4d_convolution_in_psm_loop(value: bool): + """Sets whether a 4D convolution is used for the PSM computation.""" + if not isinstance(value, bool): + raise ValueError("The value provided must be a python bool.") + global _USE_4D_CONVOLUTION + _USE_4D_CONVOLUTION = value + + +def get_use_4d_convolution_in_psm_loop() -> bool: + """Returns whether a 4D convolution is used for the PSM computation.""" + return _USE_4D_CONVOLUTION + + +def _ceil(x: int, y: int) -> int: + """Computes `ceil(x / y)` with only integer operations.""" + return - (- x // y) + + +class _ConvSpec: + """Layout specification for arrays that will be used in a convolution.""" + + def __init__(self, order: Sequence[int]): + """Initializes the array specification with the provided order.""" + self.order = tuple(order) + + def __len__(self): + return len(self.order) + + @property + def n_axis(self) -> int: + """Returns the index of the batch axis.""" + return self.order[0] + + @property + def c_axis(self) -> int: + """Returns the index of the channel axis.""" + return self.order[1] + + @property + def spatial_axes(self) -> tuple[int, ...]: + """Returns the indices of the spatial axes.""" + return self.order[2:] + + def get_n(self, shape: Shape) -> int: + """Returns the batch size of the given shape, under this spec layout.""" + return shape[self.n_axis] + + def get_c(self, shape: Shape) -> int: + """Returns the channel size of the given shape, under this spec layout.""" + return shape[self.c_axis] + + def get_spatial(self, shape: Shape) -> tuple[int, ...]: + """Returns the spatial sizes of the given shape, under this spec layout.""" + return tuple(shape[i] for i in self.spatial_axes) + + def expand_spatial_axes(self) -> Self: + """Expands the layout spatial axes by preserving `n` and `c` order.""" + n_axis = self.n_axis + sum(self.n_axis > axis for axis in self.spatial_axes) + c_axis = self.c_axis + sum(self.c_axis > axis for axis in self.spatial_axes) + spatial_axes = [] + for axis in self.spatial_axes: + spatial_axes.append(axis + sum(axis > a for a in self.spatial_axes)) + spatial_axes.append(spatial_axes[-1] + 1) + return _ConvSpec([n_axis, c_axis, *spatial_axes]) + + def swap_n_and_c(self) -> Self: + """Swaps the batch and channel indices of the layout.""" + return _ConvSpec([self.c_axis, self.n_axis, *self.spatial_axes]) + + def create_shape(self, n: T, c: T, *spatial_dims: T) -> tuple[T, ...]: + """Creates a shape according to this layout specification.""" + if len(spatial_dims) != len(self.order) - 2: + raise ValueError("Incorrect number of spatial dimensions.") + result: list[T] = [None] * len(self) # pytype: disable=annotation-type-mismatch + result[self.n_axis] = n + result[self.c_axis] = c + for ax, dim in zip(self.spatial_axes, spatial_dims): + result[ax] = dim + assert all(r is not None for r in result) + return tuple(result) + + def change_nhwc_to_ihwo(self) -> Self: + """Changes the layout from `NHWC` to `IHWO` where `I=C`, `O=N`.""" + # Change the spec: NHWC -> IHWO where I=C, O=N + order = [i - 2 if i > self.spatial_axes[1] else i for i in self.order[:4]] + return _ConvSpec(order).swap_n_and_c() + + +def _slice_array( + array: Array, + indices: Sequence[int | TracedType], + sizes: Sequence[int], +) -> Array: + """Takes a slice from the array provided.""" + if any(isinstance(x, TracedType) for x in indices): + # Any of the indices are dynamic values. + return lax.dynamic_slice_p.bind(array, *indices, slice_sizes=sizes) + else: + # All indices are static values. + index = tuple(slice(i, i + size) for i, size in zip(indices, sizes)) + return array[index] + + +def _output_spatial_shape( + inputs_spatial_shape: Shape, + kernel_spatial_shape: Shape, + spatial_strides: Shape, + padding: str | Sequence[tuple[int, int]], +) -> Shape: + """Returns the output spatial shape of the corresponding convolution.""" + if isinstance(padding, str): + if padding.lower() == "valid": + return tuple(_ceil(d - k + 1, s) for d, k, s in zip(inputs_spatial_shape, + kernel_spatial_shape, + spatial_strides)) + elif padding.lower() == "same": + return tuple(_ceil(d, s) for d, s in zip(inputs_spatial_shape, + spatial_strides)) + else: + raise ValueError(f"Unrecognized padding string {padding}!") + else: + shapes_strides_padding = zip( + inputs_spatial_shape, kernel_spatial_shape, spatial_strides, padding) + return tuple(_ceil(d + p[0] + p[1] - k + 1, s) + for d, k, s, p in shapes_strides_padding) + + +def _normalize_padding( + inputs_spatial_shape: Shape, + kernel_spatial_shape: Shape, + spatial_strides: Shape, + padding: PaddingVariants, +) -> tuple[tuple[int, int], ...]: + """Returns the padding as a tuple of pairs of integers.""" + n = len(kernel_spatial_shape) + if isinstance(padding, str): + if padding.lower() == "valid": + return ((0, 0),) * n + elif padding.lower() == "same": + # https://github.com/tensorflow/tensorflow/blob/r1.8/tensorflow/core/kernels/conv_ops.cc#L571 + output_shape = _output_spatial_shape(inputs_spatial_shape, + kernel_spatial_shape, + spatial_strides, "same") + padding = [] + for out_d, d, k, s in zip(output_shape, inputs_spatial_shape, + kernel_spatial_shape, spatial_strides): + pad = max(0, (out_d - 1) * s + k - d) + padding.append((pad // 2, pad - pad // 2)) + return tuple(padding) + else: + raise ValueError(f"Unrecognized padding: {padding}!") + elif isinstance(padding, int): + return ((padding, padding),) * n + else: + final_padding = [] + for pad in padding: + if isinstance(pad, int): + final_padding.append((pad, pad)) + else: + final_padding.append(pad) + return tuple(final_padding) + + +def _normalize_strides( + kernel_spatial_shape: Shape, + strides: int | Shape, +) -> tuple[int, ...]: + """Returns the strides as a tuple of integers.""" + n = len(kernel_spatial_shape) + if strides is None: + return (1,) * n + elif isinstance(strides, int): + return (strides,) * n + else: + assert len(strides) == n + return tuple(strides) + + +def _data_format_to_dim_numbers( + data_format: str | None, + kernel_format: str = "HWIO", +) -> lax.ConvDimensionNumbers: + """Converts the data format in dim numbers.""" + if data_format is None: + data_format = "NHWC" + if not isinstance(data_format, str): + raise ValueError("data_format must be either a python string or `None`.") + data_format = lax.conv_general_permutations([data_format, + kernel_format, + data_format]) + return lax.ConvDimensionNumbers(*data_format) + + +def _parse_simple_args( + inputs_shape: Shape, + kernel_spatial_shape: int | Shape, + strides: int | Shape = 1, + padding: PaddingVariants = "VALID", + data_format: str | None = "NHWC", + dim_numbers: DimNumbers | lax.ConvDimensionNumbers | None = None, +) -> tuple[ + tuple[int, ...], + tuple[int, ...], + tuple[tuple[int, int], ...], + lax.ConvDimensionNumbers, +]: + """Parses all convolutional arguments to a single unified format. + + Args: + inputs_shape: A sequence of ints specifying the input's shape. + kernel_spatial_shape: A sequence of ints specifying the kernel's shape. + strides: A sequence of ints specifying strides in each spatial dimension, + or a single int specifying the strides in every spatial dimension. + padding: The padding can take one of the following formats: + * str - Either 'VALID' or 'SAME' + * int - Specifies the padding on both of sides of every spatial dimension. + * sequence of ints - Specifies the padding on both sides of each spatial + dimension. + * sequence of pairs of ints - Specifies the padding on each side of each + spatial dimension. + data_format: The data format layout of the inputs. + dim_numbers: If `data_format` is `None` this can specify the layout instead. + + Returns: + A tuple of the (kernel shape, strides, padding, dim_numbers) + """ + spatial_dims = len(inputs_shape) - 2 + + if data_format is not None and dim_numbers is not None: + raise ValueError("At least one of `data_format` and `dim_numbers` " + "must be None.") + + if dim_numbers is not None: + + if not isinstance(dim_numbers, lax.ConvDimensionNumbers): + + if not isinstance(dim_numbers, (list, tuple)): + raise ValueError("The provided dim_numbers argument must be either a " + "list, tuple or lax.ConvDimensionNumbers.") + + if len(dim_numbers) != 3: + raise ValueError("When the provided dim_numbers argument is a list or " + "tuple it must have length 3, but has length " + f"{len(dim_numbers)}.") + + lax_dim_numbers = lax.ConvDimensionNumbers(*dim_numbers) + + else: + lax_dim_numbers: lax.ConvDimensionNumbers = dim_numbers + + else: + lax_dim_numbers = _data_format_to_dim_numbers(data_format) + + if isinstance(kernel_spatial_shape, int): + kernel_spatial_shape = (kernel_spatial_shape,) * spatial_dims + + if len(kernel_spatial_shape) != spatial_dims: + raise ValueError("The provided argument `kernel_spatial_shape` must have " + f"length equal to the spatial dimensions {spatial_dims} of" + f" the inputs, but got {len(kernel_spatial_shape)}.") + + inputs_spatial_shape = _ConvSpec(lax_dim_numbers.lhs_spec).get_spatial( + inputs_shape) + + kernel_spatial_shape = _ConvSpec(lax_dim_numbers.rhs_spec).get_spatial( + kernel_spatial_shape) + strides = _normalize_strides(kernel_spatial_shape, strides) + + padding = _normalize_padding( + inputs_spatial_shape, kernel_spatial_shape, strides, padding) + + return kernel_spatial_shape, strides, padding, lax_dim_numbers + + +def _num_conv_locations_full_spec( + input_spatial_shape: Shape, + kernel_spatial_shape: Shape, + spatial_strides: Shape, + spatial_padding: Sequence[tuple[int, int]], +) -> int: + """The number of convolution locations from the unified spec for arguments.""" + if len(kernel_spatial_shape) != len(input_spatial_shape): + raise ValueError("The `kernel_spatial_shape` and `input_spatial_shape` " + "must have the same number of elements, got " + f"{len(kernel_spatial_shape)} and " + f"{len(input_spatial_shape)}.") + if len(spatial_strides) != len(input_spatial_shape): + raise ValueError("The `spatial_strides` and `input_spatial_shape` " + "must have the same number of elements, got " + f"{len(spatial_strides)} and " + f"{len(input_spatial_shape)}.") + if len(spatial_padding) != len(input_spatial_shape): + raise ValueError("The `spatial_padding` and `input_spatial_shape` " + "must have the same number of elements, got " + f"{len(spatial_padding)} and " + f"{len(input_spatial_shape)}.") + + num_locations = 1 + for in_dim, k_dim, stride, padding in zip( + input_spatial_shape, kernel_spatial_shape, + spatial_strides, spatial_padding): + num_locations *= _ceil(in_dim + padding[0] + padding[1] - k_dim + 1, stride) + return num_locations + + +def num_conv_locations( + inputs_spatial_shape: Shape, + kernel_spatial_shape: int | Shape, + spatial_strides: int | Shape, + spatial_padding: str | int | Sequence[tuple[int, int]], +) -> int: + """Returns the number of convolution locations for the provided shapes.""" + inputs_spatial_shape = tuple(inputs_spatial_shape) + n = len(inputs_spatial_shape) + if isinstance(kernel_spatial_shape, int): + kernel_spatial_shape = (kernel_spatial_shape,) * n + spatial_strides = _normalize_strides(kernel_spatial_shape, spatial_strides) + spatial_padding = _normalize_padding( + inputs_spatial_shape, kernel_spatial_shape, + spatial_strides, spatial_padding) + return _num_conv_locations_full_spec( + inputs_spatial_shape, kernel_spatial_shape, + spatial_strides, spatial_padding) + + +@utils.auto_scope_function +def _the_conv4d( + lhs: Array, + lhs_spec: _ConvSpec, + rhs: Array, + rhs_spec: _ConvSpec, + pad_h: int, + pad_w: int, + stride_h: int, + stride_w: int, + per_channel: bool = False, + precision: jax.lax.Precision | None = None, +) -> Array: + """Performs a special conv4d or conv2d based on the global flag.""" + assert len(rhs_spec) == 6 + if get_use_4d_convolution_in_psm_loop(): + # Reshape lhs to 6D array - (n, extra_h, 1, extra_w, 1, c) + lhs_shape = list(lhs.shape) + lhs_shape.insert(lhs_spec.spatial_axes[1] + 1, 1) + lhs_shape.insert(lhs_spec.spatial_axes[0] + 1, 1) + lhs = jnp.reshape(lhs, lhs_shape) + # Change the spec: NHAWBC -> CHAWBN + lhs_spec = rhs_spec.swap_n_and_c() + # Change the spec: NHAWBC -> IHAWBO where I=C, O=N + rhs_spec = rhs_spec.swap_n_and_c() + dim_specs = (lhs_spec.order, rhs_spec.order, lhs_spec.order) + if per_channel: + @functools.partial(jax.vmap, + in_axes=(lhs_spec.n_axis, rhs_spec.n_axis), + out_axes=-1) + def single_conv(x, y): + return lax.conv_general_dilated( + lhs=jnp.expand_dims(x, lhs_spec.n_axis), + rhs=jnp.expand_dims(y, rhs_spec.n_axis), + window_strides=(1, 1, 1, 1), + padding=((0, pad_h), (0, 0), (0, pad_w), (0, 0)), + lhs_dilation=(1, 1, 1, 1), + rhs_dilation=(stride_h, 1, stride_w, 1), + dimension_numbers=lax.ConvDimensionNumbers(*dim_specs), + precision=precision, + ) + + result = single_conv(lhs, rhs) + assert result.shape[lhs_spec.n_axis] == 1 + result = jnp.squeeze(result, lhs_spec.n_axis) + assert result.shape[2] == 1 + assert result.shape[4] == 1 + result = jnp.squeeze(result, (2, 4)) + return result[None] + else: + result = lax.conv_general_dilated( + lhs=lhs, + rhs=rhs, + window_strides=(1, 1, 1, 1), + padding=((0, pad_h), (0, 0), (0, pad_w), (0, 0)), + lhs_dilation=(1, 1, 1, 1), + rhs_dilation=(stride_h, 1, stride_w, 1), + dimension_numbers=lax.ConvDimensionNumbers(*dim_specs), + precision=precision, + ) + # Order the result such that one of the channel dims is after spatial dims + if lhs_spec != (5, 0, 1, 2, 3, 4): + min_index = 0 if lhs_spec.n_axis < lhs_spec.c_axis else 1 + max_index = 1 - min_index + axes = list(range(6)) + if lhs_spec.order[min_index] != 0: + axes.insert(0, axes.pop(lhs_spec.order[min_index])) + if lhs_spec.order[max_index] != 5: + axes.insert(5, axes.pop(lhs_spec.order[max_index])) + result = jnp.transpose(result, axes=axes) + assert result.shape[2] == 1 + assert result.shape[4] == 1 + result = jnp.squeeze(result, (2, 4)) + return result[None, None] + else: + # Change the spec: NHWC -> CHWN + lhs_spec = lhs_spec.swap_n_and_c() + # Index rhs and remove the trivial dimensions + rhs_slice: list[slice | int] = [slice(None)] * rhs.ndim + rhs_slice[rhs_spec.spatial_axes[1]] = 0 + rhs_slice[rhs_spec.spatial_axes[3]] = 0 + rhs = rhs[tuple(rhs_slice)] + rhs_spec = rhs_spec.change_nhwc_to_ihwo() + dim_specs = (lhs_spec.order, rhs_spec.order, lhs_spec.order) + if per_channel: + vmap_single_conv = jax.vmap(lambda x, y: lax.conv_general_dilated( # pylint: disable=g-long-lambda + lhs=jnp.expand_dims(x, lhs_spec.n_axis), + rhs=jnp.expand_dims(y, rhs_spec.n_axis), + window_strides=(1, 1), + padding=((0, pad_h), (0, pad_w)), + lhs_dilation=(1, 1), + rhs_dilation=(stride_h, stride_w), + dimension_numbers=lax.ConvDimensionNumbers(*dim_specs), + precision=precision, + ), in_axes=(lhs_spec.n_axis, rhs_spec.n_axis), out_axes=-1) + result = vmap_single_conv(lhs, rhs) + assert result.shape[lhs_spec.n_axis] == 1 + result = jnp.squeeze(result, lhs_spec.n_axis) + return result[None] + else: + result = lax.conv_general_dilated( + lhs=lhs, + rhs=rhs, + window_strides=(1, 1), + padding=((0, pad_h), (0, pad_w)), + lhs_dilation=(1, 1), + rhs_dilation=(stride_h, stride_w), + dimension_numbers=lax.ConvDimensionNumbers(*dim_specs), + precision=precision, + ) + # Order the result such that one of the channel dims is after spatial dims + if lhs_spec != (3, 0, 1, 2): + min_index = 0 if lhs_spec.n_axis < lhs_spec.c_axis else 1 + max_index = 1 - min_index + axes = list(range(4)) + if lhs_spec.order[min_index] != 0: + axes.insert(0, axes.pop(lhs_spec.order[min_index])) + if lhs_spec.order[max_index] != 5: + axes.insert(5, axes.pop(lhs_spec.order[max_index])) + result = jnp.transpose(result, axes=axes) + return result[None, None] + + +def _validate_inputs_lengths( + inputs: Array, + kernel_spatial_shape: Shape, + strides: Shape, + padding: tuple[tuple[int, int], ...], +) -> None: + """Checks that the provided arguments are valid.""" + spatial_dims = inputs.ndim - 2 + if spatial_dims != 2: + raise ValueError("Currently `patches_second_moment` supports only 2D " + "convolution, hence the input is expected to have rank 4," + f" but has rank {inputs.ndim}.") + if len(kernel_spatial_shape) != spatial_dims: + raise ValueError("The argument `kernel_spatial_shape` must have length " + f"equal to the number of spatial dimensions of the input -" + f" {spatial_dims}, but instead has length " + f"{len(kernel_spatial_shape)}.") + if len(padding) != spatial_dims: + raise ValueError("The argument `padding` must have length equal to the " + "number of spatial dimensions of the input - " + f"{spatial_dims}, but instead has length " + f"{len(kernel_spatial_shape)}.") + if len(strides) != 2: + raise ValueError("The argument `strides` must have length equal to the " + "number of spatial dimensions of the input - " + f"{spatial_dims}, but instead has length " + f"{len(kernel_spatial_shape)}.") + + +@functools.partial(jax.jit, static_argnums=list(range(1, 12)), + static_argnames=( + "kernel_spatial_shape", "strides", "padding", + "data_format", "dim_numbers", "inputs_dilation", + "kernel_dilation", "feature_group_count", + "batch_group_count", "unroll_loop", "precision")) +@utils.auto_scope_function +def patches_moments_explicit( + inputs: Array, + kernel_spatial_shape: int | Shape, + strides: int | Shape = 1, + padding: PaddingVariants = "VALID", + data_format: str | None = "NHWC", + dim_numbers: DimNumbers | lax.ConvDimensionNumbers | None = None, + inputs_dilation: Sequence[int] | None = None, + kernel_dilation: Sequence[int] | None = None, + feature_group_count: int = 1, + batch_group_count: int = 1, + unroll_loop: bool = False, + precision: jax.lax.Precision | None = None, + weighting_array: Array | None = None, +) -> tuple[Array, Array]: + """The exact same functionality as :func:`~patches_moments`, but explicitly extracts the patches via :func:`jax.lax.conv_general_dilated_patches`, potentially having a higher memory usage.""" + kernel_spatial_shape, strides, padding, dim_numbers = _parse_simple_args( + inputs.shape, kernel_spatial_shape, padding=padding, strides=strides, + data_format=data_format, dim_numbers=dim_numbers) + _validate_inputs_lengths(inputs, kernel_spatial_shape, strides, padding) + + in_spec = _ConvSpec(dim_numbers.lhs_spec) + out_spec = _ConvSpec(dim_numbers.out_spec) + n = in_spec.get_n(inputs.shape) + c = in_spec.get_c(inputs.shape) + inputs_spatial_shape = in_spec.get_spatial(inputs.shape) + spec = _ConvSpec(dim_numbers.out_spec).swap_n_and_c().order + matmul_dim_numbers = lax.ConvDimensionNumbers(spec, spec, spec) + + if feature_group_count not in (1, in_spec.get_c(inputs.shape)): + raise ValueError("`patches_moments_explicit` does not support " + "`feature_group_count` different from 1 or the number of " + "channels of the inputs.") + if batch_group_count != 1: + raise ValueError("`patches_moments_explicit` does not support " + "`batch_group_count` different from 1.") + + per_channel = feature_group_count != 1 + vector_target_shape = kernel_spatial_shape + (c,) + leading_shape = kernel_spatial_shape if per_channel else vector_target_shape + matrix_target_shape = leading_shape + vector_target_shape + vector_axis = tuple(a for a in range(4) if a != out_spec.c_axis) + + # Broadcast the weighting function + if weighting_array is not None: + if weighting_array.ndim == inputs.ndim: + pass + elif weighting_array.ndim == inputs.ndim - 1: + axis = dim_numbers.lhs_spec[1] + weighting_array = jnp.expand_dims(weighting_array, axis=axis) + elif weighting_array.ndim == 1: + while weighting_array.ndim < inputs.ndim: + weighting_array = weighting_array[:, None] + else: + raise ValueError(f"`weighting_array` shape {weighting_array.shape} is " + f"not compatible with the inputs shape {inputs.shape}" + ".") + + if not per_channel: + vector_shape = (c,) + kernel_spatial_shape + matrix_shape = vector_shape + vector_shape + if weighting_array is None: + weighting_array = jnp.ones([], dtype=inputs.dtype) + + # Standard explicit patches calculation + extracted_patches = lax.conv_general_dilated_patches( + inputs, + filter_shape=kernel_spatial_shape, + window_strides=strides, + padding=padding, + lhs_dilation=inputs_dilation, + rhs_dilation=kernel_dilation, + dimension_numbers=dim_numbers, + precision=precision, + ) + + weighted_patches = extracted_patches * weighting_array + matrix_results = lax.conv_general_dilated( + extracted_patches, + weighted_patches, + window_strides=strides, + padding="VALID", + dimension_numbers=matmul_dim_numbers, + precision=precision, + ) + matrix_results = jnp.reshape(matrix_results, matrix_shape) + vector_results = jnp.reshape( + jnp.sum(weighted_patches, axis=vector_axis), vector_shape) + + if c > 1: + # The output of `conv_general_dilated_patches` is ordered `chw` + return (jnp.transpose(matrix_results, (1, 2, 0, 4, 5, 3)), + jnp.transpose(vector_results, [1, 2, 0])) + else: + return (jnp.reshape(matrix_results, matrix_target_shape), + jnp.reshape(vector_results, vector_target_shape)) + + # Loop over channels + def general_loop_body(i, image): + index = in_spec.create_shape(0, i, 0, 0) + sizes = in_spec.create_shape(n, 1, *inputs_spatial_shape) + image_channel = _slice_array(image, index, sizes) + + # Index the weighting function + if weighting_array is not None: + if weighting_array.shape[in_spec.c_axis] == 1: + wf_i = weighting_array + else: + wf_n = weighting_array[in_spec.n_axis] + wf_spatial = [weighting_array.shape[a] for a in in_spec.spatial_axes] + wf_sizes = in_spec.create_shape(wf_n, jnp.ones([]), *wf_spatial) # pytype: disable=wrong-arg-types # jnp-type + wf_i = _slice_array(weighting_array, index, wf_sizes) + else: + wf_i = None + + matrix, vector = patches_moments_explicit( + image_channel, + kernel_spatial_shape=kernel_spatial_shape, + strides=strides, + padding=padding, + data_format=None, + dim_numbers=dim_numbers, + precision=precision, + weighting_array=wf_i, + ) + return jnp.squeeze(matrix, axis=2), vector + + if unroll_loop: + results = [general_loop_body(ii, inputs) for ii in range(c)] + matrix_results, vector_results = zip(*results) + matrix_results = jnp.concatenate(matrix_results, axis=-1) + vector_results = jnp.concatenate(vector_results, axis=-1) + return matrix_results, vector_results + + def loop_cond(args): + return args[0] < c + + def loop_body(args): + + i, image, matrix_result, vector_result = args + + matrix_update, vector_update = general_loop_body(i, image) + + matrix_result = lax.dynamic_update_slice( + matrix_result, matrix_update, (0, 0, 0, 0, i)) + + vector_result = lax.dynamic_update_slice( + vector_result, vector_update, (0, 0, i)) + + return i + 1, image, matrix_result, vector_result + + init_vals = (0, inputs, + jnp.zeros(matrix_target_shape, dtype=inputs.dtype), + jnp.zeros(vector_target_shape, dtype=inputs.dtype)) + + return lax.while_loop(loop_cond, loop_body, init_vals)[-2:] # pytype: disable=bad-return-type # lax-types + + +@functools.partial(jax.jit, static_argnums=list(range(1, 12)), + static_argnames=( + "kernel_spatial_shape", "strides", "padding", + "data_format", "dim_numbers", "inputs_dilation", + "kernel_dilation", "feature_group_count", + "batch_group_count", "unroll_loop", "precision")) +@utils.auto_scope_function +def patches_moments( + inputs: Array, + kernel_spatial_shape: int | Shape, + strides: int | Shape = 1, + padding: PaddingVariants = "VALID", + data_format: str | None = "NHWC", + dim_numbers: DimNumbers | lax.ConvDimensionNumbers | None = None, + inputs_dilation: Sequence[int] | None = None, + kernel_dilation: Sequence[int] | None = None, + feature_group_count: int = 1, + batch_group_count: int = 1, + unroll_loop: bool = False, + precision: jax.lax.Precision | None = None, + weighting_array: Array | None = None, +) -> tuple[Array, Array]: + """Computes the first and second moment of the convolutional patches. + + Since the code is written to support arbitrary convolution data formats, e.g. + both NHWC and NCHW, in comments above any of the procedures is written the + simplified version of what the statements below do, if the data format + was fixed to NHWC. + + Args: + inputs: The batch of images. + kernel_spatial_shape: The spatial dimensions of the filter (int or list of + ints). + strides: The spatial dimensions of the strides (int or list of ints). + padding: The padding (str or list of pairs of ints). + data_format: The data format of the inputs (None, NHWC, NCHW). + dim_numbers: Instance of :class:`jax.lax.ConvDimensionNumbers` instead of + data_format. + inputs_dilation: An integer or sequence of integers, specifying the dilation + for the image. Currently, `patches_moments` does not support dilation, so + the only allowed values are `None, 1, (1,1)`. + kernel_dilation: An integer or sequence of integers, specifying the dilation + for the kernel. Currently, `patches_moments` does not support dilation, so + the only allowed values are `None, 1, (1,1)`. + feature_group_count: The feature grouping for grouped convolutions. + Currently, `patches_moments` supports only 1 and number of channels of the + inputs. + batch_group_count: The batch grouping for grouped convolutions. Currently, + `patches_moments` supports only 1. + unroll_loop: Whether to unroll the loop in python. + precision: In what precision to run the computation. For more details please + read Jax documentation of :func:`jax.lax.conv_general_dilated`. + weighting_array: A tensor specifying additional weighting of each element + of the moment's average. + + Returns: + The matrix of the patches' second and first moment as a pair. The tensor of + the patches' second moment has a shape `kernel_spatial_shape + (, channels) + + kernel_spatial_shape + (, channels)`. The tensor of the patches' first + moment has a shape `kernel_spatial_shape + (, channels)`. + """ + kernel_spatial_shape, strides, padding, dim_numbers = _parse_simple_args( + inputs.shape, kernel_spatial_shape, padding=padding, strides=strides, + data_format=data_format, dim_numbers=dim_numbers) + _validate_inputs_lengths(inputs, kernel_spatial_shape, strides, padding) + + # Extract useful fixed integer values from the inputs + in_spec = _ConvSpec(dim_numbers.lhs_spec) + rhs_spec = _ConvSpec(dim_numbers.rhs_spec) + inputs_spatial_shape = in_spec.get_spatial(inputs.shape) + n = in_spec.get_n(inputs.shape) + c = in_spec.get_c(inputs.shape) + in_h, in_w = inputs_spatial_shape + ker_h, ker_w = kernel_spatial_shape + pad_h, pad_w = padding + s_h, s_w = strides + + if inputs_dilation not in (None, 1, (1, 1)): + raise ValueError("`patches_second_moment` does not support input dilation.") + if kernel_dilation not in (None, 1, (1, 1)): + raise ValueError("`patches_second_moment` does not support kernel " + "dilation.") + if feature_group_count not in (1, in_spec.get_c(inputs.shape)): + raise ValueError("`patches_second_moment` does not support " + "`feature_group_count` different from 1 or the number of " + "channels of the inputs.") + if batch_group_count != 1: + raise ValueError("PSM does not support `batch_group_count` different from " + "1.") + per_channel = feature_group_count != 1 + + # Sanity check + if in_h + pad_h[0] + pad_h[1] < ker_h or in_w + pad_w[0] + pad_w[1] < ker_w: + padded_h = in_h + pad_h[0] + pad_h[1] + padded_w = in_w + pad_w[0] + pad_w[1] + raise ValueError("The provided image has spatial padded shape " + f"({padded_h}, {padded_w}) while the kernel has a larger " + f"shape ({ker_h}, {ker_w}). This means a convolution is " + "not possible.") + + # First we calculate the maximum number of times the kernel can be applied + # into the image, including the padding and ignoring the stride. + ker_max_h = in_h + pad_h[0] + pad_h[1] - ker_h + 1 + ker_max_w = in_w + pad_w[0] + pad_w[1] - ker_w + 1 + # Second we calculate the size of the image that is covered when performing + # a VALID convolution with the kernel, provided the padding. + out_h = _ceil(ker_max_h, s_h) * s_h - s_h + ker_h + out_w = _ceil(ker_max_w, s_w) * s_w - s_w + ker_w + # Finally, we potentially add extra padding on the right in order to make the + # padded image sizes divisible by their strides. This is needed so we can use + # later reshape the image into multiples of the strides, which allows us to + # execute a strided slice via XLA's dynamic slice. Note that + # in certain cases this could lead to negative padding, which is correct. + # Example: image (9, 9), kernel (2, 2), strides (2, 2), padding (0, 0) + # Then ker_max = 8, out_h = 8, padded_height = 8 and the padding is -1. + padded_h = _ceil(out_h, s_h) * s_h + padded_w = _ceil(out_w, s_w) * s_w + # Actually pad the image (extra 0 for internal padding has to be added) + extra_pad_h = (pad_h[0], padded_h - in_h - pad_h[0], 0) + extra_pad_w = (pad_w[0], padded_w - in_w - pad_w[0], 0) + spatial_padding = in_spec.create_shape( + (0, 0, 0), (0, 0, 0), extra_pad_h, extra_pad_w) + padded_image = lax.pad(inputs, jnp.asarray(0.0, dtype=inputs.dtype), + spatial_padding) + + # Reshape the input based on strides + # rhs_shape = [n, out_h // str_h, str_h, out_w // str_w, str_w, c] + rhs_spec = in_spec.expand_spatial_axes() + rhs_shape = rhs_spec.create_shape( + n, c, padded_h // s_h, s_h, padded_w // s_w, s_w) + + # sizes = (n, rhs_h, 1, rhs_w, 1, c) + rhs_h = (padded_h - ker_h) // s_h + 1 + rhs_w = (padded_w - ker_w) // s_w + 1 + sizes = rhs_spec.create_shape(n, c, rhs_h, 1, rhs_w, 1) + + # Broadcast the weighting function + if weighting_array is not None: + if weighting_array.ndim == inputs.ndim: + shape = rhs_spec.create_shape(n, c, rhs_h, 1, rhs_w, 1) + elif weighting_array.ndim == inputs.ndim - 1: + shape = rhs_spec.create_shape(n, 1, rhs_h, 1, rhs_w, 1) + elif weighting_array.ndim == 1: + shape = rhs_spec.create_shape(n, 1, 1, 1, 1, 1) + else: + raise ValueError(f"`weighting_array` shape {weighting_array.shape} is " + f"not compatible with the inputs shape {inputs.shape}" + ".") + reshaped_weighting_array = jnp.reshape(weighting_array, shape) + else: + reshaped_weighting_array = 1 + + def general_loop_body(i, image): + reshaped_image = jnp.reshape(image, rhs_shape) + + # Slice the reshaped input + iw = i % ker_w + ih = i // ker_w + + # index = (0, ih // sh, ih % sh, iw // sw, iw % sw, 0) + index = rhs_spec.create_shape( + 0, 0, ih // s_h, ih % s_h, iw // s_w, iw % s_w) + conv_rhs = _slice_array(reshaped_image, index, sizes) + conv_rhs = conv_rhs * reshaped_weighting_array + + # Compute the correct padding for the convolution + dilated_bound_h = 0 if rhs_h == 0 else (rhs_h - 1) * s_h + 1 + dilated_bound_w = 0 if rhs_w == 0 else (rhs_w - 1) * s_w + 1 + conv_pad_h = ker_h - (padded_h - dilated_bound_h + 1) + conv_pad_w = ker_w - (padded_w - dilated_bound_w + 1) + + # Compute matrix update + matrix_update = _the_conv4d( + lhs=image, + lhs_spec=in_spec, + rhs=conv_rhs, + rhs_spec=rhs_spec, + pad_h=conv_pad_h, + pad_w=conv_pad_w, + stride_h=s_h, + stride_w=s_w, + per_channel=per_channel, + precision=precision, + ) + + # Compute vector update + axis = tuple(i for i in range(len(rhs_spec)) if i != rhs_spec.c_axis) + + vector_update = jnp.sum(conv_rhs, axis=axis) + vector_update = lax.broadcast_in_dim(vector_update, (1, 1, c), (2,)) + + return ih, iw, matrix_update, vector_update + + vector_shape = kernel_spatial_shape + (c,) + leading_shape = kernel_spatial_shape if per_channel else vector_shape + matrix_shape = leading_shape + vector_shape + + if unroll_loop: + + matrix_results, vector_results = zip( + *[general_loop_body(ii, padded_image)[-2:] + for ii in range(ker_h * ker_w)]) + + matrix_results = jnp.stack(matrix_results, axis=0) + matrix_results = jnp.reshape(matrix_results, matrix_shape) + + vector_results = jnp.stack(vector_results, axis=0) + vector_results = jnp.reshape(vector_results, vector_shape) + + return matrix_results, vector_results + + else: + + def loop_cond(args): + return args[0] < ker_h * ker_w + + def loop_body(loop_inputs): + + i, image, matrix_result, vector_result = loop_inputs + ih, iw, matrix_update, vector_update = general_loop_body(i, image) + + # Update matrix value + indices = (ih, iw, 0, 0, 0) + (() if per_channel else (0,)) + matrix_result = lax.dynamic_update_slice_p.bind( + matrix_result, matrix_update, *indices) + + # Update vector value + vector_result = lax.dynamic_update_slice_p.bind( + vector_result, vector_update, ih, iw, 0) + + return i + 1, image, matrix_result, vector_result + + # Initialize loop states with zeros + matrix_init = jnp.zeros(matrix_shape, dtype=inputs.dtype) + vector_init = jnp.zeros(vector_shape, dtype=inputs.dtype) + init_vals = (0, padded_image, matrix_init, vector_init) + + return lax.while_loop(loop_cond, loop_body, init_vals)[-2:] # pytype: disable=bad-return-type # lax-types diff --git a/src/kfac_jax/_src/tag_graph_matcher.py b/src/kfac_jax/_src/tag_graph_matcher.py new file mode 100644 index 0000000000000000000000000000000000000000..4f11c617d96a983b740724ec4f8923534be619b0 --- /dev/null +++ b/src/kfac_jax/_src/tag_graph_matcher.py @@ -0,0 +1,2220 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC functionality for auto-detecting layer tags and graph matching.""" + +import collections +import dataclasses +import functools +import itertools +import pprint +from typing import Any, Callable, Mapping, Sequence, Set, TypeVar, TYPE_CHECKING + +from absl import logging +import immutabledict +import jax +import jax.extend as jex + +import jax.numpy as jnp # pylint: disable=g-import-not-at-top +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import utils +import numpy as np + +jax_version = ( + jax.__version_info__ if hasattr(jax, "__version_info__") + else tuple(map(int, jax.__version__.split(".")))) + +if jax_version >= (0, 10, 0): + DebugInfo = jex.core.DebugInfo + DropVar = jex.core.DropVar + gensym = jex.core.gensym + new_jaxpr_eqn = jex.core.new_jaxpr_eqn +else: + if jax_version >= (0, 5, 1): + DebugInfo = jax.core.DebugInfo + else: + DebugInfo = jax.core.JaxprDebugInfo # pytype: disable=module-attr + DropVar = jax.core.DropVar + gensym = jax.core.gensym + new_jaxpr_eqn = jax.core.new_jaxpr_eqn + + +HIGHER_ORDER_NAMES = ("cond", "while", "scan", "pjit", "xla_call", "xla_pmap") +ITERATIVE_HIGHER_ORDER_NAMES = ("while", "scan") + +# Types for annotation +Array = utils.Array +PyTreeDef = utils.PyTreeDef +Var = jex.core.Var +Vars = Sequence[Var] +Jaxpr = jex.core.Jaxpr +ClosedJaxpr = jex.core.ClosedJaxpr +_JAXPR_TYPES_MERGED = Jaxpr is ClosedJaxpr +JaxprEqn = jex.core.JaxprEqn +JaxprEqns = Sequence[JaxprEqn] +T = TypeVar("T") +J = TypeVar("J", Jaxpr, ClosedJaxpr) +JaxprOrClosedJaxpr = Jaxpr | ClosedJaxpr +EquivalenceFunction = Callable[[JaxprEqn, JaxprEqn], bool] +MakeVarFunc = Callable[[jax.core.AbstractValue], Var] +VarProcessor = Callable[[Vars, MakeVarFunc], tuple[Vars, JaxprEqns]] +PatternComputeFunc = Callable[[Array, Sequence[Array]], Array] +ParameterExtractorFunc = Callable[[JaxprEqns], Mapping[str, Any]] +TagCtor = Callable[[Vars, Vars, JaxprEqns, MakeVarFunc], JaxprEqn] + + +def _closed_jaxpr(jaxpr: Jaxpr, consts: Sequence[Any]) -> ClosedJaxpr: + """Attaches constant values using the cross-version constructor form. + + JAX 0.11 merged ``Jaxpr`` and ``ClosedJaxpr`` and removed support for + ``ClosedJaxpr(jaxpr=..., consts=...)`` keyword construction. The legacy + positional form remains supported in both JAX 0.10 and 0.11. + """ + return ClosedJaxpr(jaxpr, consts) + + +def _scan_num_consts(eqn: JaxprEqn) -> int: + """Returns the number of scan constants across JAX 0.10 and 0.11.""" + if "num_consts" in eqn.params: + return eqn.params["num_consts"] + # JAX 0.11 represents the (const, carry, xs) input partition as a FlatTree. + ft_in = eqn.params.get("ft_in") + if ft_in is None: + raise KeyError("scan equation has neither num_consts nor ft_in") + consts, _, _ = ft_in.unpack() + return len(consts) + + +def eval_jaxpr_eqn(eqn: JaxprEqn, in_values: list[T]) -> list[T]: + """Computes the outputs of the given Jaxpr equation.""" + + user_context = jex.source_info_util.user_context + + if TYPE_CHECKING or jax_version >= (0, 9, 2): + bind_params = eqn.primitive.get_bind_params(eqn.params) + with user_context(eqn.source_info.traceback): + output = eqn.primitive.bind(*in_values, **bind_params) + else: + subfuns, bind_params = eqn.primitive.get_bind_params(eqn.params) + with user_context(eqn.source_info.traceback): + output = eqn.primitive.bind(*subfuns, *in_values, **bind_params) + + if not isinstance(output, list): + return [output] + else: + return output + + +def reshape_equivalent( + equation1: JaxprEqn, + equation2: JaxprEqn, +) -> bool: + """Equivalence rule for :func:`~jax.numpy.reshape` primitives.""" + if not (equation1.primitive.name == "reshape" and + equation2.primitive.name == "reshape"): + raise ValueError("This is only applicable to `reshape` primitive.") + + return equation1.params["dimensions"] == equation2.params["dimensions"] + + +def broadcast_in_dim_equivalent( + equation1: JaxprEqn, + equation2: JaxprEqn, +) -> bool: + """Equivalence rule for :func:`~jax.numpy.broadcast` primitives.""" + if not (equation1.primitive.name == "broadcast_in_dim" and + equation2.primitive.name == "broadcast_in_dim"): + raise ValueError("This is only applicable to `broadcast_in_dim` primitive.") + return True + + +def conv_general_dilated_equivalent( + equation1: JaxprEqn, + equation2: JaxprEqn, +) -> bool: + """Equivalence rule for :func:`~jax.lax.conv_general_dilated` primitives.""" + if not (equation1.primitive.name == "conv_general_dilated" and + equation2.primitive.name == "conv_general_dilated"): + raise ValueError("This is only applicable to `conv_general_dilated` " + "primitive.") + params1 = equation1.params + params2 = equation2.params + for k in ("window_strides", "padding", + "lhs_dilation", "rhs_dilation"): + if len(params1[k]) != len(params2[k]): + return False + # pytype: disable=attribute-error + if (len(params1["dimension_numbers"].lhs_spec) != + len(params2["dimension_numbers"].lhs_spec)): + return False + if (len(params1["dimension_numbers"].rhs_spec) != + len(params2["dimension_numbers"].rhs_spec)): + return False + if (len(params1["dimension_numbers"].out_spec) != + len(params2["dimension_numbers"].out_spec)): + return False + if ((params1["feature_group_count"] > 1) != + (params2["feature_group_count"] > 1)): + return False + if ((params1["batch_group_count"] > 1) != + (params2["batch_group_count"] > 1)): + return False + # pytype: enable=attribute-error + return True + + +def dot_general_equivalent( + equation1: JaxprEqn, + equation2: JaxprEqn, +) -> bool: + if not (equation1.primitive.name == "dot_general" and + equation2.primitive.name == "dot_general"): + raise ValueError("This is only applicable to `dot_general_equivalent` " + "primitive.") + # We ignore precision and preferred_element_type + return (equation1.params["dimension_numbers"] == + equation2.params["dimension_numbers"]) + + +DEFAULT_SPECIAL_EQUIVALENCE_RULES = immutabledict.immutabledict({ + "reshape": reshape_equivalent, + "broadcast_in_dim": broadcast_in_dim_equivalent, + "conv_general_dilated": conv_general_dilated_equivalent, + "dot_general": dot_general_equivalent, +}) + + +class GraphMatcherComparator: + """A class to compare and determine equivalence of abstract Jax equations.""" + + def __init__( + self, + commutative_ops_names: Sequence[str] = ("add", "mul"), + special_eqn_equivalence_rules: + Mapping[str, EquivalenceFunction] = DEFAULT_SPECIAL_EQUIVALENCE_RULES, + ): + """Initializes the instance. + + Args: + commutative_ops_names: A sequence of all Jax primitive names, which are + consider commutative ops and the order of their arguments is irrelevant. + special_eqn_equivalence_rules: A mapping of a Jax primitive names to a + comparison rule, which to be used instead of the default comparator, + which looks that the whole dictionaries of extra parameters to the + primitives match. + """ + self._commutative_ops_names = set(commutative_ops_names) + self._special_eqn_equivalence_rules = dict(**special_eqn_equivalence_rules) + + @property + def commutative_ops_names(self) -> set[str]: + """The set of commutative ops.""" + return self._commutative_ops_names + + @property + def special_eqn_equivalence_rules(self) -> Mapping[str, EquivalenceFunction]: + """The special equivalence rules.""" + return self._special_eqn_equivalence_rules + + def add_commutative_op_name(self, name: str): + """Adds a name to the set of primitive ops considered to be commutative.""" + if name in self.commutative_ops_names: + raise ValueError(f"Commutative op {name!r} has already been added.") + self._commutative_ops_names.add(name) + + def add_special_equivalence_rule( + self, + name: str, + equivalence_rule: EquivalenceFunction, + ): + """Adds the special equivalence rule for ``name`` to the global store.""" + if name in self.special_eqn_equivalence_rules: + raise ValueError( + f"Special equation equivalence rule already exists for name: {name}") + self._special_eqn_equivalence_rules[name] = equivalence_rule + + def are_equivalent( + self, + equation1: JaxprEqn, + equation2: JaxprEqn, + ) -> bool: + """Returns whether the two equations are considered equivalent.""" + + if equation1.primitive.name != equation2.primitive.name: + return False + + equivalence_rule = self.special_eqn_equivalence_rules.get( + equation1.primitive.name) + + if equivalence_rule is not None: + return equivalence_rule(equation1, equation2) + + # Default comparison + return equation1.params == equation2.params + + +@dataclasses.dataclass(frozen=True) +class JaxprGraph: + """A wrapper around Jaxpr as a graph for pattern matching. + + Attributes: + name: The name for this Jaxpr graph. + closed_jaxpr: The original `ClosedJaxpr` that is being wrapped. + params_tree: The PyTreeDef of the parameter variables. + params_vars: A flat list of all the abstract parameter variables. + out_tree: The PyTreeDef of the outputs of the function. + tag_ctor: This is an optional attribute, that defines if this is used during + automatic layer tag registration, how to construct the corresponding layer + tag primitive from the subgraph matching this pattern. + losses_eqns: A tuple of all the Jaxpr equations corresponding to a loss + tag. + var_to_creation_op: A mapping of variables to the Jax equation that created + it. + manual_registrations: Any layer tag equations that have been manually + registered. + jaxpr: The underlying :class:`Jaxpr` part of ``self.closed_jaxpr``. + consts: The underlying constants part ``self.closed_jaxpr``. + outvars: The output variables of the underlying :class:`Jaxpr` part + of ``self.closed_jaxpr``. + """ + name: str + closed_jaxpr: ClosedJaxpr + params_tree: PyTreeDef + params_vars: Vars + out_tree: PyTreeDef + tag_ctor: TagCtor | None + + @property + def jaxpr(self) -> Jaxpr: + return self.closed_jaxpr.jaxpr + + @property + def consts(self) -> Sequence[Any]: + return self.closed_jaxpr.consts + + @property + def outvars(self) -> Vars: + return self.jaxpr.outvars # pytype:disable=bad-return-type + + def sub_graph_eqns( + self, + root_vars: Sequence[Var], + leaf_vars: Sequence[Var], + ) -> JaxprEqns: + """Returns the sub-graph equations between root vars and leaf vars.""" + + eqns = [] + # Extract the subgraph equations such that they both depend on root_vars and + # leaf_vars depends on them + + if any(v in self.params_vars for v in leaf_vars): + # The special case of a generic tag, where the output is a parameter + assert all(v in self.params_vars for v in leaf_vars) + return () + + to_process_eqns = [self.var_to_creation_op[v] for v in leaf_vars] + processed_vars = set() + + while to_process_eqns: + + next_eqn = to_process_eqns.pop() + eqns.append(next_eqn) + + for v in next_eqn.invars: + if (not isinstance(v, jex.core.Literal) and v not in root_vars and + v not in processed_vars and v in self.var_to_creation_op): + to_process_eqns.append(self.var_to_creation_op[v]) + processed_vars.add(v) + + return tuple(eqns) + + @functools.cached_property + def losses_eqns(self) -> tuple[tags.LossTagEqn, ...]: + # Note that this function won't look inside higher order primitives of this + # graph to find loss tags. + return tuple( + eqn for eqn in self.closed_jaxpr.jaxpr.eqns + if isinstance(eqn.primitive, tags.LossTag) + ) + + @functools.cached_property + def var_to_creation_op(self) -> immutabledict.immutabledict: + return immutabledict.immutabledict( + sum(([(var, eqn) for var in eqn.outvars] + for eqn in self.jaxpr.eqns), [])) + + @functools.cached_property + def manual_registrations(self) -> tuple[tags.LayerTagEqn, ...]: + """Returns all manually registered tags.""" + + # Note that this function won't look inside higher order primitives of this + # graph to find layer tags. + + registered_tags = [] + + for eqn in self.jaxpr.eqns: + + if isinstance(eqn.primitive, tags.LayerTag): + + for param in tags.layer_eqn_data(eqn).params: + if param not in self.params_vars: + raise ValueError("One of the parameters of the manual layer " + f"registration equation: {eqn} is not part of " + "the parameters of the global function.") + + registered_tags.append(eqn) + + return tuple(registered_tags) + + +def make_jax_graph( + func: utils.Func, + func_args: utils.FuncArgs, + params_index: int | Sequence[int], + name: str, + compute_only_loss_tags: bool, + clean_broadcasts: bool, + tag_ctor: TagCtor | None = None, +) -> JaxprGraph: + """Creates a :class:`~JaxGraph` instance from the provided function and arguments.""" + + in_tree = jax.tree_util.tree_structure(func_args) + closed_jaxpr, out_shapes = jax.make_jaxpr(func, return_shape=True)(*func_args) + + if compute_only_loss_tags: + + make_var_func = gensym() + eqns = [] + sub_graph_vars = set() + loss_tags_output_vars = [] + + for eqn in reversed(closed_jaxpr.jaxpr.eqns): + + if (isinstance(eqn.primitive, tags.LossTag) or + any(v in sub_graph_vars for v in eqn.outvars)): + + if isinstance(eqn.primitive, tags.LossTag): + + new_out_vars = [] + for v in eqn.outvars: + + if isinstance(v, DropVar): + new_out_vars.append(make_var_func(v.aval)) + else: + new_out_vars.append(v) + + loss_tags_output_vars.extend(new_out_vars[::-1]) + eqns.append(eqn.replace(outvars=new_out_vars)) + + else: + eqns.append(eqn) + + sub_graph_vars.update( + v for v in eqn.invars if not isinstance(v, jex.core.Literal) + ) + + consts_i = [ + i + for i, c in enumerate(closed_jaxpr.jaxpr.constvars) + if c in sub_graph_vars + ] + + debug_info = closed_jaxpr.jaxpr.debug_info + if debug_info is not None: + debug_info = DebugInfo( + debug_info.traced_for, + debug_info.func_src_info, + debug_info.arg_names, + tuple([f"{i}" for i in range(len(loss_tags_output_vars))]), + ) + + closed_jaxpr = _closed_jaxpr( + closed_jaxpr.jaxpr.replace( + eqns=eqns[::-1], + constvars=[closed_jaxpr.jaxpr.constvars[i] for i in consts_i], + outvars=loss_tags_output_vars[::-1], + debug_info=debug_info, + ), + [closed_jaxpr.consts[i] for i in consts_i], + ) + out_shapes = [jax.ShapeDtypeStruct(shape=v.aval.shape, dtype=v.aval.dtype) + for v in closed_jaxpr.jaxpr.outvars] # pytype:disable=attribute-error + + closed_jaxpr = clean_jaxpr(closed_jaxpr) + + if clean_broadcasts: + closed_jaxpr = merge_broadcasts_jaxpr(closed_jaxpr) + closed_jaxpr = clean_jaxpr(closed_jaxpr) + + in_vars = jax.tree_util.tree_unflatten(in_tree, closed_jaxpr.jaxpr.invars) # pytype:disable=attribute-error + + if isinstance(params_index, int): + params_vars = in_vars[params_index] + else: + params_vars = tuple(in_vars[i] for i in params_index) + + params_vars, params_tree = jax.tree_util.tree_flatten(params_vars) + + return JaxprGraph( + name=name, + closed_jaxpr=closed_jaxpr, + params_tree=params_tree, + params_vars=params_vars, + out_tree=jax.tree_util.tree_structure(out_shapes), + tag_ctor=tag_ctor + ) + + +@dataclasses.dataclass(frozen=True) +class GraphPattern: + """A graph pattern used for automatically detecting layers. + + The graph matcher needs to trace at least once the full function, which + means the caller needs to provide it with dummy arguments. The shapes of the + arguments do not matter, as the graph matcher ignores their values, however + the rank does. Especially if there is some broadcasting happening you should + register with every possible broadcast pattern. As a general advice avoid + using a shape to be 1, unless you want the pattern to specifically match + that, as some operations, like squeeze for example, can have special + behaviour then. + + Attributes: + name: The name of the pattern that is being registered to. + tag_primitive: The primitive tag to bind. + compute_func: The function that performs the computation. + parameters_extractor_func: A function that extracts from the traced Jaxpr + any parameters that are passed into the tag. + example_args: Example arguments that can be inputted into ``func``. + in_values_preprocessor: A function that can optionally modify the in_vals + passed to the tag_primitive, from those that are usually the input to + the jaxpr. + jaxpr: The underlying :class:`Jaxpr` represented by the pattern. + param_vars: The list of :class:`Var` that correspond to parameters + in the pattern. + graph: A :class:`JaxprGraph` representation of the pattern. + """ + name: str + tag_primitive: tags.LayerTag + compute_func: PatternComputeFunc + parameters_extractor_func: ParameterExtractorFunc + example_args: utils.FuncArgs + in_values_preprocessor: VarProcessor | None = None + + @property + def jaxpr(self) -> Jaxpr: + return self.graph.jaxpr + + @property + def param_vars(self) -> Vars: + return self.graph.params_vars + + @functools.cached_property + def graph(self) -> JaxprGraph: + """A :class:`JaxprGraph` representation of the pattern.""" + jnp_args = jax.tree_util.tree_map(jnp.asarray, self.example_args) + return make_jax_graph( + func=self.compute_func, + func_args=jnp_args, + params_index=1, + name=self.name, + compute_only_loss_tags=False, + clean_broadcasts=True, + ) + + def tag_ctor( + self, + in_vars: Vars, + out_vars: Vars, + graph_eqns: JaxprEqns, + make_var_func: MakeVarFunc, + ) -> JaxprEqns: + """Registers the layer tag for this graph pattern. + + Args: + in_vars: The input variables to the pattern. + out_vars: The output variables to the pattern. + graph_eqns: The real graph equations corresponding to the pattern. + make_var_func: A function to create correctly new variables. + Returns: + A sequence of any additional equations that are created from creating the + tag. + """ + assert len(out_vars) == 1 + + if self.in_values_preprocessor is not None: + in_vars, eqns = self.in_values_preprocessor(in_vars, make_var_func) + else: + eqns = [] + + new_out_vars = [make_var_func(v.aval) for v in out_vars] + + tag_eqn = new_jaxpr_eqn( + invars=[*out_vars, *in_vars], + outvars=new_out_vars, + primitive=tags.layer_tag, + params=self.parameters_extractor_func(graph_eqns), + effects=set(), + ) + + return [*eqns, tag_eqn] + + +@dataclasses.dataclass(frozen=True) +class GraphMatch: + """Represents a match of the pattern on some graph. + + Attributes: + pattern: The pattern that has been matched. + variables_map: Mapping of variables from the pattern to the original graph, + on which it has been matched. + graph_eqns: All the equations in the original graph, that correspond to + computation of the pattern. + output_var: The variable in the original graph, that correspond to the + output variable of the pattern. + param_graph_variables: All variables in the original graph, that correspond + to parameters of the pattern. + name: The name of the pattern that has been matched. + """ + pattern: GraphPattern + variables_map: Mapping[Var, Var] + graph_eqns: JaxprEqns + + @property + def name(self) -> str: + return self.pattern.name + + @functools.cached_property + def output_var(self) -> Var: + return self.variables_map[self.pattern.jaxpr.outvars[0]] + + @functools.cached_property + def param_graph_variables(self) -> Vars: + return [self.variables_map[p] for p in self.pattern.graph.params_vars] + + def create_eqns_and_update_env( + self, + env: dict[Var, Var], + make_var_func: MakeVarFunc, + ) -> JaxprEqns: + """Creates a new equations for this match and inserts output vars in the environment.""" + + in_vars = [self.variables_map[k] for k in self.pattern.graph.jaxpr.invars] + in_vars = [env.get(v, v) if isinstance(v, Var) else v for v in in_vars] + + out_vars = [self.variables_map[k] for k in self.pattern.graph.jaxpr.outvars] + out_vars = [env.get(v, v) for v in out_vars] + + eqns = self.pattern.tag_ctor( + in_vars, out_vars, self.graph_eqns, make_var_func) + + assert len(out_vars) == len(eqns[-1].outvars) + + # Reinsert the output in the environment + for k, v in zip(out_vars, eqns[-1].outvars): + env[k] = v + + return eqns + + +def match_equations( + graph: JaxprGraph, + current_variables_map: Mapping[Var, Var], + reversed_eqns_to_match: Sequence[JaxprEqn], + input_vars: Vars, + param_variables: Vars, + graph_matcher_rules: GraphMatcherComparator, + matchable_graph_params: Set[Var], +) -> dict[Var, Var] | None: + """Tries to continue matching the remaining equations to the Jaxpr graph. + + Args: + graph: The :class:`~JaxprGraph` on which we are searching for matching + equations. + current_variables_map: A mapping from a pattern variables to graph + variables, which describes what is the current partial mapping between + the pattern and the graph. + reversed_eqns_to_match: The remaining equations of the pattern that have + not yet been matched to the graph. + input_vars: The input variables of the pattern. + param_variables: The parameter variables of the pattern. + graph_matcher_rules: A :class:`~GraphMatcherRules` instance, which is used + for determining equivalence of individual Jax primitives. + matchable_graph_params: A subset of graph.params_vars consisting of + parameters that may appear in matches as parameters (not merely input + variables). + + Returns: + ``None`` if it is not possible to finish matching the remaining equations + in the graph. Otherwise, returns the full match of the pattern onto the + graph, in terms of a variable to variable mapping. + """ + + # Copy the variables mapping + current_variables_map = dict(current_variables_map) + + def add_vars_if_possible( + eqn_vars: Sequence[Var], + graph_vars: Sequence[Var] + ) -> bool: + """Tries to update the current variables map. + + If at least one of the pattern variables is a parameter, but the + corresponding graph variable is not or vise-versa, the method does not + update the current variables map and returns ``False``. Similarly, if at + least one of the graph variables is a :class:`Literal` (meaning a + constant, independent of the function inputs) and the corresponding + pattern variable is not an input to the pattern, it returns ``False``. In + all other cases it updates the map and returns ``True``. + + Args: + eqn_vars: The variables from a single equation of the pattern. + graph_vars: The variables from a corresponding equation of the graph. + + Returns: + A boolean describing whether the method succeeded to update the + current variables map. + """ + for var1, var2 in zip(eqn_vars, graph_vars): + + var2_matchable = isinstance(var2, jex.core.Var) and ( + var2 in matchable_graph_params) + + if (var1 in param_variables and not var2_matchable or + var1 not in param_variables and var2_matchable or + (isinstance(var2, jex.core.Literal) and var1 not in input_vars)): + return False + + current_variables_map.update(zip(eqn_vars, graph_vars)) + + return True + + # Loop over all remaining equations to match + for i, eqn in enumerate(reversed_eqns_to_match): + + assert all(v in current_variables_map for v in eqn.outvars) + + # Retrieve the graph equation, whose output currently corresponds to the + # first output variable of the pattern equation. + first_output_var = current_variables_map[eqn.outvars[0]] + graph_eqn = graph.var_to_creation_op.get(first_output_var) + + if graph_eqn is None: + assert first_output_var in (graph.jaxpr.invars + graph.jaxpr.constvars) + # Clearly the pattern equation is not an input or parameter + return None + + assert isinstance(graph_eqn, JaxprEqn) + + # For equations with more than one output, make sure all output variables + # in the graph are generated from the same graph equation. + for v in eqn.outvars[1:]: + if graph_eqn != graph.var_to_creation_op.get(current_variables_map[v]): + return None + + # Check that the graph and pattern equation are equivalent + if not graph_matcher_rules.are_equivalent(graph_eqn, eqn): + return None + + # Sanity check + assert len(eqn.invars) == len(graph_eqn.invars) + + if eqn.primitive.name in graph_matcher_rules.commutative_ops_names: + + # For commutative ops we search through all possible pair alignments. + # This requires a recursive solution, on top of the iterative one. + results = [] + for permutation in itertools.permutations(range(len(eqn.invars))): + + pattern_vars = [eqn.invars[j] for j in permutation] + + # Check if this ordering is feasible + if not add_vars_if_possible(pattern_vars, graph_eqn.invars): + continue + + # Recursively continue by trying to match the remaining equations. + candidate_map = match_equations( + graph=graph, + current_variables_map=current_variables_map, + reversed_eqns_to_match=reversed_eqns_to_match[i + 1:], + input_vars=input_vars, + param_variables=param_variables, + graph_matcher_rules=graph_matcher_rules, + matchable_graph_params=matchable_graph_params, + ) + + if candidate_map is not None: + # Sanity check + assert all(candidate_map[p] in matchable_graph_params + for p in param_variables) + results.append(candidate_map) + + # Return appropriately + if len(results) > 1: + raise ValueError("Found multiple branch matches in pattern at " + f"associative op {eqn.primitive.name}.") + elif len(results) == 1: + return results[0] + else: + return None + + elif not add_vars_if_possible(eqn.invars, graph_eqn.invars): + # In the case where we can't update the current variables map directly + # return + return None + + return current_variables_map + + +def match_pattern( + graph: JaxprGraph, + root_eqn: JaxprEqn, + pattern: GraphPattern, + graph_matcher_rules: GraphMatcherComparator, + matchable_graph_params: Set[Var], +) -> GraphMatch | None: + """Tries to match the ``pattern`` in the Jaxpr graph from the ``root_eqn``. + + Args: + graph: The :class:`~JaxprGraph` on which we are searching for matching + equations. + root_eqn: The equation in the graph, which is assumed to match the output + equation of the pattern. + pattern: The pattern, which we are trying to match. + graph_matcher_rules: A :class:`~GraphMatcherRules` instance, which is used + for determining equivalence of individual Jax primitives. + matchable_graph_params: A subset of graph.params_vars consisting of + parameters that may appear in matches as parameters (not merely input + variables). + + Returns: + The variable to variable mapping between the pattern and graph variable, + if the pattern can be matched to the root equation, otherwise ``None``. + """ + + # Check the number of output variables match. + if len(pattern.jaxpr.outvars) != len(root_eqn.outvars): + return None + + # Set the current variables mapping to the output variables and then try to + # check the match from there. + match_variables_map = match_equations( + graph=graph, + current_variables_map=dict(zip(pattern.jaxpr.outvars, + root_eqn.outvars)), + reversed_eqns_to_match=tuple(reversed(pattern.jaxpr.eqns)), + input_vars=pattern.jaxpr.invars, + param_variables=pattern.param_vars, + graph_matcher_rules=graph_matcher_rules, + matchable_graph_params=matchable_graph_params, + ) + + if match_variables_map is None: + return None + + # Extract all the graph equations corresponding to the pattern. + graph_eqns = [] + for k, v in match_variables_map.items(): + + if (k not in pattern.graph.jaxpr.invars and + not isinstance(v, jex.core.Literal)): + + creation_op = graph.var_to_creation_op[v] + + assert isinstance(creation_op, JaxprEqn) + + graph_eqns.append(creation_op) + + return GraphMatch( + pattern=pattern, + variables_map=match_variables_map, + graph_eqns=graph_eqns, + ) + + +def find_layer_tags_and_patterns( + graph: JaxprGraph, + eqns_for_patterns: Sequence[JaxprEqn], + graph_matcher_rules: GraphMatcherComparator, + graph_patterns: Sequence[GraphPattern], + matchable_params: Set[Var], +) -> dict[Var, GraphMatch]: + """Tries to automatically match ``patterns_to_match`` in the Jaxpr graph. + + The method returns all newly discovered matches of any pattern. Each entry has + as a key the variable of the graph corresponding to the output of the pattern, + while each value is a triple ``(pattern, match_map, eqns)`` where ``pattern`` + is the :class:`~JaxprGraph` of the pattern that has been matched, + ``match_map`` is mapping the pattern variables to the corresponding graph + variables and ``eqns`` is the sequence of all graph equations corresponding to + the pattern equations. + + Args: + graph: The :class:`~JaxprGraph` on which we are searching for matching + equations. + eqns_for_patterns: All equation that should be considered for finding + a pattern. + graph_matcher_rules: A :class:`~GraphMatcherRules` instance, which is used + for determining equivalence of individual Jax primitives. + graph_patterns: A sequence of :class:`~GraphPattern` objects, which contain + all patterns to use, in order of precedence, which to try to find in the + graph before registering a parameter with a generic layer tag. + matchable_params: A subset of graph.params_vars consisting of parameters + that may appear in matches as parameters (not merely input variables). + + Returns: + See above. + """ + + # This list keeps track to any equations that are already in a pattern and + # hence should not be part of any other. + registered_equations = [] + + # First add any manual registrations to this. + for eqn in graph.manual_registrations: + + layer_data = tags.layer_eqn_data(eqn) + + for manual_eqn in graph.sub_graph_eqns( + layer_data.inputs + layer_data.params, layer_data.outputs + ): + registered_equations.append(manual_eqn) + + matches = {} + + # Loop through all equations in reverse, and for each one check every pattern + for eqn in reversed(eqns_for_patterns): + + if eqn in registered_equations or eqn.primitive.name in HIGHER_ORDER_NAMES: + continue + + for pattern in graph_patterns: + + match = match_pattern( + graph=graph, + root_eqn=eqn, + pattern=pattern, + graph_matcher_rules=graph_matcher_rules, + matchable_graph_params=matchable_params, + ) + + if match is not None: + + # Add all the match equations to the registered equations + registered_equations.extend(match.graph_eqns) + + # Add the match to the mapping of graph matches + matches[match.output_var] = match + + break + + return matches + + +def read_env( + env: dict[jex.core.Var, T], + variables: list[jax.core.Atom], +) -> list[T]: + """Reads from the variable-to-array environment during tracing.""" + result = [] + assert isinstance(variables, list) + for v in variables: + if isinstance(v, jex.core.Literal): + # Literals are values baked into the Jaxpr + result.append(v.val) + elif isinstance(v, DropVar): + result.append(None) + else: + result.append(env[v]) + return result + + +def write_env( + env: dict[jex.core.Var, T], + variables: list[jex.core.Var], + values: list[T], +) -> None: + """Writes to the variable-to-array environment during tracing.""" + assert len(variables) == len(values) + for variables, val in zip(variables, values): + env[variables] = val + + +def to_closed_jaxpr(jaxpr: JaxprOrClosedJaxpr) -> ClosedJaxpr: + if _JAXPR_TYPES_MERGED: + # In JAX 0.11 every Jaxpr carries its own (possibly empty) const values. + # Re-wrapping it would both be redundant and risk discarding those values. + return jaxpr + if isinstance(jaxpr, Jaxpr): + return _closed_jaxpr(jaxpr, []) + return jaxpr + + +def to_jaxpr_or_closed_jaxpr(closed_jaxpr: ClosedJaxpr, original: J) -> J: + if _JAXPR_TYPES_MERGED: + return closed_jaxpr + if isinstance(original, Jaxpr): + return closed_jaxpr.jaxpr + return closed_jaxpr + + +def apply_to_higher_order_primitives( + eqn: JaxprEqn, + func: Callable[[J], J]): + """Applies `func` only to higher order Jax primitives.""" + + if eqn.primitive.name not in HIGHER_ORDER_NAMES: + return eqn + + elif eqn.primitive.name == "cond": + params = dict(**eqn.params) + params["branches"] = tuple( + func(branch) for branch in params["branches"] + ) + return eqn.replace(params=params) + + elif eqn.primitive.name == "while": + params = dict(**eqn.params) + params["body_jaxpr"] = func(params["body_jaxpr"]) + return eqn.replace(params=params) + + elif eqn.primitive.name in ("scan", "pjit"): + params = dict(**eqn.params) + params["jaxpr"] = func(params["jaxpr"]) + return eqn.replace(params=params) + + elif eqn.primitive.name in ("xla_call", "xla_pmap"): + params = dict(**eqn.params) + params["call_jaxpr"] = func(params["call_jaxpr"]) + return eqn.replace(params=params) + + else: + raise NotImplementedError() + + +def clean_jaxpr( + jaxpr: J, + preserve_tags: bool = True, + outvar_is_dep: tuple[bool, ...] | None = None, +) -> J: + """Runs dead code elimination on a Jaxpr, retaining loss and layer tags.""" + + closed_jaxpr = to_closed_jaxpr(jaxpr) + eqns = [] + + if outvar_is_dep is None: + outvar_is_dep = (True,) * len(closed_jaxpr.jaxpr.outvars) + + final_outvars = [] + dependants = set() + + for var, is_dep in zip(closed_jaxpr.jaxpr.outvars, outvar_is_dep, + strict=True): + if is_dep: + + final_outvars.append(var) + + if not isinstance(var, jex.core.Literal): + dependants.add(var) + + for eqn in reversed(closed_jaxpr.jaxpr.eqns): + + # It's much more complicated to trace dependencies through *iterative* + # higher order primitives, so we don't do it. + if eqn.primitive.name in ITERATIVE_HIGHER_ORDER_NAMES: + outvar_is_dep_for_eqn = (True,) * len(eqn.outvars) + else: + outvar_is_dep_for_eqn = tuple(var in dependants for var in eqn.outvars) + + # Note that we currently only trace dependencies into higher order + # primitives, but not *through* them. If a single output of a higher order + # primitive is a dependency, then all of its inputs are treated as such too. + eqn = apply_to_higher_order_primitives( + eqn, + functools.partial( + clean_jaxpr, + outvar_is_dep=outvar_is_dep_for_eqn, + preserve_tags=preserve_tags + ), + ) + + if eqn.primitive.name in HIGHER_ORDER_NAMES: + + params = dict(**eqn.params) + + if "out_shardings" in params: + params["out_shardings"] = utils.filter_sequence(params["out_shardings"], + outvar_is_dep_for_eqn) + if "out_layouts" in params: + params["out_layouts"] = utils.filter_sequence(params["out_layouts"], + outvar_is_dep_for_eqn) + + eqn = eqn.replace( + outvars=utils.filter_sequence(eqn.outvars, outvar_is_dep_for_eqn), + params=params, + ) + + # else: + # assert all(outvar_is_dep_for_eqn) or not any(outvar_is_dep_for_eqn) + + check = False + + for v in eqn.outvars: + if v in dependants: + dependants.remove(v) + check = True + + if isinstance(eqn.primitive, (tags.LossTag, tags.LayerTag)): + check = check or preserve_tags + + if check: + eqns.append(eqn) + new_dependants = set(v for v in eqn.invars + if not isinstance(v, jex.core.Literal)) + dependants = dependants.union(new_dependants) + + # Dependants should only be invars + dependants = dependants - set(closed_jaxpr.jaxpr.invars + + closed_jaxpr.jaxpr.constvars) + + if dependants: + raise ValueError("Something went wrong with the dead code elimination.") + + closed_jaxpr = _closed_jaxpr( + closed_jaxpr.jaxpr.replace(eqns=list(reversed(eqns)), + outvars=final_outvars), + closed_jaxpr.consts, + ) + + return to_jaxpr_or_closed_jaxpr(closed_jaxpr, jaxpr) + + +def clean_layer_tags_jaxpr( + jaxpr: J, + only_remove_auto_tags: bool = False, +) -> tuple[J, tuple[tags.LayerTagEqn | JaxprEqn, ...]]: + """Returns a Jaxpr with layer tags removed, and the layer tags (which properly refer to variables in the returned Jaxpr).""" + + closed_jaxpr = to_closed_jaxpr(jaxpr) + eqns = [] + layer_tag_eqns = [] + var_map = {} + + for eqn in closed_jaxpr.jaxpr.eqns: + + if (isinstance(eqn.primitive, tags.LayerTag) + and (not only_remove_auto_tags or (eqn.params["meta"].name is not None + and "Auto" in eqn.params["meta"].name + ) + ) + ): + for ind1, ind2 in enumerate(eqn.params["meta"].outputs_index): + var_map[eqn.outvars[ind1]] = eqn.invars[ind2] + + else: + eqns.append(eqn) + + if isinstance(eqn.primitive, tags.LayerTag): + layer_tag_eqns.append(eqn) + + def remap_input_vars( + eqns: Sequence[JaxprEqn], var_map: Mapping[jex.core.Var, jex.core.Var] + ) -> list[JaxprEqn]: + """Remaps the input variables of a JaxprEqn. + + Args: + eqns: The list of JaxprEqns to remap. + var_map: A mapping from variables to new variables. + + Returns: + A new list of JaxprEqns with remapped input variables. + """ + + eqns_new = [] + + for eqn in eqns: + + new_invars = [] + for var in eqn.invars: + if not isinstance(var, jex.core.Literal) and var in var_map.keys(): + new_invars.append(var_map[var]) + else: + new_invars.append(var) + + eqns_new.append(eqn.replace(invars=new_invars)) + + return eqns_new + + eqns_new = remap_input_vars(eqns, var_map) + layer_tag_eqns_new = remap_input_vars(layer_tag_eqns, var_map) + + closed_jaxpr = _closed_jaxpr( + closed_jaxpr.jaxpr.replace(eqns=eqns_new), + closed_jaxpr.consts, + ) + + return ( + to_jaxpr_or_closed_jaxpr(closed_jaxpr, jaxpr), + tuple(layer_tag_eqns_new), + ) + + +# Prototype for clean_jaxpr using JAX's dce_jaxpr. Doesn't work because +# dce_jaxpr will remove any equations with no used outputs, regardless of the +# dce_rule for that equation's primitive. Adding an "effect" to loss/layer +# tags also won't work, because we sometimes actually do want to remove them +# from the graph (when preserve_tags is False). +# def clean_jaxpr( +# jaxpr: J, +# preserve_tags: bool = True, +# ) -> J: +# """Runs dead code elimination on a Jaxpr, retaining loss and layer tags.""" + +# def dce_jaxpr_tag_rule( +# used_outputs: list[bool], +# eqn: JaxprEqn +# ) -> tuple[list[bool], JaxprEqn | None]: + +# assert len(used_outputs) == len(eqn.outvars) + +# if any(used_outputs) or preserve_tags: +# return [True] * len(eqn.invars), eqn +# else: +# return [False] * len(eqn.invars), None + +# closed_jaxpr = to_closed_jaxpr(jaxpr) + +# pe.dce_rules[tags.LossTag] = dce_jaxpr_tag_rule +# pe.dce_rules[tags.LayerTag] = dce_jaxpr_tag_rule + +# cleaned_jaxpr, _ = pe.dce_jaxpr( +# closed_jaxpr.jaxpr, +# used_outputs=(True,) * len(closed_jaxpr.jaxpr.outvars), +# instantiate=True) + +# pe.dce_rules.pop(tags.LossTag) +# pe.dce_rules.pop(tags.LayerTag) + +# closed_jaxpr = ClosedJaxpr( +# jaxpr=cleaned_jaxpr, +# consts=closed_jaxpr.consts, +# ) + +# return to_jaxpr_or_closed_jaxpr(closed_jaxpr, jaxpr) + + +def merge_broadcasts_jaxpr(jaxpr: J) -> J: + """Merges consecutive broadcasts in the given Jaxpr.""" + + closed_jaxpr = to_closed_jaxpr(jaxpr) + + broadcasts_outputs = {} + eqns = list() + + for eqn in closed_jaxpr.jaxpr.eqns: + + eqn = apply_to_higher_order_primitives(eqn, merge_broadcasts_jaxpr) + + # We ignore broadcasting of constants + if (eqn.primitive.name == "broadcast_in_dim" and + not all(isinstance(v, jex.core.Literal) for v in eqn.invars)): + + if eqn.invars[0] in broadcasts_outputs: + # Construct a merged equation from the previous and current one + prev_eqn = broadcasts_outputs[eqn.invars[0]] + + broadcasts_outputs[eqn.outvars[0]] = prev_eqn.replace( + params={ + "shape": eqn.params["shape"], + "broadcast_dimensions": tuple( + eqn.params["broadcast_dimensions"][d] + for d in prev_eqn.params["broadcast_dimensions"] + ), + "sharding": None, + }, + outvars=eqn.outvars, + ) + + else: + broadcasts_outputs[eqn.outvars[0]] = eqn + + if eqn.outvars[0] in closed_jaxpr.jaxpr.outvars: + # We must preserve output equations + eqns.append(broadcasts_outputs[eqn.outvars[0]]) + + else: + for v in eqn.invars: + if not isinstance(v, jex.core.Literal) and v in broadcasts_outputs: + eqns.append(broadcasts_outputs[v]) + + eqns.append(eqn) + + closed_jaxpr = _closed_jaxpr( + closed_jaxpr.jaxpr.replace(eqns=eqns), + closed_jaxpr.consts, + ) + return to_jaxpr_or_closed_jaxpr(closed_jaxpr, jaxpr) + + +def num_unique_inputs(eqns: Sequence[JaxprEqn]) -> int: + n = 0 + vars_so_far = set() + for eqn in eqns: + for v in eqn.invars: + if v not in vars_so_far: + vars_so_far.add(v) + n += 1 + for v in eqn.outvars: + vars_so_far.add(v) + return n + +# _____ _ _ _ _ +# | __ \ (_) | | | | (_) +# | |__) |___ __ _ _ ___| |_ _ __ __ _| |_ _ ___ _ __ ___ +# | _ // _ \/ _` | / __| __| '__/ _` | __| |/ _ \| '_ \/ __| +# | | \ \ __/ (_| | \__ \ |_| | | (_| | |_| | (_) | | | \__ \ +# |_| \_\___|\__, |_|___/\__|_| \__,_|\__|_|\___/|_| |_|___/ +# __/ | +# |___/ + + +def _dense( + x: Array, + params: Sequence[Array], + axes: int, + with_reshape: bool, + ) -> Array: + """Example of a dense layer function.""" + # NOTE: This function uses `tensordot` so contracts over the last + # `axes` dimensions of `x` and the first `axes` dimensions of `params`. + match params: + case [w, b]: + y = jnp.tensordot(x, w, axes=axes) + if with_reshape: + return y + b.reshape((1,) * (y.ndim - b.ndim) + b.shape) + else: + return y + b + case [w]: + return jnp.tensordot(x, w, axes=axes) + case _: + raise ValueError("Unsupported parameters list") + + +def _dense_parameter_extractor( + reversed_eqns: Sequence[JaxprEqn], + variant: str = "dense", +) -> Mapping[str, Any]: + """Extracts all parameters from the `dot_general` operator.""" + n = num_unique_inputs(reversed_eqns[::-1]) + + for eqn in reversed_eqns: + if eqn.primitive.name == "dot_general": + return dict( + meta=tags.LayerMetaData( + variant=variant, + outputs_index=(0,), + inputs_index=(1,), + params_index=tuple(i + 2 for i in range(n - 1)), + ), + **eqn.params, + ) + assert False + + +def _make_general_dense_pattern( + with_bias: bool, + with_reshape: bool, + num_repeated_axes: int, + num_in_dims: int, + num_out_dims: int, +) -> GraphPattern: + """Creates a pattern for a dense or repeated dense layer.""" + batch_dim = (2,) + repeating_dims = tuple(itertools.repeat(7, num_repeated_axes)) + + out_dims = tuple(i + 2 for i in range(num_out_dims)) + in_dims = tuple(i + 2 for i in range(num_in_dims)) + x_shape = batch_dim + repeating_dims + in_dims + weight_shape = in_dims + out_dims + p_shapes = [weight_shape, out_dims] if with_bias else [weight_shape] + + name = "dense_with_bias" if with_bias else "dense_no_bias" + name = name + ("_with_reshape" if with_reshape else "_no_reshape") + + if num_repeated_axes > 0: + name = f"repeated[{num_repeated_axes}]_{name}" + variant = "repeated_dense" + else: + variant = "dense" + + return GraphPattern( + name=name, + tag_primitive=tags.layer_tag, + compute_func=functools.partial( + _dense, axes=num_in_dims, with_reshape=with_reshape), + parameters_extractor_func=functools.partial( + _dense_parameter_extractor, variant=variant), + example_args=[np.zeros(x_shape), [np.zeros(s) for s in p_shapes]], + ) + + +def _conv2d(x: Array, params: Sequence[Array], flax_style: bool) -> Array: + """Example of a conv2d layer function.""" + + w = params[0] + + y = jax.lax.conv_general_dilated( + x, + w, + window_strides=(2, 2), + padding="SAME", + dimension_numbers=("NHWC", "HWIO", "NHWC")) + + if len(params) == 1: + # No bias + return y + + # Add bias + if flax_style: + bias = params[1] + return y + bias.reshape((1,) * (y.ndim - bias.ndim) + bias.shape) + return y + params[1][None, None, None] + + +def _conv2d_parameter_extractor( + reversed_eqns: Sequence[JaxprEqn], + variant: str = "conv2d", +) -> Mapping[str, Any]: + """Extracts all parameters from the `conv_general_dilated` operator.""" + + n = num_unique_inputs(reversed_eqns[::-1]) + + for eqn in reversed_eqns: + if eqn.primitive.name == "conv_general_dilated": + return dict( + meta=tags.LayerMetaData( + variant=variant, + outputs_index=(0,), + inputs_index=(1,), + params_index=tuple(i + 2 for i in range(n - 1)), + ), + **eqn.params, + ) + + assert False + + +def _make_conv2d_pattern( + with_bias: bool, + flax_style: bool, +) -> GraphPattern: + + x_shape = [2, 8, 8, 5] + + p_shapes = ([[3, 3, 5, 4], [4]] if with_bias else + [[3, 3, 5, 4]]) + + return GraphPattern( + name="conv2d_with_bias" if with_bias else "conv2d_no_bias", + tag_primitive=tags.layer_tag, + compute_func=functools.partial(_conv2d, flax_style=flax_style), + parameters_extractor_func=_conv2d_parameter_extractor, + example_args=[np.zeros(x_shape), [np.zeros(s) for s in p_shapes]], + ) + + +def _scale_and_shift( + x: Array, + params: Sequence[Array], + has_scale: bool, + has_shift: bool, +) -> Array: + """Example of a scale and shift function.""" + + if has_scale and has_shift: + scale, shift = params + return x * scale + shift + + elif has_scale: + [scale] = params + return x * scale + + elif has_shift: + [shift] = params + return x + shift + + else: + raise ValueError("You must have either `has_scale` or `has_shift` set " + "to True.") + + +def _scale_and_shift_parameter_extractor( + reversed_eqns: Sequence[JaxprEqn], + variant: str = "scale_and_shift", +) -> Mapping[str, Any]: + """Extracts all parameters from the scale and shift operator.""" + + has_scale = False + + has_shift = False + + for eqn in reversed_eqns: + if eqn.primitive.name == "mul": + has_scale = True + elif eqn.primitive.name == "add": + has_shift = True + + return dict( + meta=tags.LayerMetaData( + variant=variant, + outputs_index=(0,), + inputs_index=(1,), + params_index=tuple(i + 2 for i in range(has_scale + has_shift)), + ), + has_scale=has_scale, + has_shift=has_shift, + ) + + +def _make_scale_and_shift_pattern( + broadcast_ndim: int, + has_scale: bool, + has_shift: bool, + p_dim: int = 13, +) -> GraphPattern: + """Creates a scale and shift graph pattern.""" + + assert broadcast_ndim >= 0 + + assert has_scale or has_shift + + x_shape = [i + 2 for i in range(broadcast_ndim)] + [p_dim] + p_shapes = [[p_dim], [p_dim]] if (has_scale and has_shift) else [[p_dim]] + + if has_scale and has_shift: + name = f"scale_and_shift_broadcast_{broadcast_ndim}" + elif has_scale: + name = f"scale_only_broadcast_{broadcast_ndim}" + elif has_shift: + name = f"shift_only_broadcast_{broadcast_ndim}" + else: + raise ValueError("Unreachable.") + + return GraphPattern( + name=name, + tag_primitive=tags.layer_tag, + compute_func=functools.partial( + _scale_and_shift, has_scale=has_scale, has_shift=has_shift), + parameters_extractor_func=_scale_and_shift_parameter_extractor, + example_args=[np.zeros(x_shape), [np.zeros(s) for s in p_shapes]], + ) + + +def _normalization_haiku_flax( + inputs: Sequence[Array], + params: Sequence[Array], + has_scale: bool, + has_shift: bool, + has_reshape: bool, +) -> Array: + """Example of normalization as is defined in Haiku/Flax.""" + + if len(params) not in (1, 2): + raise ValueError("The inputs to the `normalization_haiku` computation must " + f"have either 1 or 2 parameters, but got {len(params)}.") + + [inputs, rsqrt_var] = inputs + + if has_scale: + scale = params[0] + if has_reshape: + scale = scale.reshape( + [1] * (inputs.ndim - scale.ndim) + list(scale.shape)) + inv = scale * rsqrt_var + else: + inv = rsqrt_var + + outputs = inputs * inv + + if has_shift: + shift = params[1] + if has_reshape: + shift = shift.reshape( + [1] * (inputs.ndim - shift.ndim) + list(shift.shape)) + return outputs + shift + return outputs + + +def _normalization_haiku_preprocessor( + in_vars: Vars, + make_var_func: MakeVarFunc, +) -> tuple[tuple[Var, ...], JaxprEqns]: + """Preprocesses the inputs to a Haiku normalization layer. + + The standard ``scale_and_shift`` represents the following canonical + computation: + y = x * scale + shift + Normalization performs a similar computation, where the `normalized_x` below + represents the standard ``x`` input to ``scale_and_shift``: + normalized_x = (x - m) / sqrt(var(x) + eps) + y = normalized_x * scale + shift + Each ``layer_tag`` represents a specific computation and hence it expects its + inputs to be in canonical form. For ``scale_and_shift`` the input must be + the array that gets multiplied by the ``scale`` before the ``shift`` addition + as shown above. However, Haiku performs normalization slightly out of order: + y = [(x - m) * scale] / sqrt(var(x) + eps) + shift + As a result, in the Jax computation graph the canonical input (normalized_x) + does not exist, because of the ordering of the multiplication and division. + To remedy this we have to add this additional function, which to be able to + compute from the variables in the Haiku normalization computation, the + canonical input to ``scale_and_shift`` tag. + + Args: + in_vars: The input variables to the pattern. + make_var_func: A function to create correctly new variables. + + Returns: + The canonical input to ``scale_and_shift`` pattern. + """ + + [in_var, rsqrt_var, *param_vars] = in_vars + + # The equation below corresponds to the computation: + # normalized_inputs = inputs * rsqrt_var + + normalized_inputs_var = make_var_func(in_var.aval) + + normalized_inputs_eqn = new_jaxpr_eqn( + invars=[in_var, rsqrt_var], + outvars=[normalized_inputs_var], + primitive=jax.lax.mul_p, + params=dict(), + effects=set(), + ) + + return (normalized_inputs_var, *param_vars), [normalized_inputs_eqn] + + +def _make_normalization_haiku_flax_pattern( + broadcast_ndim: int, + has_reshape: bool, + p_dim: int = 13, + has_shift: bool = True, +) -> GraphPattern: + """Creates a pattern for a Haiku/Flax normalization layer.""" + + assert broadcast_ndim >= 0 + + x_shape = [i + 2 for i in range(broadcast_ndim)] + [p_dim] + + example_params = [np.zeros([p_dim])] + if has_shift: + example_params.append(np.zeros([p_dim])) + + return GraphPattern( + name=f"normalization_haiku_broadcast_{broadcast_ndim}", + tag_primitive=tags.layer_tag, + compute_func=functools.partial( + _normalization_haiku_flax, + has_scale=True, + has_shift=has_shift, + has_reshape=has_reshape), + parameters_extractor_func=_scale_and_shift_parameter_extractor, + example_args=[[np.zeros(x_shape), np.zeros(x_shape)], example_params], + in_values_preprocessor=_normalization_haiku_preprocessor + ) + +# NOTE: itertools iterates the last iterator first +# i.e. [(True, False), 0, 1, 1] [(True, False), 0, 1, 2] ... +DENSE_GRAPH_PATTERNS = tuple( + _make_general_dense_pattern( + with_bias=b, + with_reshape=r, + num_repeated_axes=rep, + num_in_dims=n_ins, + num_out_dims=n_outs) + for (b, r), rep, n_ins, n_outs in itertools.product( + ((True, False), (True, True), (False, False)), + range(3), + range(1, 3), + range(1, 3) + ) +) + +NORMALIZATION_GRAPH_PATTERNS = tuple( + _make_normalization_haiku_flax_pattern( + broadcast_ndim=n, + has_reshape=r, + has_shift=s) + for n, r, s in itertools.product( + range(2), + (False, True), + (False, True), + ) +) + +DEFAULT_GRAPH_PATTERNS = DENSE_GRAPH_PATTERNS + ( + _make_conv2d_pattern(True, False), + _make_conv2d_pattern(True, True), + _make_conv2d_pattern(False, False), + _make_scale_and_shift_pattern(1, True, True), + _make_scale_and_shift_pattern(0, True, True) + ) + +DEFAULT_GRAPH_PATTERNS += NORMALIZATION_GRAPH_PATTERNS + +DEFAULT_GRAPH_PATTERNS += ( + _make_scale_and_shift_pattern(1, True, False), + _make_scale_and_shift_pattern(0, True, False), + _make_scale_and_shift_pattern(1, False, True), + _make_scale_and_shift_pattern(0, False, True), +) + + +class TagLocation: + """Represents a tag location inside a function graph.""" + + def __init__( + self, + tag_eqn: JaxprEqn, + parent_equations: Sequence[tuple[JaxprEqn, int]] = (), + ): + self.tag_eqn = tag_eqn + self.parent_equations = list(parent_equations) + + @property + def base_name(self) -> str: + meta = self.tag_eqn.params.get("meta") + assert meta is not None and isinstance(meta, tags.LayerMetaData) + assert meta.name is not None, self.tag_eqn + return meta.name + + @property + def full_name(self) -> str: + """The full name of the tag location.""" + + prefix = "" + param_vars = self.bottom_level_parameters + + for eqn, n in reversed(self.parent_equations): + + assert eqn.primitive.name in HIGHER_ORDER_NAMES + + # Prefix for this higher order primitive + prefix = prefix + f"{eqn.primitive.name}_{n}/" + + if eqn.primitive.name == "cond": + raise NotImplementedError() + + elif eqn.primitive.name == "scan": + + p_indexes = [eqn.params["jaxpr"].jaxpr.invars.index(p) + for p in param_vars] + checks = [pi < _scan_num_consts(eqn) for pi in p_indexes] + + if not (all(checks) or all(not ci for ci in checks)): + raise ValueError("Parameters inside scan of the same tag are not both" + " carry or const.") + + if all(checks): + prefix = prefix + "const/" + else: + prefix = prefix + "carry/" + + elif eqn.primitive.name == "pjit": + p_indexes = [eqn.params["jaxpr"].jaxpr.invars.index(p) + for p in param_vars] + + elif eqn.primitive.name == "while": + p_indexes = [eqn.params["body_jaxpr"].jaxpr.invars.index(p) + for p in param_vars] + + elif eqn.primitive.name in ("xla_call", "xla_pmap"): + p_indexes = [eqn.params["call_jaxpr"].invars.index(p) + for p in param_vars] + + else: + raise NotImplementedError() + + param_vars = [eqn.invars[pi] for pi in p_indexes] + + return prefix + self.base_name + + @property + def bottom_level_parameters(self) -> tuple[Var, ...]: + """The bottom level variables of the tag location.""" + return tags.layer_eqn_data(self.tag_eqn).params + + @property + def top_level_parameters(self) -> tuple[Var, ...]: + """The top level parameter variables of the tag location.""" + + param_vars = self.bottom_level_parameters + + for eqn, _ in reversed(self.parent_equations): + + assert eqn.primitive.name in HIGHER_ORDER_NAMES + + if eqn.primitive.name == "cond": + raise NotImplementedError() + + elif eqn.primitive.name in ("scan", "pjit"): + invars = eqn.params["jaxpr"].jaxpr.invars + + elif eqn.primitive.name == "while": + invars = eqn.params["body_jaxpr"].jaxpr.invars + + elif eqn.primitive.name in ("xla_call", "xla_pmap"): + invars = eqn.params["call_jaxpr"].invars + + else: + raise NotImplementedError() + + # Indices inside of the higher order primitive + p_indexes = [invars.index(p) for p in param_vars] + + # Inputs (to the higher order primitive) corresponding to those indices + param_vars = tuple(eqn.invars[pi] for pi in p_indexes) + + return param_vars + + def add_parent_eqn(self, eqn: JaxprEqn, counter: int): + assert eqn.primitive.name in HIGHER_ORDER_NAMES + self.parent_equations.append((eqn, counter)) + + +class TaggedFunction: + """Represents a function that has been processed and auto tagged.""" + + def __init__( + self, + func_graph: JaxprGraph, + tag_locations: Sequence[TagLocation], + ): + self._func_graph = func_graph + self._tag_locations = tag_locations + self._flat_func = jex.core.jaxpr_as_fun(func_graph.closed_jaxpr) + self._param_labels = self._compute_parameter_labels() + + def __call__(self, *args, **kwargs): + flat_args = jax.tree_util.tree_leaves(args) + flat_output = self._flat_func(*flat_args) + return jax.tree_util.tree_unflatten(self._func_graph.out_tree, flat_output) + + def _compute_parameter_labels(self) -> Mapping[Var, Sequence[str]]: + """Computes the parameter labels as a dict from params to strings.""" + + # Collect all registrations for every tagged parameter + tagged_params = {} + + for tag_l in self._tag_locations: + for p in tag_l.top_level_parameters: + + assert p in self._func_graph.params_vars + + if p not in tagged_params: + tagged_params[p] = [] + + tagged_params[p].append(tag_l.full_name) + + return tagged_params + + def print_parameter_tags(self): + """Prints all the parameter registrations.""" + # Print all tag parameter registrations + + labels = ["|".join(self._param_labels.get(p, ["Orphan"])) + for p in self._func_graph.params_vars] + logging.info("=" * 50) + logging.info("Graph parameter registrations:") + + for line in pprint.pformat(jax.tree_util.tree_unflatten( + self._func_graph.params_tree, labels, + )).split("\n"): + logging.info("%s", line) + + logging.info("=" * 50) + + def check_multiple_registrations(self): + for p in self._func_graph.params_vars: + if len(self._param_labels[p]) > 1: + raise ValueError(f"Parameter {p} has been registered to multiple tags: " + f"{self._param_labels[p]}.") + + +def _auto_register_tags( + graph: JaxprGraph, + graph_matcher_rules: GraphMatcherComparator, + graph_patterns: Sequence[GraphPattern], + register_orphans: bool, + register_only_until_losses: bool, + matchable_params: Set[Var], +) -> tuple[JaxprGraph, Sequence[TagLocation]]: + """Internal function for automatic registration of layer tags.""" + + higher_counters = { + "cond": 0, + "while": 0, + "scan": 0, + "pjit": 0, + "xla_call": 0, + "xla_pmap": 0, + } + + # Extract the sub-graph that leads to losses + if register_only_until_losses: + + eqns_for_registration = [] + sub_graph_vars = set() + for eqn in reversed(graph.jaxpr.eqns): + + # Note that graph.losses_eqns won't recurse into higher order primitives + # to find loss tags, so any losses defined inside such primitives will be + # effectively chopped out. + + if (eqn in graph.losses_eqns or + any(v in sub_graph_vars for v in eqn.outvars)): + + eqns_for_registration.append(eqn) + sub_graph_vars.update( + v for v in eqn.invars if not isinstance(v, jex.core.Literal)) + + eqns_for_registration = eqns_for_registration[::-1] + + else: + eqns_for_registration = graph.jaxpr.eqns + + # Count number of uses of each parameter and if it exceeds 1, we don't do any + # automatic registration. Note that we don't have to recurse into higher order + # primitives to count uses because we only care about whether there is more + # than one use of a given parameter. If there is, and they are not all inside + # of one higher order primitive, we will catch it here. If they're all in one + # higher order primitive, then we will catch that when we recursively call + # _auto_register_tags, and no registrations will happen at that level. + # Finally, the parameter in question will be seen as an orphan and registered + # as generic *only* at the top-level call of _auto_register_tags, as intended + # (since register_orphans is True only at the top-level call). + param_uses = collections.Counter() + for eqn in eqns_for_registration: + if not isinstance(eqn.primitive, tags.LayerTag): + for v in eqn.invars: + if v in graph.params_vars: + param_uses[v] += 1 + + manual_registrations = graph.manual_registrations + manually_tagged_params = set() + for eqn in manual_registrations: + for p in tags.layer_eqn_data(eqn).params: + manually_tagged_params.add(p) + + # Parameters that are eligible for auto-matching are those that have exactly + # one use in the graph and are not manually tagged. + single_use_params = {p for p in graph.params_vars if param_uses[p] <= 1} + matchable_params = (single_use_params & matchable_params) - manually_tagged_params # pylint: disable=line-too-long + + # Process all higher order primitives + eqns = [] + tag_locations = [] + for eqn in graph.jaxpr.eqns: + + if not (eqn in eqns_for_registration + and eqn.primitive.name in HIGHER_ORDER_NAMES): + + eqns.append(eqn) + + continue + + eqn_name = eqn.primitive.name + if eqn_name == "cond": + sub_jaxprs = eqn.params["branches"] + elif eqn_name == "while": + sub_jaxprs = [eqn.params["body_jaxpr"]] + elif eqn_name in ("scan", "pjit"): + sub_jaxprs = [eqn.params["jaxpr"]] + elif eqn_name in ("xla_call", "xla_pmap"): + sub_jaxprs = [eqn.params["call_jaxpr"]] + else: + raise NotImplementedError() + + final_jaxprs = [] + final_tag_locations = [] + for original_jaxpr in sub_jaxprs: + + sub_jaxpr = to_closed_jaxpr(original_jaxpr) + + sub_params_vars = [] + sub_matchable_params = set() + for outer_v, inner_v in zip(eqn.invars, sub_jaxpr.jaxpr.invars): + + if outer_v in graph.params_vars: + sub_params_vars.append(inner_v) + + if isinstance(outer_v, Var) and outer_v in matchable_params: + assert isinstance(inner_v, Var) + sub_matchable_params.add(inner_v) + + sub_graph, sub_tag_locations = _auto_register_tags( + graph=JaxprGraph( + name=graph.name + f"_{eqn_name}", + closed_jaxpr=sub_jaxpr, + params_tree=jax.tree_util.tree_structure(sub_params_vars), + params_vars=sub_params_vars, + out_tree=jax.tree_util.tree_structure(sub_jaxpr.jaxpr.outvars), + tag_ctor=None, + ), + graph_matcher_rules=graph_matcher_rules, + graph_patterns=graph_patterns, + register_orphans=False, + register_only_until_losses=False, + matchable_params=sub_matchable_params, + ) + + final_jaxprs.append( + to_jaxpr_or_closed_jaxpr(sub_graph.closed_jaxpr, original_jaxpr)) + + final_tag_locations.append(sub_tag_locations) + + if eqn_name == "cond": + if final_tag_locations[0] or final_tag_locations[1]: + # TODO(botev): We need to check each branch has identical registrations + raise NotImplementedError() + sub_tag_locations = [] + else: + # Extract the sub jaxpr parameter tag registrations and input vars + [sub_tag_locations] = final_tag_locations # pylint:disable=unbalanced-tuple-unpacking + + del final_tag_locations + + # Update the jaxpr parameter in the equation + eqn_params = dict(**eqn.params) + if eqn_name == "cond": + eqn_params["branches"] = tuple(final_jaxprs) + elif eqn_name == "while": + [eqn_params["body_jaxpr"]] = final_jaxprs # pylint:disable=unbalanced-tuple-unpacking + elif eqn_name in ("scan", "pjit"): + [eqn_params["jaxpr"]] = final_jaxprs # pylint:disable=unbalanced-tuple-unpacking + elif eqn_name in ("xla_call", "xla_pmap"): + [eqn_params["call_jaxpr"]] = final_jaxprs # pylint:disable=unbalanced-tuple-unpacking + else: + raise NotImplementedError() + + eqns.append(eqn.replace(params=eqn_params)) + + del final_jaxprs + + # Insert the sub-registrations into the tagged_params + for tag_l in sub_tag_locations: + tag_l.add_parent_eqn(eqns[-1], higher_counters[eqn_name]) + higher_counters[eqn_name] = higher_counters[eqn_name] + 1 + tag_locations.append(tag_l) + + # Make a new graph with the replaced higher order equations + mid_graph = JaxprGraph( + name=graph.name, + closed_jaxpr=_closed_jaxpr( + graph.jaxpr.replace(eqns=eqns), + graph.consts, + ), + params_tree=graph.params_tree, + params_vars=graph.params_vars, + out_tree=graph.out_tree, + tag_ctor=None, + ) + del graph + + # Find matches + matches = find_layer_tags_and_patterns( + graph=mid_graph, + eqns_for_patterns=eqns_for_registration, + graph_matcher_rules=graph_matcher_rules, + graph_patterns=graph_patterns, + matchable_params=matchable_params, + ) + + tagged_params = set() + + # Registrations in higher order primitives + for tag_l in tag_locations: + for p in tag_l.top_level_parameters: + tagged_params.add(p) + + # Manual registrations + for manual_eqn in manual_registrations: + for p in tags.layer_eqn_data(manual_eqn).params: + tagged_params.add(p) + + # Automatically detected registrations + for match in matches.values(): + for p in match.param_graph_variables: + tagged_params.add(p) + + # Create the Jaxpr with all the tag registrations + make_var_func = gensym() + eqns = list() + env = {} + pattern_counters = {} + + if register_orphans: + + for param in mid_graph.params_vars: + + if param not in tagged_params: + + orphan_p = make_var_func(param.aval) + + n = pattern_counters.get("generic", 0) + pattern_counters["generic"] = n + 1 + + eqns.append( + new_jaxpr_eqn( + invars=[param], + outvars=[orphan_p], + primitive=tags.layer_tag, + params=dict( + meta=tags.LayerMetaData( + variant="generic", + inputs_index=(), + outputs_index=(0,), + params_index=(0,), + name=f"Auto[generic({n})]", + ) + ), + effects=set(), + ) + ) + + env[param] = orphan_p + tag_locations.append(TagLocation(eqns[-1])) + + for eqn in mid_graph.jaxpr.eqns: + + invars = [env.get(v, v) if isinstance(v, Var) else v + for v in eqn.invars] + + eqns.append(eqn.replace(invars=invars)) + + if isinstance(eqn.primitive, tags.LayerTag): + + # Mark manual registrations + meta = eqns[-1].params.get("meta") + assert meta is not None and isinstance(meta, tags.LayerMetaData) + + if meta.name is None: + n = pattern_counters.get(meta.variant, 0) + pattern_counters[meta.variant] = n + 1 + meta.name = f"Manual[{meta.variant}({n})]" + + tag_locations.append(TagLocation(eqn)) + + for var in eqn.outvars: + + # Check if this is a match of a graph pattern + match = matches.get(var) + + if match is not None: + + for additional_eqn in match.create_eqns_and_update_env(env, + make_var_func): + eqns.append(additional_eqn) + + # Mark automatic registration + meta = eqns[-1].params.get("meta") + assert meta is not None and isinstance(meta, tags.LayerMetaData) + assert meta.name is None + n = pattern_counters.get(meta.variant, 0) + pattern_counters[meta.variant] = n + 1 + meta.name = (f"Auto[tag_variant={meta.variant}({n})|" + f"match_type={match.name}]") + tag_locations.append(TagLocation(eqns[-1])) + + final_outvars = [env.get(v, v) if isinstance(v, Var) else v + for v in mid_graph.jaxpr.outvars] + + final_graph = JaxprGraph( + name=mid_graph.name, + closed_jaxpr=_closed_jaxpr( + mid_graph.jaxpr.replace(eqns=eqns, outvars=final_outvars), + mid_graph.closed_jaxpr.consts, + ), + params_tree=mid_graph.params_tree, + params_vars=mid_graph.params_vars, + out_tree=mid_graph.out_tree, + tag_ctor=None, + ) + + return final_graph, tag_locations + + +def auto_register_tags( + func: utils.Func, + func_args: utils.FuncArgs, + params_index: int = 0, + register_only_generic: bool = False, + compute_only_loss_tags: bool = True, + patterns_to_skip: Sequence[str] = (), + graph_matcher_rules: GraphMatcherComparator = GraphMatcherComparator(), + graph_patterns: Sequence[GraphPattern] = DEFAULT_GRAPH_PATTERNS, +) -> TaggedFunction: + """Transforms the function by automatically registering layer tags. + + Args: + func: The original function to transform. + func_args: Example arguments to ``func`` which to be used for tracing it. + params_index: Specifies, which inputs to the function are to be considered + a parameter variable. Specifically - ``inputs[params_index]``. + register_only_generic: If ``True`` registers all parameters not already in a + layer tag with a generic tag, effectively ignoring ``graph_patterns``. + compute_only_loss_tags: If set to ``True`` (default) the resulting function + will only compute the loss tags in ``func``, not its full computation and + actual output. + patterns_to_skip: The names of any patterns from the provided list, which to + be skipped/not used during the pattern matching. + graph_matcher_rules: A :class:`~GraphMatcherRules` instance, which is used + for determining equivalence of individual Jax primitives. + graph_patterns: A sequence of :class:`~GraphPattern` objects, which contain + all patterns to use, in order of precedence, which to try to find in the + graph before registering a parameter with a generic layer tag. + Returns: + A transformed function as described above. + """ + + graph = make_jax_graph( + func=func, + func_args=func_args, + params_index=params_index, + name="main", + compute_only_loss_tags=compute_only_loss_tags, + clean_broadcasts=True, + ) + + patterns = () if register_only_generic else tuple( + pattern for pattern in graph_patterns + if pattern.name not in patterns_to_skip) + + func_graph, tagged_locations = _auto_register_tags( + graph=graph, + graph_matcher_rules=graph_matcher_rules, + graph_patterns=patterns, + register_orphans=True, + register_only_until_losses=True, + matchable_params=set(graph.params_vars), + ) + + func = TaggedFunction( + func_graph=func_graph, + tag_locations=tagged_locations, + ) + func.print_parameter_tags() + + func.check_multiple_registrations() + + return func diff --git a/src/kfac_jax/_src/tracer.py b/src/kfac_jax/_src/tracer.py new file mode 100644 index 0000000000000000000000000000000000000000..52530db52a4085f145279fe19b2681e41d8c5870 --- /dev/null +++ b/src/kfac_jax/_src/tracer.py @@ -0,0 +1,1398 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC tracing functionality for functions needed for curvature estimation.""" +from collections.abc import Iterable +import dataclasses +import functools +import itertools +from typing import Any, Callable, Generic, Sequence, TypeVar + +from absl import logging +import jax +import jax.extend as jex +import jax.numpy as jnp +from kfac_jax._src import layers_and_loss_tags as tags +from kfac_jax._src import loss_functions +from kfac_jax._src import tag_graph_matcher as tgm +from kfac_jax._src import utils +from typing_extensions import TypeAlias + +# Types for annotations +T = TypeVar("T") +Array = utils.Array +Shape = utils.Shape +Params = utils.Params +FuncArgs = utils.FuncArgs +FuncOuts = utils.FuncOuts +Var = jex.core.Var +LossFunction = loss_functions.LossFunction +LossFunctionInputs = loss_functions.LossFunctionInputs + + +ProcJaxpr: TypeAlias = "ProcessedJaxpr" +TaggedFunction = Callable[..., tuple[LossFunction, ...]] +Func = Callable[..., T] + +FunctionTransformation = Callable[..., T] +TransformedFunction = Callable[..., T] +JaxprExtractor = Callable[..., "ProcessedJaxpr"] + + +@dataclasses.dataclass(frozen=True, kw_only=True, unsafe_hash=True) +@jax.tree_util.register_pytree_node_class +class LayerVjpData(Generic[T]): + """A compact class for all data related to layer tag information during VJP.""" + primals: tags.LayerData[T] + tangents: tags.LayerData[T] + + def tree_flatten(self) -> tuple[ + tuple[tags.LayerData[T], tags.LayerData[T]], + None, + ]: + return (self.primals, self.tangents), None + + @classmethod + def tree_unflatten(cls, aux_data, children): + assert aux_data is None + primals, tangents = children + return cls(primals=primals, tangents=tangents) + + +LossTagsVjp = tuple[ + tuple[LossFunction, ...], Callable[[Sequence[LossFunctionInputs]], Params] +] +LossTagsJvp = tuple[ + tuple[LossFunction, ...], + tuple[LossFunctionInputs, ...], +] +LayerTagVjp = tuple[ + tuple[LossFunction, ...], + Callable[ + [tuple[LossFunctionInputs, ...]], + tuple[LayerVjpData[Array], ...], # pytype: disable=invalid-annotation + ], +] +LayerTagVjpAndValueAndGrad = tuple[ + tuple[LossFunction, ...], + Callable[ + [tuple[LossFunctionInputs, ...]], + tuple[LayerVjpData[Array], ...], # pytype: disable=invalid-annotation + ], + Array, + Params, +] +JaxprOrClosedJaxpr = jex.core.Jaxpr | jex.core.ClosedJaxpr + + +def shape_and_type(x: Array) -> tuple[Shape, jnp.dtype]: + """Returns the shape and type of the given array.""" + return x.shape, x.dtype + + +def make_cache_key( + func_args: FuncArgs, *args: Any +) -> tuple[utils.PyTreeDef, tuple[tuple[Shape, jnp.dtype], ...]]: + """Creates a key for caching Jax function arguments.""" + + args_flat, tree_structure = jax.tree_util.tree_flatten((func_args, args)) + + return tree_structure, tuple(map(shape_and_type, args_flat)) + + +def extract_tags( + jaxpr: jex.core.Jaxpr, +) -> tuple[tuple[tags.LayerTagEqn, ...], tuple[tags.LossTagEqn, ...]]: + """Extracts the layer and the loss tags from the given Jaxpr.""" + + return ( + tuple( + eqn for eqn in jaxpr.eqns if isinstance(eqn.primitive, tags.LayerTag) + ), + tuple( + eqn for eqn in jaxpr.eqns if isinstance(eqn.primitive, tags.LossTag) + ), + ) + + +def name_layer_tags(layer_tags: tuple[tags.LayerTagEqn, ...]) -> None: + # This adds names to all registrations when `auto_register_tags=False` + tag_counter = {} + for layer_tag in layer_tags: + meta = layer_tag.params.get("meta") + if meta is None: + raise ValueError("Layer tag %s has no meta parameter" % layer_tag) + assert isinstance(meta, tags.LayerMetaData) + if meta.name is None: + n = tag_counter.get(layer_tag.primitive.name, 0) + tag_counter[layer_tag.primitive.name] = n + 1 + meta.name = f"Manual[{layer_tag.primitive.name}|{n}]" + + +def order_layer_tags( + params_vars_flat: Sequence[Var], + layer_tags: Sequence[tags.LayerTagEqn], + allow_left_out_params: bool = False, +) -> tuple[tuple[tags.LayerTagEqn, ...], tuple[tuple[int, ...], ...]]: + """Sorts the layer tags based on the index of the parameters they contain. + + Args: + params_vars_flat: A sequence of all parameter variables. + layer_tags: A sequence of all layer tags. + allow_left_out_params: Whether to raise an error if there are any parameter + variables which are not part of a layer tag. + + Returns: + A pair of tuples ``(layer_tags, tags_indices)``, where ``layer_tags`` has + the ordered sequence of the input ``layer_tags`` and ``tags_indices`` + contains a sequence of tuples, where each tuple has the indices of the + parameters associated with the corresponding layer tag. + """ + tags_param_indices = [] + used_indices = set() + + for eqn in layer_tags: + + # Collect the equation parameter indices + tag_vars = tags.layer_eqn_data(eqn).params + vars_indices = tuple(params_vars_flat.index(v) for v in tag_vars) + meta = eqn.params.get("meta") + if meta is None or not isinstance(meta, tags.LayerMetaData): + raise ValueError(f"Layer tag {eqn} has invalid metadata.") + meta.params_canonical_order = tuple( + i for i, _ in sorted(enumerate(vars_indices), key=lambda x: x[1]) + ) + + if any(i in used_indices for i in vars_indices): + raise ValueError("Reusing variable in a second block.") + + used_indices = used_indices.union(vars_indices) + tags_param_indices.append(vars_indices) + + left_out_indices = set(range(len(params_vars_flat))) - used_indices + + if left_out_indices and not allow_left_out_params: + raise ValueError( + "The following parameter indices were not assigned a " + f"block: {left_out_indices}." + ) + + if not layer_tags: + return (), () + else: + # Sort by the vars minimum index + sorted_index_and_blocks = sorted( + zip(layer_tags, tags_param_indices), key=lambda x: min(x[1]) + ) + return tuple(zip(*sorted_index_and_blocks)) + + +class ProcessedJaxpr(utils.Finalizable): + """A wrapper around Jaxpr, with useful additional data. + + Attributes: + jaxpr: The original Jaxpr that is being wrapped. + consts: The constants returned from the tracing of the original Jaxpr. + in_tree: The PyTree structure of the inputs to the function that the + original Jaxpr has been created from. + params_index: Specifies, which inputs to the function are to be considered a + parameter variable. Specifically - ``inputs[params_index]``. + loss_tags: A tuple of all of the loss tags in the original Jaxpr. + layer_tags: A sorted tuple of all of the layer tags in the original Jaxpr. + The sorting order is based on the indices of the parameters associated + with each layer tag. + layer_indices: A sequence of tuples, where each tuple has the indices of the + parameters associated with the corresponding layer tag. + """ + + def __init__( + self, + jaxpr: jex.core.Jaxpr, + consts: list[Any], + in_tree: utils.PyTreeDef, + params_index: int, + allow_left_out_params: bool = False, + ): + """Initializes the instance. + + Args: + jaxpr: The raw Jaxpr. + consts: The constants needed for evaluation of the raw Jaxpr. + in_tree: The PyTree structure of the inputs to the function that the + ``jaxpr`` has been created from. + params_index: Specifies, which inputs to the function are to be considered + a parameter variable. Specifically - ``inputs[params_index]``. + allow_left_out_params: Whether to raise an error if any of the parameter + variables is not included in any layer tag. + """ + + super().__init__() + + self.jaxpr = jaxpr + self.consts = consts + self.in_tree = in_tree + self.params_index = params_index + + # Positional construction works before and after JAX 0.11's + # Jaxpr/ClosedJaxpr merge; keyword ``jaxpr=`` no longer does. + closed_jaxpr = jex.core.ClosedJaxpr(self.jaxpr, self.consts) + self.jaxpr, self.layer_tags = tgm.clean_layer_tags_jaxpr(closed_jaxpr) + self.jaxpr = self.jaxpr.jaxpr + + _, self.loss_tags = extract_tags(self.jaxpr) + + name_layer_tags(self.layer_tags) + + self.layer_tags, self.layer_indices = order_layer_tags( + params_vars_flat=self.params_vars_flat, + layer_tags=self.layer_tags, + allow_left_out_params=allow_left_out_params, + ) + + self.finalize() + + @property + def in_vars_flat(self) -> list[Var]: + """A flat list of all of the abstract input variables.""" + return self.jaxpr.invars + + @property + def in_vars(self) -> utils.PyTree[Var]: + """The abstract input variables, as an un-flatten structure.""" + return jax.tree_util.tree_unflatten(self.in_tree, self.in_vars_flat) + + @property + def params_vars(self) -> utils.PyTree[Var]: + """The abstract parameter variables, as an un-flatten structure.""" + return self.in_vars[self.params_index] + + @property + def params_vars_flat(self) -> list[Var]: + """A flat list of all abstract parameter variables.""" + return jax.tree_util.tree_leaves(self.params_vars) + + @property + def params_tree(self) -> utils.PyTreeDef: + """The PyTree structure of the parameter variables.""" + return jax.tree_util.tree_structure(self.params_vars) + + def log_registered_losses(self): + logging.info("Graph registered losses:") + + for loss_tag in self.loss_tags: + meta = loss_tag.params.get("meta") + assert meta is not None and isinstance(meta, tags.LossMetaData) + assert len(loss_tag.invars) == len(meta.argument_names) + + args = [] + for name, var in zip(meta.argument_names, loss_tag.invars): + args.append(f"{name}={var}") + + args_str = ", ".join(args) + + logging.info("%s(%s)", tags.loss_eqn_class_name(loss_tag), args_str) + + logging.info("=" * 50) + + def reconstruct_losses( + self, + losses_inputs: tuple[LossFunctionInputs, ...], + ) -> tuple[LossFunction, ...]: + losses = [] + + for eqn, loss_args in zip(self.loss_tags, losses_inputs): + loss: LossFunction = tags.loss_eqn_construct_loss(eqn, *loss_args) + losses.append(loss) + + return tuple(losses) + + @classmethod + def make_from_func( + cls, + func: Func[Any], + func_args: FuncArgs, + params_index: int = 0, + auto_register_tags: bool = True, + allow_left_out_params: bool = False, + **auto_registration_kwargs: Any, + ) -> ProcJaxpr: + """Constructs a :class:`~ProcessedJaxpr` from a the given function. + + Args: + func: The model function, which will be traced. + func_args: Function arguments to use for tracing. + params_index: The variables from the function arguments which are at this + index (e.g. ``func_args[params_index]``) are to be considered model + parameters. + auto_register_tags: Whether to run an automatic layer registration on the + function (e.g. :func:`~auto_register_tags`). + allow_left_out_params: If this is set to ``False`` an error would be + raised if there are any model parameters that have not be assigned to a + layer tag. + **auto_registration_kwargs: Any additional keyword arguments, to be passed + to the automatic registration pass. + + Returns: + A :class:`~ProcessedJaxpr` representing the model function. + """ + + func_args = tuple(func_args) + + if auto_register_tags: + func = tgm.auto_register_tags( + func=func, + func_args=func_args, + params_index=params_index, + **auto_registration_kwargs, + ) + + typed_jaxpr = jax.make_jaxpr(func)(*func_args) + jaxpr, consts = typed_jaxpr.jaxpr, typed_jaxpr.literals + + in_tree = jax.tree_util.tree_structure(func_args) + + processed_jaxpr = ProcessedJaxpr( + jaxpr=jaxpr, + consts=consts, + in_tree=in_tree, + params_index=params_index, + allow_left_out_params=allow_left_out_params, + ) + processed_jaxpr.log_registered_losses() + + return processed_jaxpr + + def __eq__(self, other: ProcJaxpr) -> bool: + """Compares two ProcessedJaxpr instances by tree structure.""" + + # Verify whether input trees are equivalent + if self.in_tree != other.in_tree: + return False + + # Verify whether layer indices are equivalent + if len(self.layer_indices) != len(other.layer_indices): + return False + + for ref_l_index, l_index in zip(self.layer_indices, other.layer_indices): + + if len(ref_l_index) != len(l_index): + return False + + if any(p_i != p_j for p_i, p_j in zip(ref_l_index, l_index)): + return False + + # Verify layer tags are equivalent + if len(self.layer_tags) != len(other.layer_tags): + return False + + if any( + ref_tag.primitive != tag.primitive + for ref_tag, tag in zip(self.layer_tags, other.layer_tags) + ): + return False + + # Verify whether parameter shapes are equivalent + if any( + p_i.aval.shape != p_j.aval.shape # pytype: disable=attribute-error + for p_i, p_j in zip(self.params_vars_flat, other.params_vars_flat) + ): + return False + + return True + + +def cached_transformation( + func: Func[T], + transformation: FunctionTransformation[T], + params_index: int = 0, + auto_register_tags: bool = True, + allow_left_out_params: bool = False, + allow_no_losses: bool = False, + raise_error_on_diff_jaxpr: bool = True, + **auto_registration_kwargs: Any, +) -> tuple[TransformedFunction[T], JaxprExtractor]: + """Caches ``transformation(preprocessed_jaxpr, func_args, *args)``. + + The caching mechanism uses the ``func_args`` PyTree, dtypes and shapes for + hashing. + + Args: + func: The main model function, which will be transformed. + transformation: The actual transformation of ``func``. + params_index: The variables from the function arguments which are at this + index (e.g. ``func_args[params_index]``) are to be considered model + parameters. + auto_register_tags: Whether to run an automatic layer registration on the + function (e.g. :func:`~auto_register_tags`). + allow_left_out_params: If this is set to ``False`` an error would be raised + if there are any model parameters that have not be assigned to a layer + tag. + allow_no_losses: If this is set to ``False`` an error would be raised if no + registered losses have been found when tracing the function. + raise_error_on_diff_jaxpr: Whether to raise an exception if the function has + been traced before, with different arguments, and the new Jaxpr graph + differs in more than just the shapes and dtypes of the Jaxpr equations. + **auto_registration_kwargs: Any additional keyword arguments, to be passed + to the automatic registration pass. + + Returns: + A function with a signature ``f(func_args, *args, return_only_jaxpr)`` which + evaluates the transformation of ``func`` at ``func_args``. The extra + ``args`` are any additional array arguments passed to the transformation, + while the last flag indicates whether to just return the + :class:`~ProcessedJaxpr` instead of the transformation output. Also returns + a function that returns the processed Jaxpr of `func` for a given set of + function arguments. + """ + cache = {} + + def retrieve(func_args): + + # Construct a key and check cache for hits + key = make_cache_key(func_args) + + if key not in cache: + + # Process the function + processed_jaxpr = ProcessedJaxpr.make_from_func( + func=func, + func_args=func_args, + params_index=params_index, + auto_register_tags=auto_register_tags, + allow_left_out_params=allow_left_out_params, + **auto_registration_kwargs, + ) + + if not allow_no_losses and not processed_jaxpr.loss_tags: + raise ValueError("No registered losses have been found during tracing.") + + if cache and raise_error_on_diff_jaxpr: + + # If any previous `ProcessedJaxpr` exists verify that it is equivalent + ref_jaxpr, _ = cache[next(iter(cache))] + + if ref_jaxpr != processed_jaxpr: + raise ValueError( + "The consecutive tracing of the provided function " + "yielded a non-equivalent `ProcessedJaxpr`." + ) + + f = functools.partial(transformation, processed_jaxpr) + cache[key] = (processed_jaxpr, f) + + return cache[key] + + @functools.wraps(transformation) + def wrapped_transformation(func_args: FuncArgs, *args: Any) -> T: + _, f = retrieve(func_args) + return f(func_args, *args) + + def get_processed_jaxpr(func_args: FuncArgs, *_: Any) -> ProcessedJaxpr: + closed_jaxpr, _ = retrieve(func_args) + return closed_jaxpr + + return wrapped_transformation, get_processed_jaxpr + + +def construct_compute_losses_inputs( + processed_jaxpr: ProcessedJaxpr, + primal_func_args: FuncArgs, + params_index: int, + drop_loss_tags: bool = True, +) -> Callable[ + [Params], + tuple[tuple[LossFunctionInputs, ...], tuple[LossFunctionInputs, ...]], +]: + """Constructs a function that computes the inputs to all loss tags. + + The returned function takes as input only the parameters, as specified by + ``params_index``, and returns a tuple containing the input values to the first + ``num_losses`` loss tags in the Jaxpr. This is done by iterating sequentially + over all equations in the Jaxpr, evaluating each equation, until the correct + number of loss tags have been discovered and returning the values of their + inputs. + + Args: + processed_jaxpr: The `ProcessedJaxpr` representing the function. + primal_func_args: The concrete values for the inputs to the Jaxpr. + params_index: The variables from the function arguments which are at this + index (e.g. ``func_args[params_index]``) are to be considered model + parameters. + drop_loss_tags: Whether to remove the loss tags primitive when computing the + loss functions inputs. + + Returns: + A function which computes the inputs to the first ``num_losses`` loss tags. + """ + + def forward_compute_losses( + primal_params: Params, + ) -> tuple[tuple[LossFunctionInputs, ...], tuple[LossFunctionInputs, ...]]: + """Computes and returns the inputs to the first ``num_losses`` loss tags.""" + + # Check the provided inputs match the original primals. + local_func_args = list(primal_func_args) + original_params = local_func_args[params_index] + + if not utils.abstract_objects_equal(original_params, primal_params): + raise ValueError( + "The `primal_params` should have the same abstract " + "structure as the original parameters passed in to the " + "function." + ) + + local_func_args[params_index] = primal_params + flat_args = jax.tree_util.tree_leaves(local_func_args) + + # Mapping from variable -> value + env = {} + read = functools.partial(tgm.read_env, env) + write = functools.partial(tgm.write_env, env) + + # Bind args and consts to environment + write(processed_jaxpr.jaxpr.invars, flat_args) + write(processed_jaxpr.jaxpr.constvars, processed_jaxpr.consts) + + # Loop through equations and evaluate primitives using `bind` + losses_so_far = 0 + losses_p_deps = [] + losses_inputs = [] + for eqn in processed_jaxpr.jaxpr.eqns: + + if isinstance(eqn.primitive, tags.LossTag): + assert eqn == processed_jaxpr.loss_tags[losses_so_far] + + if not drop_loss_tags: + write(eqn.outvars, tgm.eval_jaxpr_eqn(eqn, read(eqn.invars))) + + losses_inputs.append(read(eqn.invars)) + losses_p_deps.append(read(tags.loss_eqn_parameter_dependants(eqn))) + losses_so_far += 1 + + else: + write(eqn.outvars, tgm.eval_jaxpr_eqn(eqn, read(eqn.invars))) + + if losses_so_far == len(processed_jaxpr.loss_tags): + break + + return tuple(tuple(p) for p in losses_p_deps), tuple(losses_inputs) # pytype: disable=bad-return-type + + return forward_compute_losses + + +def _compute_all_losses( + p_jaxpr: ProcessedJaxpr, + primal_func_args: FuncArgs, +) -> tuple[LossFunction, ...]: + """Returns all loss functions objects.""" + if not p_jaxpr.loss_tags: + raise ValueError("The provided `ProcessedJaxpr` has no loss tags.") + + losses_func = construct_compute_losses_inputs( + processed_jaxpr=p_jaxpr, + primal_func_args=primal_func_args, + params_index=p_jaxpr.params_index, + ) + _, losses_inputs = losses_func(primal_func_args[p_jaxpr.params_index]) + return p_jaxpr.reconstruct_losses(losses_inputs) + + +def _loss_tags_vjp( + p_jaxpr: ProcessedJaxpr, + primal_func_args: FuncArgs, +) -> LossTagsVjp: + """Computes a (backward-mode) vector-Jacobian product for the vector of losses given by the loss tags. + + The function has similar interface to :func:`jax.vjp`. It takes as inputs the + concrete values of the primals at which the Jacobian will be evaluated. It + returns a pair of ``(losses, losses_vjp)``, where losses is a tuple of + :class:`~LossFunction` objects and ``vjp_func`` is a function + taking as inputs the concrete values of the tangents of the inputs for each + loss tag (corresponding to a loss object in ``losses``) and returns the + corresponding tangents of the parameters. + + Args: + p_jaxpr: The :class:``~ProcessedJaxpr`` representing the model function. + This must include at least one loss tag. + primal_func_args: The primals at which to evaluate the Jacobian. + + Returns: + The computed ``losses`` and ``losses_vjp`` pair. + """ + + if not p_jaxpr.loss_tags: + raise ValueError("The provided `ProcessedJaxpr` has no loss tags.") + + losses_func = construct_compute_losses_inputs( + processed_jaxpr=p_jaxpr, + primal_func_args=primal_func_args, + params_index=p_jaxpr.params_index, + ) + + primal_params = primal_func_args[p_jaxpr.params_index] + _, full_vjp_func, losses_inputs = jax.vjp( + losses_func, primal_params, has_aux=True + ) + + def losses_vjp_func(losses_tangents: Sequence[LossFunctionInputs]) -> Params: + """Computes the vector-Jacobian product w.r.t. the parameters. + + Args: + losses_tangents: The tangents to all loss tag's inputs. + + Returns: + The parameters' tangents, as a result of the vector-Jacobian product. + """ + + if len(losses_tangents) != len(p_jaxpr.loss_tags): + raise ValueError( + "The argument `tangents` must be a sequence of the " + "tangents to each loss tag in the same order as the " + "loss objects that have been returned. The number of " + f"loss_tags is {len(p_jaxpr.loss_tags)}, but the length " + f"of `tangents` is {len(losses_tangents)}." + ) + + for i, loss_tangents in enumerate(losses_tangents): + if not isinstance(loss_tangents, Sequence): + raise ValueError( + "Each element of the argument `tangents` must be " + f"a sequence, but tangents[{i}] has type " + f"{type(loss_tangents)}." + ) + + [params_tangents] = full_vjp_func(losses_tangents) + + return params_tangents + + return p_jaxpr.reconstruct_losses(losses_inputs), losses_vjp_func + + +def _loss_tags_jvp( + p_jaxpr: ProcessedJaxpr, + primal_func_args: FuncArgs, + params_tangents: Params, +) -> LossTagsJvp: + """Computes a (forward-mode) Jacobian-vector product for the losses given by the loss tags. + + The function has similar interface to :func:`jax.jvp`. It takes as inputs the + concrete values of the primals at which the Jacobian will be evaluated at and + the concrete values of the tangents for the **parameters**, as specified by + ``processed_jaxpr.params_index``. It returns a pair of + ``(losses, losses_tangents)``, where ``losses`` is a tuple of + :class:`~LossFunction` objects, and ``losses_tangents`` is + a tuple containing the tangents of the inputs for each loss tag (corresponding + to a loss object in ``losses``). + + Args: + p_jaxpr: The :class:`~ProcessedJaxpr` representing the model function. This + must include at least one loss tag. + primal_func_args: The primals at which to evaluate the Jacobian. + params_tangents: The vector of tangents which to multiply with the Jacobian. + + Returns: + The computed ``losses`` and ``losses_tangents`` pair. + """ + + if not p_jaxpr.loss_tags: + raise ValueError("The provided `ProcessedJaxpr` has no loss tags.") + + losses_func = construct_compute_losses_inputs( + processed_jaxpr=p_jaxpr, + primal_func_args=primal_func_args, + params_index=p_jaxpr.params_index, + ) + + primal_params = (primal_func_args[p_jaxpr.params_index],) + + tangents = (params_tangents,) + + (_, losses_tangents, losses_inputs) = jax.jvp( + losses_func, primal_params, tangents, has_aux=True + ) + + return p_jaxpr.reconstruct_losses(losses_inputs), losses_tangents + + +def _loss_tags_hvp( + processed_jaxpr: ProcessedJaxpr, + primal_func_args: FuncArgs, + params_tangents: Params, +) -> tuple[Params, tuple[LossFunction, ...]]: + """Computes a Hessian-vector product of the function w.r.t. all loss tags. + + The function takes as inputs the concrete values of the primals for the + function arguments at which the Hessian will be evaluated at and the concrete + values of the tangents for the **parameters**, as specified by + ``processed_jaxpr.params_index``. It returns the product of the Hessian with + this tangents via backward-over-forward mode. + + Args: + processed_jaxpr: The :class:`~ProcessedJaxpr` representing the model + function. This must include at least one loss tag. + primal_func_args: The primals at which to evaluate the Hessian. + params_tangents: The vector of tangents which to multiply with the Hessian. + + Returns: + The parameter-structured vector representing the Hessian-vector product and + the resulting :class:`~LossFunction` objects that correspond to every + loss tag. + """ + + if not processed_jaxpr.loss_tags: + raise ValueError("The provided `ProcessedJaxpr` has no loss tags.") + + losses_func = construct_compute_losses_inputs( + processed_jaxpr=processed_jaxpr, + primal_func_args=primal_func_args, + params_index=processed_jaxpr.params_index, + ) + + def losses_sum(param_primals: Params) -> Array: + # This computes the sum of losses evaluated. Makes it easier because we can + # now use jax.grad rather than jax.vjp for taking derivatives. + _, losses_inputs = losses_func(param_primals) + losses = processed_jaxpr.reconstruct_losses(losses_inputs) + return sum(jnp.sum(loss.evaluate()) for loss in losses) + + # Directional derivative function + df_dot_dv = lambda p: (jax.jvp(losses_sum, [p], [params_tangents])[1]) + hvp = jax.grad(df_dot_dv)(primal_func_args[processed_jaxpr.params_index]) + + _, losses_inputs = losses_func(primal_func_args[processed_jaxpr.params_index]) + return hvp, processed_jaxpr.reconstruct_losses(losses_inputs) + + +@jax.tree_util.register_dataclass +@dataclasses.dataclass(frozen=True) +class VarMap: + """A mapping from jaxpr variables to values. + + Variables in a jaxpr are not ordered, and thus ``dict[Var, ...]`` cannot be + passed to PyTree APIs. This class works around that by indexing the dict + on the ID of each variable instead of the variable itself. + """ + id_to_var: dict[int, Var] = dataclasses.field( + default_factory=dict, metadata=dict(static=True) + ) + var_to_val: dict[int, Any] = dataclasses.field(default_factory=dict) + + def __contains__(self, var: Var) -> bool: + return id(var) in self.id_to_var + + def __getitem__(self, var: Var) -> Any: + return self.var_to_val[id(var)] + + def get(self, var: Var) -> Any: + return self.var_to_val[id(var)] + + def update(self, it: Iterable[tuple[Var, Any]]) -> None: + for var, val in it: + self.id_to_var[id(var)] = var + self.var_to_val[id(var)] = val + + @classmethod + def create(cls, it: Iterable[tuple[Var, Any]]) -> "VarMap": + self = cls() + self.update(it) + return self + + +def _layer_tag_vjp( + processed_jaxpr: ProcessedJaxpr, + primal_func_args: FuncArgs, +) -> LayerTagVjp: + """Computes primal values and tangents w.r.t. all layer tags. + + The returned function has similar interface to :func:`jax.vjp`. It takes as + inputs the concrete values of the primals at which the Jacobian will be + evaluated. It returns a pair of ``(losses, vjp_func)``, where losses is a + tuple of :class:`~LossFunction` objects and ``vjp_func`` is a function taking + as inputs the concrete values of the tangents of the inputs for each loss tag + (corresponding to a loss object in ``losses``) and returns a list of + quantities computed for each layer tag in ``processed_jaxpr``. Each entry of + the list is a :class:`~LayerVjpData` with the keys ``"primals", "tangents"`` + mapping to LayerData objects that each separate the primals (or tangents) into + inputs, outputs, and parameters. + + Args: + processed_jaxpr: The :class:`~ProcessedJaxpr` representing the model + function. This must include at least one loss tag. + primal_func_args: The primals at which to evaluate the Jacobian. + + Returns: + The computed ``losses`` and ``vjp_func`` pair. + """ + layer_tag_invars = jax.tree_util.tree_leaves( + [tag.invars for tag in processed_jaxpr.layer_tags] + ) + # We exclude literals because they are not hashable. Their values will be read + # properly for the primals due to the special handling in tgm.read_env. For + # the tangents they won't matter, since only layer inputs (which together with + # the layer outputs and params form the layer tag invars) can be literals, and + # we explicitly exclude the inputs from the tangents (since their computation + # isn't generally correct anyway). + layer_tag_invars = list(set(v for v in layer_tag_invars + if not isinstance(v, jex.core.Literal))) + + def forward() -> tuple[Array, ...]: + """Computes the values of all inputs to all **layer** tags.""" + + own_func_args = primal_func_args + + # Mapping from variable -> value + env = {} + read = functools.partial(tgm.read_env, env) + write = functools.partial(tgm.write_env, env) + + # Bind args and consts to environment + write( + processed_jaxpr.jaxpr.invars, jax.tree_util.tree_leaves(own_func_args) + ) + write(processed_jaxpr.jaxpr.constvars, processed_jaxpr.consts) + + # Loop through equations and evaluate them + num_losses_passed = 0 + for eqn in processed_jaxpr.jaxpr.eqns: + if isinstance(eqn.primitive, tags.LossTag): + num_losses_passed += 1 + if num_losses_passed == len(processed_jaxpr.loss_tags): + break + else: + write(eqn.outvars, tgm.eval_jaxpr_eqn(eqn, read(eqn.invars))) + + assert num_losses_passed == len(processed_jaxpr.loss_tags) + + return tuple(read(layer_tag_invars)) + + def forward_aux( + aux: dict[Var, Array], + ) -> tuple[tuple[LossFunctionInputs, ...], tuple[LossFunctionInputs, ...]]: + """Computes the inputs and kwargs of all **loss** tags. + + Args: + aux: A mapping from an Jaxpr variable to an additional auxiliary value. + For each variable in this mapping, we add to the value computed during + standard evaluation the auxiliary value. This is done in order to be + able to compute gradients wrt all intermediate expressions corresponding + to the Jaxpr variables in this mapping + + Returns: + The pair of ``(losses_inputs, losses_kwargs)`` where ``losses_inputs`` + is a tuple of the input values for each loss tag, and ``losses_kwargs`` + is a tuple of the kwargs values of each loss tag. + """ + + own_func_args = primal_func_args + + # Mapping from variable -> value + env: dict[jex.core.Var, Array] = {} + read = functools.partial(tgm.read_env, env) + + def write(variables: list[jex.core.Var], values: list[Array]) -> None: + + tgm.write_env(env, variables, values) + + for v in variables: + if not isinstance(v, jex.core.Literal) and v in aux: + env[v] = env[v] + aux[v] + + # Bind args and consts to environment + write( + processed_jaxpr.jaxpr.invars, jax.tree_util.tree_leaves(own_func_args) + ) + + write(processed_jaxpr.jaxpr.constvars, processed_jaxpr.consts) + + # Loop through equations and evaluate primitives using `bind` + num_losses_passed = 0 + losses_p_dependants = [] + losses_inputs_values = [] + + for eqn in processed_jaxpr.jaxpr.eqns: + + input_values = read(eqn.invars) + + if isinstance(eqn.primitive, tags.LossTag): + + loss: LossFunction = tags.loss_eqn_construct_loss(eqn, *input_values) + + losses_p_dependants.append(loss.parameter_dependants) + losses_inputs_values.append(tuple(input_values)) + + num_losses_passed += 1 + + if num_losses_passed == len(processed_jaxpr.loss_tags): + break + + else: + write(eqn.outvars, tgm.eval_jaxpr_eqn(eqn, input_values)) + + assert num_losses_passed == len(processed_jaxpr.loss_tags) + + # Read the inputs to the loss functions, but also return the target values + return tuple(losses_p_dependants), tuple(losses_inputs_values) + + # First compute the primal values for the inputs to all layer tags + layer_input_values = forward() + primals = zip(layer_tag_invars, layer_input_values) + + # Update with the values of all parameters, which are inputs to the function + primals = itertools.chain( + primals, + zip( + processed_jaxpr.jaxpr.invars, + jax.tree_util.tree_leaves(primal_func_args), + ), + ) + + primals_dict = VarMap.create(primals) + + # Create auxiliary values all equal to zero. + aux_values = jax.tree_util.tree_map(jnp.zeros_like, layer_input_values) + + # Create a mapping from all layer tag inputs to the zero values + aux_dict = VarMap.create(zip(layer_tag_invars, aux_values)) + + # These values would now allow us to compute gradients wrt the layer tags + # inputs, which are intermediate expressions in the Jaxpr. + _, aux_vjp, losses_inputs = jax.vjp(forward_aux, aux_dict, has_aux=True) + + # Compute the actual loss objects. + losses: list[LossFunction] = [ + tags.loss_eqn_construct_loss(tag, *inputs) + for tag, inputs in zip(processed_jaxpr.loss_tags, losses_inputs) + ] + + def vjp_func( + tangents: tuple[LossFunctionInputs, ...], # pytype: disable=invalid-annotation + ) -> tuple[LayerVjpData[Array], ...]: # pytype: disable=invalid-annotation + """Computes a (reverse-mode) vector-Jacobian product w.r.t. all layer tags. + + Args: + tangents: The concrete tangent values for the tangents of the inputs to + all **loss** tags. + + Returns: + A tuple containing both the primal and tangent values for the inputs to + all **layer** tags. The values are provided as a dictionary with keys: + ``inputs, outputs, params, outputs_tangent, params_tangent``. + """ + [tangents_dict] = aux_vjp(tangents) + + read_primals = functools.partial(tgm.read_env, primals_dict) + read_tangents = functools.partial(tgm.read_env, tangents_dict) + layers_info = [] + + for tag in processed_jaxpr.layer_tags: + + primals = read_primals(tag.invars) # pytype: disable=wrong-arg-types + tangents = read_tangents(tag.invars) + + # The input tangents could be potentially wrong, so we don't include them. + # This is because the one of the invars could be a Literal, which breaks + # the above logic, or one of the invars could be shared with another + # layer, so that both layers would get the sum of the input tangents for + # that variable. There is also the possibility of "input preprocessing", + # which affects the tag inputs but not the corresponding layer in the + # graph. + layers_info.append(LayerVjpData( + primals=tag.primitive.layer_data(primals, tag.params), + tangents=tag.primitive.layer_data(tangents, tag.params, + exclude_inputs=True), + )) + + return tuple(layers_info) + + return tuple(losses), vjp_func + + +def _layer_tag_vjp_and_value_and_grad_from_surrogate( + processed_jaxpr: ProcessedJaxpr, + primal_func_args: FuncArgs, +) -> LayerTagVjpAndValueAndGrad: + """Transposes one tagged forward for KFAC and ordinary gradients. + + The source function returns ``(primal_loss, gradient_surrogate)``. Its + surrogate has the desired first-order parameter gradient, while the first + result is the loss value reported by the optimizer. Artificial injections at + every layer operand (including every owned model parameter) are exposed to + one VJP. Exact-Fisher seeds and the surrogate gradient therefore reuse the + same model primal evaluation, while curvature retains the established + aux-only transpose structure. This is only first-order reverse mode: unlike + transforming a custom-JVP or an already transposed value-and-grad graph, it + cannot introduce Hessian-vector work. + """ + layer_tag_invars = jax.tree_util.tree_leaves( + [tag.invars for tag in processed_jaxpr.layer_tags] + ) + layer_tag_invars = list( + set( + var + for var in layer_tag_invars + if not isinstance(var, jex.core.Literal) + ) + ) + def forward_aux(aux: VarMap): + env: dict[jex.core.Var, Array] = {} + read = functools.partial(tgm.read_env, env) + + def write(variables: list[jex.core.Var], values: list[Array]) -> None: + tgm.write_env(env, variables, values) + for var in variables: + if not isinstance(var, jex.core.Literal) and var in aux: + env[var] = env[var] + aux[var] + + write( + processed_jaxpr.jaxpr.invars, + jax.tree_util.tree_leaves(primal_func_args), + ) + write(processed_jaxpr.jaxpr.constvars, processed_jaxpr.consts) + + losses_inputs_values = [] + losses_p_dependants = [] + for eqn in processed_jaxpr.jaxpr.eqns: + input_values = read(eqn.invars) + if isinstance(eqn.primitive, tags.LossTag): + loss = tags.loss_eqn_construct_loss(eqn, *input_values) + losses_inputs_values.append(tuple(input_values)) + losses_p_dependants.append(loss.parameter_dependants) + write(eqn.outvars, tgm.eval_jaxpr_eqn(eqn, input_values)) + + if len(losses_inputs_values) != len(processed_jaxpr.loss_tags): + raise ValueError( + "The shared-forward interpreter did not encounter every " + "registered loss." + ) + forward_outputs = tuple(read(processed_jaxpr.jaxpr.outvars)) + if len(forward_outputs) != 2: + raise ValueError( + "The shared-forward source must return exactly " + "`(primal_loss, gradient_surrogate)`; got " + f"{len(forward_outputs)} flattened outputs." + ) + primal_loss, gradient_surrogate = forward_outputs + for name, value in ( + ("primal_loss", primal_loss), + ("gradient_surrogate", gradient_surrogate), + ): + if value.shape: + raise ValueError( + f"The shared-forward `{name}` must be scalar; got shape " + f"{value.shape}." + ) + return ( + (tuple(losses_p_dependants), gradient_surrogate), + ( + tuple(losses_inputs_values), + tuple(read(layer_tag_invars)), + primal_loss, + gradient_surrogate, + ), + ) + + aux_values = tuple( + jnp.zeros(var.aval.shape, var.aval.dtype) for var in layer_tag_invars + ) + aux_dict = VarMap.create(zip(layer_tag_invars, aux_values)) + _, aux_vjp, forward_values = jax.vjp( + forward_aux, + aux_dict, + has_aux=True, + ) + ( + losses_inputs, + layer_input_values, + primal_loss, + gradient_surrogate, + ) = forward_values + + primals = itertools.chain( + zip(layer_tag_invars, layer_input_values), + zip( + processed_jaxpr.jaxpr.invars, + jax.tree_util.tree_leaves(primal_func_args), + ), + ) + primals_dict = VarMap.create(primals) + losses = tuple( + tags.loss_eqn_construct_loss(tag, *inputs) + for tag, inputs in zip(processed_jaxpr.loss_tags, losses_inputs) + ) + + def losses_vjp( + tangents: tuple[LossFunctionInputs, ...], + ) -> tuple[LayerVjpData[Array], ...]: + [tangents_dict] = aux_vjp( + (tangents, jnp.zeros_like(gradient_surrogate)) + ) + read_primals = functools.partial(tgm.read_env, primals_dict) + read_tangents = functools.partial(tgm.read_env, tangents_dict) + return tuple( + LayerVjpData( + primals=tag.primitive.layer_data( + read_primals(tag.invars), tag.params + ), + tangents=tag.primitive.layer_data( + read_tangents(tag.invars), + tag.params, + exclude_inputs=True, + ), + ) + for tag in processed_jaxpr.layer_tags + ) + + zero_loss_tangents = tuple( + jax.tree_util.tree_map(jnp.zeros_like, loss.parameter_dependants) + for loss in losses + ) + [gradient_tangents_dict] = aux_vjp( + (zero_loss_tangents, jnp.ones_like(gradient_surrogate)) + ) + grads = jax.tree_util.tree_unflatten( + processed_jaxpr.params_tree, + tuple( + gradient_tangents_dict[var] + for var in processed_jaxpr.params_vars_flat + ), + ) + return losses, losses_vjp, primal_loss, grads + + +def compute_all_losses( + func: Func[Any], + params_index: int = 0, +) -> tuple[TransformedFunction[tuple[LossFunction, ...]], JaxprExtractor]: + """Creates a function that when called, returns all loss objects. + + The returned function takes as inputs the concrete values of the primals for + which to compute the loss objects. + + Args: + func: The model function, which must include at least one loss registration. + params_index: The variables from the function arguments which are at this + index (e.g. `func_args[params_index]`) are to be considered model + parameters. + + Returns: + A function that computes all loss objects with signature + `Callable[[FuncArgs], tuple[LossFunction, ...]]`, and a function that + returns the processed Jaxpr of `func` for a given set of function arguments. + """ + # Note that this function is independent of any layer tags, hence we can avoid + # calling the auto registration. + return cached_transformation( + func=func, + transformation=_compute_all_losses, + verifier=lambda: None, + params_index=params_index, + auto_register_tags=False, + allow_left_out_params=True, + ) + + +def loss_tags_vjp( + func: Func[Any], + params_index: int = 0, +) -> tuple[TransformedFunction[LossTagsVjp], JaxprExtractor]: + """Creates a function for the vector-Jacobian product w.r.t. all loss tags. + + The returned function has a similar interface to :func:`jax.vjp`. It takes as + inputs the concrete values of the primals at which the Jacobian will be + evaluated. It returns a pair ``(losses, losses_vjp)``, where losses is a + tuple of :class:`~LossFunction` objects and ``vjp_func`` is a function taking + as inputs the concrete values of the tangents of the inputs for each loss tag + (corresponding to a loss object in ``losses``) and returns the corresponding + tangents of the parameters. + + Args: + func: The model function, which must include at least one loss registration. + params_index: The variables from the function arguments which are at this + index (e.g. `func_args[params_index]`) are to be considered model + parameters. + + Returns: + A function that computes the vector-Jacobian product with signature + `Callable[[FuncArgs], LossTagsVjp]`, and a function that returns the + processed Jaxpr of `func` for a given set of function arguments. + """ + # Note that this function is independent of any layer tags, hence we can avoid + # calling the auto registration. + return cached_transformation( + func=func, + transformation=_loss_tags_vjp, + verifier=lambda: None, + params_index=params_index, + auto_register_tags=False, + allow_left_out_params=True, + ) + + +def loss_tags_jvp( + func: Func[Any], + params_index: int = 0, +) -> tuple[TransformedFunction[LossTagsJvp], JaxprExtractor]: + """Creates a function for the Jacobian-vector product w.r.t. all loss tags. + + The returned function has a similar interface to :func:`jax.jvp`. It takes as + inputs the concrete values of the primals at which the Jacobian will be + evaluated at and the concrete values of the tangents for the **parameters**, + as specified by ``processed_jaxpr.params_index``. It returns a pair + ``(losses, losses_tangents)``, where ``losses`` is a tuple of + :class:`~LossFunction` objects, and ``losses_tangents`` is a tuple containing + the tangents of the inputs for each loss tag (corresponding to a loss object + in ``losses``). + + Args: + func: The model function, which must include at least one loss registration. + params_index: The variables from the function arguments which are at this + index (e.g. `func_args[params_index]`) are to be considered model + parameters. + + Returns: + A function that computes the Jacobian-vector product with signature + `Callable[[FuncArgs, Params], LossTagsVjp]`, and a function that returns + the processed Jaxpr of `func` for a given set of function arguments. + """ + # Note that this function is independent of any layer tags, hence we can avoid + # calling the auto registration. + return cached_transformation( + func=func, + transformation=_loss_tags_jvp, + verifier=lambda: None, + params_index=params_index, + auto_register_tags=False, + allow_left_out_params=True, + ) + + +def loss_tags_hvp( + func: Func[Any], + params_index: int = 0, +) -> tuple[ + TransformedFunction[tuple[Params, tuple[LossFunction, ...]]], JaxprExtractor +]: + """Creates a function for the Hessian-vector product w.r.t. all loss tags. + + The returned function takes as inputs the concrete values of the primals for + the function arguments at which the Hessian will be evaluated at and the + concrete values of the tangents for the **parameters**, as specified by + ``processed_jaxpr.params_index``. It returns the product of the Hessian with + these tangents via backward-over-forward mode autodiff. + + Args: + func: The model function, which must include at least one loss registration. + params_index: The variables from the function arguments which are at this + index (e.g. `func_args[params_index]`) are to be considered model + parameters. + + Returns: + A function that computes the Hessian-vector product and also returns all + losses, with signature `Callable[[FuncArgs, Params], + tuple[LossTagsVjp, tuple[LossFunction, ...]]`, and a function that returns + the processed Jaxpr of `func` for a given set of function arguments. + """ + # Note that this function is independent of any layer tags, hence we can avoid + # calling the auto registration. + return cached_transformation( + func=func, + transformation=_loss_tags_hvp, + verifier=lambda: None, + params_index=params_index, + auto_register_tags=False, + allow_left_out_params=True, + ) + + +def layer_tags_vjp( + func: Func[Any], + params_index: int = 0, + auto_register_tags: bool = True, + raise_error_on_diff_jaxpr: bool = True, + **auto_registration_kwargs, +) -> tuple[TransformedFunction[LayerTagVjp], JaxprExtractor]: + """Creates a function for primal values and tangents w.r.t. all layer tags. + + The returned function has a similar interface to :func:`jax.vjp`. It takes as + inputs the concrete values of the primals at which the Jacobian will be + evaluated. It returns a pair ``(losses, vjp_func)``, where ``losses`` is a + tuple of :class:`~LossFunction` objects and ``vjp_func`` is a function taking + as inputs the concrete values of the tangents of the inputs for each loss tag + (corresponding to a loss object in ``losses``) and returns a list of + quantities computed for each layer tag in ``processed_jaxpr``. Each entry of + the list is a :class:`~LayerVjpData` with the keys ``"primals", "tangents"`` + mapping to LayerData objects that each separate the primals (or tangents) into + inputs, outputs, and parameters. + + Args: + func: The model function, which must include at least one loss registration. + params_index: The variables from the function arguments which are at this + index (e.g. ``func_args[params_index]``) are to be considered model + parameters. + auto_register_tags: Whether to run an automatic layer registration on the + function (e.g. :func:`~auto_register_tags`). + raise_error_on_diff_jaxpr: When tracing with different arguments, if the + returned jaxpr has a different graph will raise an exception. + **auto_registration_kwargs: Any additional keyword arguments, to be passed + to the automatic registration pass. + + Returns: + Returns the above described function, and a function that returns the + processed Jaxpr of `func` for a given set of function arguments. + """ + + return cached_transformation( + func=func, + transformation=_layer_tag_vjp, + params_index=params_index, + auto_register_tags=auto_register_tags, + allow_left_out_params=False, + raise_error_on_diff_jaxpr=raise_error_on_diff_jaxpr, + **auto_registration_kwargs, + ) + + +def layer_tags_vjp_and_value_and_grad( + value_func: Func[Any], + params_index: int = 0, + auto_register_tags: bool = True, + raise_error_on_diff_jaxpr: bool = True, + **auto_registration_kwargs, +) -> tuple[ + TransformedFunction[LayerTagVjpAndValueAndGrad], + JaxprExtractor, +]: + """Creates a layer VJP that also returns loss and surrogate gradient. + + ``value_func`` must return exactly ``(primal_loss, gradient_surrogate)``. + Both are scalars, and the gradient of the second result with respect to the + parameters must equal the training gradient associated with the first. The + transformed callable receives the ordinary function arguments and executes + the tagged source graph once for exact-Fisher statistics, the reported loss, + and the ordinary parameter gradient. + """ + + # The ordinary layer-VJP auto-tagger deliberately discards the function's + # declared outputs. This transform must retain + # `(primal_loss, gradient_surrogate)`, while the matcher itself still stops + # registering layers at the loss tags. + auto_registration_kwargs["compute_only_loss_tags"] = False + return cached_transformation( + func=value_func, + transformation=_layer_tag_vjp_and_value_and_grad_from_surrogate, + params_index=params_index, + auto_register_tags=auto_register_tags, + allow_left_out_params=False, + raise_error_on_diff_jaxpr=raise_error_on_diff_jaxpr, + **auto_registration_kwargs, + ) diff --git a/src/kfac_jax/_src/utils/__init__.py b/src/kfac_jax/_src/utils/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..f0102c0f2052df58b4901b310144bfa6fea4f808 --- /dev/null +++ b/src/kfac_jax/_src/utils/__init__.py @@ -0,0 +1,150 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC related utility classes and functions.""" + +from kfac_jax._src.utils import accumulators +from kfac_jax._src.utils import math +from kfac_jax._src.utils import misc +from kfac_jax._src.utils import parallel +from kfac_jax._src.utils import staging +from kfac_jax._src.utils import types + +# types +Array = types.Array +PRNGKey = types.PRNGKey +Scalar = types.Scalar +Numeric = types.Numeric +Shape = types.Shape +DType = types.DType +PyTree = types.PyTree +ArrayTree = types.ArrayTree +TArrayTree = types.TArrayTree +Params = types.Params +Batch = types.Batch +FuncState = types.FuncState +FuncAux = types.FuncAux +PyTreeDef = types.PyTreeDef +FuncArgs = types.FuncArgs +FuncOuts = types.FuncOuts +Func = types.Func +ValueFunc = types.ValueFunc +ValueAndGradFunc = types.ValueAndGradFunc +AssumedFuncOutput = types.AssumedFuncOutput +ScheduleType = types.ScheduleType +tree_is_empty = types.tree_is_empty +abstract_objects_equal = types.abstract_objects_equal +get_float_dtype_and_check_consistency = ( + types.get_float_dtype_and_check_consistency) +del types + +# misc +deserialize_state_tree = misc.deserialize_state_tree +serialize_state_tree = misc.serialize_state_tree +to_tuple_or_repeat = misc.to_tuple_or_repeat +filter_sequence = misc.filter_sequence +first_dim_is_size = misc.first_dim_is_size +fake_element_from_iterator = misc.fake_element_from_iterator +default_batch_size_extractor = misc.default_batch_size_extractor +auto_scope_function = misc.auto_scope_function +auto_scope_method = misc.auto_scope_method +register_state_class = misc.register_state_class +replace_char = misc.replace_char +call_func_with_conditional_kwargs = misc.call_func_with_conditional_kwargs +Finalizable = misc.Finalizable +State = misc.State +rearrange = misc.rearrange +del misc + +# parallel +in_pmap = parallel.in_pmap +wrap_if_pmap = parallel.wrap_if_pmap +pmean_if_pmap = parallel.pmean_if_pmap +psum_if_pmap = parallel.psum_if_pmap +pmap_mean = parallel.pmap_mean +pmap_sum = parallel.pmap_sum +using_legacy_pmap = parallel.using_legacy_pmap +get_device_n_contents = parallel.get_device_n_contents +get_first = parallel.get_first +get_mean = parallel.get_mean +get_sum = parallel.get_sum +broadcast_all_local_devices = parallel.broadcast_all_local_devices +pmap_zeros_like = parallel.pmap_zeros_like +jit_zeros_like = parallel.jit_zeros_like +replicate_all_local_devices = parallel.replicate_all_local_devices +make_different_rng_key_on_all_devices = ( + parallel.make_different_rng_key_on_all_devices) +p_split = parallel.p_split +p_split_num = parallel.p_split_num +host_sync = parallel.host_sync +host_all_gather = parallel.host_all_gather +host_mean = parallel.host_mean +pmap_sync_and_divide_value = parallel.pmap_sync_and_divide_value +jit_sync_and_divide_value = parallel.jit_sync_and_divide_value +copy_array = parallel.copy_array +copy_obj = parallel.copy_obj +pmap_copy_obj = parallel.pmap_copy_obj +distribute_thunks = parallel.distribute_thunks +del parallel + +# math +set_special_case_zero_inv = math.set_special_case_zero_inv +get_special_case_zero_inv = math.get_special_case_zero_inv +set_use_cholesky_inversion = math.set_use_cholesky_inversion +get_use_cholesky_inversion = math.get_use_cholesky_inversion +product = math.product +outer_product = math.outer_product +scalar_mul = math.scalar_mul +scalar_div = math.scalar_div +weighted_sum_of_objects = math.weighted_sum_of_objects +sum_of_objects = math.sum_objects +pytree_size = math.pytree_size +inner_product = math.inner_product +symmetric_matrix_inner_products = math.symmetric_matrix_inner_products +asymmetric_matrix_inner_products = math.asymmetric_matrix_inner_products +matrix_of_inner_products = math.matrix_of_inner_products +vector_of_inner_products = math.vector_of_inner_products +block_permuted = math.block_permuted +norm = math.norm +squared_norm = math.squared_norm +per_parameter_norm = math.per_parameter_norm +psd_inv = math.psd_inv +psd_solve = math.psd_solve +psd_solve_maybe_zero_last_idx = math.psd_solve_maybe_zero_last_idx +pi_adjusted_kronecker_factors = math.pi_adjusted_kronecker_factors +pi_adjusted_kronecker_inverse = math.pi_adjusted_kronecker_inverse +kronecker_product_axis_mul_v = math.kronecker_product_axis_mul_v +kronecker_eigen_basis_axis_mul_v = math.kronecker_eigen_basis_axis_mul_v +kronecker_product_mul_v = math.kronecker_product_mul_v +kronecker_eigen_basis_mul_v = math.kronecker_eigen_basis_mul_v +safe_psd_eigh = math.safe_psd_eigh +tnt_scale = math.tnt_scale +loop_and_parallelize_average = math.loop_and_parallelize_average +psd_matrix_norm = math.psd_matrix_norm +invert_psd_matrices = math.invert_psd_matrices +inverse_sqrt_psd_matrices = math.inverse_sqrt_psd_matrices +stable_sqrt = math.stable_sqrt +cosine_similarity = math.cosine_similarity + +del math + +# accumulators +default_add_function = accumulators.default_add_function +WeightedMovingAverage = accumulators.WeightedMovingAverage +MultiChunkAccumulator = accumulators.MultiChunkAccumulator +del accumulators + +# staged +staged = staging.staged +WithStagedMethods = staging.WithStagedMethods +del staging diff --git a/src/kfac_jax/_src/utils/accumulators.py b/src/kfac_jax/_src/utils/accumulators.py new file mode 100644 index 0000000000000000000000000000000000000000..be575e5ba753c5936dfb7428cc58ef1e72be0315 --- /dev/null +++ b/src/kfac_jax/_src/utils/accumulators.py @@ -0,0 +1,297 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC for accumulating statistics.""" +from typing import Any, Callable, Generic + +import jax +import jax.numpy as jnp + +from kfac_jax._src.utils import misc +from kfac_jax._src.utils import parallel +from kfac_jax._src.utils import types + +Array = types.Array +Numeric = types.Numeric +Shape = types.Shape +DType = types.DType +ArrayTree = types.ArrayTree +TArrayTree = types.TArrayTree + +AddFunction = Callable[[TArrayTree, TArrayTree, Numeric, Numeric], TArrayTree] + + +def default_add_function( + obj1: TArrayTree, + obj2: TArrayTree, + coeff1: Numeric, + coeff2: Numeric +) -> TArrayTree: + + return jax.tree_util.tree_map( + lambda x, y: coeff1 * x + coeff2 * y, obj1, obj2) + + +@misc.register_state_class +class WeightedMovingAverage(Generic[TArrayTree], misc.State): + """A wrapped class for an arbitrary weighted moving average.""" + + weight: Numeric + value: TArrayTree | None + + @property + def ndim(self) -> int: + assert self.value is not None + return self.value.ndim # pytype: disable=attribute-error + + @property + def shape(self) -> Shape: + assert self.value is not None + return self.value.shape # pytype: disable=attribute-error + + @property + def dtype(self) -> DType: + assert self.value is not None + return self.value.dtype # pytype: disable=attribute-error + + def update( + self, + value: TArrayTree, + old_weight_multiplier: Numeric, + new_weight: Numeric, + add_function: AddFunction = default_add_function, + ): + """Updates the underlying array and weight accordingly.""" + + assert self.value is not None + + # A negative value of new_weight means we should only update the value + # (with -new_weight) and not the total running weight. This roughly + # corresponds to summation instead of averaging, and is useful in a few + # contexts. + self.weight = old_weight_multiplier * self.weight + jax.nn.relu(new_weight) + eta_for_old = jax.nn.relu(new_weight) / self.weight + eta_for_new = jnp.abs(new_weight) / self.weight + + self.value = add_function(self.value, value, 1.0 - eta_for_old, eta_for_new) + + def sync(self, pmap_axis_name: str | None): + """Syncs the underlying array across devices.""" + + if self.value is None: + raise ValueError("`_value` has not been set yet.") + + self.value = parallel.pmean_if_pmap(self.value, pmap_axis_name) + + def clear(self, value_to_none: bool = False): + """Resets the weighted average.""" + + self.weight = jnp.zeros_like(self.weight) + self.value = None if value_to_none else jnp.zeros_like(self.value) + + def value_and_clear(self) -> TArrayTree: + """Retrieves the value of the weighted average and clears it.""" + + value = self.value + self.clear() + + assert value is not None + return value + + @classmethod + def zeros_array( + cls, + shape: Shape, + dtype: DType | None = None, + ) -> "WeightedMovingAverage[Array]": + """Initializes a `WeightedMovingAverage` with a single array of zeros.""" + + return cls( # pytype: disable=wrong-keyword-args + weight=jnp.zeros([], dtype=dtype), + value=jnp.zeros(shape, dtype=dtype), + ) + + @classmethod + def zeros_like(cls, value: TArrayTree) -> "WeightedMovingAverage[TArrayTree]": + """Initializes a `WeightedMovingAverage` with zeros structure like `value`.""" + + return cls( # pytype: disable=wrong-keyword-args + weight=jnp.array( + 0.0, dtype=types.get_float_dtype_and_check_consistency(value) + ), + value=jax.tree_util.tree_map(jnp.zeros_like, value), + ) + + +class MultiChunkAccumulator(Generic[TArrayTree]): + """Statistics accumulation, abstracted over multiple chunks.""" + + def __init__( + self, + init_obj_value: TArrayTree | None, + weight: Numeric, + multi_device: bool, + ): + """Initializes an accumulator instance with the provided object and counter. + + Args: + init_obj_value: The initial value of the accumulator. + weight: The initial weight, which specifies how many samples are assumed + to have been already counted in the initial value of the accumulator. + multi_device: Whether the objects that are accumulated are outputs of a + multi-device computation (e.g. `jax.pmap`). + """ + self._accumulator = init_obj_value + self._weight = weight + self._multi_device = multi_device + + @property + def accumulator(self) -> TArrayTree | None: + """The current value of the underlying not-normalized accumulator.""" + return self._accumulator + + @property + def weight(self) -> Numeric | None: + """The current normalization weight of the underlying accumulator.""" + return self._weight + + @property + def multi_device(self) -> bool: + """Whether the accumulator is the output of a multi-device computation.""" + return self._multi_device + + @property + def value(self) -> TArrayTree | None: + """The current normalized value of the accumulator.""" + + if types.tree_is_empty(self.accumulator): + return self.accumulator + + if self._multi_device: + return parallel.pmap_sync_and_divide_value(self.accumulator, self.weight) + else: + return parallel.jit_sync_and_divide_value(self.accumulator, self.weight) + + def clear(self) -> None: + """Sets the underlying accumulator and weight to `None`.""" + self._accumulator = None + self._weight = None + + def value_and_clear(self) -> TArrayTree | None: + """Retrieves the normalized value of the accumulator and clears it.""" + + value = self.value + self.clear() + + return value + + def add(self, value_obj: TArrayTree, weight: Numeric = 1): + """Adds an element to the moving average and the max. + + The exact update equation for the statistics are: + raw_value_t = raw_value_{t-1} + value_obj * weight + weight_t = weight_{t-1} + weight + + Args: + value_obj: The value of the object, which scaled by `weight` will be added + to the accumulator. + weight: The relative weight of the `value_obj`. + """ + + value_obj = jax.tree_util.tree_map(lambda x: x * weight, value_obj) + + if self._accumulator is None: + + self._accumulator = value_obj + + if isinstance(weight, types.SCALAR_TYPES): + self._weight = jnp.full_like(self._weight, weight) + + elif not isinstance(weight, jax.Array): + raise ValueError("`weight` should be an instance of float, int or " + "jax.Array.") + + elif self._weight.shape != weight.shape: # pytype: disable=attribute-error # numpy-scalars + raise ValueError("If `weight` is an `jnp.ndarray` then should have the " + "same shape as the weight of the accumulator.") + else: + self._weight = weight + + return + + if not types.tree_is_empty(self._accumulator): + + if types.tree_is_empty(value_obj): + raise ValueError("The provided `value_obj` has an empty PyTree " + "structure, but the accumulator has been initialized " + "with a non-empty PyTree object.") + + self._accumulator = jax.tree_util.tree_map( + jnp.add, self._accumulator, value_obj) + + elif not types.tree_is_empty(value_obj): + + raise ValueError("The provided `value_obj` has a non-empty PyTree " + "structure, but the accumulator has been initialized " + "with an empty PyTree object.") + + self._weight = self._weight + weight + + @classmethod + def zeros_like( + cls, + obj: TArrayTree, + multi_device: bool + ) -> "MultiChunkAccumulator[TArrayTree]": + """Creates a zero initialized accumulator as `obj`.""" + + if multi_device: + value = (parallel.pmap_zeros_like(obj) + if not types.tree_is_empty(obj) else obj) + weight = parallel.replicate_all_local_devices( + jnp.zeros([], dtype=jnp.int32)) + else: + value = (parallel.jit_zeros_like(obj) + if not types.tree_is_empty(obj) else obj) + weight = jnp.zeros([], dtype=jnp.int32) + + return cls(value, weight, multi_device) + + @classmethod + def empty(cls, multi_device: bool) -> "MultiChunkAccumulator[Any]": + """Creates an empty accumulator.""" + + weight = jnp.zeros([], dtype=jnp.int32) + + if multi_device: + weight = parallel.replicate_all_local_devices(weight) + + return cls(None, weight, multi_device) + + def __repr__(self): + return (f"{self.__class__.__name__}({self._accumulator!r}, " + f"{self._weight!r}, {self._multi_device})") + + def copy(self): + """Returns a copy of the PyTree structure (but not the JAX arrays).""" + + (flattened, structure) = jax.tree_util.tree_flatten(self) + + return jax.tree_util.tree_unflatten(structure, flattened) + + +jax.tree_util.register_pytree_node( + MultiChunkAccumulator, + lambda x: ((x.accumulator, x.weight), (x.multi_device,)), + lambda fixed, arrays: MultiChunkAccumulator(*arrays, *fixed) +) diff --git a/src/kfac_jax/_src/utils/math.py b/src/kfac_jax/_src/utils/math.py new file mode 100644 index 0000000000000000000000000000000000000000..d47af37cffb1feebdf99cb8cc895d86416404a7d --- /dev/null +++ b/src/kfac_jax/_src/utils/math.py @@ -0,0 +1,1212 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC utilities for various mathematical operations.""" +import functools +import string +from typing import Callable, Iterable, Sequence, TypeVar + +import jax +from jax import lax +from jax.experimental.sparse import linalg as experimental_splinalg +import jax.numpy as jnp +from jax.scipy import linalg +from kfac_jax._src.utils import types +import numpy as np +import optax +import tree + + +Array = types.Array +Numeric = types.Numeric +PRNGKey = types.PRNGKey +ArrayTree = types.ArrayTree +TArrayTree = types.TArrayTree +TNumeric = TypeVar("TNumeric", bound=Numeric) + +_ALPHABET = string.ascii_lowercase + +# If true we use a special case formula for when a block has one or more zero +# factors. +_SPECIAL_CASE_ZERO_INV: bool = True + +# Cholesky inverses are deterministic on GPUs and somewhat faster to compute, +# but tend to perform a bit worse and not tolerate very low damping values. +# If disabling Cholesky inverses and using GPUs, one must make sure to set +# distributed_inverses=True, since otherwise the devices can potentially become +# out of sync, which is a silent and very serious failure +_USE_CHOLESKY_INVERSION: bool = False + + +def set_special_case_zero_inv(value: bool): + """Sets whether `pi_adjusted_inverse` handles zero and nan matrices.""" + global _SPECIAL_CASE_ZERO_INV + _SPECIAL_CASE_ZERO_INV = value + + +def get_special_case_zero_inv() -> bool: + """Returns whether `pi_adjusted_inverse` handles zero and nan matrices.""" + return _SPECIAL_CASE_ZERO_INV + + +def set_use_cholesky_inversion(value: bool): + """Sets whether `pi_adjusted_inverse` handles zero and nan matrices.""" + global _USE_CHOLESKY_INVERSION + _USE_CHOLESKY_INVERSION = value + + +def get_use_cholesky_inversion() -> bool: + """Returns whether `pi_adjusted_inverse` handles zero and nan matrices.""" + return _USE_CHOLESKY_INVERSION + + +def product(iterable_object: Iterable[TNumeric]) -> TNumeric: + """Computes the product of all elements in the iterable.""" + x = 1 + + for element in iterable_object: + x = x * element + + return x + + +def outer_product(*arrays: Array) -> Array: + """Computes the outer product of an arbitrary number of vectors.""" + if not all(a.ndim == 1 for a in arrays): + raise ValueError("All arrays must be vectors.") + in_str = ",".join(_ALPHABET[:len(arrays)]) + out_str = _ALPHABET[:len(arrays)] + return jnp.einsum(f"{in_str}->{out_str}", *arrays) + + +def scalar_mul(obj: TArrayTree, scalar: Numeric) -> TArrayTree: + """Multiplies all PyTree leaves of the object by the provided scalar.""" + # The check below is in its current form because of how `jax.jit` tracing + # mechanism work. If we use `scalar == 1` and `scalar` is an array, inside a + # `jit` context, jax will raise an error, since you are not allowed to use + # abstract values in concrete boolean statements, like native python + # if/while/for constructs. + if isinstance(scalar, types.SCALAR_TYPES) and scalar == 1.0: + return obj + + return jax.tree_util.tree_map(lambda x: x * scalar, obj) + + +def scalar_div(obj: TArrayTree, scalar: Numeric) -> TArrayTree: + """Divides all PyTree leaves of the object by the provided scalar.""" + # The check below is in its current form because of how `jax.jit` tracing + # mechanism work. If we use `scalar == 1` and `scalar` is an array, inside a + # `jit` context, jax will raise an error, since you are not allowed to use + # abstract values in concrete boolean statements, like native python + # if/while/for constructs. + if isinstance(scalar, types.SCALAR_TYPES) and scalar == 1.0: + return obj + + return jax.tree_util.tree_map(lambda x: x / scalar, obj) + + +def weighted_sum_of_objects( + objects: Sequence[TArrayTree], + coefficients: Sequence[Numeric], +) -> TArrayTree: + """Computes a weighted sum of the objects'. + + The function computes `sum_i coefficients[i] * objects[i]`. All objects must + have the same PyTree structure, and PyTree leaves in equivalent positions must + have the same shape. + + Args: + objects: The sequence of objects to be summed together. + coefficients: The coefficients corresponding to each object instance. + + Returns: + An object, representing the weighted sum, of the same type as the inputs. + """ + if len(objects) != len(coefficients): + raise ValueError("The number of coefficients must equal the number of " + "objects.") + if not objects: + raise ValueError("The objects' sequences can not be empty.") + + accumulator = scalar_mul(objects[0], coefficients[0]) + + for o_i, c_i in zip(objects[1:], coefficients[1:]): + if not types.abstract_objects_equal(accumulator, o_i): + raise ValueError("One or more objects do not have equivalent abstract " + "structure.") + accumulator = jax.tree_util.tree_map( + jnp.add, accumulator, scalar_mul(o_i, c_i)) + + return accumulator + + +def sum_objects(objects: Sequence[TArrayTree]) -> TArrayTree: + return weighted_sum_of_objects(objects, [1] * len(objects)) + + +def pytree_size(pytree): + """Computes total size of pytree leaves.""" + return jax.tree_util.tree_reduce( + lambda x, y: x + y, jax.tree_util.tree_map(jnp.size, pytree), 0 + ) + + +def _inner_product_float64(obj1: ArrayTree, obj2: ArrayTree) -> Array: + """Computes inner product explicitly in float64 precision.""" + + raise NotImplementedError() + + # This function isn't currently working due to a break in + # jax.experimental.enable_x64. + + # def array_ip(x, y): + # x = jnp.array(jnp.reshape(x, [-1]), dtype=jnp.float64) + # y = jnp.array(jnp.reshape(y, [-1]), dtype=jnp.float64) + # return jnp.dot(x, y, precision=lax.Precision.HIGHEST) + + # original_dtype = types.get_float_dtype_and_check_consistency((obj1, obj2)) + + # with jax.experimental.enable_x64(): + + # elements_inner_products = jax.tree_util.tree_map(array_ip, obj1, obj2) + + # flat_list = jax.tree_util.tree_leaves(elements_inner_products) + # result = flat_list[0] + + # for element_ip in flat_list[1:]: + # result = result + element_ip + + # return jnp.array(result, dtype=original_dtype) + + +def inner_product( + obj1: ArrayTree, + obj2: ArrayTree, + in_float64: bool = False +) -> Array: + """Computes the inner product ``. + + To compute the inner product, each of the two input objects is assumed to + represent a vector by flattening and concatenating all of their PyTree leaves. + Objects `obj1` and `obj2` must have the same PyTree structure, and PyTree + leaves in equivalent positions must have the same shape. + + Args: + obj1: The first object representing a vector. + obj2: The second object representing a vector. + in_float64: Whether to compute the inner product explicitly in `float64` + precision. If this is set to `True` the computation will be in double + precision regardless of whether `float64` has been enabled in Jax. + + Returns: + The scalar value of the inner product. + """ + if not types.abstract_objects_equal(obj1, obj2, check_dtype=False): + raise ValueError("The objects do not have identical abstract structure.") + + if in_float64: + return _inner_product_float64(obj1, obj2) + + elements_product = jax.tree_util.tree_map( + lambda x, y: jnp.sum(x * y), obj1, obj2) + + return sum(jax.tree_util.tree_leaves(elements_product)) + + +def symmetric_matrix_inner_products( + vectors1: Sequence[ArrayTree], + vectors2: Sequence[ArrayTree], + ip_function: Callable[[ArrayTree, ArrayTree], Array] = inner_product, +) -> Array: + """Computes a matrix of the inner products between the two sequences. + + Note that this function assumes that the output matrix is symmetric (up to + numerical precision), and if this happens not to be the case, it won't + actually compute the matrix of inner products in the expected way. If the + output matrix is not symmetric, use `asymmetric_matrix_inner_products`instead. + + Args: + vectors1: A sequence of identically structured PyTrees, each one + representing a single vector. + vectors2: A sequence of identically structured PyTrees, each one + representing a single vector. + ip_function: A callable which computes the inner product between PyTrees. + Defaults to the standard dot-product. + + Returns: + A symmetric matrix `m` with elements `m[i, j] = ` + for `i >= j`. + """ + if len(vectors1) != len(vectors2): + raise ValueError("The two sequences should have the same length.") + + m = [[] for _ in vectors1] + for i, v_i in enumerate(vectors1): + for j, v_j in enumerate(vectors2): + if j < i: + m[i].append(m[j][i]) + else: + m[i].append(ip_function(v_i, v_j)) + + return jnp.asarray(m) + + +def asymmetric_matrix_inner_products( + vectors1: Sequence[ArrayTree], + vectors2: Sequence[ArrayTree], + ip_function: Callable[[ArrayTree, ArrayTree], Array] = inner_product, +) -> Array: + """Computes a matrix of the inner products between the two sequences. + + Unlike `symmetric_matrix_inner_products`, this function doesn't assume that + the output matrix is symmetric, and therefore can be slower when it is indeed + symmetric. + + Args: + vectors1: A sequence of identically structured PyTrees, each one + representing a single vector. + vectors2: A sequence of identically structured PyTrees, each one + representing a single vector. + ip_function: A callable which computes the inner product between PyTrees. + Defaults to the standard dot-product. + + Returns: + A matrix `m` with elements `m[i, j]`. + """ + if len(vectors1) != len(vectors2): + raise ValueError("The two sequences should have the same length.") + + m = [[] for _ in vectors1] + for i, v_i in enumerate(vectors1): + for v_j in vectors2: + m[i].append(ip_function(v_i, v_j)) + + return jnp.asarray(m) + + +def matrix_of_inner_products( + vectors: Sequence[ArrayTree], + ip_function: Callable[[ArrayTree, ArrayTree], Array] = inner_product, +) -> Array: + """Computes the matrix of inner products of the sequence of vectors. + + Args: + vectors: A sequence of identically structured PyTrees, each one representing + a single vector. + ip_function: A callable which computes the inner product between PyTrees. + Defaults to the standard dot-product. + + Returns: + A matrix `m` with elements `m[i, j] = `. + """ + return symmetric_matrix_inner_products(vectors, vectors, + ip_function=ip_function) + + +def vector_of_inner_products( + base: ArrayTree, + vectors: Sequence[ArrayTree], + ip_function: Callable[[ArrayTree, ArrayTree], Array] = inner_product, +) -> Array: + """Computes a vector of inner products with base. + + Args: + base: A PyTree representing the base vector. + vectors: A sequence of identically structured PyTrees, each one representing + a single vector. + ip_function: A callable which computes the inner product between PyTrees. + Defaults to the standard dot-product. + + Returns: + A vector `v` with elements `v[i] = `. + """ + v = [] + for v_i in vectors: + v.append(ip_function(v_i, base)) + + return jnp.asarray(v) + + +def block_permuted( + matrix: Array, + block_sizes: Sequence[int], + block_order: Sequence[int], +) -> Array: + """Permutes whole blocks of the input matrix. + + Given a square matrix, this function splits it into blocks, each one having + a size defined in `block_sizes` and permutes them, both in rows and + columns. The permutation sends to the `i` slot the `block_order[i]` block of + the input matrix. Example: + matrix = [[A_0, B_0, C_0], [A_1, B_1, C_1], [A_2, B_2, C_2]] + block_order = [2, 0, 1] + => [[C_2, A_2, B_2], [C_0, A_0, B_0], [C_1, A_1, B_1]] + + Args: + matrix: The matrix, whose blocks will be permuted. + block_sizes: A sequences of each block's size. + block_order: A sequence of the order of the blocks. + + Returns: + The resulting matrix after permuting the blocks. + """ + if len(block_sizes) != len(block_order): + raise ValueError( + f"The length of `block_sizes` (=={len(block_sizes)} " + f"and `block_order` (=={len(block_order)}) must be " + "the same.") + + if all(i == j for i, j in enumerate(block_order)): + return matrix + + indices = np.cumsum(block_sizes)[:-1] + blocks = [jnp.split(row, indices, 1) for row in jnp.split(matrix, indices, 0)] + reordered_blocks = [[blocks[i][j] for j in block_order] for i in block_order] + + return jnp.block(reordered_blocks) + + +def squared_norm(obj: ArrayTree) -> Array: + """Computes the squared Euclidean norm of the provided PyTree object.""" + elements_squared_norm = jax.tree_util.tree_map( + lambda x: jnp.sum(jnp.square(x)), obj) + + return sum(jax.tree_util.tree_leaves(elements_squared_norm)) + + +def norm(obj: ArrayTree) -> Array: + """Computes the Euclidean norm of the provided PyTree object.""" + elements_squared_norm = jax.tree_util.tree_map( + lambda x: jnp.sum(jnp.square(x)), obj) + + return jnp.sqrt(sum(jax.tree_util.tree_leaves(elements_squared_norm))) + + +def per_parameter_norm(obj: ArrayTree, key_prefix: str) -> ArrayTree: + + per_param_norm = jax.tree_util.tree_map(jnp.linalg.norm, obj) + per_param_norm = tree.flatten_with_path(per_param_norm) + + return { + key_prefix + "(" + "/".join(k) + ")": v for k, v in per_param_norm + } + + +def psd_inv(matrix: Array) -> Array: + """Computes the inverse of `matrix`, which is assumed PSD.""" + + if matrix.shape[:1] != matrix.shape[1:]: + raise ValueError(f"Expected square matrix, but got shape {matrix.shape}.") + + if get_use_cholesky_inversion(): + identity = jnp.eye(matrix.shape[0], dtype=matrix.dtype) + return linalg.solve(matrix, identity, assume_a="pos") + else: + # Cuda's LU solver will go into an infinite loop if the matrix has NaNs or + # possibly Infs, so we need to check for that before calling it. + return lax.cond( + jnp.logical_or(jnp.any(jnp.isnan(matrix)), jnp.any(jnp.isinf(matrix))), + lambda: jnp.full(matrix.shape, jnp.nan, dtype=matrix.dtype), + lambda: linalg.inv(matrix), + ) + + +def psd_solve(matrix: Array, vector: Array) -> Array: + """Computes the solution of `matrix * x = vector`, for a PSD `matrix`.""" + + if matrix.shape[:1] != matrix.shape[1:]: + raise ValueError(f"Expected square matrix, but got shape {matrix.shape}.") + + if get_use_cholesky_inversion(): + return linalg.solve(matrix, vector, assume_a="pos") + else: + # Cuda's LU solver will go into an infinite loop if the matrix has NaNs or + # possibly Infs, so we need to check for that before calling it. + return lax.cond( + jnp.logical_or(jnp.any(jnp.isnan(matrix)), jnp.any(jnp.isinf(matrix))), + lambda: jnp.full(vector.shape, jnp.nan, dtype=vector.dtype), + lambda: linalg.solve(matrix, vector), + ) + + +def psd_solve_without_last_idx(a: Array, b: Array) -> Array: + sub_a = a[..., :-1, :-1] + sub_b = b[..., :-1] + sub_x = psd_solve(sub_a, sub_b) + return jnp.concatenate([sub_x, jnp.zeros_like(b[..., :1])], axis=-1) + + +def psd_solve_maybe_zero_last_idx(a: Array, b: Array) -> Array: + # Check the last column and row are zero. + check = jnp.logical_and(jnp.all(a[..., -1] == 0), jnp.all(a[..., -1, :] == 0)) + return jax.lax.cond(check, psd_solve_without_last_idx, psd_solve, a, b) + + +def psd_matrix_norm( + matrix: Array, + norm_type: str = "avg_diag", + method_2norm: str = "lobpcg", + rng_key: PRNGKey | None = None +) -> Numeric: + """Computes one of several different matrix norms for PSD matrices. + + NOTE: not all the functions options provided here are actually norms, but most + are. + + Args: + matrix: a square matrix represented as a 2D array, a 1D vector giving the + diagonal, or a 0D scalar (which gets interpreted as a 1x1 matrix). Must be + positive semi-definite (PSD). + norm_type: a string specifying the type of matrix norm. Can be "2_norm" for + the matrix 2-norm aka the spectral norm, "avg_diag" for the average of + diagonal entries, "1_norm" for the matrix 1-norm, or "avg_fro" for the + Frobenius norm divided by the square root of the number of rows. + method_2norm: a string specifying the method used to compute 2-norms. Can + be "lobpcg" (recommended) or "power_iteration". + rng_key: an optional JAX PRNGKey key to used initialize the lobpcg method + for computing the 2-norm. + + Returns: + A 0D scalar giving the requested norm. + """ + + if norm_type == "2_norm": + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.max(matrix) + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + + if method_2norm == "lobpcg": + + if rng_key is None: + rng_key = jax.random.PRNGKey(123) + + v = jax.random.normal(rng_key, shape=[matrix.shape[0], 1]) + + return experimental_splinalg.lobpcg_standard( + matrix, v, m=300, tol=1e-8)[0][0] + + elif method_2norm == "power_iteration": + + return float(optax.power_iteration( + matrix, num_iters=300, error_tolerance=1e-7)[0]) + + else: + raise ValueError(f"Unrecognized method string: '{norm_type}'") + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "avg_diag": + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.sum(matrix) / matrix.shape[0] + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + return jnp.trace(matrix) / matrix.shape[0] + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "median_diag": + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.median(matrix) + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + return jnp.median(jnp.diag(matrix)) + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "trace": + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.sum(matrix) + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + return jnp.trace(matrix) + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "median_eig": + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.median(matrix) + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + # call safe_psd_eigh instead? + s, _ = jnp.linalg.eigh(matrix) + return jnp.median(s) + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "one_over_dim": # this isn't a norm + + if matrix.ndim == 0: + return 1.0 + + elif matrix.ndim == 1: + return 1.0 / matrix.shape[0] + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + return 1.0 / matrix.shape[0] + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "1_norm": # equiv to inf norm for symmetric matrices + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.max(matrix) + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + return jnp.linalg.norm(matrix, ord=1) + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "avg_fro": + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.linalg.norm(matrix) / jnp.sqrt(matrix.shape[0]) + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + return jnp.linalg.norm(matrix) / jnp.sqrt(matrix.shape[0]) + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + elif norm_type == "fro": + + if matrix.ndim == 0: + return matrix + + elif matrix.ndim == 1: + return jnp.linalg.norm(matrix) + + elif matrix.ndim == 2 and matrix.shape[0] == matrix.shape[1]: + return jnp.linalg.norm(matrix) + + else: + raise ValueError(f"Unsupported shape for factor array: {matrix.shape}") + + raise ValueError(f"Unrecognized norm type: '{norm_type}'") + + +def pi_adjusted_kronecker_factors( + *factors: Array, + damping: Numeric +) -> tuple[Array, ...]: + """Computes Kronecker factors with pi-adjusted factored damping. + + The `f1 kron f2 kron ... kron fn + damping * I` is not a Kronecker product + in general, because of the added identity. [1] proposed a pi-adjusted factored + damping approach to approximate it as a Kronecker product. [2] generalized + this approach from two to tree factors, and [3] generalized it to arbitrary + numbers of factors. This function implements the generalized approach. + + [1] - https://arxiv.org/abs/1503.05671 + [2] - https://openreview.net/forum?id=SkkTMpjex + [3] - https://ui.adsabs.harvard.edu/abs/2021arXiv210602925R/abstract + + Args: + *factors: A list of factors represented as 2D arrays, vectors (which are + interpreted as representing the diagonal of a matrix) or scalars (which + are interpreted as being a 1x1 matrix). All factors must be PSD. + damping: The weight of the identity added to the Kronecker product. + + Returns: + A list of factors with the same length as `factors`, and with the same + corresponding representations, whose Kronecker product approximates + `(f1 kron f2 kron ... kron fn) + damping * I` according to the + pi-adjusted factored-damping approach. + """ + + # The implementation writes each single factor as `c_i u_i`, where the matrix + # `u_i` is such that `trace(u_i) / dim(u_i) = 1`. We then factor out all the + # scalar factors `c_i` into a single overall scaling coefficient and + # distribute the damping to each single non-scalar factor `u_i` equally. + + norm_type = "avg_diag" + + norms = jnp.array([psd_matrix_norm(f, norm_type=norm_type) for f in factors]) + + k = len(factors) + + # TODO(jamesmartens,botev): consider making the use of special behavior for + # scalar factors a module-level configurable option. One can argue that scalar + # factors should behave the same as non-scalar factors for the sake of + # consistent behavior as the layer widths shrink to 1. + + def regular_case() -> tuple[Array, ...]: + + num_non_scalars = sum(1 if f.size != 1 else 0 for f in factors) + + # Compute the normalized factors `u_i`, such that Trace(u_i) / dim(u_i) = 1 + us = [fi / ni for fi, ni in zip(factors, norms)] + + if num_non_scalars != 0: + + # Distribute c and damping/c among k factors, where c = jnp.prod(norms), + # satisfying kron(factors) = c * kron(us). + + # NOTE: c_k (geometric mean of norms) can also be calculated by + # c ** (1/k) = jnp.prod(norms) ** (1 / len(norms)), but this alternative + # can make the result zero due to the multiplication of (potentially) + # small values, i.e. jnp.prod(norms). + c_k = jnp.exp(jnp.mean(jnp.log(norms))) + + d_k = jnp.power(damping, 1.0 / k) / c_k + + if k > num_non_scalars: + + c_non_scalar = c_k ** (float(k) / num_non_scalars) + + # We distribute the damping only inside the non-scalar factors + d_hat = jnp.power(damping, 1.0 / num_non_scalars) / c_non_scalar + + else: + d_hat = d_k + + else: + + # This could cause under/overflow, but it's unavoidable here. + c = jnp.prod(jnp.array(norms)) + + # In the case where all factors are scalar we need to add the damping and + # then take the k-th root + c_k = jnp.power(c + damping, 1.0 / k) + + u_hats = [] + + for u in us: + if u.size == 1: # scalar case + u_hat = jnp.ones_like(u) # damping not used in the scalar factors + + elif u.ndim == 2: + u_hat = u + d_hat * jnp.eye(u.shape[0], dtype=u.dtype) + + else: # diagonal case + assert u.ndim == 1 + u_hat = u + d_hat + + u_hats.append(u_hat * c_k) + + return tuple(u_hats) + + def zero_case() -> tuple[Array, ...]: + + # In the special case where for some reason one of the factors is zero, then + # the we write each factor as `damping^(1/k) * I`. + + c_k = jnp.power(damping, 1.0 / k) + + return tuple( + c_k * (jnp.eye(fi.shape[0], dtype=fi.dtype) if fi.ndim == 2 else + jnp.ones_like(fi)) + for fi in factors + ) + + if get_special_case_zero_inv(): + return lax.cond(jnp.greater(jnp.min(norms), 0.0), regular_case, zero_case) + + else: + return regular_case() + + +def invert_psd_matrices( + matrices: ArrayTree +) -> ArrayTree: + """Inverts a PyTree of matrices. + + Args: + matrices: A PyTree of 2D arrays, vectors (which are interpreted as + representing the diagonal of a matrix) or scalars (which are interpreted + as being a 1x1 matrix) representing the matrices to be inverted. All + matrices must be PSD. + + Returns: + A PyTree of matrices giving the inverses of the corresponding matrices + passed as arguments (with the same respective representations). + """ + + def invert_psd_matrix(m): + + if m.ndim == 2: + return psd_inv(m) + + assert m.ndim <= 1 + return 1.0 / m + + return jax.tree_util.tree_map(invert_psd_matrix, matrices) + + +def inverse_sqrt_psd_matrices(matrices: ArrayTree) -> ArrayTree: + + def inverse_sqrt_psd_matrix(m): + + if m.ndim == 2: + # Check copy.bara.sky before changing the next line: + return qr_pth_inv_root.qr_pth_inv_root(4, m, cholesky_qr=True) + + assert m.ndim <= 1 + return 1.0 / jnp.sqrt(m) + + return jax.tree_util.tree_map(inverse_sqrt_psd_matrix, matrices) + + +def pi_adjusted_kronecker_inverse( + *factors: Array, + damping: Numeric, +) -> tuple[Array, ...]: + """Computes pi-adjusted factored damping inverses. + + The inverse of `(f1 kron f2 kron ... kron fn) + damping * I` is not Kronecker + factored in general, because of the added identity. [1] proposed a pi-adjusted + factored damping approach to approximate the inverse as a Kronecker product. + [2] generalized this approach from two to tree factors, and [3] generalized it + to arbitrary numbers of factors. This function implements the generalized + approach. + + [1] - https://arxiv.org/abs/1503.05671 + [2] - https://openreview.net/forum?id=SkkTMpjex + [3] - https://ui.adsabs.harvard.edu/abs/2021arXiv210602925R/abstract + + Args: + *factors: A list of factors represented as 2D arrays, vectors (which are + interpreted as representing the diagonal of a matrix) or scalars (which + are interpreted as being a 1x1 matrix). All factors must be PSD. + damping: The weight of the identity added to the Kronecker product. + + Returns: + A list of factors with the same length as `factors`, and with the same + corresponding representations, whose Kronecker product approximates the + inverse of `(f1 kron f2 kron ... kron fn) + damping * I` according to the + pi-adjusted factored-damping approach. + """ + + return invert_psd_matrices( + pi_adjusted_kronecker_factors(*factors, damping=damping)) + + +def kronecker_product_axis_mul_v( + factors: Sequence[Array], + v: Array, + axis_groups: Sequence[Sequence[int]] | None = None, + transpose: bool | Sequence[bool] = False, +): + """Computes ``kron(*factors) rvec(v)`` where ``rvec`` is row-wise vectorization. + + Args: + factors: The sequence of factors forming the Kronecker product. Must be + square 2D arrays or `None`, which is interpreted as identity. + v: A tensor whose vectorization will be multiplied by the Kronecker product. + axis_groups: A list whose i-th element is a sequence of consecutive integers + specifying the axes of the input tensor ``v`` that correspond to the i-th + Kronecker factor. Passing ``None`` is equivalent to passing + ``[[0],[1],[2],...]``. + transpose: A single boolean or a sequence of booleans. If it is a sequence, + each element specifies if the corresponding factor should be transposed. + If it is a single boolean, specifies if all factors should be transposed. + + Returns: + The result, shaped as a tensor, of multiplying the vectorization of the + input tensor by the Kronecker-factored matrix. + """ + if axis_groups is None: + axis_groups = tuple((i,) for i in range(v.ndim)) + else: + axis_groups = tuple(tuple(group) for group in axis_groups) + + # Sanity checks + if sum(axis_groups, ()) != tuple(range(v.ndim)): + raise ValueError(f"The `axis_groups={axis_groups}` are either not in " + f"consecutive order or do not cover exactly the axis of " + f"the input `v`..") + if len(factors) != len(axis_groups): + raise ValueError("The number of factors provided must be equal to the " + "number of axis groups provided.") + + if isinstance(transpose, bool): + transpose = [transpose] * len(factors) + + elif len(transpose) != len(factors): + raise ValueError("The length of the transpose sequence must match the " + "number of factors.") + + factor_strs = ["yz" if t else "zy" for t in transpose] + general_str = _ALPHABET[:v.ndim] + + result = v + for group, factor, f_str in zip(axis_groups, factors, factor_strs): + + if factor is None: + continue + + # This flattens all axis in `group` of `result` into a single one. + shape = v.shape[:min(group)] + (-1,) + v.shape[max(group) + 1:] + vector = result.reshape(shape) + + # This contracts `result` with `factor` along the single axis. + vector_str = general_str[:min(group)] + "y" + general_str[max(group) + 1:] + result_str = vector_str.replace("y", "z") + einsum_str = f"{f_str},{vector_str}->{result_str}" + r_next = jnp.einsum(einsum_str, factor, vector) + + # This reshapes back to the original shape. + result = r_next.reshape(v.shape) + + return result + + +def kronecker_eigen_basis_axis_mul_v( + q_factors: Sequence[Array], + eigenvalues: Array, + v: Array, + axis_groups: Sequence[Sequence[int]] | None = None, +): + """Computes a matrix-vector product in a Kronecker product eigen-basis. + + The function computes: + ``kron(*q_factors) diag(eigenvalues) kron(*q_factors)^T rvec(v)`` + + where all variables are appropriately sized matrices and ``rvec`` is + row-wise vectorization. The computation is related to the usual Kronecker + product ``kron(*factors) rvec(v)``, if ``factors`` are all symmetric PSD + matrices and ``q_factors`` are the matrices of eigenvectors of ``factors`` and + ``eigenvalues`` is the kronecker product of the eigenvalues of ``factors``. + However, the function does not assume that its inputs are of this form. + + Args: + q_factors: A sequence of the orthonormal basis of eigenvectors of each + Kronecker factor. + eigenvalues: A tensor containing the eigenvalues (e.g. the Kronecker product + of eigenvalues of all factors). + v: The input vector as a tensor. + axis_groups: A list whose i-th element is a sequence of consecutive integers + specifying the axes of the input tensor ``v`` that correspond to the i-th + Kronecker factor. Passing ``None`` is equivalent to passing + ``[[0],[1],[2],...]``. + + Returns: + The result of multiplying the input vector by the Kronecker product of the + factors, shaped as a tensor. + """ + q_proj_v = kronecker_product_axis_mul_v(q_factors, v, axis_groups, True) + + if eigenvalues.shape != q_proj_v.shape: + raise ValueError("The eigenvalues array should have the same shape as the " + "projection of `v` onto `kron(*factors)`.") + + eig_weighted_v = eigenvalues * q_proj_v + + return kronecker_product_axis_mul_v(q_factors, eig_weighted_v, axis_groups) + + +def kronecker_product_mul_v( + a: Array, + b: Array, + v: Array, + a_is_symmetric: bool, +) -> Array: + """Computes `unvec[(a kron b) vec(v)]` for correctly sized input matrices.""" + del a_is_symmetric # not used + return kronecker_product_axis_mul_v([b, a], v) + + +def kronecker_eigen_basis_mul_v( + q_a: Array, + q_b: Array, + eigenvalues: Array, + v: Array, +) -> Array: + """Computes a matrix-vector product in a Kronecker product eigen-basis. + + The function computes: + `(q_a kron q_b) diagonal(eigenvalues) (q_a kron q_b)^T vec(v)` + + where all variables are appropriately sized matrices. The computation is + related to the usual Kronecker product `(a kron b) vec(v)`, if `a` and `b` are + symmetric matrices and `q_a` and `q_b` are the matrices of eigenvectors of `a` + and `b` and `eigenvalues` is the outer product of the eigenvalues of `a` and + `b`. However, the function does not assume anything about the `eigenvalues` + and allows for any dense matrix. + + Args: + q_a: An orthonormal basis for eigenvectors of the first Kronecker factor. + q_b: An orthonormal basis for eigenvectors of the second Kronecker factor. + eigenvalues: A matrix containing the eigenvalues (e.g. the product of + eigenvalues of both factors). + v: The input vector as a matrix. + + Returns: + The result of the matrix-vector product. + """ + return kronecker_eigen_basis_axis_mul_v([q_b, q_a], eigenvalues, v) + + +def _host_eigh(x: Array, *_) -> tuple[Array, Array]: + """This calls the CPU numpy function for eigh.""" + + shape_s = jax.ShapeDtypeStruct(x.shape[:-1], x.dtype) + shape_q = jax.ShapeDtypeStruct(x.shape, x.dtype) + + return jax.pure_callback(np.linalg.eigh, (shape_s, shape_q), x) + + +def _eigh( + x: Array, + force_on_host: bool = False, +) -> tuple[Array, Array]: + """Computes eigenvectors and eigenvalues, with optionally offloading to cpu.""" + + if force_on_host: + return _host_eigh(x) + + s, q = jnp.linalg.eigh(x) + + # Recently with CUDA 11.7 there is a bug in cuSOLVER which makes the eigh + # implementation unstable sometimes on GPUs. + return jax.lax.cond( + jnp.any(jnp.isnan(s)), + _host_eigh, + lambda *args: args[1:], + x, s, q + ) + + +def safe_psd_eigh( + x: Array, + force_on_host: bool = False, +) -> tuple[Array, Array]: + """Computes the eigenvalue decomposition for a PSD matrix. + + The function is similar to `jax.numpy.linalg.eigh`, but it clips the returned + eigenvalues to always be non-negative, which we know mathematically holds for + PSD matrices, but due to numerical errors `jax.numpy.linalg.eigh` could return + negative values. + + Args: + x: The input matrix, assumed to be PSD. + force_on_host: If `True` will perform the computation on the host CPU. + + Returns: + A pair of (eigenvalues, eigenvectors) arrays. + """ + + d = x.shape[0] + + # Here we are handling the case of NaNs separately, because in some versions + # of cuda and cudablas they can cause a runtime error. + s, q = lax.cond( + jnp.any(jnp.isnan(x)), + lambda _: (jnp.full([d], jnp.nan, dtype=x.dtype), # pylint: disable=g-long-lambda + jnp.full([d, d], jnp.nan, dtype=x.dtype)), + functools.partial(_eigh, force_on_host=force_on_host), + x, + ) + + # The matrix is PSD by construction, but numerical inaccuracies can produce + # slightly negative eigenvalues. Hence, clip at zero. + return jnp.clip(s, min=0.0), q + + +def tnt_scale(factors: Sequence[Array]) -> Numeric: + """Computes the correct scaling factor for a TNT factorization.""" + + if len(factors) == 1: + return 1.0 + + # These should be the same values + zs = jnp.asarray([jnp.trace(factor) for factor in factors]) + + # We want to compute geometric_mean(zs) ** -(num_factors - 1) + + # deal with the case where one of the factors is zero + mean_log = lax.select( + jnp.greater(jnp.min(zs), 0.0), + jnp.mean(jnp.log(zs)), + 0.0, + ) + + return jnp.exp(-(len(factors) - 1) * mean_log) + + +def loop_and_parallelize_average( + func: Callable[..., ArrayTree], + max_parallel_size: int, +) -> Callable[..., ArrayTree]: + """Returns a function that computes the average of `func` over any arguments. + + The returned function is mathematically equivalent to + jnp.mean(jax.vmap(func)(*args), axis=0). + However, naively using the above code could lead to prohibitively large memory + usage, as it scales linearly with the leading axis size of `args`, because of + `jax.vmap`. To amortize the memory cost, if the leading axis has size larger + than `max_parallel_size`, we call multiple times `vmap` in a loop via `scan` + by splitting the arguments to multiple chunks. This allows to trade off memory + usage for the cost of compute time. + + Args: + func: A function that computes a singleton output. + max_parallel_size: The maximum number of elements that are allowed to be + part of a single call to `jax.vmap`. + + Returns: + A function that computes the averaged output of `func` over the leading + axis of its arguments. + """ + vmap_fn = jax.vmap(func) + + @functools.wraps(func) + def average_func(*args) -> ArrayTree: + + lead_axis_sizes = set(x.shape[0] for x in jax.tree_util.tree_leaves(args)) + + if not lead_axis_sizes: + raise ValueError("You must pass in at least one argument with a PyTree " + "leaf node.") + + elif len(lead_axis_sizes) != 1: + raise ValueError(f"Inconsistent leading axis sizes seen: " + f"{lead_axis_sizes!r}.") + + leading_size = next(iter(lead_axis_sizes)) + + singleton_args = jax.tree_util.tree_map(lambda _x: _x[0], args) + _, output_tree = jax.make_jaxpr(func, return_shape=True)(*singleton_args) + + singleton_size = sum(x.size for x in jax.tree_util.tree_leaves(output_tree)) + output_size = singleton_size * leading_size + + # Compute the loop size and any remainder size + if max_parallel_size is None or output_size <= max_parallel_size: + + parallel_size = leading_size + + else: + parallel_size = max( + min(max_parallel_size // singleton_size, leading_size), 1) + + # The arguments have to be split into chunks along their leading axis, + # however since `jax.scan` does not support inputs with different size, + # if the leading axis is not divisible by the parallel_size, we need to + # separately compute the values for the last remaining arguments chunks. + num_parallel_chunks = leading_size // parallel_size + remainder_size = leading_size % parallel_size + all_chunks_size = leading_size - remainder_size + + # Index to get the loop arguments + loop_args = jax.tree_util.tree_map(lambda x: x[:all_chunks_size], args) + + if num_parallel_chunks == 1: + averaged_value = jnp.mean(vmap_fn(*loop_args), axis=0) + + else: + + def scan_fn(accumulator, args_): + + vmap_value = vmap_fn(*args_) + + avg_value = jax.tree_util.tree_map( + lambda x: jnp.mean(x, axis=0), vmap_value) + + return jax.tree_util.tree_map(jnp.add, accumulator, avg_value), None + + loop_shape = (num_parallel_chunks, parallel_size) + + loop_args = jax.tree_util.tree_map( + lambda x: x.reshape(loop_shape + x.shape[1:]), + loop_args) + + summed_value, _ = jax.lax.scan( + scan_fn, + init=jax.tree_util.tree_map( + jnp.zeros_like, output_tree), + xs=loop_args) + + averaged_value = scalar_div(summed_value, num_parallel_chunks) + + if remainder_size == 0: + return averaged_value + + # Index to get the remainder arguments + remainder_args = jax.tree_util.tree_map(lambda x: x[all_chunks_size:], args) + remainder_value = jnp.mean(vmap_fn(*remainder_args), axis=0) + + avg_weight = all_chunks_size / leading_size + remainder_weight = remainder_size / leading_size + + return weighted_sum_of_objects( + [averaged_value, remainder_value], [avg_weight, remainder_weight]) + + return average_func + + +@functools.partial(jax.custom_jvp, nondiff_argnums=(1,)) +def _sqrt_bound_derivative( + x: jax.Array, + max_gradient: float | jax.Array, +) -> jax.Array: + """Computes a square root with a gradient clipped at `max_gradient`.""" + del max_gradient # unused + return jnp.sqrt(x) + + +def _stable_sqrt_fwd( + max_gradient: float | jax.Array, + primals: tuple[jax.Array], # pylint: disable=g-one-element-tuple + tangents: tuple[jax.Array], # pylint: disable=g-one-element-tuple +) -> tuple[jax.Array, jax.Array]: + """Forward mode autodiff of square-root.""" + (x,) = primals + x_pre = jnp.maximum(x, 1 / (4 * max_gradient**2)) + + _, tangent = jax.jvp(jnp.sqrt, (x_pre,), tangents) + + return jnp.sqrt(x), tangent + + +_sqrt_bound_derivative.defjvp(_stable_sqrt_fwd) + +stable_sqrt = functools.partial(_sqrt_bound_derivative, max_gradient=1000.0) + + +def cosine_similarity(v1: ArrayTree, v2: ArrayTree) -> Array: + """Computes the cosine similarity between flattened pytrees.""" + return inner_product(v1, v2) / (norm(v1) * norm(v2)) diff --git a/src/kfac_jax/_src/utils/misc.py b/src/kfac_jax/_src/utils/misc.py new file mode 100644 index 0000000000000000000000000000000000000000..e09a0ecaf94e3e7e0c635f3e787ca519c9565449 --- /dev/null +++ b/src/kfac_jax/_src/utils/misc.py @@ -0,0 +1,395 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC related utility classes and functions.""" +import abc +import dataclasses +import functools +import inspect +import itertools +from typing import Any, Callable, Iterator, Sequence, TypeVar + +import jax +import jax.numpy as jnp +from kfac_jax._src.utils import types + +T = types.T +S = TypeVar("S", bound=Sequence) +Array = types.Array +Numeric = types.Numeric +ArrayTree = types.ArrayTree +TArrayTree = types.TArrayTree +StateType = TypeVar("StateType") +StateTree = types.PyTree["State"] + + +STATE_CLASSES_SERIALIZATION_DICT = {} + + +def fake_element_from_iterator( + iterator: Iterator[TArrayTree], +) -> tuple[TArrayTree, Iterator[TArrayTree]]: + """Returns a zeroed-out initial element of the iterator "non-destructively". + + This function mutates the input iterator, hence after calling this function + it will be advanced by one. An equivalent to the original iterator (e.g. not + advanced by one) is returned as the second element of the returned pair. The + advised usage of the function is: + `fake_element, iterator = fake_element_from_iterator(iterator)` + + Args: + iterator: A PyTree iterator. Must yield at least one element. + + Returns: + A pair `(element, output_iterator)` where `element` is a zeroed-out version + of the first element of the iterator, and `output_iterator` is an + equivalent iterator to the input one. + """ + init_element = next(iterator) + fake_element = jax.tree_util.tree_map(jnp.zeros_like, init_element) + def equivalent_iterator() -> Iterator[ArrayTree]: + yield init_element + # For some reason unknown to us, "yield from" can fail in certain + # circumstances + while True: + yield next(iterator) + return fake_element, equivalent_iterator() + + +def filter_sequence( + unfiltered_sequence: S, + bool_sequence: Sequence[bool] +) -> S: + + filtered = itertools.compress(unfiltered_sequence, bool_sequence) + + return tuple(filtered) if isinstance( + unfiltered_sequence, tuple) else list(filtered) + + +def to_tuple_or_repeat( + x: Numeric | Sequence[Numeric], + length: int, +) -> tuple[Numeric, ...]: + """Converts `x` to a tuple of fixed length. + + If `x` is an array, it is split along its last axis to a tuple (assumed to + have `x.shape[-1] == length`). If it is a scalar, the scalar is repeated + `length` times into a tuple, and if it is a list or a tuple it is just + verified that its length is the same. + + Args: + x: The input array, scalar, list or tuple. + length: The length of the returned tuple. + + Returns: + A tuple constructed by either replicating or splitting `x`. + """ + if isinstance(x, jnp.ndarray) and x.size > 1: # pytype: disable=attribute-error + assert x.shape[-1] == length # pytype: disable=attribute-error + return tuple(x[..., i] for i in range(length)) + elif isinstance(x, (list, tuple)): + assert len(x) == length + return tuple(x) + elif isinstance(x, (int, float, jnp.ndarray)): + return (x,) * length + else: + raise ValueError(f"Unrecognized type for `x` - {type(x)}.") + + +def first_dim_is_size(size: int, *args: Array) -> bool: + """Checks that each element of `args` has first axis size equal to `size`.""" + return all(arg.shape[0] == size for arg in args) + + +def rearrange(x: Array, spec: str) -> Array: + """Rearranges the array according to the given spec, equivalent to https://einops.rocks/api/rearrange.""" + + in_str, out_str = spec.split("->") + in_str = in_str.replace(" ", "") + out_str = out_str.replace(" ", "") + + assert len(in_str) == x.ndim + + # Reorder + order = [] + for axis in out_str.replace("(", "").replace(")", ""): + if axis != "1": + order.append(in_str.index(axis)) + x = jnp.transpose(x, order) + + assert len(order) == x.ndim + + # Reshape + shape = [] + open_bracket = False + size = None + i = 0 + for s in out_str: + if s == "1": + shape.append(1) + elif s == "(": + size = 1 + open_bracket = True + elif s == ")": + open_bracket = False + shape.append(size) + size = None + elif open_bracket: + size *= x.shape[i] + i += 1 + else: + shape.append(x.shape[i]) + i += 1 + + return x.reshape(shape) + + +class State(abc.ABC): + """Abstract class for state classes.""" + + @classmethod + def field_names(cls) -> tuple[str, ...]: + return tuple(field.name for field in dataclasses.fields(cls)) # pytype: disable=wrong-arg-types + + @classmethod + def field_types(cls) -> dict[str, type[Any]]: + return {field.name: field.type for field in dataclasses.fields(cls)} # pytype: disable=wrong-arg-types + + @property + def field_values(self) -> tuple[ArrayTree, ...]: + return tuple(getattr(self, name) for name in self.field_names()) + + def copy(self: StateType) -> StateType: + """Returns a copy of the PyTree structure (but not the JAX arrays).""" + (flattened, structure) = jax.tree_util.tree_flatten(self) + return jax.tree_util.tree_unflatten(structure, flattened) + + def tree_flatten(self) -> tuple[tuple[ArrayTree, ...], tuple[str, ...]]: + return self.field_values, self.field_names() + + @classmethod + def tree_unflatten( + cls, + aux_data: tuple[str, ...], + children: tuple[ArrayTree, ...], + ): + return cls(**dict(zip(aux_data, children))) + + def __repr__(self) -> str: + return (f"{self.__class__.__name__}(" + + ",".join(f"{name}={v!r}" for name, v in self.field_values) + + ")") + + +def register_state_class(class_type: type[Any]) -> type[Any]: + """Extended dataclass decorator, which also registers the class as a PyTree. + + The function is equivalent to `dataclasses.dataclass`, but additionally + registers the `class_type` as a PyTree. This is done by setting the PyTree + nodes of all `dataclasses.fields` of the class. + + Args: + class_type: The class type to transform. + + Returns: + The transformed `class_type` which is now a dataclass and also registered as + a PyTree. + """ + if not issubclass(class_type, State): + raise ValueError( + f"Class {class_type} is not a subclass of kfac_jax.utils.State." + ) + + class_type = dataclasses.dataclass(class_type) + class_type = jax.tree_util.register_pytree_node_class(class_type) + class_name = f"{class_type.__module__}.{class_type.__qualname__}" + STATE_CLASSES_SERIALIZATION_DICT[class_name] = class_type + return class_type + + +def serialize_state_tree(instance: StateTree) -> ArrayTree: + """Returns a recursively constructed dictionary of the state.""" + if isinstance(instance, State): + result_dict = {name: serialize_state_tree(getattr(instance, name)) + for name in instance.field_names()} + cls = instance.__class__ + result_dict["__class__"] = f"{cls.__module__}.{cls.__qualname__}" + return result_dict + + elif isinstance(instance, list): + return [serialize_state_tree(v) for v in instance] + + elif isinstance(instance, tuple): + return tuple(serialize_state_tree(v) for v in instance) + + elif isinstance(instance, set): + return set(serialize_state_tree(v) for v in instance) + + elif isinstance(instance, dict): + return {k: serialize_state_tree(v) for k, v in instance.items()} + + else: + return instance # pytype: disable=bad-return-type + + +def deserialize_state_tree(representation: ArrayTree) -> StateTree: + """Returns the state class using a recursively constructed.""" + if isinstance(representation, list): + return [deserialize_state_tree(v) for v in representation] + + elif isinstance(representation, tuple): + return tuple(deserialize_state_tree(v) for v in representation) + + elif isinstance(representation, set): + return set(deserialize_state_tree(v) for v in representation) + + elif isinstance(representation, dict): + if "__class__" not in representation: + return {k: deserialize_state_tree(v) for k, v in representation.items()} + + class_name = representation.pop("__class__") + if class_name not in STATE_CLASSES_SERIALIZATION_DICT: + raise ValueError(f"Did not find how to reconstruct class {class_name}.") + + dict_rep = deserialize_state_tree(representation) + return STATE_CLASSES_SERIALIZATION_DICT[class_name](**dict_rep) + + else: + return representation + + +class Finalizable(abc.ABC): + """A mixin for classes that can "finalize" their attributes. + + The class provides the function `finalize` which freezes all attributes of the + instance after its call. Any attributes assignment thereafter will raise an + error. All subclasses must always call `super().__init__()` for the mixin to + function properly, and they must set any attributes before any call to + `finalize` has happened. + """ + + def __init__( + self, + forbid_setting_attributes_after_finalize: bool = True, + excluded_attribute_names: Sequence[str] = (), + **parent_kwargs: Any, + ): + """Initializes the instance. + + Args: + forbid_setting_attributes_after_finalize: If `True`, trying to set + attributes (via direct obj.attr = ...) after `finalize` was called on + the instance will raise an error. If `False`, this is not checked. + excluded_attribute_names: When `forbid_setting_attributes_after_finalize` + is set to `True` this specifies any attributes names that can still be + set. + **parent_kwargs: Any keyword arguments to be passed to any parent class. + """ + self._finalized = False + self._forbid_setting_attributes = forbid_setting_attributes_after_finalize + excluded_attribute_names = set(excluded_attribute_names) + excluded_attribute_names.add("_forbid_setting_attributes") + self._excluded_attribute_names = frozenset(excluded_attribute_names) + super().__init__(**parent_kwargs) + + @property + def finalized(self) -> bool: + """Whether the object has already been finalized.""" + return self._finalized # pytype: disable=attribute-error + + def finalize(self, *args: Any, **kwargs: Any): + """Finalizes the object, after which no attributes can be set.""" + + if self.finalized: + raise ValueError("Object has already been finalized.") + + self._finalize(*args, **kwargs) + self._finalized = True + + def _finalize(self, *args: Any, **kwargs: Any): + """Any logic that a child class needs to do during the finalization.""" + + def unlock_attributes(self): + self._forbid_setting_attributes = False + + def lock_attributes(self): + self._forbid_setting_attributes = True + + def __setattr__(self, name: str, value: Any): + + # We have to use default values here because __setattr__ will be called + # before the attributes are initialized. i.e. the constructor will call this + # function as it initializes _finalized etc. + is_finalized = getattr(self, "_finalized", False) + has_name_attr = hasattr(self, name) + set_attribute_forbidden = getattr(self, "_forbid_setting_attributes", True) + name_is_excluded = name in getattr(self, "_excluded_attribute_names", ()) + + if is_finalized and not has_name_attr: + raise AttributeError( + "Can't create new attributes after finalization. Attempted to create " + f"{name}." + ) + + elif not is_finalized or not set_attribute_forbidden or name_is_excluded: + super().__setattr__(name, value) + + else: + raise AttributeError("Can't set attributes after finalization.") + + +def auto_scope_method(method: Callable[..., T]) -> Callable[..., T]: + """Wraps the method call to have automatically generated Jax name scope.""" + @functools.wraps(method) + def wrapped(instance, *args, **kwargs): + class_name = type(instance).__name__ + method_name = method.__name__ + if method_name.startswith("_"): + method_name = method_name[1:] + with jax.named_scope(f"{class_name}_{method_name}"): + return method(instance, *args, **kwargs) + + return wrapped + + +def auto_scope_function(func: Callable[..., T]) -> Callable[..., T]: + """Wraps the function call to have automatically generated Jax name scope.""" + @functools.wraps(func) + def wrapped(*args, **kwargs): + with jax.named_scope(func.__name__): + return func(*args, **kwargs) + + return wrapped + + +def default_batch_size_extractor(batch: types.Batch) -> int: + """Computes the batch size as the size of axis `0` of the first element.""" + return jax.tree_util.tree_leaves(batch)[0].shape[0] + + +def replace_char(original: str, new_str: str, index: int) -> str: + """Replaces the character at a given location.""" + return original[:index] + new_str + original[index + 1 :] + + +def call_func_with_conditional_kwargs( + func: Callable[..., T], + *func_args: Any, + **kwargs: Any, +) -> T: + + sig = inspect.signature(func) + func_kwargs = {k: v for k, v in kwargs.items() if k in sig.parameters} + + return func(*func_args, **func_kwargs) diff --git a/src/kfac_jax/_src/utils/parallel.py b/src/kfac_jax/_src/utils/parallel.py new file mode 100644 index 0000000000000000000000000000000000000000..625e948f292d4ccbbaf7c7b7d47a66a2a16170b9 --- /dev/null +++ b/src/kfac_jax/_src/utils/parallel.py @@ -0,0 +1,407 @@ +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC utilities for multi-device execution.""" +import functools +from typing import Any, Callable, Sequence + +import jax +from jax import lax +import jax.numpy as jnp +from kfac_jax._src.utils import types + +try: + # JAX v0.10.0 or newer + from jax.extend.core import unsafe_get_axis_names_DO_NOT_USE # pylint: disable=g-import-not-at-top +except ImportError: + # JAX v0.9.2 or older + from jax.core import unsafe_get_axis_names_DO_NOT_USE # pylint: disable=g-import-not-at-top + +jax_version = ( + jax.__version_info__ if hasattr(jax, "__version_info__") + else tuple(map(int, jax.__version__.split(".")))) + + +Array = types.Array +Numeric = types.Numeric +PRNGKey = types.PRNGKey +TArrayTree = types.TArrayTree + + +def _axis_name_tuple(axis_name): + if axis_name is None: + return () + if isinstance(axis_name, tuple): + return axis_name + return (axis_name,) + + +def in_pmap(axis_name: str | tuple[str, ...] | None) -> bool: + """Returns whether we are in a pmap with the given axis name.""" + + if axis_name is None: + return False + + axis_names = unsafe_get_axis_names_DO_NOT_USE() + requested = _axis_name_tuple(axis_name) + + if all(name in axis_names for name in requested): + return True + + if len(axis_names) > 0: + raise ValueError( + f"In pmap with axis names {axis_names}, but wrong axis name " + f"({axis_name}) was provided. This is likely a bug." + ) + + return False + + +def wrap_if_pmap( + p_func: Callable[[TArrayTree, str], TArrayTree], +) -> Callable[[TArrayTree, str | None], TArrayTree]: + """Wraps `p_func` to be executed only when inside a `jax.pmap` context.""" + + @functools.wraps(p_func) + def p_func_if_pmap(obj: TArrayTree, axis_name: str | None) -> TArrayTree: + return p_func(obj, axis_name) if in_pmap(axis_name) else obj + + return p_func_if_pmap + + +# TODO(jamesmartens,botev): We no longer use wrap_if_pmap in the below +# definitions since it doesn't seem to transmit type info properly. Investigate? +def pmean_if_pmap(obj: TArrayTree, axis_name: str | None) -> TArrayTree: + return lax.pmean(obj, axis_name) if in_pmap(axis_name) else obj + + +def psum_if_pmap(obj: TArrayTree, axis_name: str | None) -> TArrayTree: + return lax.psum(obj, axis_name) if in_pmap(axis_name) else obj + + +pmap_mean = jax.pmap(lambda x: lax.pmean(x, "i"), axis_name="i") +pmap_sum = jax.pmap(lambda x: lax.psum(x, "i"), axis_name="i") + + +def is_scalar(x: Any) -> bool: + return isinstance(x, (float, int)) or ( + isinstance(x, jax.Array) and not x.shape + ) + + +def using_legacy_pmap() -> bool: + """Returns whether the legacy pmap is being used.""" + return False + + +def get_device_n_contents(obj: TArrayTree, n: int) -> TArrayTree: + """Gets the contents from pmap output for device n.""" + + def _get_device_n_contents(value: Numeric) -> Numeric: + + if is_scalar(value): + return value + + if using_legacy_pmap(): + return value[n] + + assert isinstance(value, jax.Array) + + if isinstance(value.sharding, jax.sharding.SingleDeviceSharding): + return value[n] + + assert isinstance(value.sharding, jax.NamedSharding) + + shard_data = value.addressable_shards[n].data + if value.sharding.spec[0] is None: + return shard_data + + return shard_data.squeeze(0) + + return jax.tree_util.tree_map(_get_device_n_contents, obj) + + +def get_first(obj: TArrayTree) -> TArrayTree: + return get_device_n_contents(obj, 0) + + +def get_mean(obj: TArrayTree) -> TArrayTree: + """Returns the average of `obj` over different devices.""" + return get_first(pmap_mean(obj)) + + +def get_sum(obj: TArrayTree) -> TArrayTree: + """Returns the sum of `obj` over different devices.""" + return get_first(pmap_sum(obj)) + + +_broadcast_all_local_devices_legacy = jax.pmap(lambda x: x) +_broadcast_all_local_devices_cache: dict[ + str | None, Callable[[TArrayTree], TArrayTree] +] = {} + + +def broadcast_all_local_devices( + obj: TArrayTree, axis_name: str | None = None +) -> TArrayTree: + """Broadcasts `obj` to all local Jax devices. + + Args: + obj: A pytree to broadcast. + axis_name: Optional axis name for the pmap. + + Returns: + The broadcasted pytree. + """ + if types.tree_is_empty(obj): + return obj + + # When no axis_name provided, use legacy pmap. + if axis_name is None: + return _broadcast_all_local_devices_legacy(obj) + + devices = jax.local_devices() + mesh = jax.sharding.Mesh(devices, (axis_name,)) + sharding = jax.NamedSharding(mesh, jax.sharding.PartitionSpec(axis_name)) + + def _broadcast_with_axis(x): + return jax.device_put(x, sharding) + + return jax.tree_util.tree_map(_broadcast_with_axis, obj) + + +pmap_zeros_like = jax.pmap(lambda x: jax.tree_util.tree_map(jnp.zeros_like, x)) +jit_zeros_like = jax.jit(lambda x: jax.tree_util.tree_map(jnp.zeros_like, x)) + + +def replicate_all_local_devices( + obj: TArrayTree, axis_name: str | None = None +) -> TArrayTree: + """Replicates `obj` to all local Jax devices. + + Args: + obj: A pytree to replicate. + axis_name: Optional axis name for sharding. When the result will be passed + to a pmap with a specific axis_name, this should match to avoid mesh + sharding mismatches. + + Returns: + The replicated pytree. + """ + if types.tree_is_empty(obj): + return obj + + devices = jax.local_devices() + + # When no axis_name is provided, use the original device_put_replicated. + if axis_name is None: + return jax.device_put_replicated(obj, devices=devices) + + mesh = jax.sharding.Mesh(devices, (axis_name,)) + sharding = jax.NamedSharding(mesh, jax.P(axis_name)) + + def _replicate_with_axis(x): + # Stack to add the device dimension, then device_put with sharding. + stacked = jnp.stack([x] * len(devices)) + return jax.device_put(stacked, sharding) + + return jax.tree_util.tree_map(_replicate_with_axis, obj) + + +def make_different_rng_key_on_all_devices(rng: PRNGKey) -> PRNGKey: + """Makes a different PRNG for all Jax devices and processes.""" + + rng = jax.random.fold_in(rng, jax.process_index()) + rng = jax.random.split(rng, jax.local_device_count()) + + return broadcast_all_local_devices(rng) + + +p_split = jax.pmap(lambda key: tuple(jax.random.split(key))) + +p_split_num = jax.pmap(lambda key, num: tuple(jax.random.split(key, num)), + static_broadcasted_argnums=1) + + +default_device_sync = None + + +def host_sync( + obj: TArrayTree, + sync_op: Callable[[TArrayTree, str], TArrayTree], +) -> TArrayTree: + """Syncs `obj` across multiple hosts with the operation `sync_op`.""" + + # The implementation here is to use the pmap syncing mechanisms but with only + # the default device of each host. Technically we could do this with all + # the devices on each host, but that would possibly be wasteful. + + if jax.process_count() > 1: + + # We set default_device_sync here because calling jax.local_devices during + # the library import stage will break JAX. + + global default_device_sync + + if default_device_sync is None: + + default_devices = [jax.local_devices(process_index=p_idx)[0] + for p_idx in range(jax.process_count())] + + default_device_sync = jax.pmap(lambda x, sync_op: sync_op(x, "i"), + devices=default_devices, + axis_name="i", + static_broadcasted_argnums=1) + + obj = jax.tree_util.tree_map(lambda x: jnp.expand_dims(x, axis=0), obj) + + return get_first(default_device_sync(obj, sync_op)) + + return obj + + +def host_all_gather(x: TArrayTree) -> TArrayTree: + """Gathers on every host the values of the PyTree leaves `x`.""" + return host_sync(x, lax.all_gather) + + +def host_mean(x: TArrayTree) -> TArrayTree: + """Computes the mean of the PyTree leaves of `x` over multiple hosts.""" + return host_sync(x, lax.pmean) + + +def sync_and_divide_value( + value: TArrayTree, + counter: Numeric, + axis_name: str | None = None, +) -> TArrayTree: + """Computes the mean of `value` over all hosts and divides it by `counter`.""" + value = jax.tree_util.tree_map(lambda x: x / counter, value) + return pmean_if_pmap(value, axis_name) + + +jit_sync_and_divide_value = jax.jit(sync_and_divide_value) +pmap_sync_and_divide_value = jax.pmap( + functools.partial(sync_and_divide_value, axis_name="i"), + axis_name="i", +) + + +# We might be able to change this to "return jnp.array(x)" in newer JAX +# versions. Or maybe we can use jnp.copy now? +def copy_array(x: Array) -> Array: + """Copies a Jax array so that it can be donated freely.""" + return x + jnp.zeros_like(x) + + +copy_obj = jax.jit(lambda x: jax.tree_util.tree_map(copy_array, x)) +_pmap_copy_obj = jax.pmap(copy_obj) + + +def pmap_copy_obj(x: TArrayTree | None) -> TArrayTree | None: + + # pmap will fail to work if passed a totally empty tree + if x is None: + return None + + if types.tree_is_empty(x): + # this does a shallow copy of the tree similar to .copy(): + (flattened, structure) = jax.tree_util.tree_flatten(x) + return jax.tree_util.tree_unflatten(structure, flattened) + + return _pmap_copy_obj(x) + + +def distribute_thunks( + thunks: Sequence[Callable[[], TArrayTree]], + pmap_axis_name: str, + ) -> TArrayTree: + """Distributes the computation of a list of thunks over the pmapped devices. + + Given a list of thunks, this function distributes their computation over the + devices of the current pmap in a round-robin fashion, syncronizes the results + across devices, and then returns them as a sequence of PyTrees. + + Note that this function is meant to be used in a compiled context, and may + call ``thunk[i]()`` several times for each i, with all but one call getting + "optimized away" by XLA. + + Args: + thunks: A sequence of callables performing the desired computations. Each + callable must take zero arguments and return a PyTree of JAX arrays. As + with callables passed to (most) standard JAX API functions, these need to + be stateless and free of side effects. The output of each callable must be + the same regardless of the device it is executed on. + pmap_axis_name: The name of the pmap axis to use. + + Returns: + A sequence of PyTrees that are the output of the corresponding element of + ``thunks``. + """ + + # The strategy here is to make a callable for each device which executes only + # the thunks i such that i % total_devices == device_index, and returns a tree + # of zeros for the remaining thunks. We then do a lax.switch over these based + # on device_index, and return psum over these. Note that the more obvious way + # of doing this, which is to perform a psum over the output of a sequence of + # lax.cond calls (with one for each thunk), won't work in general. This is + # because in order to save memory, XLA will sometimes elect to execute these + # conds sequentially instead of in parallel. + + if not in_pmap(pmap_axis_name): + raise ValueError(f"Provided pmap_axis_name {pmap_axis_name} is not a valid " + "pmap axis in current pmap (or this function was not " + "called in a pmap).") + + assert pmap_axis_name is not None + + axis_names = _axis_name_tuple(pmap_axis_name) + total_devices = lax.psum(1, axis_name=pmap_axis_name) # returns a constant + if len(axis_names) == 1: + current_device_index = lax.axis_index(axis_names[0]) + else: + # Linearise the multi-axis shard_map index so distributed thunk work is + # spread over the full data mesh, not just one named axis. + current_device_index = 0 + stride = 1 + for axis in reversed(axis_names): + current_device_index = current_device_index + lax.axis_index(axis) * stride + stride = stride * lax.psum(1, axis_name=axis) + + # This should get optimized away by XLA since we don't use the values: + dummy_output_trees = tuple(thunk() for thunk in thunks) + + def make_branch(device_index): + + def branch(): + """Execute only thunks i such that i % total_devices == device_index.""" + + outs = [] + for i in range(len(thunks)): + + if i % total_devices == device_index: + outs.append(thunks[i]()) + else: + outs.append( + jax.tree_util.tree_map(jnp.zeros_like, dummy_output_trees[i])) + + return tuple(outs) + + return branch + + branches = tuple(make_branch(device_index) + for device_index in range(total_devices)) + + output_trees = jax.lax.switch(current_device_index, branches) + + return jax.lax.psum(output_trees, axis_name=pmap_axis_name) diff --git a/src/kfac_jax/_src/utils/staging.py b/src/kfac_jax/_src/utils/staging.py new file mode 100644 index 0000000000000000000000000000000000000000..960236a51363492b94d74876aa2bb15d61b121f6 --- /dev/null +++ b/src/kfac_jax/_src/utils/staging.py @@ -0,0 +1,385 @@ +# Modifications copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC JAX staging utilities.""" + +import functools +import inspect +import operator +from typing import Any, Callable, Sequence + +import jax +from jax import lax +import jax.numpy as jnp + +from kfac_jax._src.utils import misc +from kfac_jax._src.utils import parallel +from kfac_jax._src.utils import types + + +TArrayTree = types.TArrayTree + + +class WithStagedMethods(misc.Finalizable): + """An mixin for classes which can have staged/compiled methods.""" + + class StagingContext: + """A context manager for handling methods that are staged/compiled.""" + + def __init__(self, wsm_instance: "WithStagedMethods"): + """Initializes the context manager. + + Args: + wsm_instance: The corresponding `WithStagedMethods` instance. + """ + self._wsm_instance = wsm_instance + + def __enter__(self): + """Enters the staging context.""" + if self._wsm_instance._in_staging: + raise RuntimeError("Cannot enter staging context while already in " + "staging context.") + self._wsm_instance._in_staging = True + + def __exit__(self, *_): + """Exits the staging context.""" + assert self._wsm_instance._in_staging, "Exiting while not in staging." + self._wsm_instance._in_staging = False + + def __init__( + self, + multi_device: bool = False, + pmap_axis_name: str | None = None, + debug: bool = False, + **parent_kwargs: Any, + ): + """Initializes the instance. + + Args: + multi_device: Whether any of decorated staged methods are to be run on a + single or multiple devices. If this is set to `True` than any call + would internally be delegated to `jax.pmap` and otherwise to `jax.jit`. + pmap_axis_name: The name of the pmap axis to use when running on + multiple devices. This is required if `multi_device=True`. + debug: If this is set `True` than any call to a stage method would + directly call the method and would not stage/compile it. + **parent_kwargs: Any additional keyword arguments for the parent class. + """ + if "excluded_attribute_names" in parent_kwargs: + parent_kwargs["excluded_attribute_names"] = ( + ("_in_staging",) + tuple(parent_kwargs["excluded_attribute_names"])) + else: + parent_kwargs["excluded_attribute_names"] = ("_in_staging",) + + super().__init__(**parent_kwargs) + + def _valid_axis_name(name): + return ( + isinstance(name, str) + or ( + isinstance(name, tuple) + and len(name) > 0 + and all(isinstance(axis, str) for axis in name) + ) + ) + + if multi_device and not _valid_axis_name(pmap_axis_name): + raise ValueError( + "When `multi_device=True` you must pass a string axis name or " + "a tuple of string axis names for `pmap_axis_name`." + ) + + self._multi_device = multi_device + self._pmap_axis_name = pmap_axis_name + self._debug = debug + self._in_staging = False + + @property + def multi_device(self) -> bool: + """Indicates whether staged method will be run across multiple devices.""" + return self._multi_device + + @property + def pmap_axis_name(self): + """The name of the `jax.pmap` axis to use for staged methods.""" + if self.debug: + return None + return self._pmap_axis_name + + @property + def debug(self) -> bool: + """Whether staged methods would be run in 'debug' mode.""" + return self._debug + + @property + def in_staging(self) -> bool: + """Whether we are in a staging context while compiling staged methods.""" + return self._in_staging + + def staging_context(self) -> StagingContext: + """Returns a staging context manager, linked to this instance.""" + return self.StagingContext(self) + + def get_first(self, obj: TArrayTree) -> TArrayTree: + """Indexes the `obj` PyTree leaves over leading axis if `multi_device`.""" + return parallel.get_first(obj) if self.multi_device else obj + + def copy_obj(self, obj: TArrayTree | None) -> TArrayTree | None: + """Copies the object.""" + if self.multi_device: + return parallel.pmap_copy_obj(obj) + else: + return parallel.copy_obj(obj) + + def replicate(self, obj: TArrayTree) -> TArrayTree: + """Replicates the object to all local devices if `multi_device`.""" + if self.multi_device: + return parallel.replicate_all_local_devices(obj) + else: + return obj + + def pmean_if_pmap_wrapper( + self, + func: Callable[..., TArrayTree], + ) -> Callable[..., TArrayTree]: + """Wraps a function to perform a pmean if `multi_device`.""" + if self.multi_device and not self.debug: + return lambda *args, **kwargs: lax.pmean( + func(*args, **kwargs), self.pmap_axis_name + ) + else: + return func + + +def staged( + method: Callable[..., TArrayTree], + static_argnums: int | Sequence[int] | None = None, + donate_argnums: int | Sequence[int] | None = None, +) -> Callable[..., TArrayTree]: + """Makes the instance method staged. + + This decorator **should** only be applied to instance methods of classes that + inherit from the `WithStagedMethods` class. The decorator makes the decorated + method staged, which is equivalent to `jax.jit` if `instance.multi_device` is + `False` and to `jax.pmap` otherwise. + + Note that the point of this abstraction around JAX's compilation is to make + sure that jitting/pmapping is only done once, so that if we are already in a + compiled/staged method, we won't initiate a second nested compilation when + calling into second staged method. + + Note that when specifying static and donated argunms, the `self` reference + **must not** be counted. Example: + + @functools.partial(staged, donate_argunms=0) + def try(self, x): + ... + + then `instance.try(x)` is equivalent to + `jax.jit(instance.try, donate_argnums=0)(x)` if `instance.multi_device` is + `False` and to `jax.pmap(instance.try, donate_argnums=0)(x)` otherwise. + + Args: + method: The method to be transformed into a staged method. + static_argnums: The static argument numbers, as defined in `jax.jit/pmap`. + donate_argnums: The donated argument numbers, as defined in + `jax.jit/pmap`. + + Returns: + The transformed method, which will now be a staged function. + """ + + if isinstance(static_argnums, int): + static_argnums = (static_argnums,) + + # This is needed because of b/147015762 + if donate_argnums is None: + donate_argnums = () + if isinstance(donate_argnums, int): + donate_argnums = (donate_argnums,) + else: + donate_argnums: tuple[int, ...] = tuple(donate_argnums) + + original_static_argnums = static_argnums or () + + # shift static_argnums by 1 and include instance (self) + static_argnums = (0,) + tuple(i + 1 for i in (static_argnums or ())) + # shift donate_argnums by 1 and include state + donate_argnums = tuple(i + 1 for i in donate_argnums) + + pmap_funcs = {} + jitted_func = jax.jit(method, + static_argnums=static_argnums, + donate_argnums=donate_argnums) + + @functools.wraps(method) + def decorated( + instance: "WithStagedMethods", + *args: Any, + **kwargs: Any + ) -> TArrayTree: + + sig = inspect.signature(method) + bound_args = sig.bind(instance, *args, **kwargs) + bound_args.apply_defaults() + args, kwargs = bound_args.args[1:], bound_args.kwargs + + if instance.in_staging: + return method(instance, *args, **kwargs) + + with instance.staging_context(): + + if instance.multi_device and instance.debug: + # In this case we want to call `method` once for each device index. + # Note that this might not always produce sensible behavior, and will + # depend on the details of the method and if it has side effects on the + # state of the class. Note that pmean operations won't happen, since the + # actual output of pmapped methods won't be numerically correct. + + bcast_argnums = [ + i for i in range(len(args)) if (i in original_static_argnums + or parallel.is_scalar(args[i]))] + + outs = [] + non_bcast_args = [args[i] if i not in bcast_argnums else None + for i in range(len(args))] + + for i in range(jax.local_device_count()): + + non_bcast_args_i = jax.tree_util.tree_map( + operator.itemgetter(i), non_bcast_args) + + args_i = [ + non_bcast_args_i[j] if j not in bcast_argnums else args[j] + for j in range(len(args)) + ] + + kwargs_i = jax.tree_util.tree_map(operator.itemgetter(i), kwargs) + + with jax.disable_jit(): + outs.append(method(instance, *args_i, **kwargs_i)) + + outs = jax.tree_util.tree_map(lambda *args_: jnp.stack(args_), *outs) + + elif instance.debug: + with jax.disable_jit(): + outs = method(instance, *args, **kwargs) + + elif instance.multi_device: + + # shard_map handles the parallelism (consumes NamedSharded inputs + # natively, with no + # device_put_replicated/all-to-all boilerplate); jit handles + # caching and donation (buffer reuse for opt state — critical to + # avoid per-step state all-to-all). + from jax.sharding import ( + Mesh as _Mesh, NamedSharding as _NS, + PartitionSpec as _P, + SingleDeviceSharding as _SDS, + ) + + static_idx = set(original_static_argnums) + dynamic_idx = [i for i in range(len(args)) if i not in static_idx] + # `donate_argnums` (outer closure) has already been shifted +1 + # for instance/self at line ~206. Un-shift to recover the + # user-supplied indices over `args` (excluding self), then map + # to dynamic-only positions (jit donate_argnums positional). + _orig_donate = tuple(i - 1 for i in donate_argnums if i >= 1) + donate_dynamic = tuple( + dynamic_idx.index(i) for i in _orig_donate if i in dynamic_idx + ) + + axis_names = instance.pmap_axis_name + axis_names = axis_names if isinstance(axis_names, tuple) else (axis_names,) + + def _spec_for(leaf): + if parallel.is_scalar(leaf): + return _P() + sh = getattr(leaf, "sharding", None) + if isinstance(sh, _NS): + return sh.spec + if isinstance(sh, _SDS): + return _P() + return _P() + + def _mesh_from_dynamic_args(dynamic_args_): + for leaf in jax.tree_util.tree_leaves(dynamic_args_): + sh = getattr(leaf, "sharding", None) + if isinstance(sh, _NS): + mesh_ = sh.mesh + if all(axis in mesh_.axis_names for axis in axis_names): + return mesh_ + if len(axis_names) == 1: + return _Mesh(jax.local_devices(), axis_names) + raise ValueError( + "multi-axis KFAC staging needs at least one NamedSharded " + f"dynamic argument on a mesh containing axes {axis_names!r}" + ) + + # `args` order signature is fixed; cache by static-arg values + # plus dynamic-arg shardings (per-leaf shapes/dtypes are + # handled by jit's own cache). + dynamic_args = tuple(args[i] for i in dynamic_idx) + mesh = _mesh_from_dynamic_args(dynamic_args) + in_specs = tuple( + jax.tree_util.tree_map(_spec_for, a) for a in dynamic_args + ) + cache_key = ( + instance.pmap_axis_name, + tuple(mesh.axis_names), + tuple(mesh.shape.items()), + tuple((i, args[i]) for i in sorted(static_idx)), + str(in_specs), + ) + func = pmap_funcs.get(cache_key) + if func is None: + # Closure-capture instance + static args; the inner function + # only takes dynamic args (so jit's static_argnums isn't + # needed — the static values are baked into the closure). + static_kv = {i: args[i] for i in static_idx} + + def _bound_method(*dyn_args, + _instance=instance, + _static_kv=static_kv, + _kwargs=kwargs, + _n=len(args)): + full = [] + di = 0 + for i in range(_n): + if i in static_idx: + full.append(_static_kv[i]) + else: + full.append(dyn_args[di]) + di += 1 + return method(_instance, *full, **_kwargs) + + sharded = jax.shard_map( + _bound_method, mesh=mesh, + in_specs=in_specs, out_specs=_P(), + check_vma=False, + ) + func = jax.jit(sharded, donate_argnums=donate_dynamic) + + pmap_funcs[cache_key] = func + + outs = func(*dynamic_args) + + else: + outs = jitted_func(instance, *args, **kwargs) + + return outs + + return decorated diff --git a/src/kfac_jax/_src/utils/types.py b/src/kfac_jax/_src/utils/types.py new file mode 100644 index 0000000000000000000000000000000000000000..9d42353ccee5a94cf70da05c79c469d2187f312b --- /dev/null +++ b/src/kfac_jax/_src/utils/types.py @@ -0,0 +1,80 @@ +# Modifications copyright (c) 2026 Simulacra Research Inc. +# SPDX-License-Identifier: Apache-2.0 + +# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""K-FAC annotation types and general tree operations.""" +from typing import Callable, Mapping, Sequence, TypeVar + +import jax +import jax.numpy as jnp + +# Types for annotation +T = TypeVar("T") +Array = jax.Array +PRNGKey = Array +Scalar = float | int +Numeric = Array | Scalar +Shape = tuple[int, ...] +DType = jax.typing.DTypeLike +PyTree = T | Sequence["PyTree[T]"] | Mapping[str, "PyTree[T]"] +ArrayTree = PyTree[Array] +TArrayTree = TypeVar("TArrayTree", bound=ArrayTree) +Params = TypeVar("Params", bound=ArrayTree) +Batch = TypeVar("Batch", bound=ArrayTree) +FuncState = TypeVar("FuncState", bound=ArrayTree) +FuncAux = dict[str, ArrayTree] +PyTreeDef = jax.tree_util.PyTreeDef +FuncArgs = Sequence[ArrayTree] +FuncOuts = Array | tuple[Array, FuncAux] +Func = Callable[..., FuncOuts] +ValueFunc = Callable[..., Array] +ValueAndGradFunc = Callable[..., tuple[Array, Params]] +AssumedFuncOutput = (Array | tuple[Array, FuncAux] | + tuple[Array, tuple[FuncState, FuncAux]]) +SCALAR_TYPES = (float, int) +ScheduleType = ( + Callable[[Numeric, Numeric | None], Numeric] | + Callable[[Numeric], Numeric] + ) + + +def tree_is_empty(obj: ArrayTree) -> bool: + """Returns whether the given PyTree is empty.""" + return not jax.tree_util.tree_leaves(obj) + + +def abstract_objects_equal( + obj1: ArrayTree, + obj2: ArrayTree, + check_dtype: bool = True +) -> bool: + """`True` if the objects have the same PyTree structure, shapes and dtypes.""" + return (jax.tree_util.tree_structure(obj1) == + jax.tree_util.tree_structure(obj2) and + all(e1.shape == e2.shape and (e1.dtype == e2.dtype or not check_dtype) + for e1, e2 in zip(jax.tree_util.tree_leaves(obj1), + jax.tree_util.tree_leaves(obj2)))) + + +def get_float_dtype_and_check_consistency(obj: ArrayTree) -> DType | None: + """Checks that all leaves have the same float dtype, and returns this.""" + + leaves = jax.tree_util.tree_leaves(obj) + + for leaf in leaves: + if leaf.dtype != jnp.float32: + raise ValueError("Only float32 values are supported.") + + return jnp.dtype(jnp.float32) if leaves else None diff --git a/src/kfac_jax/py.typed b/src/kfac_jax/py.typed new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/third_party/folx/LICENSE b/third_party/folx/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..22aed37e650bbf933b6983cda9c2c5db65dcdd04 --- /dev/null +++ b/third_party/folx/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) Microsoft Corporation. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/third_party/jax/LICENSE b/third_party/jax/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..d645695673349e3947e8e5ae42332d0ac3164cd7 --- /dev/null +++ b/third_party/jax/LICENSE @@ -0,0 +1,202 @@ + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/third_party/kfac_jax/LICENSE b/third_party/kfac_jax/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..d645695673349e3947e8e5ae42332d0ac3164cd7 --- /dev/null +++ b/third_party/kfac_jax/LICENSE @@ -0,0 +1,202 @@ + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/weights/hamiltonzero_v1.eqx b/weights/hamiltonzero_v1.eqx new file mode 100644 index 0000000000000000000000000000000000000000..dd18576ec4c52bac7b72c756064c3710d49bc54d --- /dev/null +++ b/weights/hamiltonzero_v1.eqx @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2dc415ef30cbe9398b84fc1fb6a9c6bedd80171ed5fb2893fa1019eedaae28ee +size 2190177792