LightPFN 1.0.0: public model card, report, provenance; remove the private 0.1.0 package files
Browse files- LICENSE +202 -202
- LightPFN_report.pdf +2 -2
- NOTICE +9 -4
- README.md +64 -90
- artifacts.json +0 -52
- dependency_licenses.json +0 -471
- dist/lightpfn-0.1.0-py3-none-any.whl +0 -0
- dist/lightpfn-0.1.0.tar.gz +0 -3
- installed_verification.json +0 -17
- package/.gitignore +0 -53
- package/LICENSE +0 -202
- package/NOTICE +0 -6
- package/PKG-INFO +0 -123
- package/docs/DEPENDENCIES.md +0 -28
- package/docs/PACKAGE_README.md +0 -92
- package/lightpfn/__init__.py +0 -25
- package/lightpfn/checkpoint.py +0 -107
- package/lightpfn/device.py +0 -51
- package/lightpfn/model/__init__.py +0 -1
- package/lightpfn/model/layers.py +0 -190
- package/lightpfn/model/lightpfn.py +0 -490
- package/lightpfn/pretrained.json +0 -4
- package/lightpfn/sklearn.py +0 -280
- package/lightpfn/vulkan/__init__.py +0 -40
- package/lightpfn/vulkan/engine.py +0 -210
- package/lightpfn/vulkan/kernels.py +0 -457
- package/lightpfn/vulkan/model.py +0 -548
- package/pyproject.toml +0 -43
- package/tests/test_release.py +0 -94
- package/tests/test_sklearn_api.py +0 -124
- provenance.json +20 -18
LICENSE
CHANGED
|
@@ -1,202 +1,202 @@
|
|
| 1 |
-
|
| 2 |
-
Apache License
|
| 3 |
-
Version 2.0, January 2004
|
| 4 |
-
http://www.apache.org/licenses/
|
| 5 |
-
|
| 6 |
-
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 7 |
-
|
| 8 |
-
1. Definitions.
|
| 9 |
-
|
| 10 |
-
"License" shall mean the terms and conditions for use, reproduction,
|
| 11 |
-
and distribution as defined by Sections 1 through 9 of this document.
|
| 12 |
-
|
| 13 |
-
"Licensor" shall mean the copyright owner or entity authorized by
|
| 14 |
-
the copyright owner that is granting the License.
|
| 15 |
-
|
| 16 |
-
"Legal Entity" shall mean the union of the acting entity and all
|
| 17 |
-
other entities that control, are controlled by, or are under common
|
| 18 |
-
control with that entity. For the purposes of this definition,
|
| 19 |
-
"control" means (i) the power, direct or indirect, to cause the
|
| 20 |
-
direction or management of such entity, whether by contract or
|
| 21 |
-
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 22 |
-
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 23 |
-
|
| 24 |
-
"You" (or "Your") shall mean an individual or Legal Entity
|
| 25 |
-
exercising permissions granted by this License.
|
| 26 |
-
|
| 27 |
-
"Source" form shall mean the preferred form for making modifications,
|
| 28 |
-
including but not limited to software source code, documentation
|
| 29 |
-
source, and configuration files.
|
| 30 |
-
|
| 31 |
-
"Object" form shall mean any form resulting from mechanical
|
| 32 |
-
transformation or translation of a Source form, including but
|
| 33 |
-
not limited to compiled object code, generated documentation,
|
| 34 |
-
and conversions to other media types.
|
| 35 |
-
|
| 36 |
-
"Work" shall mean the work of authorship, whether in Source or
|
| 37 |
-
Object form, made available under the License, as indicated by a
|
| 38 |
-
copyright notice that is included in or attached to the work
|
| 39 |
-
(an example is provided in the Appendix below).
|
| 40 |
-
|
| 41 |
-
"Derivative Works" shall mean any work, whether in Source or Object
|
| 42 |
-
form, that is based on (or derived from) the Work and for which the
|
| 43 |
-
editorial revisions, annotations, elaborations, or other modifications
|
| 44 |
-
represent, as a whole, an original work of authorship. For the purposes
|
| 45 |
-
of this License, Derivative Works shall not include works that remain
|
| 46 |
-
separable from, or merely link (or bind by name) to the interfaces of,
|
| 47 |
-
the Work and Derivative Works thereof.
|
| 48 |
-
|
| 49 |
-
"Contribution" shall mean any work of authorship, including
|
| 50 |
-
the original version of the Work and any modifications or additions
|
| 51 |
-
to that Work or Derivative Works thereof, that is intentionally
|
| 52 |
-
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 53 |
-
or by an individual or Legal Entity authorized to submit on behalf of
|
| 54 |
-
the copyright owner. For the purposes of this definition, "submitted"
|
| 55 |
-
means any form of electronic, verbal, or written communication sent
|
| 56 |
-
to the Licensor or its representatives, including but not limited to
|
| 57 |
-
communication on electronic mailing lists, source code control systems,
|
| 58 |
-
and issue tracking systems that are managed by, or on behalf of, the
|
| 59 |
-
Licensor for the purpose of discussing and improving the Work, but
|
| 60 |
-
excluding communication that is conspicuously marked or otherwise
|
| 61 |
-
designated in writing by the copyright owner as "Not a Contribution."
|
| 62 |
-
|
| 63 |
-
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 64 |
-
on behalf of whom a Contribution has been received by Licensor and
|
| 65 |
-
subsequently incorporated within the Work.
|
| 66 |
-
|
| 67 |
-
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 68 |
-
this License, each Contributor hereby grants to You a perpetual,
|
| 69 |
-
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 70 |
-
copyright license to reproduce, prepare Derivative Works of,
|
| 71 |
-
publicly display, publicly perform, sublicense, and distribute the
|
| 72 |
-
Work and such Derivative Works in Source or Object form.
|
| 73 |
-
|
| 74 |
-
3. Grant of Patent License. Subject to the terms and conditions of
|
| 75 |
-
this License, each Contributor hereby grants to You a perpetual,
|
| 76 |
-
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 77 |
-
(except as stated in this section) patent license to make, have made,
|
| 78 |
-
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 79 |
-
where such license applies only to those patent claims licensable
|
| 80 |
-
by such Contributor that are necessarily infringed by their
|
| 81 |
-
Contribution(s) alone or by combination of their Contribution(s)
|
| 82 |
-
with the Work to which such Contribution(s) was submitted. If You
|
| 83 |
-
institute patent litigation against any entity (including a
|
| 84 |
-
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 85 |
-
or a Contribution incorporated within the Work constitutes direct
|
| 86 |
-
or contributory patent infringement, then any patent licenses
|
| 87 |
-
granted to You under this License for that Work shall terminate
|
| 88 |
-
as of the date such litigation is filed.
|
| 89 |
-
|
| 90 |
-
4. Redistribution. You may reproduce and distribute copies of the
|
| 91 |
-
Work or Derivative Works thereof in any medium, with or without
|
| 92 |
-
modifications, and in Source or Object form, provided that You
|
| 93 |
-
meet the following conditions:
|
| 94 |
-
|
| 95 |
-
(a) You must give any other recipients of the Work or
|
| 96 |
-
Derivative Works a copy of this License; and
|
| 97 |
-
|
| 98 |
-
(b) You must cause any modified files to carry prominent notices
|
| 99 |
-
stating that You changed the files; and
|
| 100 |
-
|
| 101 |
-
(c) You must retain, in the Source form of any Derivative Works
|
| 102 |
-
that You distribute, all copyright, patent, trademark, and
|
| 103 |
-
attribution notices from the Source form of the Work,
|
| 104 |
-
excluding those notices that do not pertain to any part of
|
| 105 |
-
the Derivative Works; and
|
| 106 |
-
|
| 107 |
-
(d) If the Work includes a "NOTICE" text file as part of its
|
| 108 |
-
distribution, then any Derivative Works that You distribute must
|
| 109 |
-
include a readable copy of the attribution notices contained
|
| 110 |
-
within such NOTICE file, excluding those notices that do not
|
| 111 |
-
pertain to any part of the Derivative Works, in at least one
|
| 112 |
-
of the following places: within a NOTICE text file distributed
|
| 113 |
-
as part of the Derivative Works; within the Source form or
|
| 114 |
-
documentation, if provided along with the Derivative Works; or,
|
| 115 |
-
within a display generated by the Derivative Works, if and
|
| 116 |
-
wherever such third-party notices normally appear. The contents
|
| 117 |
-
of the NOTICE file are for informational purposes only and
|
| 118 |
-
do not modify the License. You may add Your own attribution
|
| 119 |
-
notices within Derivative Works that You distribute, alongside
|
| 120 |
-
or as an addendum to the NOTICE text from the Work, provided
|
| 121 |
-
that such additional attribution notices cannot be construed
|
| 122 |
-
as modifying the License.
|
| 123 |
-
|
| 124 |
-
You may add Your own copyright statement to Your modifications and
|
| 125 |
-
may provide additional or different license terms and conditions
|
| 126 |
-
for use, reproduction, or distribution of Your modifications, or
|
| 127 |
-
for any such Derivative Works as a whole, provided Your use,
|
| 128 |
-
reproduction, and distribution of the Work otherwise complies with
|
| 129 |
-
the conditions stated in this License.
|
| 130 |
-
|
| 131 |
-
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 132 |
-
any Contribution intentionally submitted for inclusion in the Work
|
| 133 |
-
by You to the Licensor shall be under the terms and conditions of
|
| 134 |
-
this License, without any additional terms or conditions.
|
| 135 |
-
Notwithstanding the above, nothing herein shall supersede or modify
|
| 136 |
-
the terms of any separate license agreement you may have executed
|
| 137 |
-
with Licensor regarding such Contributions.
|
| 138 |
-
|
| 139 |
-
6. Trademarks. This License does not grant permission to use the trade
|
| 140 |
-
names, trademarks, service marks, or product names of the Licensor,
|
| 141 |
-
except as required for reasonable and customary use in describing the
|
| 142 |
-
origin of the Work and reproducing the content of the NOTICE file.
|
| 143 |
-
|
| 144 |
-
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 145 |
-
agreed to in writing, Licensor provides the Work (and each
|
| 146 |
-
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 147 |
-
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 148 |
-
implied, including, without limitation, any warranties or conditions
|
| 149 |
-
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 150 |
-
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 151 |
-
appropriateness of using or redistributing the Work and assume any
|
| 152 |
-
risks associated with Your exercise of permissions under this License.
|
| 153 |
-
|
| 154 |
-
8. Limitation of Liability. In no event and under no legal theory,
|
| 155 |
-
whether in tort (including negligence), contract, or otherwise,
|
| 156 |
-
unless required by applicable law (such as deliberate and grossly
|
| 157 |
-
negligent acts) or agreed to in writing, shall any Contributor be
|
| 158 |
-
liable to You for damages, including any direct, indirect, special,
|
| 159 |
-
incidental, or consequential damages of any character arising as a
|
| 160 |
-
result of this License or out of the use or inability to use the
|
| 161 |
-
Work (including but not limited to damages for loss of goodwill,
|
| 162 |
-
work stoppage, computer failure or malfunction, or any and all
|
| 163 |
-
other commercial damages or losses), even if such Contributor
|
| 164 |
-
has been advised of the possibility of such damages.
|
| 165 |
-
|
| 166 |
-
9. Accepting Warranty or Additional Liability. While redistributing
|
| 167 |
-
the Work or Derivative Works thereof, You may choose to offer,
|
| 168 |
-
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 169 |
-
or other liability obligations and/or rights consistent with this
|
| 170 |
-
License. However, in accepting such obligations, You may act only
|
| 171 |
-
on Your own behalf and on Your sole responsibility, not on behalf
|
| 172 |
-
of any other Contributor, and only if You agree to indemnify,
|
| 173 |
-
defend, and hold each Contributor harmless for any liability
|
| 174 |
-
incurred by, or claims asserted against, such Contributor by reason
|
| 175 |
-
of your accepting any such warranty or additional liability.
|
| 176 |
-
|
| 177 |
-
END OF TERMS AND CONDITIONS
|
| 178 |
-
|
| 179 |
-
APPENDIX: How to apply the Apache License to your work.
|
| 180 |
-
|
| 181 |
-
To apply the Apache License to your work, attach the following
|
| 182 |
-
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 183 |
-
replaced with your own identifying information. (Don't include
|
| 184 |
-
the brackets!) The text should be enclosed in the appropriate
|
| 185 |
-
comment syntax for the file format. We also recommend that a
|
| 186 |
-
file or class name and description of purpose be included on the
|
| 187 |
-
same "printed page" as the copyright notice for easier
|
| 188 |
-
identification within third-party archives.
|
| 189 |
-
|
| 190 |
-
Copyright [yyyy] [name of copyright owner]
|
| 191 |
-
|
| 192 |
-
Licensed under the Apache License, Version 2.0 (the "License");
|
| 193 |
-
you may not use this file except in compliance with the License.
|
| 194 |
-
You may obtain a copy of the License at
|
| 195 |
-
|
| 196 |
-
http://www.apache.org/licenses/LICENSE-2.0
|
| 197 |
-
|
| 198 |
-
Unless required by applicable law or agreed to in writing, software
|
| 199 |
-
distributed under the License is distributed on an "AS IS" BASIS,
|
| 200 |
-
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 201 |
-
See the License for the specific language governing permissions and
|
| 202 |
-
limitations under the License.
|
|
|
|
| 1 |
+
|
| 2 |
+
Apache License
|
| 3 |
+
Version 2.0, January 2004
|
| 4 |
+
http://www.apache.org/licenses/
|
| 5 |
+
|
| 6 |
+
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 7 |
+
|
| 8 |
+
1. Definitions.
|
| 9 |
+
|
| 10 |
+
"License" shall mean the terms and conditions for use, reproduction,
|
| 11 |
+
and distribution as defined by Sections 1 through 9 of this document.
|
| 12 |
+
|
| 13 |
+
"Licensor" shall mean the copyright owner or entity authorized by
|
| 14 |
+
the copyright owner that is granting the License.
|
| 15 |
+
|
| 16 |
+
"Legal Entity" shall mean the union of the acting entity and all
|
| 17 |
+
other entities that control, are controlled by, or are under common
|
| 18 |
+
control with that entity. For the purposes of this definition,
|
| 19 |
+
"control" means (i) the power, direct or indirect, to cause the
|
| 20 |
+
direction or management of such entity, whether by contract or
|
| 21 |
+
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 22 |
+
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 23 |
+
|
| 24 |
+
"You" (or "Your") shall mean an individual or Legal Entity
|
| 25 |
+
exercising permissions granted by this License.
|
| 26 |
+
|
| 27 |
+
"Source" form shall mean the preferred form for making modifications,
|
| 28 |
+
including but not limited to software source code, documentation
|
| 29 |
+
source, and configuration files.
|
| 30 |
+
|
| 31 |
+
"Object" form shall mean any form resulting from mechanical
|
| 32 |
+
transformation or translation of a Source form, including but
|
| 33 |
+
not limited to compiled object code, generated documentation,
|
| 34 |
+
and conversions to other media types.
|
| 35 |
+
|
| 36 |
+
"Work" shall mean the work of authorship, whether in Source or
|
| 37 |
+
Object form, made available under the License, as indicated by a
|
| 38 |
+
copyright notice that is included in or attached to the work
|
| 39 |
+
(an example is provided in the Appendix below).
|
| 40 |
+
|
| 41 |
+
"Derivative Works" shall mean any work, whether in Source or Object
|
| 42 |
+
form, that is based on (or derived from) the Work and for which the
|
| 43 |
+
editorial revisions, annotations, elaborations, or other modifications
|
| 44 |
+
represent, as a whole, an original work of authorship. For the purposes
|
| 45 |
+
of this License, Derivative Works shall not include works that remain
|
| 46 |
+
separable from, or merely link (or bind by name) to the interfaces of,
|
| 47 |
+
the Work and Derivative Works thereof.
|
| 48 |
+
|
| 49 |
+
"Contribution" shall mean any work of authorship, including
|
| 50 |
+
the original version of the Work and any modifications or additions
|
| 51 |
+
to that Work or Derivative Works thereof, that is intentionally
|
| 52 |
+
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 53 |
+
or by an individual or Legal Entity authorized to submit on behalf of
|
| 54 |
+
the copyright owner. For the purposes of this definition, "submitted"
|
| 55 |
+
means any form of electronic, verbal, or written communication sent
|
| 56 |
+
to the Licensor or its representatives, including but not limited to
|
| 57 |
+
communication on electronic mailing lists, source code control systems,
|
| 58 |
+
and issue tracking systems that are managed by, or on behalf of, the
|
| 59 |
+
Licensor for the purpose of discussing and improving the Work, but
|
| 60 |
+
excluding communication that is conspicuously marked or otherwise
|
| 61 |
+
designated in writing by the copyright owner as "Not a Contribution."
|
| 62 |
+
|
| 63 |
+
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 64 |
+
on behalf of whom a Contribution has been received by Licensor and
|
| 65 |
+
subsequently incorporated within the Work.
|
| 66 |
+
|
| 67 |
+
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 68 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 69 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 70 |
+
copyright license to reproduce, prepare Derivative Works of,
|
| 71 |
+
publicly display, publicly perform, sublicense, and distribute the
|
| 72 |
+
Work and such Derivative Works in Source or Object form.
|
| 73 |
+
|
| 74 |
+
3. Grant of Patent License. Subject to the terms and conditions of
|
| 75 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 76 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 77 |
+
(except as stated in this section) patent license to make, have made,
|
| 78 |
+
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 79 |
+
where such license applies only to those patent claims licensable
|
| 80 |
+
by such Contributor that are necessarily infringed by their
|
| 81 |
+
Contribution(s) alone or by combination of their Contribution(s)
|
| 82 |
+
with the Work to which such Contribution(s) was submitted. If You
|
| 83 |
+
institute patent litigation against any entity (including a
|
| 84 |
+
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 85 |
+
or a Contribution incorporated within the Work constitutes direct
|
| 86 |
+
or contributory patent infringement, then any patent licenses
|
| 87 |
+
granted to You under this License for that Work shall terminate
|
| 88 |
+
as of the date such litigation is filed.
|
| 89 |
+
|
| 90 |
+
4. Redistribution. You may reproduce and distribute copies of the
|
| 91 |
+
Work or Derivative Works thereof in any medium, with or without
|
| 92 |
+
modifications, and in Source or Object form, provided that You
|
| 93 |
+
meet the following conditions:
|
| 94 |
+
|
| 95 |
+
(a) You must give any other recipients of the Work or
|
| 96 |
+
Derivative Works a copy of this License; and
|
| 97 |
+
|
| 98 |
+
(b) You must cause any modified files to carry prominent notices
|
| 99 |
+
stating that You changed the files; and
|
| 100 |
+
|
| 101 |
+
(c) You must retain, in the Source form of any Derivative Works
|
| 102 |
+
that You distribute, all copyright, patent, trademark, and
|
| 103 |
+
attribution notices from the Source form of the Work,
|
| 104 |
+
excluding those notices that do not pertain to any part of
|
| 105 |
+
the Derivative Works; and
|
| 106 |
+
|
| 107 |
+
(d) If the Work includes a "NOTICE" text file as part of its
|
| 108 |
+
distribution, then any Derivative Works that You distribute must
|
| 109 |
+
include a readable copy of the attribution notices contained
|
| 110 |
+
within such NOTICE file, excluding those notices that do not
|
| 111 |
+
pertain to any part of the Derivative Works, in at least one
|
| 112 |
+
of the following places: within a NOTICE text file distributed
|
| 113 |
+
as part of the Derivative Works; within the Source form or
|
| 114 |
+
documentation, if provided along with the Derivative Works; or,
|
| 115 |
+
within a display generated by the Derivative Works, if and
|
| 116 |
+
wherever such third-party notices normally appear. The contents
|
| 117 |
+
of the NOTICE file are for informational purposes only and
|
| 118 |
+
do not modify the License. You may add Your own attribution
|
| 119 |
+
notices within Derivative Works that You distribute, alongside
|
| 120 |
+
or as an addendum to the NOTICE text from the Work, provided
|
| 121 |
+
that such additional attribution notices cannot be construed
|
| 122 |
+
as modifying the License.
|
| 123 |
+
|
| 124 |
+
You may add Your own copyright statement to Your modifications and
|
| 125 |
+
may provide additional or different license terms and conditions
|
| 126 |
+
for use, reproduction, or distribution of Your modifications, or
|
| 127 |
+
for any such Derivative Works as a whole, provided Your use,
|
| 128 |
+
reproduction, and distribution of the Work otherwise complies with
|
| 129 |
+
the conditions stated in this License.
|
| 130 |
+
|
| 131 |
+
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 132 |
+
any Contribution intentionally submitted for inclusion in the Work
|
| 133 |
+
by You to the Licensor shall be under the terms and conditions of
|
| 134 |
+
this License, without any additional terms or conditions.
|
| 135 |
+
Notwithstanding the above, nothing herein shall supersede or modify
|
| 136 |
+
the terms of any separate license agreement you may have executed
|
| 137 |
+
with Licensor regarding such Contributions.
|
| 138 |
+
|
| 139 |
+
6. Trademarks. This License does not grant permission to use the trade
|
| 140 |
+
names, trademarks, service marks, or product names of the Licensor,
|
| 141 |
+
except as required for reasonable and customary use in describing the
|
| 142 |
+
origin of the Work and reproducing the content of the NOTICE file.
|
| 143 |
+
|
| 144 |
+
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 145 |
+
agreed to in writing, Licensor provides the Work (and each
|
| 146 |
+
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 147 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 148 |
+
implied, including, without limitation, any warranties or conditions
|
| 149 |
+
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 150 |
+
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 151 |
+
appropriateness of using or redistributing the Work and assume any
|
| 152 |
+
risks associated with Your exercise of permissions under this License.
|
| 153 |
+
|
| 154 |
+
8. Limitation of Liability. In no event and under no legal theory,
|
| 155 |
+
whether in tort (including negligence), contract, or otherwise,
|
| 156 |
+
unless required by applicable law (such as deliberate and grossly
|
| 157 |
+
negligent acts) or agreed to in writing, shall any Contributor be
|
| 158 |
+
liable to You for damages, including any direct, indirect, special,
|
| 159 |
+
incidental, or consequential damages of any character arising as a
|
| 160 |
+
result of this License or out of the use or inability to use the
|
| 161 |
+
Work (including but not limited to damages for loss of goodwill,
|
| 162 |
+
work stoppage, computer failure or malfunction, or any and all
|
| 163 |
+
other commercial damages or losses), even if such Contributor
|
| 164 |
+
has been advised of the possibility of such damages.
|
| 165 |
+
|
| 166 |
+
9. Accepting Warranty or Additional Liability. While redistributing
|
| 167 |
+
the Work or Derivative Works thereof, You may choose to offer,
|
| 168 |
+
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 169 |
+
or other liability obligations and/or rights consistent with this
|
| 170 |
+
License. However, in accepting such obligations, You may act only
|
| 171 |
+
on Your own behalf and on Your sole responsibility, not on behalf
|
| 172 |
+
of any other Contributor, and only if You agree to indemnify,
|
| 173 |
+
defend, and hold each Contributor harmless for any liability
|
| 174 |
+
incurred by, or claims asserted against, such Contributor by reason
|
| 175 |
+
of your accepting any such warranty or additional liability.
|
| 176 |
+
|
| 177 |
+
END OF TERMS AND CONDITIONS
|
| 178 |
+
|
| 179 |
+
APPENDIX: How to apply the Apache License to your work.
|
| 180 |
+
|
| 181 |
+
To apply the Apache License to your work, attach the following
|
| 182 |
+
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 183 |
+
replaced with your own identifying information. (Don't include
|
| 184 |
+
the brackets!) The text should be enclosed in the appropriate
|
| 185 |
+
comment syntax for the file format. We also recommend that a
|
| 186 |
+
file or class name and description of purpose be included on the
|
| 187 |
+
same "printed page" as the copyright notice for easier
|
| 188 |
+
identification within third-party archives.
|
| 189 |
+
|
| 190 |
+
Copyright [yyyy] [name of copyright owner]
|
| 191 |
+
|
| 192 |
+
Licensed under the Apache License, Version 2.0 (the "License");
|
| 193 |
+
you may not use this file except in compliance with the License.
|
| 194 |
+
You may obtain a copy of the License at
|
| 195 |
+
|
| 196 |
+
http://www.apache.org/licenses/LICENSE-2.0
|
| 197 |
+
|
| 198 |
+
Unless required by applicable law or agreed to in writing, software
|
| 199 |
+
distributed under the License is distributed on an "AS IS" BASIS,
|
| 200 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 201 |
+
See the License for the specific language governing permissions and
|
| 202 |
+
limitations under the License.
|
LightPFN_report.pdf
CHANGED
|
@@ -1,3 +1,3 @@
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:
|
| 3 |
-
size
|
|
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:274be6425d074c984ce49cee56f40a0801421017a2851febe916ac7bbc4718f4
|
| 3 |
+
size 331223
|
NOTICE
CHANGED
|
@@ -1,6 +1,11 @@
|
|
| 1 |
LightPFN
|
| 2 |
-
Copyright 2026 Giorgio
|
| 3 |
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
LightPFN
|
| 2 |
+
Copyright 2026 Giorgio Ottoboni
|
| 3 |
|
| 4 |
+
This product includes software and model weights licensed under the Apache License, Version 2.0.
|
| 5 |
+
|
| 6 |
+
The model was trained from scratch on synthetic data only, without distillation and without weights or
|
| 7 |
+
outputs of other tabular foundation models.
|
| 8 |
+
|
| 9 |
+
The optional training code (lightpfn.prior) uses the structural causal model prior of TabICL
|
| 10 |
+
(https://github.com/soda-inria/tabicl, BSD-3-Clause, Copyright (c) 2025, Soda team @ Inria) as an installed
|
| 11 |
+
dependency; TabICL code is not redistributed in this repository or in the published wheel.
|
README.md
CHANGED
|
@@ -4,122 +4,96 @@ library_name: lightpfn
|
|
| 4 |
pipeline_tag: tabular-classification
|
| 5 |
tags:
|
| 6 |
- tabular
|
| 7 |
-
- classification
|
| 8 |
- in-context-learning
|
|
|
|
|
|
|
| 9 |
- synthetic-data
|
| 10 |
-
-
|
|
|
|
| 11 |
- safetensors
|
| 12 |
---
|
| 13 |
|
| 14 |
-
# LightPFN
|
| 15 |
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
|
|
|
| 24 |
|
| 25 |
-
|
| 26 |
-
models were used in training. External tabular models were evaluation baselines
|
| 27 |
-
only. The release excludes their code/weights, synthetic priors and training
|
| 28 |
-
dependencies.
|
| 29 |
-
|
| 30 |
-
## Installation
|
| 31 |
-
|
| 32 |
-
Authenticate with an account authorized to read this private repository:
|
| 33 |
|
| 34 |
```bash
|
| 35 |
-
|
| 36 |
-
hf download ueuegio/LightPFN dist/lightpfn-0.1.0-py3-none-any.whl --local-dir .
|
| 37 |
-
pip install "./dist/lightpfn-0.1.0-py3-none-any.whl[sklearn,hf]"
|
| 38 |
```
|
| 39 |
|
| 40 |
-
Install `huggingface_hub` first if the `hf` command is unavailable. For CPU-only
|
| 41 |
-
use, install torch from PyTorch's official CPU wheel index before the package.
|
| 42 |
-
An optional Vulkan extra provides AMD/Intel/NVIDIA GPU inference without CUDA:
|
| 43 |
-
`pip install "./dist/lightpfn-0.1.0-py3-none-any.whl[sklearn,hf,vulkan]"`.
|
| 44 |
-
|
| 45 |
```python
|
| 46 |
from lightpfn import LightPFNClassifier
|
| 47 |
|
| 48 |
-
clf = LightPFNClassifier(
|
| 49 |
-
clf.fit(X_train, y_train)
|
| 50 |
proba = clf.predict_proba(X_test)
|
| 51 |
-
pred = clf.predict(X_test)
|
| 52 |
```
|
| 53 |
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
Inputs are dense numeric tables; missing values can be NaN. Encode string
|
| 58 |
-
categories before fitting. Native categorical handling is not implemented.
|
| 59 |
-
The deprecated `cat` mask does not change the computation and warns when used
|
| 60 |
-
to flag categories.
|
| 61 |
-
|
| 62 |
-
The base package needs only torch and numpy; sklearn, Hub/safetensors and Vulkan
|
| 63 |
-
are optional extras. The inference package imports no prior or external foundation
|
| 64 |
-
model. The root `model.safetensors` stores tensors and `config.json` stores the
|
| 65 |
-
architecture. Legacy tensor/dict `.pt` files load only with `weights_only=True`.
|
| 66 |
-
No unrestricted pickle fallback or remote-code loading is provided.
|
| 67 |
-
|
| 68 |
-
## Training
|
| 69 |
-
|
| 70 |
-
Architecture: B4, summary-token rows, three row-refinement rounds, seven ICL
|
| 71 |
-
blocks, retrieval decoder. The synthetic mix is graph v3 90% / rule v1 10%,
|
| 72 |
-
without augmentation. Model size was fixed below 5M throughout development.
|
| 73 |
|
| 74 |
-
|
| 75 |
-
RTX 5090 GPUs; tables of at most 2,048 rows.
|
| 76 |
-
- Long-context stage: 39,250 steps, LR 1e-4 to 1e-5, tables up to 60,000 rows.
|
| 77 |
-
- Released weights: EMA `final_long/ema_step039250.pt`, converted losslessly to
|
| 78 |
-
FP32 safetensors. Training optimizer/state is excluded.
|
| 79 |
|
| 80 |
-
|
| 81 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 82 |
|
| 83 |
## Evaluation
|
| 84 |
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
the finalists. These are local experiments, **not an official TabArena submission**.
|
| 88 |
-
GBDT comparisons use library defaults, not tuned TabArena configurations.
|
| 89 |
|
| 90 |
-
|
|
| 91 |
|---|---|
|
| 92 |
-
|
|
| 93 |
-
|
|
| 94 |
-
| TabArena
|
| 95 |
-
|
|
| 96 |
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
TabArena-lite and full results should not be conflated. The full report and
|
| 100 |
-
evaluation outputs are in the source repository.
|
| 101 |
|
| 102 |
## Limitations
|
| 103 |
|
| 104 |
-
Classification
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
##
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
pipeline_tag: tabular-classification
|
| 5 |
tags:
|
| 6 |
- tabular
|
| 7 |
+
- tabular-classification
|
| 8 |
- in-context-learning
|
| 9 |
+
- prior-data-fitted-network
|
| 10 |
+
- foundation-model
|
| 11 |
- synthetic-data
|
| 12 |
+
- scikit-learn
|
| 13 |
+
- vulkan
|
| 14 |
- safetensors
|
| 15 |
---
|
| 16 |
|
| 17 |
+
# LightPFN
|
| 18 |
|
| 19 |
+
LightPFN is a small tabular foundation model for classification: a 4,603,088-parameter in-context learner
|
| 20 |
+
pretrained only on synthetic data. `fit` stores the training set as context and `predict_proba` answers in one
|
| 21 |
+
forward pass, with no training on your data and no hyperparameters to tune. It runs on CPU, CUDA, ROCm and, through
|
| 22 |
+
its own Vulkan kernels, on AMD, Intel and NVIDIA GPUs.
|
| 23 |
|
| 24 |
+
- Code, documentation and training pipeline: [github.com/GioOtto/LightPFN](https://github.com/GioOtto/LightPFN)
|
| 25 |
+
- Package: [pypi.org/project/LightPFN](https://pypi.org/project/LightPFN/)
|
| 26 |
+
- Technical report: [LightPFN_report.pdf](./LightPFN_report.pdf)
|
| 27 |
+
- License: Apache 2.0, code and weights
|
| 28 |
|
| 29 |
+
## Usage
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
|
| 31 |
```bash
|
| 32 |
+
pip install LightPFN
|
|
|
|
|
|
|
| 33 |
```
|
| 34 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 35 |
```python
|
| 36 |
from lightpfn import LightPFNClassifier
|
| 37 |
|
| 38 |
+
clf = LightPFNClassifier(n_estimators=4, random_state=0)
|
| 39 |
+
clf.fit(X_train, y_train) # downloads these weights at a pinned commit on the first fit
|
| 40 |
proba = clf.predict_proba(X_test)
|
|
|
|
| 41 |
```
|
| 42 |
|
| 43 |
+
`X` can be a NumPy array or a pandas DataFrame with numeric, categorical, string and boolean columns and missing
|
| 44 |
+
values. `device="auto"` uses a CUDA or ROCm GPU, else a Vulkan GPU (`pip install "LightPFN[vulkan]"`), else the CPU.
|
| 45 |
+
The [user guide](https://github.com/GioOtto/LightPFN/blob/main/docs/en/GUIDE.md) covers every option.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 46 |
|
| 47 |
+
## Model
|
|
|
|
|
|
|
|
|
|
|
|
|
| 48 |
|
| 49 |
+
| | |
|
| 50 |
+
|---|---|
|
| 51 |
+
| Task | classification, 2 to 10 classes |
|
| 52 |
+
| Parameters | 4,603,088 (18 MB, float32) |
|
| 53 |
+
| Architecture | cell embedding (value, rank, missing flag), two induced column stages, row refinement with four summary tokens, row compression, seven in-context blocks, retrieval decoder |
|
| 54 |
+
| Pretraining data | synthetic only: 90% structural causal graph prior, 10% rule prior (XOR, parity, lookup, trees); 7.68 million task draws from 4.03 million distinct tasks |
|
| 55 |
+
| Training | 120,000 steps of 64 tasks on two RTX 5090 (tables up to 2,048 rows), then 39,250 steps on tables up to 60,000 rows |
|
| 56 |
+
| Not used in training | real datasets, distillation, weights or outputs of other tabular foundation models |
|
| 57 |
|
| 58 |
## Evaluation
|
| 59 |
|
| 60 |
+
The official TabArena-Lite evaluation includes default, tuned and ensembled baselines. Other evaluations
|
| 61 |
+
use default baselines; AUC differences have 95% paired bootstrap intervals.
|
|
|
|
|
|
|
| 62 |
|
| 63 |
+
| Benchmark | Result |
|
| 64 |
|---|---|
|
| 65 |
+
| Official TabArena-Lite pipeline, 38 classification datasets, default configuration, four estimators | Elo 1420 (+67 / -66), 25th of 99 methods, 38 successful tasks, none imputed; above GBDT point estimates, intervals overlap tuned CatBoost; author-run, pending maintainer full-benchmark verification |
|
| 66 |
+
| 55 OpenML-CC18 datasets outside TabArena (at most 1,000 rows, one estimator) | mean AUC 0.911, +0.86 [0.42, 1.40] over CatBoost, +1.8 to +2.1 over LightGBM, XGBoost and random forest |
|
| 67 |
+
| TabArena, 38 classification tasks, official splits, first repeat (our harness), four estimators | mean AUC 0.858, rank 2.50 of 7, behind TabICLv2 (0.864), lower error than CatBoost on 76% of tasks |
|
| 68 |
+
| 13 OpenML datasets of 50,000 to 2.2M rows, 10,000 to 100,000 training rows | mean AUC difference from CatBoost between -0.34 and +0.28 points, intervals include zero |
|
| 69 |
|
| 70 |
+
Full tables, timings and the official TabArena results:
|
| 71 |
+
[docs/en/RESULTS.md](https://github.com/GioOtto/LightPFN/blob/main/docs/en/RESULTS.md).
|
|
|
|
|
|
|
| 72 |
|
| 73 |
## Limitations
|
| 74 |
|
| 75 |
+
- Classification only (2 to 10 classes); regression is planned for version 2.
|
| 76 |
+
- Categorical columns are read as ordinal codes. On tables dominated by high-cardinality categorical columns
|
| 77 |
+
CatBoost is ahead; native categorical handling is planned for version 2.
|
| 78 |
+
- Above 20,000 training rows each estimator reads a stratified subsample (`max_context`).
|
| 79 |
+
- CPU time grows with the context length; on large tables a GPU is much faster.
|
| 80 |
+
|
| 81 |
+
## Files
|
| 82 |
+
|
| 83 |
+
| File | Content |
|
| 84 |
+
|---|---|
|
| 85 |
+
| `model.safetensors`, `config.json` | weights (float32) and architecture of the released model |
|
| 86 |
+
| `provenance.json` | hashes of the weights and of the source checkpoint |
|
| 87 |
+
| `LightPFN_report.pdf` | technical report |
|
| 88 |
+
| `LICENSE`, `NOTICE` | Apache License 2.0 and attribution notice |
|
| 89 |
+
|
| 90 |
+
## Citation
|
| 91 |
+
|
| 92 |
+
```bibtex
|
| 93 |
+
@techreport{ottoboni2026lightpfn,
|
| 94 |
+
title = {A Sling Against Giants: {LightPFN}, a 4.6M-parameter tabular in-context classifier designed to stay small},
|
| 95 |
+
author = {Ottoboni, Giorgio},
|
| 96 |
+
year = {2026},
|
| 97 |
+
url = {https://github.com/GioOtto/LightPFN}
|
| 98 |
+
}
|
| 99 |
+
```
|
artifacts.json
DELETED
|
@@ -1,52 +0,0 @@
|
|
| 1 |
-
{
|
| 2 |
-
"lightpfn-0.1.0-py3-none-any.whl": {
|
| 3 |
-
"sha256": "c97936566a9d82a012d865d40111263f4fea9f21eb8a44840ac4e040b5f55604",
|
| 4 |
-
"bytes": 46760,
|
| 5 |
-
"files": [
|
| 6 |
-
"lightpfn/__init__.py",
|
| 7 |
-
"lightpfn/checkpoint.py",
|
| 8 |
-
"lightpfn/device.py",
|
| 9 |
-
"lightpfn/pretrained.json",
|
| 10 |
-
"lightpfn/sklearn.py",
|
| 11 |
-
"lightpfn/model/__init__.py",
|
| 12 |
-
"lightpfn/model/layers.py",
|
| 13 |
-
"lightpfn/model/lightpfn.py",
|
| 14 |
-
"lightpfn/vulkan/__init__.py",
|
| 15 |
-
"lightpfn/vulkan/engine.py",
|
| 16 |
-
"lightpfn/vulkan/kernels.py",
|
| 17 |
-
"lightpfn/vulkan/model.py",
|
| 18 |
-
"lightpfn-0.1.0.dist-info/METADATA",
|
| 19 |
-
"lightpfn-0.1.0.dist-info/WHEEL",
|
| 20 |
-
"lightpfn-0.1.0.dist-info/licenses/LICENSE",
|
| 21 |
-
"lightpfn-0.1.0.dist-info/licenses/NOTICE",
|
| 22 |
-
"lightpfn-0.1.0.dist-info/RECORD"
|
| 23 |
-
]
|
| 24 |
-
},
|
| 25 |
-
"lightpfn-0.1.0.tar.gz": {
|
| 26 |
-
"sha256": "3a9606af85ad6d178963717cb7ad0dec432ff44ea7bea000396bb34c0ff2ab5d",
|
| 27 |
-
"bytes": 45222,
|
| 28 |
-
"files": [
|
| 29 |
-
"lightpfn-0.1.0/docs/DEPENDENCIES.md",
|
| 30 |
-
"lightpfn-0.1.0/lightpfn/__init__.py",
|
| 31 |
-
"lightpfn-0.1.0/lightpfn/checkpoint.py",
|
| 32 |
-
"lightpfn-0.1.0/lightpfn/device.py",
|
| 33 |
-
"lightpfn-0.1.0/lightpfn/pretrained.json",
|
| 34 |
-
"lightpfn-0.1.0/lightpfn/sklearn.py",
|
| 35 |
-
"lightpfn-0.1.0/lightpfn/model/__init__.py",
|
| 36 |
-
"lightpfn-0.1.0/lightpfn/model/layers.py",
|
| 37 |
-
"lightpfn-0.1.0/lightpfn/model/lightpfn.py",
|
| 38 |
-
"lightpfn-0.1.0/lightpfn/vulkan/__init__.py",
|
| 39 |
-
"lightpfn-0.1.0/lightpfn/vulkan/engine.py",
|
| 40 |
-
"lightpfn-0.1.0/lightpfn/vulkan/kernels.py",
|
| 41 |
-
"lightpfn-0.1.0/lightpfn/vulkan/model.py",
|
| 42 |
-
"lightpfn-0.1.0/tests/test_release.py",
|
| 43 |
-
"lightpfn-0.1.0/tests/test_sklearn_api.py",
|
| 44 |
-
"lightpfn-0.1.0/.gitignore",
|
| 45 |
-
"lightpfn-0.1.0/LICENSE",
|
| 46 |
-
"lightpfn-0.1.0/NOTICE",
|
| 47 |
-
"lightpfn-0.1.0/pyproject.toml",
|
| 48 |
-
"lightpfn-0.1.0/docs/PACKAGE_README.md",
|
| 49 |
-
"lightpfn-0.1.0/PKG-INFO"
|
| 50 |
-
]
|
| 51 |
-
}
|
| 52 |
-
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
dependency_licenses.json
DELETED
|
@@ -1,471 +0,0 @@
|
|
| 1 |
-
{
|
| 2 |
-
"anyio": {
|
| 3 |
-
"version": "4.15.1",
|
| 4 |
-
"license_expression": "MIT",
|
| 5 |
-
"license_summary": "The MIT License (MIT)",
|
| 6 |
-
"license_classifiers": [],
|
| 7 |
-
"license_files": [
|
| 8 |
-
"anyio-4.15.1.dist-info/licenses/LICENSE"
|
| 9 |
-
]
|
| 10 |
-
},
|
| 11 |
-
"cffi": {
|
| 12 |
-
"version": "2.1.1",
|
| 13 |
-
"license_expression": "MIT-0",
|
| 14 |
-
"license_summary": "Except when otherwise stated (look for LICENSE files in directories or",
|
| 15 |
-
"license_classifiers": [],
|
| 16 |
-
"license_files": [
|
| 17 |
-
"cffi-2.1.1.dist-info/licenses/LICENSE"
|
| 18 |
-
]
|
| 19 |
-
},
|
| 20 |
-
"click": {
|
| 21 |
-
"version": "8.5.0",
|
| 22 |
-
"license_expression": "BSD-3-Clause",
|
| 23 |
-
"license_summary": "Copyright 2014 Pallets",
|
| 24 |
-
"license_classifiers": [],
|
| 25 |
-
"license_files": [
|
| 26 |
-
"click-8.5.0.dist-info/licenses/LICENSE.txt"
|
| 27 |
-
]
|
| 28 |
-
},
|
| 29 |
-
"cloudpickle": {
|
| 30 |
-
"version": "3.1.2",
|
| 31 |
-
"license_expression": null,
|
| 32 |
-
"license_summary": "BSD-3-Clause",
|
| 33 |
-
"license_classifiers": [
|
| 34 |
-
"License :: OSI Approved :: BSD License"
|
| 35 |
-
],
|
| 36 |
-
"license_files": [
|
| 37 |
-
"cloudpickle-3.1.2.dist-info/licenses/LICENSE"
|
| 38 |
-
]
|
| 39 |
-
},
|
| 40 |
-
"colorama": {
|
| 41 |
-
"version": "0.4.6",
|
| 42 |
-
"license_expression": null,
|
| 43 |
-
"license_summary": "Copyright (c) 2010 Jonathan Hartley",
|
| 44 |
-
"license_classifiers": [
|
| 45 |
-
"License :: OSI Approved :: BSD License"
|
| 46 |
-
],
|
| 47 |
-
"license_files": [
|
| 48 |
-
"colorama-0.4.6.dist-info/licenses/LICENSE.txt"
|
| 49 |
-
]
|
| 50 |
-
},
|
| 51 |
-
"filelock": {
|
| 52 |
-
"version": "3.32.3",
|
| 53 |
-
"license_expression": "MIT",
|
| 54 |
-
"license_summary": "MIT License",
|
| 55 |
-
"license_classifiers": [
|
| 56 |
-
"License :: OSI Approved :: MIT License"
|
| 57 |
-
],
|
| 58 |
-
"license_files": [
|
| 59 |
-
"filelock-3.32.3.dist-info/licenses/LICENSE"
|
| 60 |
-
]
|
| 61 |
-
},
|
| 62 |
-
"fsspec": {
|
| 63 |
-
"version": "2026.7.0",
|
| 64 |
-
"license_expression": "BSD-3-Clause",
|
| 65 |
-
"license_summary": "BSD 3-Clause License",
|
| 66 |
-
"license_classifiers": [],
|
| 67 |
-
"license_files": [
|
| 68 |
-
"fsspec-2026.7.0.dist-info/licenses/LICENSE"
|
| 69 |
-
]
|
| 70 |
-
},
|
| 71 |
-
"h11": {
|
| 72 |
-
"version": "0.16.0",
|
| 73 |
-
"license_expression": null,
|
| 74 |
-
"license_summary": "MIT",
|
| 75 |
-
"license_classifiers": [
|
| 76 |
-
"License :: OSI Approved :: MIT License"
|
| 77 |
-
],
|
| 78 |
-
"license_files": [
|
| 79 |
-
"h11-0.16.0.dist-info/licenses/LICENSE.txt"
|
| 80 |
-
]
|
| 81 |
-
},
|
| 82 |
-
"hf-xet": {
|
| 83 |
-
"version": "1.6.0",
|
| 84 |
-
"license_expression": "Apache-2.0",
|
| 85 |
-
"license_summary": "Apache License",
|
| 86 |
-
"license_classifiers": [
|
| 87 |
-
"License :: OSI Approved :: Apache Software License"
|
| 88 |
-
],
|
| 89 |
-
"license_files": [
|
| 90 |
-
"hf_xet-1.6.0.dist-info/licenses/LICENSE"
|
| 91 |
-
]
|
| 92 |
-
},
|
| 93 |
-
"httpcore2": {
|
| 94 |
-
"version": "2.13.1",
|
| 95 |
-
"license_expression": "BSD-3-Clause",
|
| 96 |
-
"license_summary": "Copyright \u00a9 2026 to present Pydantic Services Inc. and individual contributors.",
|
| 97 |
-
"license_classifiers": [
|
| 98 |
-
"License :: OSI Approved :: BSD License"
|
| 99 |
-
],
|
| 100 |
-
"license_files": [
|
| 101 |
-
"httpcore2-2.13.1.dist-info/licenses/LICENSE.md"
|
| 102 |
-
]
|
| 103 |
-
},
|
| 104 |
-
"httpx2": {
|
| 105 |
-
"version": "2.13.1",
|
| 106 |
-
"license_expression": "BSD-3-Clause",
|
| 107 |
-
"license_summary": "Copyright \u00a9 2026 to present Pydantic Services Inc. and individual contributors.",
|
| 108 |
-
"license_classifiers": [
|
| 109 |
-
"License :: OSI Approved :: BSD License"
|
| 110 |
-
],
|
| 111 |
-
"license_files": [
|
| 112 |
-
"httpx2-2.13.1.dist-info/licenses/LICENSE.md"
|
| 113 |
-
]
|
| 114 |
-
},
|
| 115 |
-
"huggingface-hub": {
|
| 116 |
-
"version": "2.1.1",
|
| 117 |
-
"license_expression": null,
|
| 118 |
-
"license_summary": "Apache-2.0",
|
| 119 |
-
"license_classifiers": [
|
| 120 |
-
"License :: OSI Approved :: Apache Software License"
|
| 121 |
-
],
|
| 122 |
-
"license_files": [
|
| 123 |
-
"huggingface_hub-2.1.1.dist-info/licenses/LICENSE"
|
| 124 |
-
]
|
| 125 |
-
},
|
| 126 |
-
"idna": {
|
| 127 |
-
"version": "3.20",
|
| 128 |
-
"license_expression": "BSD-3-Clause",
|
| 129 |
-
"license_summary": "BSD 3-Clause License",
|
| 130 |
-
"license_classifiers": [],
|
| 131 |
-
"license_files": [
|
| 132 |
-
"idna-3.20.dist-info/licenses/LICENSE.md"
|
| 133 |
-
]
|
| 134 |
-
},
|
| 135 |
-
"jinja2": {
|
| 136 |
-
"version": "3.1.6",
|
| 137 |
-
"license_expression": null,
|
| 138 |
-
"license_summary": "Copyright 2007 Pallets",
|
| 139 |
-
"license_classifiers": [
|
| 140 |
-
"License :: OSI Approved :: BSD License"
|
| 141 |
-
],
|
| 142 |
-
"license_files": [
|
| 143 |
-
"jinja2-3.1.6.dist-info/licenses/LICENSE.txt"
|
| 144 |
-
]
|
| 145 |
-
},
|
| 146 |
-
"joblib": {
|
| 147 |
-
"version": "1.6.0",
|
| 148 |
-
"license_expression": "BSD-3-Clause",
|
| 149 |
-
"license_summary": "BSD 3-Clause License",
|
| 150 |
-
"license_classifiers": [],
|
| 151 |
-
"license_files": [
|
| 152 |
-
"joblib-1.6.0.dist-info/licenses/LICENSE.txt"
|
| 153 |
-
]
|
| 154 |
-
},
|
| 155 |
-
"markupsafe": {
|
| 156 |
-
"version": "3.0.4",
|
| 157 |
-
"license_expression": "BSD-3-Clause",
|
| 158 |
-
"license_summary": "Copyright 2010 Pallets",
|
| 159 |
-
"license_classifiers": [],
|
| 160 |
-
"license_files": [
|
| 161 |
-
"markupsafe-3.0.4.dist-info/licenses/LICENSE.txt"
|
| 162 |
-
]
|
| 163 |
-
},
|
| 164 |
-
"mpmath": {
|
| 165 |
-
"version": "1.3.0",
|
| 166 |
-
"license_expression": null,
|
| 167 |
-
"license_summary": "BSD",
|
| 168 |
-
"license_classifiers": [
|
| 169 |
-
"License :: OSI Approved :: BSD License"
|
| 170 |
-
],
|
| 171 |
-
"license_files": [
|
| 172 |
-
"mpmath-1.3.0.dist-info/LICENSE"
|
| 173 |
-
]
|
| 174 |
-
},
|
| 175 |
-
"narwhals": {
|
| 176 |
-
"version": "2.27.0",
|
| 177 |
-
"license_expression": "MIT",
|
| 178 |
-
"license_summary": "MIT License",
|
| 179 |
-
"license_classifiers": [],
|
| 180 |
-
"license_files": [
|
| 181 |
-
"narwhals-2.27.0.dist-info/licenses/LICENSE.md"
|
| 182 |
-
]
|
| 183 |
-
},
|
| 184 |
-
"networkx": {
|
| 185 |
-
"version": "3.6.1",
|
| 186 |
-
"license_expression": "BSD-3-Clause",
|
| 187 |
-
"license_summary": "NetworkX is distributed with the 3-clause BSD license.",
|
| 188 |
-
"license_classifiers": [],
|
| 189 |
-
"license_files": [
|
| 190 |
-
"networkx-3.6.1.dist-info/licenses/LICENSE.txt"
|
| 191 |
-
]
|
| 192 |
-
},
|
| 193 |
-
"numpy": {
|
| 194 |
-
"version": "2.5.3",
|
| 195 |
-
"license_expression": "BSD-3-Clause AND 0BSD AND MIT AND Zlib AND CC0-1.0",
|
| 196 |
-
"license_summary": "Copyright (c) 2005-2025, NumPy Developers.",
|
| 197 |
-
"license_classifiers": [],
|
| 198 |
-
"license_files": [
|
| 199 |
-
"numpy-2.5.3.dist-info/licenses/LICENSE.txt",
|
| 200 |
-
"numpy-2.5.3.dist-info/licenses/numpy/_core/include/numpy/libdivide/LICENSE.txt",
|
| 201 |
-
"numpy-2.5.3.dist-info/licenses/numpy/_core/src/common/pythoncapi-compat/COPYING",
|
| 202 |
-
"numpy-2.5.3.dist-info/licenses/numpy/_core/src/highway/LICENSE",
|
| 203 |
-
"numpy-2.5.3.dist-info/licenses/numpy/_core/src/multiarray/dragon4_LICENSE.txt",
|
| 204 |
-
"numpy-2.5.3.dist-info/licenses/numpy/_core/src/npysort/x86-simd-sort/LICENSE.md",
|
| 205 |
-
"numpy-2.5.3.dist-info/licenses/numpy/_core/src/umath/svml/LICENSE",
|
| 206 |
-
"numpy-2.5.3.dist-info/licenses/numpy/fft/pocketfft/LICENSE.md",
|
| 207 |
-
"numpy-2.5.3.dist-info/licenses/numpy/linalg/lapack_lite/LICENSE.txt",
|
| 208 |
-
"numpy-2.5.3.dist-info/licenses/numpy/ma/LICENSE",
|
| 209 |
-
"numpy-2.5.3.dist-info/licenses/numpy/random/LICENSE.md",
|
| 210 |
-
"numpy-2.5.3.dist-info/licenses/numpy/random/src/distributions/LICENSE.md",
|
| 211 |
-
"numpy-2.5.3.dist-info/licenses/numpy/random/src/mt19937/LICENSE.md",
|
| 212 |
-
"numpy-2.5.3.dist-info/licenses/numpy/random/src/pcg64/LICENSE.md",
|
| 213 |
-
"numpy-2.5.3.dist-info/licenses/numpy/random/src/philox/LICENSE.md",
|
| 214 |
-
"numpy-2.5.3.dist-info/licenses/numpy/random/src/sfc64/LICENSE.md",
|
| 215 |
-
"numpy-2.5.3.dist-info/licenses/numpy/random/src/splitmix64/LICENSE.md"
|
| 216 |
-
]
|
| 217 |
-
},
|
| 218 |
-
"packaging": {
|
| 219 |
-
"version": "26.3",
|
| 220 |
-
"license_expression": "Apache-2.0 OR BSD-2-Clause",
|
| 221 |
-
"license_summary": "This software is made available under the terms of *either* of the licenses",
|
| 222 |
-
"license_classifiers": [],
|
| 223 |
-
"license_files": [
|
| 224 |
-
"packaging-26.3.dist-info/licenses/LICENSE",
|
| 225 |
-
"packaging-26.3.dist-info/licenses/LICENSE.APACHE",
|
| 226 |
-
"packaging-26.3.dist-info/licenses/LICENSE.BSD"
|
| 227 |
-
]
|
| 228 |
-
},
|
| 229 |
-
"pycparser": {
|
| 230 |
-
"version": "3.0",
|
| 231 |
-
"license_expression": "BSD-3-Clause",
|
| 232 |
-
"license_summary": "pycparser -- A C parser in Python",
|
| 233 |
-
"license_classifiers": [],
|
| 234 |
-
"license_files": [
|
| 235 |
-
"pycparser-3.0.dist-info/licenses/LICENSE"
|
| 236 |
-
]
|
| 237 |
-
},
|
| 238 |
-
"pyyaml": {
|
| 239 |
-
"version": "6.0.3",
|
| 240 |
-
"license_expression": null,
|
| 241 |
-
"license_summary": "MIT",
|
| 242 |
-
"license_classifiers": [
|
| 243 |
-
"License :: OSI Approved :: MIT License"
|
| 244 |
-
],
|
| 245 |
-
"license_files": [
|
| 246 |
-
"pyyaml-6.0.3.dist-info/licenses/LICENSE"
|
| 247 |
-
]
|
| 248 |
-
},
|
| 249 |
-
"rendercanvas": {
|
| 250 |
-
"version": "2.7.2",
|
| 251 |
-
"license_expression": null,
|
| 252 |
-
"license_summary": "BSD 2-Clause License",
|
| 253 |
-
"license_classifiers": [],
|
| 254 |
-
"license_files": [
|
| 255 |
-
"rendercanvas-2.7.2.dist-info/licenses/LICENSE"
|
| 256 |
-
]
|
| 257 |
-
},
|
| 258 |
-
"safetensors": {
|
| 259 |
-
"version": "0.8.0",
|
| 260 |
-
"license_expression": null,
|
| 261 |
-
"license_summary": "Apache License",
|
| 262 |
-
"license_classifiers": [
|
| 263 |
-
"License :: OSI Approved :: Apache Software License"
|
| 264 |
-
],
|
| 265 |
-
"license_files": [
|
| 266 |
-
"safetensors-0.8.0.dist-info/licenses/LICENSE"
|
| 267 |
-
]
|
| 268 |
-
},
|
| 269 |
-
"scikit-learn": {
|
| 270 |
-
"version": "1.9.1",
|
| 271 |
-
"license_expression": "BSD-3-Clause",
|
| 272 |
-
"license_summary": "BSD 3-Clause License",
|
| 273 |
-
"license_classifiers": [],
|
| 274 |
-
"license_files": [
|
| 275 |
-
"scikit_learn-1.9.1.dist-info/licenses/COPYING"
|
| 276 |
-
]
|
| 277 |
-
},
|
| 278 |
-
"scipy": {
|
| 279 |
-
"version": "1.18.1",
|
| 280 |
-
"license_expression": null,
|
| 281 |
-
"license_summary": "Copyright (c) 2001-2002 Enthought, Inc. 2003, SciPy Developers.",
|
| 282 |
-
"license_classifiers": [
|
| 283 |
-
"License :: OSI Approved :: BSD License"
|
| 284 |
-
],
|
| 285 |
-
"license_files": [
|
| 286 |
-
"scipy-1.18.1.dist-info/LICENSE.txt"
|
| 287 |
-
]
|
| 288 |
-
},
|
| 289 |
-
"setuptools": {
|
| 290 |
-
"version": "83.0.0",
|
| 291 |
-
"license_expression": "MIT",
|
| 292 |
-
"license_summary": "GNU LESSER GENERAL PUBLIC LICENSE",
|
| 293 |
-
"license_classifiers": [],
|
| 294 |
-
"license_files": [
|
| 295 |
-
"setuptools/_vendor/autocommand-2.2.2.dist-info/LICENSE",
|
| 296 |
-
"setuptools/_vendor/backports.tarfile-1.2.0.dist-info/LICENSE",
|
| 297 |
-
"setuptools/_vendor/importlib_metadata-8.7.1.dist-info/licenses/LICENSE",
|
| 298 |
-
"setuptools/_vendor/jaraco.text-4.0.0.dist-info/LICENSE",
|
| 299 |
-
"setuptools/_vendor/jaraco_context-6.1.0.dist-info/licenses/LICENSE",
|
| 300 |
-
"setuptools/_vendor/jaraco_functools-4.4.0.dist-info/licenses/LICENSE",
|
| 301 |
-
"setuptools/_vendor/more_itertools-10.8.0.dist-info/licenses/LICENSE",
|
| 302 |
-
"setuptools/_vendor/packaging-26.0.dist-info/licenses/LICENSE",
|
| 303 |
-
"setuptools/_vendor/packaging-26.0.dist-info/licenses/LICENSE.APACHE",
|
| 304 |
-
"setuptools/_vendor/packaging-26.0.dist-info/licenses/LICENSE.BSD",
|
| 305 |
-
"setuptools/_vendor/platformdirs-4.4.0.dist-info/licenses/LICENSE",
|
| 306 |
-
"setuptools/_vendor/tomli-2.4.0.dist-info/licenses/LICENSE",
|
| 307 |
-
"setuptools/_vendor/wheel-0.46.3.dist-info/licenses/LICENSE.txt",
|
| 308 |
-
"setuptools/_vendor/zipp-3.23.0.dist-info/licenses/LICENSE"
|
| 309 |
-
]
|
| 310 |
-
},
|
| 311 |
-
"sympy": {
|
| 312 |
-
"version": "1.14.0",
|
| 313 |
-
"license_expression": null,
|
| 314 |
-
"license_summary": "BSD",
|
| 315 |
-
"license_classifiers": [
|
| 316 |
-
"License :: OSI Approved :: BSD License"
|
| 317 |
-
],
|
| 318 |
-
"license_files": [
|
| 319 |
-
"sympy-1.14.0.dist-info/licenses/AUTHORS",
|
| 320 |
-
"sympy-1.14.0.dist-info/licenses/LICENSE"
|
| 321 |
-
]
|
| 322 |
-
},
|
| 323 |
-
"threadpoolctl": {
|
| 324 |
-
"version": "3.7.0",
|
| 325 |
-
"license_expression": "BSD-3-Clause",
|
| 326 |
-
"license_summary": "Copyright (c) 2019, threadpoolctl contributors",
|
| 327 |
-
"license_classifiers": [],
|
| 328 |
-
"license_files": [
|
| 329 |
-
"threadpoolctl-3.7.0.dist-info/licenses/LICENSE"
|
| 330 |
-
]
|
| 331 |
-
},
|
| 332 |
-
"torch": {
|
| 333 |
-
"version": "2.14.1+cpu",
|
| 334 |
-
"license_expression": "Apache-2.0 AND Apache-2.0 WITH LLVM-exception AND BSD-2-Clause AND BSD-3-Clause AND BSL-1.0 AND MIT",
|
| 335 |
-
"license_summary": "From PyTorch:",
|
| 336 |
-
"license_classifiers": [],
|
| 337 |
-
"license_files": [
|
| 338 |
-
"torch-2.14.1+cpu.dist-info/licenses/LICENSE",
|
| 339 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/FP16/LICENSE",
|
| 340 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/FXdiv/LICENSE",
|
| 341 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/NNPACK/LICENSE",
|
| 342 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/NVTX/LICENSE.txt",
|
| 343 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/NVTX/docs/LICENSE.txt",
|
| 344 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/NVTX/python/LICENSE.txt",
|
| 345 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/NVTX/rust/LICENSE",
|
| 346 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/NVTX/tools/docs/github-markdown-css/license",
|
| 347 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/VulkanMemoryAllocator/LICENSE.txt",
|
| 348 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/XNNPACK/LICENSE",
|
| 349 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/aiter/3rdparty/composable_kernel/LICENSE",
|
| 350 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/aiter/3rdparty/composable_kernel/docs/license.rst",
|
| 351 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/aiter/LICENSE",
|
| 352 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/benchmark/LICENSE",
|
| 353 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/composable_kernel/LICENSE",
|
| 354 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/composable_kernel/docs/license.rst",
|
| 355 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/cpp-httplib/LICENSE",
|
| 356 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/cpuinfo/LICENSE",
|
| 357 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/cpuinfo/deps/clog/LICENSE",
|
| 358 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/cudnn_frontend/LICENSE.txt",
|
| 359 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/cutlass/LICENSE.txt",
|
| 360 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/cutlass/python/LICENSE.txt",
|
| 361 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/LICENSE",
|
| 362 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/composable_kernel/LICENSE",
|
| 363 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/composable_kernel/docs/license.rst",
|
| 364 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/cpuinfo/LICENSE",
|
| 365 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/cpuinfo/deps/clog/LICENSE",
|
| 366 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/cutlass/LICENSE.txt",
|
| 367 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/cutlass/python/LICENSE.txt",
|
| 368 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/googletest/LICENSE",
|
| 369 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/external/hipify_torch/LICENSE.txt",
|
| 370 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/fbgemm_gpu/docs/src/general/License.rst",
|
| 371 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/fbgemm_gpu/experimental/hstu/LICENSE",
|
| 372 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/fbgemm_gpu/src/quantize_ops/mx/LICENSE",
|
| 373 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fbgemm/fbgemm_gpu/test/quantize/mx/LICENSE",
|
| 374 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/LICENSE",
|
| 375 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/csrc/composable_kernel/LICENSE",
|
| 376 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/csrc/composable_kernel/docs/license.rst",
|
| 377 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/csrc/cutlass/LICENSE.txt",
|
| 378 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/csrc/cutlass/python/LICENSE.txt",
|
| 379 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/flash_attn/cute/LICENSE",
|
| 380 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/third_party/aiter/3rdparty/composable_kernel/LICENSE",
|
| 381 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/third_party/aiter/3rdparty/composable_kernel/docs/license.rst",
|
| 382 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flash-attention/third_party/aiter/LICENSE",
|
| 383 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flatbuffers/LICENSE",
|
| 384 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flatbuffers/dart/LICENSE",
|
| 385 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/flatbuffers/swift/LICENSE",
|
| 386 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/fmt/LICENSE",
|
| 387 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/gemmlowp/gemmlowp/LICENSE",
|
| 388 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/gloo/LICENSE",
|
| 389 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/googletest/LICENSE",
|
| 390 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/ideep/LICENSE",
|
| 391 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/ideep/mkl-dnn/LICENSE",
|
| 392 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/ideep/mkl-dnn/third_party/gtest/LICENSE",
|
| 393 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/ideep/mkl-dnn/third_party/opencl/LICENSE",
|
| 394 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/kineto/LICENSE",
|
| 395 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/kineto/libkineto/third_party/dynolog_headers/LICENSE",
|
| 396 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/kineto/libkineto/third_party/fmt/LICENSE",
|
| 397 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/kineto/libkineto/third_party/googletest/LICENSE",
|
| 398 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/llvm-openmp/LICENSE.txt",
|
| 399 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mimalloc/LICENSE",
|
| 400 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/miniz-3.0.2/LICENSE",
|
| 401 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/LICENSE",
|
| 402 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/external/composable_kernel/LICENSE",
|
| 403 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/external/composable_kernel/docs/license.rst",
|
| 404 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/external/cutlass/LICENSE.txt",
|
| 405 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/external/cutlass/python/LICENSE.txt",
|
| 406 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/external/googletest/LICENSE",
|
| 407 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/external/hipify_torch/LICENSE.txt",
|
| 408 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/mslk/mslk/attention/flash_attn/LICENSE",
|
| 409 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/onnx/LICENSE",
|
| 410 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/onnx/third_party/pybind11/LICENSE",
|
| 411 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/perfetto/LICENSE",
|
| 412 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/protobuf/LICENSE",
|
| 413 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/protobuf/third_party/benchmark/LICENSE",
|
| 414 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/protobuf/third_party/googletest/LICENSE",
|
| 415 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/protobuf/third_party/googletest/googlemock/LICENSE",
|
| 416 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/protobuf/third_party/googletest/googlemock/scripts/generator/LICENSE",
|
| 417 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/protobuf/third_party/googletest/googletest/LICENSE",
|
| 418 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/protobuf/third_party/utf8_range/LICENSE",
|
| 419 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/psimd/LICENSE",
|
| 420 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/pthreadpool/LICENSE",
|
| 421 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/pybind11/LICENSE",
|
| 422 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/python-peachpy/LICENSE.rst",
|
| 423 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/sleef/LICENSE.txt",
|
| 424 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/LICENSE.txt",
|
| 425 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/googletest/LICENSE",
|
| 426 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/googletest/googlemock/LICENSE",
|
| 427 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/googletest/googlemock/scripts/generator/LICENSE",
|
| 428 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/googletest/googletest/LICENSE",
|
| 429 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/libnop/LICENSE",
|
| 430 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/libuv/LICENSE",
|
| 431 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/pybind11/LICENSE",
|
| 432 |
-
"torch-2.14.1+cpu.dist-info/licenses/third_party/tensorpipe/third_party/pybind11/tools/clang/LICENSE.TXT"
|
| 433 |
-
]
|
| 434 |
-
},
|
| 435 |
-
"tqdm": {
|
| 436 |
-
"version": "4.70.1",
|
| 437 |
-
"license_expression": null,
|
| 438 |
-
"license_summary": "MPL-2.0 AND MIT",
|
| 439 |
-
"license_classifiers": [],
|
| 440 |
-
"license_files": [
|
| 441 |
-
"tqdm-4.70.1.dist-info/licenses/LICENCE"
|
| 442 |
-
]
|
| 443 |
-
},
|
| 444 |
-
"truststore": {
|
| 445 |
-
"version": "0.10.4",
|
| 446 |
-
"license_expression": "MIT",
|
| 447 |
-
"license_summary": "The MIT License (MIT)",
|
| 448 |
-
"license_classifiers": [],
|
| 449 |
-
"license_files": [
|
| 450 |
-
"truststore-0.10.4.dist-info/licenses/LICENSE"
|
| 451 |
-
]
|
| 452 |
-
},
|
| 453 |
-
"typing-extensions": {
|
| 454 |
-
"version": "4.16.0",
|
| 455 |
-
"license_expression": "PSF-2.0",
|
| 456 |
-
"license_summary": "A. HISTORY OF THE SOFTWARE",
|
| 457 |
-
"license_classifiers": [],
|
| 458 |
-
"license_files": [
|
| 459 |
-
"typing_extensions-4.16.0.dist-info/licenses/LICENSE"
|
| 460 |
-
]
|
| 461 |
-
},
|
| 462 |
-
"wgpu": {
|
| 463 |
-
"version": "0.32.0",
|
| 464 |
-
"license_expression": null,
|
| 465 |
-
"license_summary": "BSD 2-Clause License",
|
| 466 |
-
"license_classifiers": [],
|
| 467 |
-
"license_files": [
|
| 468 |
-
"wgpu-0.32.0.dist-info/licenses/LICENSE"
|
| 469 |
-
]
|
| 470 |
-
}
|
| 471 |
-
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
dist/lightpfn-0.1.0-py3-none-any.whl
DELETED
|
Binary file (46.8 kB)
|
|
|
dist/lightpfn-0.1.0.tar.gz
DELETED
|
@@ -1,3 +0,0 @@
|
|
| 1 |
-
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:3a9606af85ad6d178963717cb7ad0dec432ff44ea7bea000396bb34c0ff2ab5d
|
| 3 |
-
size 45222
|
|
|
|
|
|
|
|
|
|
|
|
installed_verification.json
DELETED
|
@@ -1,17 +0,0 @@
|
|
| 1 |
-
{
|
| 2 |
-
"package_version": "0.1.0",
|
| 3 |
-
"model": {
|
| 4 |
-
"repo_id": "ueuegio/LightPFN",
|
| 5 |
-
"revision": "bd389ab59a89dd0e05c9ecb7c642c08ee52e9637"
|
| 6 |
-
},
|
| 7 |
-
"installed_import": true,
|
| 8 |
-
"hub_download": true,
|
| 9 |
-
"offline_network_blocked": true,
|
| 10 |
-
"estimators": 4,
|
| 11 |
-
"max_probability_difference": 0.0,
|
| 12 |
-
"probability_shape": [
|
| 13 |
-
40,
|
| 14 |
-
3
|
| 15 |
-
],
|
| 16 |
-
"finite": true
|
| 17 |
-
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/.gitignore
DELETED
|
@@ -1,53 +0,0 @@
|
|
| 1 |
-
__pycache__/
|
| 2 |
-
*.pyc
|
| 3 |
-
dist/
|
| 4 |
-
build/
|
| 5 |
-
*.egg-info/
|
| 6 |
-
.pytest_cache/
|
| 7 |
-
runs/release/staging/
|
| 8 |
-
runs/release/install_check/
|
| 9 |
-
*.safetensors
|
| 10 |
-
|
| 11 |
-
# datasets, synthetic pools, OpenML cache and TabICL offload files stay on the data disk
|
| 12 |
-
data/
|
| 13 |
-
|
| 14 |
-
# model checkpoints
|
| 15 |
-
*.pt
|
| 16 |
-
*.pt.tmp
|
| 17 |
-
|
| 18 |
-
# scratch and local copies
|
| 19 |
-
runs/snapshots/
|
| 20 |
-
runs/codex_tmp/
|
| 21 |
-
|
| 22 |
-
# full Codex transcripts (the prompts and final answers are tracked)
|
| 23 |
-
runs/reviews/*.log
|
| 24 |
-
runs/reviews/*.err
|
| 25 |
-
|
| 26 |
-
# Codex patch dumps (the history is in git)
|
| 27 |
-
runs/reviews/*.patch
|
| 28 |
-
|
| 29 |
-
# bulky analysis intermediates
|
| 30 |
-
runs/analysis_r1/*.npz
|
| 31 |
-
runs/analysis_r1/prior_feature_stats.csv
|
| 32 |
-
runs/analysis_r1/prior_geometry_all.csv
|
| 33 |
-
runs/analysis_r1/matched_n128_feature_stats.csv
|
| 34 |
-
|
| 35 |
-
# thermal watchdog output (the script is tracked)
|
| 36 |
-
runs/thermal/temps.csv
|
| 37 |
-
runs/thermal/watchdog.log
|
| 38 |
-
runs/thermal/STOP
|
| 39 |
-
|
| 40 |
-
# benchmark scratch (temporary runs, pools, test dirs)
|
| 41 |
-
runs/bench_train/pytest_tmp/
|
| 42 |
-
runs/bench_train/test_tmp/
|
| 43 |
-
runs/bench_train/pools/
|
| 44 |
-
runs/bench_train/runs/
|
| 45 |
-
runs/bench_train/plan_*.json
|
| 46 |
-
|
| 47 |
-
# LaTeX build files (the PDF is tracked)
|
| 48 |
-
docs/report/*.aux
|
| 49 |
-
docs/report/*.log
|
| 50 |
-
docs/report/*.out
|
| 51 |
-
|
| 52 |
-
# lock files of the evaluation queue
|
| 53 |
-
runs/**/.eval.lock
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/LICENSE
DELETED
|
@@ -1,202 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
Apache License
|
| 3 |
-
Version 2.0, January 2004
|
| 4 |
-
http://www.apache.org/licenses/
|
| 5 |
-
|
| 6 |
-
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 7 |
-
|
| 8 |
-
1. Definitions.
|
| 9 |
-
|
| 10 |
-
"License" shall mean the terms and conditions for use, reproduction,
|
| 11 |
-
and distribution as defined by Sections 1 through 9 of this document.
|
| 12 |
-
|
| 13 |
-
"Licensor" shall mean the copyright owner or entity authorized by
|
| 14 |
-
the copyright owner that is granting the License.
|
| 15 |
-
|
| 16 |
-
"Legal Entity" shall mean the union of the acting entity and all
|
| 17 |
-
other entities that control, are controlled by, or are under common
|
| 18 |
-
control with that entity. For the purposes of this definition,
|
| 19 |
-
"control" means (i) the power, direct or indirect, to cause the
|
| 20 |
-
direction or management of such entity, whether by contract or
|
| 21 |
-
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 22 |
-
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 23 |
-
|
| 24 |
-
"You" (or "Your") shall mean an individual or Legal Entity
|
| 25 |
-
exercising permissions granted by this License.
|
| 26 |
-
|
| 27 |
-
"Source" form shall mean the preferred form for making modifications,
|
| 28 |
-
including but not limited to software source code, documentation
|
| 29 |
-
source, and configuration files.
|
| 30 |
-
|
| 31 |
-
"Object" form shall mean any form resulting from mechanical
|
| 32 |
-
transformation or translation of a Source form, including but
|
| 33 |
-
not limited to compiled object code, generated documentation,
|
| 34 |
-
and conversions to other media types.
|
| 35 |
-
|
| 36 |
-
"Work" shall mean the work of authorship, whether in Source or
|
| 37 |
-
Object form, made available under the License, as indicated by a
|
| 38 |
-
copyright notice that is included in or attached to the work
|
| 39 |
-
(an example is provided in the Appendix below).
|
| 40 |
-
|
| 41 |
-
"Derivative Works" shall mean any work, whether in Source or Object
|
| 42 |
-
form, that is based on (or derived from) the Work and for which the
|
| 43 |
-
editorial revisions, annotations, elaborations, or other modifications
|
| 44 |
-
represent, as a whole, an original work of authorship. For the purposes
|
| 45 |
-
of this License, Derivative Works shall not include works that remain
|
| 46 |
-
separable from, or merely link (or bind by name) to the interfaces of,
|
| 47 |
-
the Work and Derivative Works thereof.
|
| 48 |
-
|
| 49 |
-
"Contribution" shall mean any work of authorship, including
|
| 50 |
-
the original version of the Work and any modifications or additions
|
| 51 |
-
to that Work or Derivative Works thereof, that is intentionally
|
| 52 |
-
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 53 |
-
or by an individual or Legal Entity authorized to submit on behalf of
|
| 54 |
-
the copyright owner. For the purposes of this definition, "submitted"
|
| 55 |
-
means any form of electronic, verbal, or written communication sent
|
| 56 |
-
to the Licensor or its representatives, including but not limited to
|
| 57 |
-
communication on electronic mailing lists, source code control systems,
|
| 58 |
-
and issue tracking systems that are managed by, or on behalf of, the
|
| 59 |
-
Licensor for the purpose of discussing and improving the Work, but
|
| 60 |
-
excluding communication that is conspicuously marked or otherwise
|
| 61 |
-
designated in writing by the copyright owner as "Not a Contribution."
|
| 62 |
-
|
| 63 |
-
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 64 |
-
on behalf of whom a Contribution has been received by Licensor and
|
| 65 |
-
subsequently incorporated within the Work.
|
| 66 |
-
|
| 67 |
-
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 68 |
-
this License, each Contributor hereby grants to You a perpetual,
|
| 69 |
-
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 70 |
-
copyright license to reproduce, prepare Derivative Works of,
|
| 71 |
-
publicly display, publicly perform, sublicense, and distribute the
|
| 72 |
-
Work and such Derivative Works in Source or Object form.
|
| 73 |
-
|
| 74 |
-
3. Grant of Patent License. Subject to the terms and conditions of
|
| 75 |
-
this License, each Contributor hereby grants to You a perpetual,
|
| 76 |
-
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 77 |
-
(except as stated in this section) patent license to make, have made,
|
| 78 |
-
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 79 |
-
where such license applies only to those patent claims licensable
|
| 80 |
-
by such Contributor that are necessarily infringed by their
|
| 81 |
-
Contribution(s) alone or by combination of their Contribution(s)
|
| 82 |
-
with the Work to which such Contribution(s) was submitted. If You
|
| 83 |
-
institute patent litigation against any entity (including a
|
| 84 |
-
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 85 |
-
or a Contribution incorporated within the Work constitutes direct
|
| 86 |
-
or contributory patent infringement, then any patent licenses
|
| 87 |
-
granted to You under this License for that Work shall terminate
|
| 88 |
-
as of the date such litigation is filed.
|
| 89 |
-
|
| 90 |
-
4. Redistribution. You may reproduce and distribute copies of the
|
| 91 |
-
Work or Derivative Works thereof in any medium, with or without
|
| 92 |
-
modifications, and in Source or Object form, provided that You
|
| 93 |
-
meet the following conditions:
|
| 94 |
-
|
| 95 |
-
(a) You must give any other recipients of the Work or
|
| 96 |
-
Derivative Works a copy of this License; and
|
| 97 |
-
|
| 98 |
-
(b) You must cause any modified files to carry prominent notices
|
| 99 |
-
stating that You changed the files; and
|
| 100 |
-
|
| 101 |
-
(c) You must retain, in the Source form of any Derivative Works
|
| 102 |
-
that You distribute, all copyright, patent, trademark, and
|
| 103 |
-
attribution notices from the Source form of the Work,
|
| 104 |
-
excluding those notices that do not pertain to any part of
|
| 105 |
-
the Derivative Works; and
|
| 106 |
-
|
| 107 |
-
(d) If the Work includes a "NOTICE" text file as part of its
|
| 108 |
-
distribution, then any Derivative Works that You distribute must
|
| 109 |
-
include a readable copy of the attribution notices contained
|
| 110 |
-
within such NOTICE file, excluding those notices that do not
|
| 111 |
-
pertain to any part of the Derivative Works, in at least one
|
| 112 |
-
of the following places: within a NOTICE text file distributed
|
| 113 |
-
as part of the Derivative Works; within the Source form or
|
| 114 |
-
documentation, if provided along with the Derivative Works; or,
|
| 115 |
-
within a display generated by the Derivative Works, if and
|
| 116 |
-
wherever such third-party notices normally appear. The contents
|
| 117 |
-
of the NOTICE file are for informational purposes only and
|
| 118 |
-
do not modify the License. You may add Your own attribution
|
| 119 |
-
notices within Derivative Works that You distribute, alongside
|
| 120 |
-
or as an addendum to the NOTICE text from the Work, provided
|
| 121 |
-
that such additional attribution notices cannot be construed
|
| 122 |
-
as modifying the License.
|
| 123 |
-
|
| 124 |
-
You may add Your own copyright statement to Your modifications and
|
| 125 |
-
may provide additional or different license terms and conditions
|
| 126 |
-
for use, reproduction, or distribution of Your modifications, or
|
| 127 |
-
for any such Derivative Works as a whole, provided Your use,
|
| 128 |
-
reproduction, and distribution of the Work otherwise complies with
|
| 129 |
-
the conditions stated in this License.
|
| 130 |
-
|
| 131 |
-
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 132 |
-
any Contribution intentionally submitted for inclusion in the Work
|
| 133 |
-
by You to the Licensor shall be under the terms and conditions of
|
| 134 |
-
this License, without any additional terms or conditions.
|
| 135 |
-
Notwithstanding the above, nothing herein shall supersede or modify
|
| 136 |
-
the terms of any separate license agreement you may have executed
|
| 137 |
-
with Licensor regarding such Contributions.
|
| 138 |
-
|
| 139 |
-
6. Trademarks. This License does not grant permission to use the trade
|
| 140 |
-
names, trademarks, service marks, or product names of the Licensor,
|
| 141 |
-
except as required for reasonable and customary use in describing the
|
| 142 |
-
origin of the Work and reproducing the content of the NOTICE file.
|
| 143 |
-
|
| 144 |
-
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 145 |
-
agreed to in writing, Licensor provides the Work (and each
|
| 146 |
-
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 147 |
-
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 148 |
-
implied, including, without limitation, any warranties or conditions
|
| 149 |
-
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 150 |
-
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 151 |
-
appropriateness of using or redistributing the Work and assume any
|
| 152 |
-
risks associated with Your exercise of permissions under this License.
|
| 153 |
-
|
| 154 |
-
8. Limitation of Liability. In no event and under no legal theory,
|
| 155 |
-
whether in tort (including negligence), contract, or otherwise,
|
| 156 |
-
unless required by applicable law (such as deliberate and grossly
|
| 157 |
-
negligent acts) or agreed to in writing, shall any Contributor be
|
| 158 |
-
liable to You for damages, including any direct, indirect, special,
|
| 159 |
-
incidental, or consequential damages of any character arising as a
|
| 160 |
-
result of this License or out of the use or inability to use the
|
| 161 |
-
Work (including but not limited to damages for loss of goodwill,
|
| 162 |
-
work stoppage, computer failure or malfunction, or any and all
|
| 163 |
-
other commercial damages or losses), even if such Contributor
|
| 164 |
-
has been advised of the possibility of such damages.
|
| 165 |
-
|
| 166 |
-
9. Accepting Warranty or Additional Liability. While redistributing
|
| 167 |
-
the Work or Derivative Works thereof, You may choose to offer,
|
| 168 |
-
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 169 |
-
or other liability obligations and/or rights consistent with this
|
| 170 |
-
License. However, in accepting such obligations, You may act only
|
| 171 |
-
on Your own behalf and on Your sole responsibility, not on behalf
|
| 172 |
-
of any other Contributor, and only if You agree to indemnify,
|
| 173 |
-
defend, and hold each Contributor harmless for any liability
|
| 174 |
-
incurred by, or claims asserted against, such Contributor by reason
|
| 175 |
-
of your accepting any such warranty or additional liability.
|
| 176 |
-
|
| 177 |
-
END OF TERMS AND CONDITIONS
|
| 178 |
-
|
| 179 |
-
APPENDIX: How to apply the Apache License to your work.
|
| 180 |
-
|
| 181 |
-
To apply the Apache License to your work, attach the following
|
| 182 |
-
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 183 |
-
replaced with your own identifying information. (Don't include
|
| 184 |
-
the brackets!) The text should be enclosed in the appropriate
|
| 185 |
-
comment syntax for the file format. We also recommend that a
|
| 186 |
-
file or class name and description of purpose be included on the
|
| 187 |
-
same "printed page" as the copyright notice for easier
|
| 188 |
-
identification within third-party archives.
|
| 189 |
-
|
| 190 |
-
Copyright [yyyy] [name of copyright owner]
|
| 191 |
-
|
| 192 |
-
Licensed under the Apache License, Version 2.0 (the "License");
|
| 193 |
-
you may not use this file except in compliance with the License.
|
| 194 |
-
You may obtain a copy of the License at
|
| 195 |
-
|
| 196 |
-
http://www.apache.org/licenses/LICENSE-2.0
|
| 197 |
-
|
| 198 |
-
Unless required by applicable law or agreed to in writing, software
|
| 199 |
-
distributed under the License is distributed on an "AS IS" BASIS,
|
| 200 |
-
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 201 |
-
See the License for the specific language governing permissions and
|
| 202 |
-
limitations under the License.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/NOTICE
DELETED
|
@@ -1,6 +0,0 @@
|
|
| 1 |
-
LightPFN
|
| 2 |
-
Copyright 2026 Giorgio (GioOtto)
|
| 3 |
-
|
| 4 |
-
Code and released model weights are licensed under the Apache License, Version 2.0.
|
| 5 |
-
The model was trained from scratch using synthetic data, without distillation
|
| 6 |
-
or training on weights or outputs from other tabular foundation models.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/PKG-INFO
DELETED
|
@@ -1,123 +0,0 @@
|
|
| 1 |
-
Metadata-Version: 2.5
|
| 2 |
-
Name: lightpfn
|
| 3 |
-
Version: 0.1.0
|
| 4 |
-
Summary: A compact tabular classifier pretrained only on synthetic data
|
| 5 |
-
Project-URL: Repository, https://github.com/GioOtto/gioPFN
|
| 6 |
-
Project-URL: Models, https://huggingface.co/ueuegio/LightPFN
|
| 7 |
-
Author-email: Giorgio <247403232+GioOtto@users.noreply.github.com>
|
| 8 |
-
License-Expression: Apache-2.0
|
| 9 |
-
License-File: LICENSE
|
| 10 |
-
License-File: NOTICE
|
| 11 |
-
Classifier: Development Status :: 3 - Alpha
|
| 12 |
-
Classifier: Intended Audience :: Science/Research
|
| 13 |
-
Classifier: Operating System :: OS Independent
|
| 14 |
-
Classifier: Programming Language :: Python :: 3
|
| 15 |
-
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
| 16 |
-
Requires-Python: >=3.10
|
| 17 |
-
Requires-Dist: numpy>=1.24
|
| 18 |
-
Requires-Dist: torch>=2.6
|
| 19 |
-
Provides-Extra: dev
|
| 20 |
-
Requires-Dist: build>=1.2; extra == 'dev'
|
| 21 |
-
Requires-Dist: hatchling>=1.27; extra == 'dev'
|
| 22 |
-
Requires-Dist: pytest>=8; extra == 'dev'
|
| 23 |
-
Provides-Extra: hf
|
| 24 |
-
Requires-Dist: huggingface-hub>=0.27; extra == 'hf'
|
| 25 |
-
Requires-Dist: safetensors>=0.5; extra == 'hf'
|
| 26 |
-
Provides-Extra: sklearn
|
| 27 |
-
Requires-Dist: scikit-learn>=1.6; extra == 'sklearn'
|
| 28 |
-
Provides-Extra: vulkan
|
| 29 |
-
Requires-Dist: wgpu<0.33,>=0.32; extra == 'vulkan'
|
| 30 |
-
Description-Content-Type: text/markdown
|
| 31 |
-
|
| 32 |
-
# LightPFN
|
| 33 |
-
|
| 34 |
-
A **4,603,088-parameter** in-context classifier pretrained from scratch only on
|
| 35 |
-
synthetic tabular tasks. `fit()` builds an inference context; no gradient training
|
| 36 |
-
is performed on your dataset. Code and released weights: **Apache 2.0**.
|
| 37 |
-
|
| 38 |
-
Python 3.10+; CPU, PyTorch CUDA/ROCm, and optional Vulkan inference.
|
| 39 |
-
The base package requires only **torch and numpy**. Synthetic priors, training
|
| 40 |
-
code, evaluation harnesses, and third-party foundation models are excluded from
|
| 41 |
-
both distribution artifacts.
|
| 42 |
-
|
| 43 |
-
## Install
|
| 44 |
-
|
| 45 |
-
Install a supplied wheel, adding the sklearn and Hugging Face extras:
|
| 46 |
-
|
| 47 |
-
```bash
|
| 48 |
-
pip install "./lightpfn-0.1.0-py3-none-any.whl[sklearn,hf]"
|
| 49 |
-
hf auth login
|
| 50 |
-
```
|
| 51 |
-
|
| 52 |
-
The initial Hugging Face repository is private: an authorized account/token is
|
| 53 |
-
required for downloads. Authentication is handled by `huggingface_hub`, using
|
| 54 |
-
its cached login or `HF_TOKEN`; credentials are never stored in the package.
|
| 55 |
-
The package is not yet published on PyPI.
|
| 56 |
-
|
| 57 |
-
From the source tree: `pip install ".[sklearn,hf]"`. For Vulkan add the `vulkan`
|
| 58 |
-
extra. For a CPU installation, install a CPU PyTorch wheel from the official
|
| 59 |
-
PyTorch index first; installing torch from PyPI may pull GPU runtime packages.
|
| 60 |
-
|
| 61 |
-
## Use
|
| 62 |
-
|
| 63 |
-
```python
|
| 64 |
-
from lightpfn import LightPFNClassifier
|
| 65 |
-
|
| 66 |
-
clf = LightPFNClassifier(device="cpu", n_estimators=4, random_state=0)
|
| 67 |
-
clf.fit(X_train, y_train)
|
| 68 |
-
probabilities = clf.predict_proba(X_test) # columns follow clf.classes_
|
| 69 |
-
predictions = clf.predict(X_test)
|
| 70 |
-
```
|
| 71 |
-
|
| 72 |
-
The constructor is cheap. The first `fit()` loads the released model from an
|
| 73 |
-
immutable Hugging Face commit, and subsequent use benefits from the HF cache.
|
| 74 |
-
`local_files_only=True` prohibits network downloads. `checkpoint=` accepts a
|
| 75 |
-
local safetensors directory/file or a compatible legacy `.pt` checkpoint.
|
| 76 |
-
`repo_id=` and `revision=` override the released model; custom repositories
|
| 77 |
-
require an explicit revision. No remote Python code is loaded.
|
| 78 |
-
|
| 79 |
-
The estimator inherits `ClassifierMixin` and `BaseEstimator`, supports
|
| 80 |
-
`get_params`, `set_params`, cloning, `Pipeline`, `GridSearchCV`, `score`, feature
|
| 81 |
-
name validation and `n_features_in_`. Use `random_state` or the legacy `seed`.
|
| 82 |
-
Standard estimator checks are exercised with local weights. Two strict query
|
| 83 |
-
invariance checks have documented FP32 tolerances: changing row order or batch
|
| 84 |
-
size can change probabilities around 1e-7. For cloning/model selection prefer
|
| 85 |
-
`checkpoint=` or the default pretrained model; sklearn's parameter hash check
|
| 86 |
-
does not reliably compare raw torch modules because it hashes storage identity.
|
| 87 |
-
Input must be a dense numeric table; NaN values are supported, infinity and
|
| 88 |
-
sparse tables are rejected. Encode strings/categories before fitting, for
|
| 89 |
-
example using sklearn's `OrdinalEncoder` with missing/unknown values mapped to
|
| 90 |
-
NaN. Native categorical semantics are not implemented. The deprecated `cat`
|
| 91 |
-
mask does not change predictions and warns when categorical columns are marked.
|
| 92 |
-
|
| 93 |
-
The model was trained for **2-10 classes**. A single-class fit returns its constant
|
| 94 |
-
class probability. For more than 10 classes, use a different classifier.
|
| 95 |
-
Above `max_context` (default 20,000), each member uses a stratified subset that
|
| 96 |
-
keeps every class. Larger contexts increase compute and memory substantially.
|
| 97 |
-
`n_threads` sets PyTorch's process-wide CPU thread count. On Vulkan, `auto`
|
| 98 |
-
falls back to CPU if the required buffers exceed adapter limits.
|
| 99 |
-
|
| 100 |
-
Low-level use with only torch/numpy:
|
| 101 |
-
|
| 102 |
-
```python
|
| 103 |
-
from lightpfn import load_model
|
| 104 |
-
|
| 105 |
-
model = load_model("inference.pt") # restricted weights_only=True
|
| 106 |
-
```
|
| 107 |
-
|
| 108 |
-
To load/export safetensors, install the `hf` extra and use `load_model(folder)`
|
| 109 |
-
or `save_model(model, folder)`. The release contains **safetensors + JSON**;
|
| 110 |
-
unrestricted pickle loading is never used by the inference package.
|
| 111 |
-
|
| 112 |
-
## Evidence and limitations
|
| 113 |
-
|
| 114 |
-
The final `final_long` checkpoint follows 120,000 base steps and 39,250
|
| 115 |
-
long-context steps. Selection used held-out synthetic tasks, probes and 55
|
| 116 |
-
OpenML-CC18 tasks outside TabArena. TabArena was inspected only for finalists.
|
| 117 |
-
On 38 classification tasks using official splits, first repeat, four estimators
|
| 118 |
-
achieved mean AUC 0.8581 and lower task error than default CatBoost on 76% of
|
| 119 |
-
tasks. This is a local evaluation, not an official TabArena submission, and does
|
| 120 |
-
not establish superiority over tuned GBDTs. Categorical-heavy datasets remain
|
| 121 |
-
a weakness; regression is unsupported.
|
| 122 |
-
|
| 123 |
-
The full experiment history and technical report are in the source repository.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/docs/DEPENDENCIES.md
DELETED
|
@@ -1,28 +0,0 @@
|
|
| 1 |
-
# Inference dependency licenses
|
| 2 |
-
|
| 3 |
-
LightPFN's code and released weights use Apache-2.0. Dependency packages are
|
| 4 |
-
installed separately and retain their upstream licenses; their bundled notices
|
| 5 |
-
and licenses must remain intact if you redistribute those packages.
|
| 6 |
-
|
| 7 |
-
| Dependency | Role | Upstream license |
|
| 8 |
-
|---|---|---|
|
| 9 |
-
| numpy | Required, numeric arrays | BSD-3-Clause; wheels also include permissive third-party notices |
|
| 10 |
-
| torch | Required, model execution and restricted local loader | BSD-3-Clause; tested CPU wheel includes Apache-2.0, LLVM exception, BSD, BSL-1.0, MIT notices |
|
| 11 |
-
| scikit-learn | Optional `sklearn` extra | BSD-3-Clause |
|
| 12 |
-
| huggingface_hub | Optional `hf` extra | Apache-2.0 |
|
| 13 |
-
| safetensors | Optional `hf` extra | Apache-2.0 |
|
| 14 |
-
| wgpu | Optional `vulkan` extra | BSD-2-Clause; native wgpu components have their own permissive notices |
|
| 15 |
-
|
| 16 |
-
These upstream licenses permit use with this Apache-2.0 package. They do not
|
| 17 |
-
relicense the dependencies as Apache-2.0. No other tabular foundation model is
|
| 18 |
-
a dependency of either the wheel or source distribution. Synthetic generation,
|
| 19 |
-
training and benchmark modules remain in the research repository and are not
|
| 20 |
-
distributed by this inference package.
|
| 21 |
-
|
| 22 |
-
`runs/release/dependency_licenses.json` records the resolved inference dependency
|
| 23 |
-
closure, versions and license metadata from the environment used to validate
|
| 24 |
-
the release. Future dependency resolutions should be audited separately.
|
| 25 |
-
The 35-package validation environment uses permissive BSD/MIT/Apache/PSF
|
| 26 |
-
licenses, together with tqdm's MPL-2.0/MIT license. Its rendercanvas license was
|
| 27 |
-
verified from the installed BSD-2-Clause license file; scipy retains its BSD
|
| 28 |
-
license and the third-party notices bundled in its wheel.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/docs/PACKAGE_README.md
DELETED
|
@@ -1,92 +0,0 @@
|
|
| 1 |
-
# LightPFN
|
| 2 |
-
|
| 3 |
-
A **4,603,088-parameter** in-context classifier pretrained from scratch only on
|
| 4 |
-
synthetic tabular tasks. `fit()` builds an inference context; no gradient training
|
| 5 |
-
is performed on your dataset. Code and released weights: **Apache 2.0**.
|
| 6 |
-
|
| 7 |
-
Python 3.10+; CPU, PyTorch CUDA/ROCm, and optional Vulkan inference.
|
| 8 |
-
The base package requires only **torch and numpy**. Synthetic priors, training
|
| 9 |
-
code, evaluation harnesses, and third-party foundation models are excluded from
|
| 10 |
-
both distribution artifacts.
|
| 11 |
-
|
| 12 |
-
## Install
|
| 13 |
-
|
| 14 |
-
Install a supplied wheel, adding the sklearn and Hugging Face extras:
|
| 15 |
-
|
| 16 |
-
```bash
|
| 17 |
-
pip install "./lightpfn-0.1.0-py3-none-any.whl[sklearn,hf]"
|
| 18 |
-
hf auth login
|
| 19 |
-
```
|
| 20 |
-
|
| 21 |
-
The initial Hugging Face repository is private: an authorized account/token is
|
| 22 |
-
required for downloads. Authentication is handled by `huggingface_hub`, using
|
| 23 |
-
its cached login or `HF_TOKEN`; credentials are never stored in the package.
|
| 24 |
-
The package is not yet published on PyPI.
|
| 25 |
-
|
| 26 |
-
From the source tree: `pip install ".[sklearn,hf]"`. For Vulkan add the `vulkan`
|
| 27 |
-
extra. For a CPU installation, install a CPU PyTorch wheel from the official
|
| 28 |
-
PyTorch index first; installing torch from PyPI may pull GPU runtime packages.
|
| 29 |
-
|
| 30 |
-
## Use
|
| 31 |
-
|
| 32 |
-
```python
|
| 33 |
-
from lightpfn import LightPFNClassifier
|
| 34 |
-
|
| 35 |
-
clf = LightPFNClassifier(device="cpu", n_estimators=4, random_state=0)
|
| 36 |
-
clf.fit(X_train, y_train)
|
| 37 |
-
probabilities = clf.predict_proba(X_test) # columns follow clf.classes_
|
| 38 |
-
predictions = clf.predict(X_test)
|
| 39 |
-
```
|
| 40 |
-
|
| 41 |
-
The constructor is cheap. The first `fit()` loads the released model from an
|
| 42 |
-
immutable Hugging Face commit, and subsequent use benefits from the HF cache.
|
| 43 |
-
`local_files_only=True` prohibits network downloads. `checkpoint=` accepts a
|
| 44 |
-
local safetensors directory/file or a compatible legacy `.pt` checkpoint.
|
| 45 |
-
`repo_id=` and `revision=` override the released model; custom repositories
|
| 46 |
-
require an explicit revision. No remote Python code is loaded.
|
| 47 |
-
|
| 48 |
-
The estimator inherits `ClassifierMixin` and `BaseEstimator`, supports
|
| 49 |
-
`get_params`, `set_params`, cloning, `Pipeline`, `GridSearchCV`, `score`, feature
|
| 50 |
-
name validation and `n_features_in_`. Use `random_state` or the legacy `seed`.
|
| 51 |
-
Standard estimator checks are exercised with local weights. Two strict query
|
| 52 |
-
invariance checks have documented FP32 tolerances: changing row order or batch
|
| 53 |
-
size can change probabilities around 1e-7. For cloning/model selection prefer
|
| 54 |
-
`checkpoint=` or the default pretrained model; sklearn's parameter hash check
|
| 55 |
-
does not reliably compare raw torch modules because it hashes storage identity.
|
| 56 |
-
Input must be a dense numeric table; NaN values are supported, infinity and
|
| 57 |
-
sparse tables are rejected. Encode strings/categories before fitting, for
|
| 58 |
-
example using sklearn's `OrdinalEncoder` with missing/unknown values mapped to
|
| 59 |
-
NaN. Native categorical semantics are not implemented. The deprecated `cat`
|
| 60 |
-
mask does not change predictions and warns when categorical columns are marked.
|
| 61 |
-
|
| 62 |
-
The model was trained for **2-10 classes**. A single-class fit returns its constant
|
| 63 |
-
class probability. For more than 10 classes, use a different classifier.
|
| 64 |
-
Above `max_context` (default 20,000), each member uses a stratified subset that
|
| 65 |
-
keeps every class. Larger contexts increase compute and memory substantially.
|
| 66 |
-
`n_threads` sets PyTorch's process-wide CPU thread count. On Vulkan, `auto`
|
| 67 |
-
falls back to CPU if the required buffers exceed adapter limits.
|
| 68 |
-
|
| 69 |
-
Low-level use with only torch/numpy:
|
| 70 |
-
|
| 71 |
-
```python
|
| 72 |
-
from lightpfn import load_model
|
| 73 |
-
|
| 74 |
-
model = load_model("inference.pt") # restricted weights_only=True
|
| 75 |
-
```
|
| 76 |
-
|
| 77 |
-
To load/export safetensors, install the `hf` extra and use `load_model(folder)`
|
| 78 |
-
or `save_model(model, folder)`. The release contains **safetensors + JSON**;
|
| 79 |
-
unrestricted pickle loading is never used by the inference package.
|
| 80 |
-
|
| 81 |
-
## Evidence and limitations
|
| 82 |
-
|
| 83 |
-
The final `final_long` checkpoint follows 120,000 base steps and 39,250
|
| 84 |
-
long-context steps. Selection used held-out synthetic tasks, probes and 55
|
| 85 |
-
OpenML-CC18 tasks outside TabArena. TabArena was inspected only for finalists.
|
| 86 |
-
On 38 classification tasks using official splits, first repeat, four estimators
|
| 87 |
-
achieved mean AUC 0.8581 and lower task error than default CatBoost on 76% of
|
| 88 |
-
tasks. This is a local evaluation, not an official TabArena submission, and does
|
| 89 |
-
not establish superiority over tuned GBDTs. Categorical-heavy datasets remain
|
| 90 |
-
a weakness; regression is unsupported.
|
| 91 |
-
|
| 92 |
-
The full experiment history and technical report are in the source repository.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/__init__.py
DELETED
|
@@ -1,25 +0,0 @@
|
|
| 1 |
-
"""LightPFN: compact tabular inference, pretrained only on synthetic data."""
|
| 2 |
-
|
| 3 |
-
__version__ = "0.1.0"
|
| 4 |
-
__all__ = ["Config", "LightPFN", "LightPFNClassifier", "load_model", "load_pretrained", "save_model", "__version__"]
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def __getattr__(name):
|
| 8 |
-
if name in ("Config", "LightPFN"):
|
| 9 |
-
from lightpfn.model import lightpfn
|
| 10 |
-
value = getattr(lightpfn, name)
|
| 11 |
-
elif name == "LightPFNClassifier":
|
| 12 |
-
try:
|
| 13 |
-
from lightpfn.sklearn import LightPFNClassifier
|
| 14 |
-
except ModuleNotFoundError as exc:
|
| 15 |
-
if exc.name != "sklearn":
|
| 16 |
-
raise
|
| 17 |
-
raise ImportError("LightPFNClassifier requires `pip install lightpfn[sklearn]`.") from exc
|
| 18 |
-
value = LightPFNClassifier
|
| 19 |
-
elif name in ("load_model", "load_pretrained", "save_model"):
|
| 20 |
-
from lightpfn import checkpoint
|
| 21 |
-
value = getattr(checkpoint, name)
|
| 22 |
-
else:
|
| 23 |
-
raise AttributeError(name)
|
| 24 |
-
globals()[name] = value
|
| 25 |
-
return value
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/checkpoint.py
DELETED
|
@@ -1,107 +0,0 @@
|
|
| 1 |
-
"""Safe local checkpoints and revision-pinned Hugging Face downloads.
|
| 2 |
-
|
| 3 |
-
The release format is model.safetensors plus config.json. Legacy tensor/dict
|
| 4 |
-
checkpoints remain supported using PyTorch's restricted weights-only loader.
|
| 5 |
-
There is deliberately no fallback to unrestricted pickle.
|
| 6 |
-
"""
|
| 7 |
-
|
| 8 |
-
import json
|
| 9 |
-
from collections.abc import Mapping
|
| 10 |
-
from dataclasses import asdict
|
| 11 |
-
from importlib.resources import files
|
| 12 |
-
from pathlib import Path
|
| 13 |
-
|
| 14 |
-
import torch
|
| 15 |
-
|
| 16 |
-
from lightpfn.model.lightpfn import Config, LightPFN
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
def _safetensors():
|
| 20 |
-
try:
|
| 21 |
-
from safetensors.torch import load_file, save_file
|
| 22 |
-
except ImportError as exc:
|
| 23 |
-
raise ImportError("Safetensors weights require `pip install lightpfn[hf]`.") from exc
|
| 24 |
-
return load_file, save_file
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
def _config(values):
|
| 28 |
-
if not isinstance(values, Mapping):
|
| 29 |
-
raise ValueError("Checkpoint config must be a mapping of architecture parameters.")
|
| 30 |
-
values = dict(values)
|
| 31 |
-
if "group_offsets" in values:
|
| 32 |
-
values["group_offsets"] = tuple(values["group_offsets"])
|
| 33 |
-
return Config(**values)
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
def load_model(path, device="cpu"):
|
| 37 |
-
"""Load a safetensors directory/file or a weights-only compatible .pt file.
|
| 38 |
-
|
| 39 |
-
Safetensors files need config.json in the same directory. Full training
|
| 40 |
-
checkpoints containing arbitrary Python objects are intentionally rejected.
|
| 41 |
-
"""
|
| 42 |
-
path = Path(path)
|
| 43 |
-
if path.is_dir():
|
| 44 |
-
path = path / "model.safetensors"
|
| 45 |
-
if path.suffix == ".safetensors":
|
| 46 |
-
load_file, _ = _safetensors()
|
| 47 |
-
config = json.loads(path.with_name("config.json").read_text(encoding="utf-8"))
|
| 48 |
-
state = load_file(str(path), device="cpu")
|
| 49 |
-
else:
|
| 50 |
-
checkpoint = torch.load(path, map_location="cpu", weights_only=True)
|
| 51 |
-
if not isinstance(checkpoint, Mapping) or "config" not in checkpoint:
|
| 52 |
-
raise ValueError("Checkpoint must contain 'config' and 'ema' or 'model' weights.")
|
| 53 |
-
config = checkpoint["config"]
|
| 54 |
-
state = checkpoint.get("ema", checkpoint.get("model"))
|
| 55 |
-
if not isinstance(state, Mapping) or not state or not all(
|
| 56 |
-
isinstance(key, str) and isinstance(value, torch.Tensor) for key, value in state.items()
|
| 57 |
-
):
|
| 58 |
-
raise ValueError("Checkpoint weights must be a nonempty tensor state dictionary.")
|
| 59 |
-
model = LightPFN(_config(config))
|
| 60 |
-
model.load_state_dict(state, strict=True)
|
| 61 |
-
return model.to(device).eval()
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
def save_model(model, directory):
|
| 65 |
-
"""Export unfused inference weights and JSON config, excluding optimizer state."""
|
| 66 |
-
if getattr(model, "is_folded", False):
|
| 67 |
-
raise ValueError("Export the original model; folded weights use a different layout.")
|
| 68 |
-
_, save_file = _safetensors()
|
| 69 |
-
directory = Path(directory)
|
| 70 |
-
directory.mkdir(parents=True, exist_ok=True)
|
| 71 |
-
state = {key: value.detach().cpu().contiguous().clone() for key, value in model.state_dict().items()}
|
| 72 |
-
save_file(state, str(directory / "model.safetensors"), metadata={"format": "pt"})
|
| 73 |
-
(directory / "config.json").write_text(json.dumps(asdict(model.cfg), indent=2) + "\n", encoding="utf-8")
|
| 74 |
-
return directory
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
def pretrained_spec():
|
| 78 |
-
"""Return the packaged model repository and immutable revision."""
|
| 79 |
-
return json.loads(files("lightpfn").joinpath("pretrained.json").read_text(encoding="utf-8"))
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
def load_pretrained(*, repo_id=None, revision=None, device="cpu", cache_dir=None, local_files_only=False):
|
| 83 |
-
"""Load the released model, or an explicit repo and revision using cached HF authentication.
|
| 84 |
-
|
| 85 |
-
A custom repository requires a revision. No remote Python code is executed.
|
| 86 |
-
For private repositories authenticate with `hf auth login` or HF_TOKEN.
|
| 87 |
-
"""
|
| 88 |
-
try:
|
| 89 |
-
from huggingface_hub import hf_hub_download
|
| 90 |
-
except ImportError as exc:
|
| 91 |
-
raise ImportError("Hugging Face downloads require `pip install lightpfn[hf]`.") from exc
|
| 92 |
-
spec = pretrained_spec()
|
| 93 |
-
repo_id = spec["repo_id"] if repo_id is None else repo_id
|
| 94 |
-
if revision is None:
|
| 95 |
-
if repo_id != spec["repo_id"]:
|
| 96 |
-
raise ValueError("Provide revision when using a custom Hugging Face repository.")
|
| 97 |
-
revision = spec["revision"]
|
| 98 |
-
if not revision:
|
| 99 |
-
raise ValueError("No pretrained revision configured; provide a checkpoint or explicit revision.")
|
| 100 |
-
options = dict(repo_id=repo_id, revision=revision, cache_dir=cache_dir, local_files_only=local_files_only)
|
| 101 |
-
config_path = Path(hf_hub_download(filename="config.json", **options))
|
| 102 |
-
# Use the resolved commit from the snapshot path for both files, even when a
|
| 103 |
-
# caller explicitly chose a mutable branch or tag.
|
| 104 |
-
resolved = config_path.parent.name
|
| 105 |
-
options["revision"] = resolved
|
| 106 |
-
weights_path = Path(hf_hub_download(filename="model.safetensors", **options))
|
| 107 |
-
return load_model(weights_path, device=device)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/device.py
DELETED
|
@@ -1,51 +0,0 @@
|
|
| 1 |
-
"""Inference device: "auto" takes torch's GPU backend when there is one (CUDA on NVIDIA, also ROCm builds of
|
| 2 |
-
torch, which use the same "cuda" device), otherwise a GPU with a Vulkan driver (AMD, Intel or NVIDIA,
|
| 3 |
-
through wgpu), otherwise the CPU. Each can be asked for explicitly: "cuda[:i]", "vulkan[:i]" (index into
|
| 4 |
-
lightpfn.vulkan.adapters()), "cpu". The environment variable LIGHTPFN_DEVICE replaces "auto" with any of
|
| 5 |
-
these, e.g. LIGHTPFN_DEVICE=cpu to keep a GPU free.
|
| 6 |
-
"""
|
| 7 |
-
|
| 8 |
-
import os
|
| 9 |
-
|
| 10 |
-
import torch
|
| 11 |
-
|
| 12 |
-
DEVICE_ENV = "LIGHTPFN_DEVICE"
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
def resolve_device(device="auto"):
|
| 16 |
-
"""The device string LightPFNClassifier runs on: "cuda[:i]", "vulkan[:i]", "cpu" (or "mps")."""
|
| 17 |
-
device = "auto" if device is None else str(device).strip().lower()
|
| 18 |
-
if device == "auto":
|
| 19 |
-
device = os.environ.get(DEVICE_ENV, "").strip().lower() or "auto"
|
| 20 |
-
if device == "auto":
|
| 21 |
-
if torch.cuda.is_available():
|
| 22 |
-
return "cuda"
|
| 23 |
-
from lightpfn import vulkan
|
| 24 |
-
|
| 25 |
-
return "vulkan" if vulkan.is_available() else "cpu"
|
| 26 |
-
kind, _, index = device.partition(":")
|
| 27 |
-
if ":" in device and not index:
|
| 28 |
-
raise ValueError(f"device {device!r}: missing index after ':'")
|
| 29 |
-
if index and not index.isdigit():
|
| 30 |
-
raise ValueError(f"device {device!r}: the index after ':' must be a number")
|
| 31 |
-
if kind == "cpu":
|
| 32 |
-
return "cpu"
|
| 33 |
-
if kind == "cuda":
|
| 34 |
-
if not torch.cuda.is_available():
|
| 35 |
-
raise RuntimeError("device 'cuda' requested, but this torch build sees no CUDA/ROCm GPU")
|
| 36 |
-
return device
|
| 37 |
-
if kind == "vulkan":
|
| 38 |
-
from lightpfn import vulkan
|
| 39 |
-
|
| 40 |
-
if not vulkan.is_available(int(index) if index else None):
|
| 41 |
-
raise RuntimeError(f"device {device!r} requested, but no Vulkan adapter was found: it needs a GPU "
|
| 42 |
-
"driver with Vulkan and `pip install wgpu` (adapters: lightpfn.vulkan.adapters())")
|
| 43 |
-
return device
|
| 44 |
-
if kind == "mps":
|
| 45 |
-
return device
|
| 46 |
-
raise ValueError(f"unknown device {device!r}: use 'auto', 'cpu', 'cuda[:i]' or 'vulkan[:i]'")
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
def kind(device):
|
| 50 |
-
""""cpu", "cuda", "vulkan" or "mps" of a resolved device string."""
|
| 51 |
-
return str(device).partition(":")[0]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/model/__init__.py
DELETED
|
@@ -1 +0,0 @@
|
|
| 1 |
-
"""LightPFN architecture; the inference package does not import synthetic priors."""
|
|
|
|
|
|
package/lightpfn/model/layers.py
DELETED
|
@@ -1,190 +0,0 @@
|
|
| 1 |
-
"""Building blocks shared by the three stages of the model.
|
| 2 |
-
|
| 3 |
-
Tensors are (batch, sequence, dim) unless a name says otherwise. Attention runs through
|
| 4 |
-
F.scaled_dot_product_attention, so the same code uses flash kernels on GPU and fused CPU kernels.
|
| 5 |
-
"""
|
| 6 |
-
|
| 7 |
-
import math
|
| 8 |
-
|
| 9 |
-
import torch
|
| 10 |
-
import torch.nn.functional as F
|
| 11 |
-
from torch import nn
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
class MLP(nn.Module):
|
| 15 |
-
"""Two-layer GELU feed-forward block; the output projection starts at zero so every residual
|
| 16 |
-
block starts as the identity (as in TabPFN-3)."""
|
| 17 |
-
|
| 18 |
-
def __init__(self, dim, hidden):
|
| 19 |
-
super().__init__()
|
| 20 |
-
self.fc1 = nn.Linear(dim, hidden, bias=False)
|
| 21 |
-
self.fc2 = nn.Linear(hidden, dim, bias=False)
|
| 22 |
-
nn.init.zeros_(self.fc2.weight)
|
| 23 |
-
|
| 24 |
-
def forward(self, x):
|
| 25 |
-
return self.fc2(F.gelu(self.fc1(x)))
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
class SoftmaxScaling(nn.Module):
|
| 29 |
-
"""Scales attention queries as a function of the number of keys n (TabPFN-3 SoftmaxScalingMLP):
|
| 30 |
-
q * base(log n) * (1 + tanh(mod(q))). Keeps attention sharp when the context grows far beyond
|
| 31 |
-
the lengths seen in training. base starts at 1 so the layer starts as a no-op."""
|
| 32 |
-
|
| 33 |
-
def __init__(self, n_heads, head_dim, hidden=64):
|
| 34 |
-
super().__init__()
|
| 35 |
-
self.n_heads, self.head_dim = n_heads, head_dim
|
| 36 |
-
self.base = nn.Sequential(nn.Linear(1, hidden), nn.GELU(), nn.Linear(hidden, n_heads * head_dim))
|
| 37 |
-
self.mod = nn.Sequential(nn.Linear(head_dim, hidden), nn.GELU(), nn.Linear(hidden, head_dim))
|
| 38 |
-
nn.init.zeros_(self.base[2].weight)
|
| 39 |
-
nn.init.ones_(self.base[2].bias)
|
| 40 |
-
nn.init.zeros_(self.mod[2].weight)
|
| 41 |
-
nn.init.zeros_(self.mod[2].bias)
|
| 42 |
-
|
| 43 |
-
def forward(self, q, n):
|
| 44 |
-
"""q: (B, H, L, D) queries, n: number of keys."""
|
| 45 |
-
logn = torch.full((1, 1), math.log(max(n, 2)), device=q.device, dtype=q.dtype)
|
| 46 |
-
base = self.base(logn).view(1, self.n_heads, 1, self.head_dim)
|
| 47 |
-
return q * base * (1 + torch.tanh(self.mod(q)))
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
# CUDA attention kernels put the batch on a grid dimension of at most 65535 blocks: a larger batch fails
|
| 51 |
-
# with "invalid argument" (the row stages of a 2-task group of 50k-row tables have 100k rows in the batch).
|
| 52 |
-
SDPA_MAX_BATCH = 32768
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
def sdpa(q, k, v, mask=None):
|
| 56 |
-
"""F.scaled_dot_product_attention in chunks of at most SDPA_MAX_BATCH along the batch dim (a size-1
|
| 57 |
-
batch dim of k, v or mask broadcasts)."""
|
| 58 |
-
B = q.shape[0]
|
| 59 |
-
if B <= SDPA_MAX_BATCH:
|
| 60 |
-
return F.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
| 61 |
-
out = []
|
| 62 |
-
for i in range(0, B, SDPA_MAX_BATCH):
|
| 63 |
-
part = [t if t is None or t.shape[0] == 1 else t[i:i + SDPA_MAX_BATCH] for t in (k, v, mask)]
|
| 64 |
-
out.append(F.scaled_dot_product_attention(q[i:i + SDPA_MAX_BATCH], *part[:2], attn_mask=part[2]))
|
| 65 |
-
return torch.cat(out)
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
def rope(x, pos, base=10000.0, interleaved=False):
|
| 69 |
-
"""Rotary position embedding (rotate-half form) on the last dim of x: (..., L, D) with
|
| 70 |
-
positions pos: (L,). interleaved=True rotates the pairs (2i, 2i + 1) instead of (i, i + D/2),
|
| 71 |
-
with one complex multiply: the same function for weights passed through interleave_rotary()."""
|
| 72 |
-
d = x.shape[-1]
|
| 73 |
-
inv_freq = base ** (-torch.arange(0, d, 2, device=x.device, dtype=torch.float32) / d)
|
| 74 |
-
ang = pos.to(torch.float32)[:, None] * inv_freq[None] # (L, D/2)
|
| 75 |
-
if interleaved and x.dtype in (torch.float32, torch.float64):
|
| 76 |
-
rot = torch.complex(ang.cos(), ang.sin()).to(torch.complex64 if x.dtype == torch.float32 else torch.complex128)
|
| 77 |
-
return torch.view_as_real(torch.view_as_complex(x.unflatten(-1, (d // 2, 2))) * rot).flatten(-2)
|
| 78 |
-
cos, sin = ang.cos().to(x.dtype), ang.sin().to(x.dtype)
|
| 79 |
-
if interleaved: # reduced precision (autocast): the same arithmetic as below on adjacent pairs
|
| 80 |
-
x1, x2 = x[..., 0::2], x[..., 1::2]
|
| 81 |
-
return torch.stack([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1).flatten(-2)
|
| 82 |
-
x1, x2 = x[..., : d // 2], x[..., d // 2 :]
|
| 83 |
-
return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1)
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
class Attention(nn.Module):
|
| 87 |
-
"""Multi-head attention with separate query and key/value inputs.
|
| 88 |
-
|
| 89 |
-
kv_heads_query < n_heads lets a set of queries use only the first key/value heads (multi-query
|
| 90 |
-
attention for test rows, as in TabPFN-3.5): keys/values of the other heads never need to be
|
| 91 |
-
cached for prediction."""
|
| 92 |
-
|
| 93 |
-
def __init__(self, dim, n_heads, scaling=False):
|
| 94 |
-
super().__init__()
|
| 95 |
-
assert dim % n_heads == 0
|
| 96 |
-
self.n_heads, self.head_dim = n_heads, dim // n_heads
|
| 97 |
-
self.q = nn.Linear(dim, dim, bias=False)
|
| 98 |
-
self.kv = nn.Linear(dim, 2 * dim, bias=False)
|
| 99 |
-
self.out = nn.Linear(dim, dim, bias=False)
|
| 100 |
-
nn.init.zeros_(self.out.weight)
|
| 101 |
-
self.scaling = SoftmaxScaling(n_heads, self.head_dim) if scaling else None
|
| 102 |
-
|
| 103 |
-
def split(self, x):
|
| 104 |
-
B, L, _ = x.shape
|
| 105 |
-
return x.view(B, L, self.n_heads, self.head_dim).transpose(1, 2) # (B, H, L, D)
|
| 106 |
-
|
| 107 |
-
def keys_values(self, x):
|
| 108 |
-
k, v = self.kv(x).chunk(2, dim=-1)
|
| 109 |
-
return self.split(k), self.split(v)
|
| 110 |
-
|
| 111 |
-
def attend(self, x, k, v, mask=None, q_rope=None, kv_heads=None):
|
| 112 |
-
"""x: (B, Lq, dim) queries; k, v: (B, Hkv, Lk, D). kv_heads=h: all query heads use the
|
| 113 |
-
first h key/value heads."""
|
| 114 |
-
q = self.split(self.q(x))
|
| 115 |
-
if q_rope is not None:
|
| 116 |
-
q = q_rope(q)
|
| 117 |
-
if self.scaling is not None:
|
| 118 |
-
q = self.scaling(q, k.shape[2])
|
| 119 |
-
if kv_heads is not None:
|
| 120 |
-
k = k[:, :kv_heads].repeat_interleave(self.n_heads // kv_heads, dim=1)
|
| 121 |
-
v = v[:, :kv_heads].repeat_interleave(self.n_heads // kv_heads, dim=1)
|
| 122 |
-
o = sdpa(q, k, v, mask)
|
| 123 |
-
B, H, L, D = o.shape
|
| 124 |
-
return self.out(o.transpose(1, 2).reshape(B, L, H * D))
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
class Block(nn.Module):
|
| 128 |
-
"""Pre-norm residual block: attention of x on a key/value sequence, then an MLP."""
|
| 129 |
-
|
| 130 |
-
def __init__(self, dim, n_heads, ff_factor=2, scaling=False, ff_hidden=None):
|
| 131 |
-
super().__init__()
|
| 132 |
-
self.norm_q = nn.RMSNorm(dim)
|
| 133 |
-
self.norm_kv = nn.RMSNorm(dim)
|
| 134 |
-
self.norm_ff = nn.RMSNorm(dim)
|
| 135 |
-
self.attn = Attention(dim, n_heads, scaling)
|
| 136 |
-
self.mlp = MLP(dim, dim * ff_factor if ff_hidden is None else ff_hidden)
|
| 137 |
-
|
| 138 |
-
def keys_values(self, kv_input):
|
| 139 |
-
return self.attn.keys_values(self.norm_kv(kv_input))
|
| 140 |
-
|
| 141 |
-
def forward(self, x, k, v, mask=None, q_rope=None, kv_heads=None):
|
| 142 |
-
x = x + self.attn.attend(self.norm_q(x), k, v, mask, q_rope, kv_heads)
|
| 143 |
-
return x + self.mlp(self.norm_ff(x))
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
class FoldedRMSNorm(nn.Module):
|
| 147 |
-
"""RMSNorm whose weight was folded into the linear layer that reads it (inference only): one
|
| 148 |
-
reduction and one multiply instead of the unfused CPU kernels. Same eps as nn.RMSNorm (eps=None:
|
| 149 |
-
that of the accumulation dtype, float32 for reduced precision), reduction in at least float32."""
|
| 150 |
-
|
| 151 |
-
def __init__(self, dim, eps=None):
|
| 152 |
-
super().__init__()
|
| 153 |
-
self.dim, self.eps = dim, eps
|
| 154 |
-
|
| 155 |
-
def forward(self, x):
|
| 156 |
-
acc = torch.promote_types(x.dtype, torch.float32)
|
| 157 |
-
eps = torch.finfo(acc).eps if self.eps is None else self.eps
|
| 158 |
-
ms = torch.linalg.vector_norm(x, dim=-1, keepdim=True, dtype=acc).square_().div_(self.dim)
|
| 159 |
-
return x * ms.add_(eps).rsqrt_().to(x.dtype)
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
@torch.no_grad()
|
| 163 |
-
def fold_norm(norm, *linears):
|
| 164 |
-
"""Moves the weight of an RMSNorm into the linear layers fed by it; returns the weightless norm."""
|
| 165 |
-
for lin in linears:
|
| 166 |
-
lin.weight.mul_(norm.weight)
|
| 167 |
-
return FoldedRMSNorm(norm.weight.numel(), norm.eps)
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
@torch.no_grad()
|
| 171 |
-
def interleave_rotary(block):
|
| 172 |
-
"""Reorders the query and key dims of each head of a block from (i, i + D/2) pairs to adjacent
|
| 173 |
-
(2i, 2i + 1) pairs: q.k is unchanged (same permutation on both), and rope(..., interleaved=True)
|
| 174 |
-
rotates the same pairs."""
|
| 175 |
-
a = block.attn
|
| 176 |
-
H, D = a.n_heads, a.head_dim
|
| 177 |
-
pairs = torch.stack([torch.arange(D // 2), torch.arange(D // 2) + D // 2], -1).flatten()
|
| 178 |
-
rows = (torch.arange(H)[:, None] * D + pairs).flatten().to(a.q.weight.device)
|
| 179 |
-
a.q.weight.copy_(a.q.weight[rows])
|
| 180 |
-
a.kv.weight[: H * D].copy_(a.kv.weight[rows])
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
class OrthogonalEmbedding(nn.Embedding):
|
| 184 |
-
"""Label embedding with orthonormal initial rows (TabPFN-3 TrainableOrthogonalEmbedding)."""
|
| 185 |
-
|
| 186 |
-
def __init__(self, n, dim):
|
| 187 |
-
super().__init__(n, dim)
|
| 188 |
-
with torch.no_grad():
|
| 189 |
-
q, _ = torch.linalg.qr(torch.randn(dim, min(n, dim)))
|
| 190 |
-
self.weight[: q.shape[1]] = q.T
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/model/lightpfn.py
DELETED
|
@@ -1,490 +0,0 @@
|
|
| 1 |
-
"""LightPFN classifier: cell embedding -> column stage -> row stage -> in-context learning -> decoder.
|
| 2 |
-
|
| 3 |
-
Design notes and sources are in docs/design.md. In short:
|
| 4 |
-
cells per-column statistics of the training rows turn each value into a z-score, an ECDF rank
|
| 5 |
-
and a NaN flag; cells are grouped with circular feature shifts and embedded with learnable
|
| 6 |
-
Fourier features (TabPFN-3.5), which avoids the low-rank collapse of a scalar linear
|
| 7 |
-
embedding (LimiX-2M)
|
| 8 |
-
column per-column ISAB whose inducing points attend to the training cells only, with the label
|
| 9 |
-
embedding added to the training cells (target-aware, TabICLv2)
|
| 10 |
-
row per-row transformer over the features with CLS tokens and RoPE; the CLS outputs are
|
| 11 |
-
concatenated into the row representation (TabICLv2)
|
| 12 |
-
ICL transformer over rows: training rows (plus learned thinking rows) attend to each other,
|
| 13 |
-
test rows attend to the training rows only, with multi-query attention (TabPFN-3.5) and
|
| 14 |
-
learned log-n softmax scaling (TabPFN-3)
|
| 15 |
-
decoder attention of test rows over the one-hot training labels (TabPFN-3 retrieval decoder)
|
| 16 |
-
|
| 17 |
-
`forward` runs `encode` (everything that depends on the training rows, returned as a Context) and
|
| 18 |
-
then `predict_logits` (test rows only), so the cached inference path is the training path.
|
| 19 |
-
"""
|
| 20 |
-
|
| 21 |
-
import copy
|
| 22 |
-
from dataclasses import dataclass, field
|
| 23 |
-
|
| 24 |
-
import torch
|
| 25 |
-
import torch.nn.functional as F
|
| 26 |
-
from torch import nn
|
| 27 |
-
|
| 28 |
-
from lightpfn.model.layers import Block, OrthogonalEmbedding, SoftmaxScaling, fold_norm, interleave_rotary, rope
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
@dataclass
|
| 32 |
-
class Config:
|
| 33 |
-
"""Opt-in architecture ablations; default parameters/state names are r2-exact.
|
| 34 |
-
|
| 35 |
-
Budget recipes at the default widths (trainable parameters, buffers excluded):
|
| 36 |
-
B1 cell_embed='rbf', icl_ff_reallocation=4: 4,777,024.
|
| 37 |
-
B2 ccmm=True, icl_ff_reallocation=5: 4,776,784.
|
| 38 |
-
B3 row_mode='summary': 4,777,072 (no extra parameters).
|
| 39 |
-
B4 row_refine=True, row_mode='summary', icl_drop_blocks=1: 4,603,088.
|
| 40 |
-
One last-MLP hidden unit costs 2*icl_dim parameters (512 at default width).
|
| 41 |
-
B4 uses R=3 independent rounds plus final broadcast: +372,032 parameters,
|
| 42 |
-
funded by removing one 546,016-parameter ICL block (8 -> 7). Other depths/
|
| 43 |
-
widths require an explicit budget choice; there is no automatic reallocation.
|
| 44 |
-
In particular, row_refine does not implicitly change row_mode or ICL depth.
|
| 45 |
-
"""
|
| 46 |
-
|
| 47 |
-
col_dim: int = 64
|
| 48 |
-
col_blocks: int = 2
|
| 49 |
-
col_heads: int = 4
|
| 50 |
-
n_inducing: int = 64
|
| 51 |
-
row_blocks: int = 3
|
| 52 |
-
row_heads: int = 4
|
| 53 |
-
n_cls: int = 4 # ICL width = n_cls * col_dim
|
| 54 |
-
icl_blocks: int = 8
|
| 55 |
-
icl_heads: int = 8
|
| 56 |
-
icl_kv_heads_test: int = 1
|
| 57 |
-
n_thinking: int = 16
|
| 58 |
-
ff_factor: int = 2
|
| 59 |
-
n_freq: int = 16
|
| 60 |
-
n_ecdf_freq: int = 4
|
| 61 |
-
group_offsets: tuple = (0, 1, 3)
|
| 62 |
-
label_slots: int = 16 # max classes; training maps classes to random slots so all are trained
|
| 63 |
-
decoder_heads: int = 4
|
| 64 |
-
rope_base: float = 10000.0
|
| 65 |
-
cell_embed: str = "fourier" # B1: "rbf", 64 fixed uniform kernels, sigma=1
|
| 66 |
-
ccmm: bool = False # B2: training-only mask token and rank-bin readout
|
| 67 |
-
ccmm_mask_fraction: float = 0.15 # observed test cells, per task; structured masks round up
|
| 68 |
-
row_mode: str = "self_attention" # B3: "summary"
|
| 69 |
-
row_refine: bool = False # B4: refinement -> independent column stage -> compression
|
| 70 |
-
row_refine_rounds: int = 3
|
| 71 |
-
icl_drop_blocks: int = 0 # explicit budget reallocation, e.g. 1 for B4
|
| 72 |
-
icl_ff_reallocation: int = 0 # hidden units removed from the LAST retained ICL MLP
|
| 73 |
-
|
| 74 |
-
def __post_init__(self):
|
| 75 |
-
if self.cell_embed not in ("fourier", "rbf"):
|
| 76 |
-
raise ValueError("cell_embed must be fourier or rbf")
|
| 77 |
-
if self.row_mode not in ("self_attention", "summary"):
|
| 78 |
-
raise ValueError("row_mode must be self_attention or summary")
|
| 79 |
-
if not 0 <= self.ccmm_mask_fraction <= 1:
|
| 80 |
-
raise ValueError("ccmm_mask_fraction must be in [0, 1]")
|
| 81 |
-
if self.row_refine_rounds < 1:
|
| 82 |
-
raise ValueError("row_refine_rounds must be positive")
|
| 83 |
-
if not 0 <= self.icl_drop_blocks < self.icl_blocks:
|
| 84 |
-
raise ValueError("icl_drop_blocks must leave at least one ICL block")
|
| 85 |
-
if not 0 <= self.icl_ff_reallocation < self.icl_dim * self.ff_factor:
|
| 86 |
-
raise ValueError("icl_ff_reallocation must leave a nonempty ICL MLP")
|
| 87 |
-
|
| 88 |
-
@property
|
| 89 |
-
def icl_dim(self):
|
| 90 |
-
return self.n_cls * self.col_dim
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
@dataclass
|
| 94 |
-
class Context:
|
| 95 |
-
"""Everything prediction needs from the training rows."""
|
| 96 |
-
|
| 97 |
-
stats: dict
|
| 98 |
-
col_kv: list
|
| 99 |
-
icl_kv: list
|
| 100 |
-
dec_k: torch.Tensor
|
| 101 |
-
y: torch.Tensor
|
| 102 |
-
slots: torch.Tensor
|
| 103 |
-
d: torch.Tensor | None
|
| 104 |
-
n_classes: int
|
| 105 |
-
extra: dict = field(default_factory=dict)
|
| 106 |
-
col_refine_kv: list | None = None
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
class CellEmbedder(nn.Module):
|
| 110 |
-
def __init__(self, cfg):
|
| 111 |
-
super().__init__()
|
| 112 |
-
G = len(cfg.group_offsets)
|
| 113 |
-
self.offsets = cfg.group_offsets
|
| 114 |
-
self.n_ecdf_freq = cfg.n_ecdf_freq
|
| 115 |
-
self.cell_embed = cfg.cell_embed
|
| 116 |
-
if cfg.cell_embed == "fourier":
|
| 117 |
-
self.freq = nn.Parameter(torch.randn(G, cfg.n_freq) * 2.0)
|
| 118 |
-
self.fourier = nn.Linear(2 * cfg.n_freq, cfg.col_dim, bias=False)
|
| 119 |
-
else:
|
| 120 |
-
# RaBEL's verified sweep favors 64 uniform kernels and fixed sigma=1.
|
| 121 |
-
# The range is adapted to LightPFN's already soft-clipped z scores.
|
| 122 |
-
self.register_buffer("rbf_centers", torch.linspace(-5.0, 5.0, 64))
|
| 123 |
-
self.rbf = nn.Linear(64, cfg.col_dim, bias=False)
|
| 124 |
-
self.meta = nn.Linear(G * (2 + 2 * cfg.n_ecdf_freq), cfg.col_dim, bias=False)
|
| 125 |
-
self.norm = nn.LayerNorm(cfg.col_dim)
|
| 126 |
-
|
| 127 |
-
@staticmethod
|
| 128 |
-
@torch.no_grad()
|
| 129 |
-
def stats(x_train):
|
| 130 |
-
"""Column statistics of the training rows (NaN ignored). x_train: (B, n, m)."""
|
| 131 |
-
x = x_train.float()
|
| 132 |
-
valid = ~torch.isnan(x)
|
| 133 |
-
cnt = valid.sum(1) # (B, m)
|
| 134 |
-
x0 = torch.where(valid, x, 0.0)
|
| 135 |
-
mean = x0.sum(1) / cnt.clamp(min=1)
|
| 136 |
-
var = (torch.where(valid, x - mean[:, None], 0.0) ** 2).sum(1) / (cnt - 1).clamp(min=1)
|
| 137 |
-
std = var.sqrt()
|
| 138 |
-
srt = torch.where(valid, x, float("inf")).transpose(1, 2).sort(-1).values.contiguous() # (B, m, n)
|
| 139 |
-
return dict(mean=mean, std=std, sorted=srt, cnt=cnt)
|
| 140 |
-
|
| 141 |
-
@staticmethod
|
| 142 |
-
def normalize(x, st):
|
| 143 |
-
"""Returns z-score (soft-clipped), ECDF mid-rank in [0, 1] and NaN flag, each (B, n, m)."""
|
| 144 |
-
x = x.float()
|
| 145 |
-
nan = torch.isnan(x)
|
| 146 |
-
z = (x - st["mean"][:, None]) / (st["std"][:, None] + 1e-6)
|
| 147 |
-
z = torch.where(nan | (st["std"][:, None] == 0), 0.0, 5.0 * torch.tanh(z / 5.0))
|
| 148 |
-
q = torch.where(nan, 0.0, x).transpose(1, 2).contiguous() # (B, m, n)
|
| 149 |
-
lo = torch.searchsorted(st["sorted"], q, right=False)
|
| 150 |
-
hi = torch.searchsorted(st["sorted"], q, right=True)
|
| 151 |
-
cnt = st["cnt"][..., None]
|
| 152 |
-
r = ((lo + hi).float() / 2 / cnt.clamp(min=1)).clamp(0, 1)
|
| 153 |
-
r = torch.where(cnt > 0, r, 0.5).transpose(1, 2)
|
| 154 |
-
r = torch.where(nan, 0.5, r)
|
| 155 |
-
return z, r, nan.float()
|
| 156 |
-
|
| 157 |
-
def group_index(self, d, B, m, device):
|
| 158 |
-
j = torch.arange(m, device=device)[None]
|
| 159 |
-
d = torch.full((B, 1), m, device=device) if d is None else d.to(device)[:, None]
|
| 160 |
-
return torch.stack([torch.where(j < d, (j + o) % d, j) for o in self.offsets], -1) # (B, m, G)
|
| 161 |
-
|
| 162 |
-
def grouped(self, x, st, d=None, mask=None, mask_token=None):
|
| 163 |
-
"""z, ECDF rank and NaN flag of every cell and of its group neighbors, each (B, n, m, G)."""
|
| 164 |
-
B, n, m = x.shape
|
| 165 |
-
if mask is None:
|
| 166 |
-
z, r, nan = self.normalize(x, st)
|
| 167 |
-
else:
|
| 168 |
-
if mask.shape != x.shape or mask.dtype != torch.bool or mask_token is None:
|
| 169 |
-
raise ValueError("cell masking needs a bool mask matching x and a mask token")
|
| 170 |
-
# Neutralize BEFORE circular grouping: otherwise a hidden value leaks into
|
| 171 |
-
# the embeddings of its neighbors through the grouped z/ECDF/NaN features.
|
| 172 |
-
z, r, nan = self.normalize(x.masked_fill(mask, 0.0), st)
|
| 173 |
-
z, r, nan = (t.masked_fill(mask, 0.0) for t in (z, r, nan))
|
| 174 |
-
idx = self.group_index(d, B, m, x.device)
|
| 175 |
-
G = idx.shape[-1]
|
| 176 |
-
gidx = idx.view(B, 1, m * G).expand(B, n, m * G)
|
| 177 |
-
return tuple(torch.gather(t, 2, gidx).view(B, n, m, G) for t in (z, r, nan))
|
| 178 |
-
|
| 179 |
-
def forward(self, x, st, d=None, mask=None, mask_token=None):
|
| 180 |
-
out = self.embed(*self.grouped(x, st, d, mask, mask_token))
|
| 181 |
-
if mask is not None:
|
| 182 |
-
out = torch.where(mask[..., None], mask_token.to(out.dtype), out)
|
| 183 |
-
return out
|
| 184 |
-
|
| 185 |
-
def embed(self, z, r, nan):
|
| 186 |
-
"""Cell embeddings (B, n, m, E) from grouped features; cells are independent, so any slice
|
| 187 |
-
of rows or columns of the features gives the same slice of the embeddings."""
|
| 188 |
-
if self.cell_embed == "fourier":
|
| 189 |
-
ang = z[..., None] * self.freq # (B, n, m, G, F), fp32
|
| 190 |
-
four = torch.cat([ang.sin(), ang.cos()], -1).sum(-2) # (B, n, m, 2F)
|
| 191 |
-
value = self.fourier(four)
|
| 192 |
-
else:
|
| 193 |
-
kernels = torch.exp(-0.5 * (z[..., None] - self.rbf_centers).square())
|
| 194 |
-
value = self.rbf(kernels.sum(-2))
|
| 195 |
-
k = torch.pi * 2.0 ** torch.arange(self.n_ecdf_freq, device=z.device)
|
| 196 |
-
rang = r[..., None] * k
|
| 197 |
-
meta = torch.cat([z, nan, rang.sin().flatten(-2), rang.cos().flatten(-2)], -1)
|
| 198 |
-
return self.norm(value + self.meta(meta)) # (B, n, m, E)
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
class ColumnStage(nn.Module):
|
| 202 |
-
"""ISAB over the rows of each column. Inducing points read the training cells only, so their
|
| 203 |
-
states (col_kv) are all that test rows need."""
|
| 204 |
-
|
| 205 |
-
def __init__(self, cfg):
|
| 206 |
-
super().__init__()
|
| 207 |
-
E = cfg.col_dim
|
| 208 |
-
self.y_emb = OrthogonalEmbedding(cfg.label_slots, E)
|
| 209 |
-
self.inducing = nn.ParameterList(nn.Parameter(torch.randn(cfg.n_inducing, E) * 0.02) for _ in range(cfg.col_blocks))
|
| 210 |
-
self.ind_blocks = nn.ModuleList(Block(E, cfg.col_heads, cfg.ff_factor) for _ in range(cfg.col_blocks))
|
| 211 |
-
self.cell_blocks = nn.ModuleList(Block(E, cfg.col_heads, cfg.ff_factor) for _ in range(cfg.col_blocks))
|
| 212 |
-
|
| 213 |
-
def context(self, cells, y_slots):
|
| 214 |
-
B, n, m, E = cells.shape
|
| 215 |
-
x = (cells + self.y_emb(y_slots)[:, :, None]).transpose(1, 2).reshape(B * m, n, E)
|
| 216 |
-
col_kv = []
|
| 217 |
-
for ind, ib, cb in zip(self.inducing, self.ind_blocks, self.cell_blocks):
|
| 218 |
-
h = ib(ind.expand(B * m, -1, -1), *ib.keys_values(x))
|
| 219 |
-
kh, vh = cb.keys_values(h)
|
| 220 |
-
x = cb(x, kh, vh)
|
| 221 |
-
col_kv.append((kh, vh))
|
| 222 |
-
return x.view(B, m, n, E).transpose(1, 2), col_kv
|
| 223 |
-
|
| 224 |
-
def query(self, cells, col_kv):
|
| 225 |
-
B, n, m, E = cells.shape
|
| 226 |
-
x = cells.transpose(1, 2).reshape(B * m, n, E)
|
| 227 |
-
for cb, (kh, vh) in zip(self.cell_blocks, col_kv):
|
| 228 |
-
x = cb(x, kh, vh)
|
| 229 |
-
return x.view(B, m, n, E).transpose(1, 2)
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
class RowStage(nn.Module):
|
| 233 |
-
"""Transformer over the features of each row; CLS tokens (no rotation) collect the row."""
|
| 234 |
-
|
| 235 |
-
def __init__(self, cfg):
|
| 236 |
-
super().__init__()
|
| 237 |
-
self.n_cls, self.rope_base = cfg.n_cls, cfg.rope_base
|
| 238 |
-
self.rope_interleaved = False # set by LightPFN.folded()
|
| 239 |
-
self.mode = cfg.row_mode
|
| 240 |
-
self.cls = nn.Parameter(torch.randn(cfg.n_cls, cfg.col_dim) * 0.02)
|
| 241 |
-
self.blocks = nn.ModuleList(Block(cfg.col_dim, cfg.row_heads, cfg.ff_factor) for _ in range(cfg.row_blocks))
|
| 242 |
-
|
| 243 |
-
def _rope(self, t):
|
| 244 |
-
C = self.n_cls
|
| 245 |
-
pos = torch.arange(t.shape[2] - C, device=t.device)
|
| 246 |
-
return torch.cat([t[:, :, :C], rope(t[:, :, C:], pos, self.rope_base, self.rope_interleaved)], dim=2)
|
| 247 |
-
|
| 248 |
-
def forward(self, cells, d=None, has_padding=None, return_cells=False):
|
| 249 |
-
B, n, m, E = cells.shape
|
| 250 |
-
C = self.n_cls
|
| 251 |
-
x = torch.cat([self.cls.to(cells.dtype).expand(B * n, C, E), cells.reshape(B * n, m, E)], dim=1)
|
| 252 |
-
mask = None
|
| 253 |
-
# Training supplies the exact decision from CPU metadata. Other callers
|
| 254 |
-
# retain the original automatic padding detection.
|
| 255 |
-
if d is not None and (bool((d < m).any()) if has_padding is None else has_padding):
|
| 256 |
-
valid = torch.arange(m, device=cells.device)[None] < d.to(cells.device)[:, None] # (B, m)
|
| 257 |
-
valid = torch.cat([torch.ones(B, C, dtype=torch.bool, device=cells.device), valid], 1)
|
| 258 |
-
mask = valid.repeat_interleave(n, 0)[:, None, None, :] # (B*n, 1, 1, C+m)
|
| 259 |
-
for blk in self.blocks:
|
| 260 |
-
if self.mode == "self_attention":
|
| 261 |
-
k, v = blk.keys_values(x)
|
| 262 |
-
x = blk(x, self._rope(k), v, mask, q_rope=self._rope)
|
| 263 |
-
else:
|
| 264 |
-
# Only C queries: cells stay fixed; summaries also attend to each other.
|
| 265 |
-
k, v = blk.keys_values(x)
|
| 266 |
-
summary = blk(x[:, :C], self._rope(k), v, mask)
|
| 267 |
-
x = torch.cat([summary, x[:, C:]], dim=1)
|
| 268 |
-
rows = x[:, :C].reshape(B, n, C * E)
|
| 269 |
-
if return_cells:
|
| 270 |
-
return rows, x[:, C:].reshape(B, n, m, E)
|
| 271 |
-
return rows
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
class RowRefinement(nn.Module):
|
| 275 |
-
"""Temporary summaries, independent broadcast/gather weights per round, final broadcast.
|
| 276 |
-
|
| 277 |
-
Each row is processed independently. Feature attention is O(m*K), K=4;
|
| 278 |
-
summaries are discarded before the second column stage and final compression.
|
| 279 |
-
"""
|
| 280 |
-
|
| 281 |
-
def __init__(self, cfg):
|
| 282 |
-
super().__init__()
|
| 283 |
-
self.rope_base = cfg.rope_base
|
| 284 |
-
self.rope_interleaved = False # set by LightPFN.folded()
|
| 285 |
-
self.summary = nn.Parameter(torch.randn(4, cfg.col_dim) * 0.02)
|
| 286 |
-
self.broadcast = nn.ModuleList(
|
| 287 |
-
Block(cfg.col_dim, cfg.row_heads, cfg.ff_factor) for _ in range(cfg.row_refine_rounds + 1))
|
| 288 |
-
self.gather = nn.ModuleList(
|
| 289 |
-
Block(cfg.col_dim, cfg.row_heads, cfg.ff_factor) for _ in range(cfg.row_refine_rounds))
|
| 290 |
-
|
| 291 |
-
def forward(self, cells, d=None, has_padding=None):
|
| 292 |
-
B, n, m, E = cells.shape
|
| 293 |
-
x = cells.reshape(B * n, m, E)
|
| 294 |
-
s = self.summary.to(cells.dtype).expand(B * n, -1, -1)
|
| 295 |
-
mask = None
|
| 296 |
-
if d is not None and (bool((d < m).any()) if has_padding is None else has_padding):
|
| 297 |
-
valid = torch.arange(m, device=cells.device)[None] < d.to(cells.device)[:, None]
|
| 298 |
-
mask = valid.repeat_interleave(n, 0)[:, None, None, :]
|
| 299 |
-
pos = torch.arange(m, device=cells.device)
|
| 300 |
-
for i, broadcast in enumerate(self.broadcast):
|
| 301 |
-
x = broadcast(x, *broadcast.keys_values(s),
|
| 302 |
-
q_rope=lambda q: rope(q, pos, self.rope_base, self.rope_interleaved))
|
| 303 |
-
if i < len(self.gather):
|
| 304 |
-
gather = self.gather[i]
|
| 305 |
-
k, v = gather.keys_values(x)
|
| 306 |
-
s = gather(s, rope(k, pos, self.rope_base, self.rope_interleaved), v, mask)
|
| 307 |
-
return x.reshape(B, n, m, E)
|
| 308 |
-
|
| 309 |
-
|
| 310 |
-
class ICLStage(nn.Module):
|
| 311 |
-
def __init__(self, cfg):
|
| 312 |
-
super().__init__()
|
| 313 |
-
D = cfg.icl_dim
|
| 314 |
-
self.kv_heads_test = cfg.icl_kv_heads_test
|
| 315 |
-
self.y_emb = OrthogonalEmbedding(cfg.label_slots, D)
|
| 316 |
-
self.thinking = nn.Parameter(torch.randn(cfg.n_thinking, D) * 0.02)
|
| 317 |
-
depth = cfg.icl_blocks - cfg.icl_drop_blocks
|
| 318 |
-
self.blocks = nn.ModuleList(
|
| 319 |
-
Block(D, cfg.icl_heads, cfg.ff_factor, scaling=True,
|
| 320 |
-
ff_hidden=D * cfg.ff_factor - (cfg.icl_ff_reallocation if i == depth - 1 else 0))
|
| 321 |
-
for i in range(depth))
|
| 322 |
-
self.norm = nn.RMSNorm(D)
|
| 323 |
-
|
| 324 |
-
def context(self, r, y_slots):
|
| 325 |
-
B = r.shape[0]
|
| 326 |
-
T = self.thinking.shape[0]
|
| 327 |
-
x = torch.cat([self.thinking.to(r.dtype).expand(B, -1, -1), r + self.y_emb(y_slots)], dim=1)
|
| 328 |
-
icl_kv = []
|
| 329 |
-
for blk in self.blocks:
|
| 330 |
-
k, v = blk.keys_values(x)
|
| 331 |
-
x = blk(x, k, v)
|
| 332 |
-
if not self.training: # the cache keeps only the heads test rows read
|
| 333 |
-
k, v = k[:, : self.kv_heads_test].contiguous(), v[:, : self.kv_heads_test].contiguous()
|
| 334 |
-
icl_kv.append((k, v))
|
| 335 |
-
return self.norm(x[:, T:]), icl_kv
|
| 336 |
-
|
| 337 |
-
def query(self, r, icl_kv):
|
| 338 |
-
x = r
|
| 339 |
-
for blk, (k, v) in zip(self.blocks, icl_kv):
|
| 340 |
-
x = blk(x, k, v, kv_heads=self.kv_heads_test)
|
| 341 |
-
return self.norm(x)
|
| 342 |
-
|
| 343 |
-
|
| 344 |
-
class RetrievalDecoder(nn.Module):
|
| 345 |
-
"""p(class | test row) = attention-weighted average of the one-hot training labels, averaged
|
| 346 |
-
over heads; logits are its log. Any number of classes up to the head dim."""
|
| 347 |
-
|
| 348 |
-
def __init__(self, cfg):
|
| 349 |
-
super().__init__()
|
| 350 |
-
D, H = cfg.icl_dim, cfg.decoder_heads
|
| 351 |
-
self.n_heads, self.head_dim = H, D // H
|
| 352 |
-
assert cfg.label_slots <= self.head_dim
|
| 353 |
-
self.q = nn.Linear(D, D, bias=False)
|
| 354 |
-
self.k = nn.Linear(D, D, bias=False)
|
| 355 |
-
self.scaling = SoftmaxScaling(H, self.head_dim)
|
| 356 |
-
|
| 357 |
-
def keys(self, h_train):
|
| 358 |
-
B, N, _ = h_train.shape
|
| 359 |
-
return self.k(h_train).view(B, N, self.n_heads, self.head_dim).transpose(1, 2)
|
| 360 |
-
|
| 361 |
-
def forward(self, k, y, h_test, n_classes):
|
| 362 |
-
B, M, _ = h_test.shape
|
| 363 |
-
q = self.q(h_test).view(B, M, self.n_heads, self.head_dim).transpose(1, 2)
|
| 364 |
-
q = self.scaling(q, k.shape[2])
|
| 365 |
-
# one-hot values padded to the head dim so fused attention kernels apply
|
| 366 |
-
v = F.one_hot(y, self.head_dim).to(q.dtype)[:, None].expand(-1, self.n_heads, -1, -1)
|
| 367 |
-
p = F.scaled_dot_product_attention(q, k.to(q.dtype), v).float().mean(1)[..., :n_classes]
|
| 368 |
-
return torch.log(p.clamp(min=1e-5) + 3e-5)
|
| 369 |
-
|
| 370 |
-
|
| 371 |
-
class LightPFN(nn.Module):
|
| 372 |
-
def __init__(self, cfg=None):
|
| 373 |
-
super().__init__()
|
| 374 |
-
self.cfg = cfg = cfg or Config()
|
| 375 |
-
self.cells = CellEmbedder(cfg)
|
| 376 |
-
self.col = ColumnStage(cfg)
|
| 377 |
-
self.row = RowStage(cfg)
|
| 378 |
-
self.icl = ICLStage(cfg)
|
| 379 |
-
self.decoder = RetrievalDecoder(cfg)
|
| 380 |
-
if cfg.row_refine:
|
| 381 |
-
self.refine = RowRefinement(cfg)
|
| 382 |
-
self.col_refine = ColumnStage(cfg) # independent, target-aware, train-only cache
|
| 383 |
-
if cfg.ccmm:
|
| 384 |
-
self.ccmm_mask_token = nn.Parameter(torch.randn(cfg.col_dim) * 0.02)
|
| 385 |
-
self.ccmm_head = nn.Sequential(nn.LayerNorm(cfg.col_dim), nn.Linear(cfg.col_dim, 32))
|
| 386 |
-
|
| 387 |
-
@torch.no_grad()
|
| 388 |
-
def folded(self):
|
| 389 |
-
"""A copy for prediction only that computes the same function faster: RMSNorm weights folded
|
| 390 |
-
into the linear layers that read them, and the query/key dims of the rotary blocks interleaved
|
| 391 |
-
so RoPE is one complex multiply. Outputs match the original up to float rounding; the copy is
|
| 392 |
-
not meant for training or for saving as a checkpoint."""
|
| 393 |
-
if getattr(self, "is_folded", False):
|
| 394 |
-
return self
|
| 395 |
-
m = copy.deepcopy(self).eval()
|
| 396 |
-
m.is_folded = True
|
| 397 |
-
for blk in (b for b in m.modules() if isinstance(b, Block)):
|
| 398 |
-
blk.norm_q = fold_norm(blk.norm_q, blk.attn.q)
|
| 399 |
-
blk.norm_kv = fold_norm(blk.norm_kv, blk.attn.kv)
|
| 400 |
-
blk.norm_ff = fold_norm(blk.norm_ff, blk.mlp.fc1)
|
| 401 |
-
m.icl.norm = fold_norm(m.icl.norm, m.decoder.q, m.decoder.k) # its output only feeds the decoder
|
| 402 |
-
rotary = [m.row] + ([m.refine] if self.cfg.row_refine else [])
|
| 403 |
-
for stage in rotary:
|
| 404 |
-
for blk in (b for b in stage.modules() if isinstance(b, Block)):
|
| 405 |
-
interleave_rotary(blk)
|
| 406 |
-
stage.rope_interleaved = True
|
| 407 |
-
return m
|
| 408 |
-
|
| 409 |
-
def encode(self, X_train, y_train, d=None, slots=None, n_classes=None, has_padding=None, chunk_cells=None):
|
| 410 |
-
"""X_train: (B, n, m) float with NaN, y_train: (B, n) labels 0..C-1, d: (B,) true feature
|
| 411 |
-
counts when features are zero-padded, slots: (B, label_slots) class -> label slot map.
|
| 412 |
-
n_classes: total number of classes C. The default, max(y_train) + 1, is wrong when the
|
| 413 |
-
highest classes have no training row: callers that know C must pass it (the sklearn
|
| 414 |
-
wrapper and the training loop do). chunk_cells (inference): run the cell stages on about
|
| 415 |
-
that many cells at a time, which keeps them in the CPU cache; the result is the same."""
|
| 416 |
-
B = X_train.shape[0]
|
| 417 |
-
n_classes = int(y_train.max()) + 1 if n_classes is None else n_classes
|
| 418 |
-
if n_classes > self.cfg.label_slots:
|
| 419 |
-
raise ValueError(f"{n_classes} classes, the model supports {self.cfg.label_slots}")
|
| 420 |
-
if slots is None:
|
| 421 |
-
slots = torch.arange(self.cfg.label_slots, device=X_train.device).expand(B, -1)
|
| 422 |
-
y_slots = torch.gather(slots, 1, y_train)
|
| 423 |
-
stats = self.cells.stats(X_train)
|
| 424 |
-
if chunk_cells is not None:
|
| 425 |
-
if torch.is_grad_enabled():
|
| 426 |
-
raise RuntimeError("chunk_cells is for inference: it writes into a shared buffer that autograd cannot track")
|
| 427 |
-
rows, col_kv, col_refine_kv = self._train_rows_chunked(X_train, stats, y_slots, d, has_padding, chunk_cells)
|
| 428 |
-
else:
|
| 429 |
-
col, col_kv = self.col.context(self.cells(X_train, stats, d), y_slots)
|
| 430 |
-
col_refine_kv = None
|
| 431 |
-
if self.cfg.row_refine:
|
| 432 |
-
col = self.refine(col, d, has_padding)
|
| 433 |
-
col, col_refine_kv = self.col_refine.context(col, y_slots)
|
| 434 |
-
rows = self.row(col, d, has_padding)
|
| 435 |
-
h, icl_kv = self.icl.context(rows, y_slots)
|
| 436 |
-
return Context(stats, col_kv, icl_kv, self.decoder.keys(h), y_train, slots, d, n_classes,
|
| 437 |
-
extra=dict(has_padding=has_padding), col_refine_kv=col_refine_kv)
|
| 438 |
-
|
| 439 |
-
def _train_rows_chunked(self, X, stats, y_slots, d, has_padding, chunk_cells):
|
| 440 |
-
"""The cell stages of encode by pieces: column stages on groups of whole columns, row stages
|
| 441 |
-
on groups of whole rows, all writing into one (B, n, m, E) buffer."""
|
| 442 |
-
B, n, m = X.shape
|
| 443 |
-
grouped = self.cells.grouped(X, stats, d)
|
| 444 |
-
cols, rows = max(1, chunk_cells // (B * n)), max(1, chunk_cells // (B * m))
|
| 445 |
-
buf = None # (B, n, m, E), allocated with the dtype of the first column-stage output
|
| 446 |
-
|
| 447 |
-
def column_stage(stage, cells_of):
|
| 448 |
-
nonlocal buf
|
| 449 |
-
parts = []
|
| 450 |
-
for j in range(0, m, cols):
|
| 451 |
-
out, kv = stage.context(cells_of(j, j + cols), y_slots)
|
| 452 |
-
if buf is None:
|
| 453 |
-
buf = out.new_empty(B, n, m, out.shape[-1])
|
| 454 |
-
buf[:, :, j : j + cols] = out
|
| 455 |
-
parts.append(kv)
|
| 456 |
-
# (B * columns, ...) caches of each piece, merged in the (batch, column) order of one call
|
| 457 |
-
return [tuple(torch.cat([p[i][t].unflatten(0, (B, -1)) for p in parts], 1).flatten(0, 1) for t in (0, 1))
|
| 458 |
-
for i in range(len(parts[0]))]
|
| 459 |
-
|
| 460 |
-
col_kv = column_stage(self.col, lambda a, b: self.cells.embed(*(t[:, :, a:b] for t in grouped)))
|
| 461 |
-
col_refine_kv = None
|
| 462 |
-
if self.cfg.row_refine:
|
| 463 |
-
for i in range(0, n, rows):
|
| 464 |
-
buf[:, i : i + rows] = self.refine(buf[:, i : i + rows], d, has_padding)
|
| 465 |
-
col_refine_kv = column_stage(self.col_refine, lambda a, b: buf[:, :, a:b])
|
| 466 |
-
out = torch.cat([self.row(buf[:, i : i + rows], d, has_padding) for i in range(0, n, rows)], 1)
|
| 467 |
-
return out, col_kv, col_refine_kv
|
| 468 |
-
|
| 469 |
-
def _test_rows(self, ctx, X_test):
|
| 470 |
-
col = self.col.query(self.cells(X_test, ctx.stats, ctx.d), ctx.col_kv)
|
| 471 |
-
if self.cfg.row_refine:
|
| 472 |
-
col = self.refine(col, ctx.d, ctx.extra.get("has_padding"))
|
| 473 |
-
col = self.col_refine.query(col, ctx.col_refine_kv)
|
| 474 |
-
return self.row(col, ctx.d, ctx.extra.get("has_padding"))
|
| 475 |
-
|
| 476 |
-
def predict_logits(self, ctx, X_test, chunk_cells=None):
|
| 477 |
-
"""Logits (B, M, n_classes) for test rows; rows are independent, so X_test can be chunked.
|
| 478 |
-
chunk_cells: run the cell stages on groups of about that many cells (same result)."""
|
| 479 |
-
if chunk_cells is None:
|
| 480 |
-
rows = self._test_rows(ctx, X_test)
|
| 481 |
-
else:
|
| 482 |
-
step = max(1, chunk_cells // (X_test.shape[0] * X_test.shape[2]))
|
| 483 |
-
rows = torch.cat([self._test_rows(ctx, X_test[:, i : i + step]) for i in range(0, X_test.shape[1], step)], 1)
|
| 484 |
-
h = self.icl.query(rows, ctx.icl_kv)
|
| 485 |
-
return self.decoder(ctx.dec_k, ctx.y, h, ctx.n_classes)
|
| 486 |
-
|
| 487 |
-
def forward(self, X, y_train, d=None, slots=None, n_classes=None, has_padding=None):
|
| 488 |
-
n_train = y_train.shape[1]
|
| 489 |
-
ctx = self.encode(X[:, :n_train], y_train, d, slots, n_classes, has_padding)
|
| 490 |
-
return self.predict_logits(ctx, X[:, n_train:])
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/pretrained.json
DELETED
|
@@ -1,4 +0,0 @@
|
|
| 1 |
-
{
|
| 2 |
-
"repo_id": "ueuegio/LightPFN",
|
| 3 |
-
"revision": "bd389ab59a89dd0e05c9ecb7c642c08ee52e9637"
|
| 4 |
-
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/sklearn.py
DELETED
|
@@ -1,280 +0,0 @@
|
|
| 1 |
-
"""scikit-learn style wrapper: fit() encodes the training set once (column statistics, inducing
|
| 2 |
-
states, ICL key/value cache), predict_proba() only runs the test rows, in chunks.
|
| 3 |
-
|
| 4 |
-
Estimators beyond the first use a random feature order and a random class -> label-slot map;
|
| 5 |
-
their probabilities are averaged. Training sets larger than max_context are subsampled per
|
| 6 |
-
estimator with stratification, so every class (even one with a single row) stays in the context.
|
| 7 |
-
"""
|
| 8 |
-
|
| 9 |
-
import copy
|
| 10 |
-
import warnings
|
| 11 |
-
from numbers import Integral
|
| 12 |
-
|
| 13 |
-
import numpy as np
|
| 14 |
-
import torch
|
| 15 |
-
from sklearn.base import BaseEstimator, ClassifierMixin
|
| 16 |
-
from sklearn.utils.multiclass import check_classification_targets
|
| 17 |
-
from sklearn.utils.validation import check_is_fitted, validate_data
|
| 18 |
-
|
| 19 |
-
from lightpfn.checkpoint import load_model, load_pretrained
|
| 20 |
-
from lightpfn.device import kind, resolve_device
|
| 21 |
-
from lightpfn.model.lightpfn import Config, LightPFN
|
| 22 |
-
|
| 23 |
-
# Inference pieces, in cells (rows x features, times the estimators of a batch); measured on an 8-core
|
| 24 |
-
# CPU and an RTX 5090 (runs/bench_infer/).
|
| 25 |
-
CPU_CHUNK_CELLS_PER_THREAD = 1024
|
| 26 |
-
GPU_CHUNK_CELLS = 1 << 20
|
| 27 |
-
CPU_BATCH_CELLS = 0 # one estimator at a time: batching saves nothing once pieces fit the cache
|
| 28 |
-
GPU_BATCH_CELLS = 1 << 22
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
def stratified_subsample(rng, y, n, min_per_class=5):
|
| 32 |
-
"""Exactly n of the len(y) > n rows: class proportions kept, and min(count, min_per_class) rows
|
| 33 |
-
of every class, taken from the largest classes (with fewer classes guaranteed if the minimums
|
| 34 |
-
alone would exceed n)."""
|
| 35 |
-
classes, counts = np.unique(y, return_counts=True)
|
| 36 |
-
if n < len(classes):
|
| 37 |
-
raise ValueError("max_context must be at least the number of classes.")
|
| 38 |
-
mins = np.minimum(counts, min_per_class)
|
| 39 |
-
if mins.sum() > n:
|
| 40 |
-
mins = np.minimum(counts, max(1, n // len(classes)))
|
| 41 |
-
take = np.maximum(np.floor(counts * n / len(y)).astype(int), mins)
|
| 42 |
-
while take.sum() > n:
|
| 43 |
-
take[np.argmax(take - mins)] -= 1
|
| 44 |
-
idx = np.concatenate([rng.choice(np.flatnonzero(y == c), size=k, replace=False) for c, k in zip(classes, take)])
|
| 45 |
-
if len(idx) < n:
|
| 46 |
-
rest = np.setdiff1d(np.arange(len(y)), idx)
|
| 47 |
-
idx = np.concatenate([idx, rng.choice(rest, size=n - len(idx), replace=False)])
|
| 48 |
-
return np.sort(idx)
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
class LightPFNClassifier(ClassifierMixin, BaseEstimator):
|
| 52 |
-
"""device: "auto" (default: CUDA/ROCm through torch, else a Vulkan GPU of any vendor, else the CPU; the
|
| 53 |
-
environment variable LIGHTPFN_DEVICE can replace it), or explicitly "cuda[:i]", "vulkan[:i]", "cpu".
|
| 54 |
-
On "vulkan" the network runs on lightpfn.vulkan (wgpu); a training set too large for the GPU's buffers
|
| 55 |
-
falls back to the CPU with a warning when the device was chosen by "auto", and raises otherwise.
|
| 56 |
-
Weights and devices are initialized on fit, so construction and sklearn clone
|
| 57 |
-
perform no downloads or GPU work. Without model/checkpoint the pinned release
|
| 58 |
-
is downloaded from Hugging Face (install the hf extra and authenticate while
|
| 59 |
-
the repository is private). Input must be numeric; NaN is supported. Encode
|
| 60 |
-
strings/categories before fit. The model was trained for 2-10 classes.
|
| 61 |
-
|
| 62 |
-
n_estimators averages feature/label permutations; max_context bounds the
|
| 63 |
-
stratified context per member. chunk_rows bounds prediction batches;
|
| 64 |
-
chunk_cells/batch_cells control cache blocking and estimator batching.
|
| 65 |
-
n_threads sets PyTorch's process-wide CPU thread count. seed is the original
|
| 66 |
-
RNG parameter; random_state is its sklearn alias (use one of them).
|
| 67 |
-
fold enables equivalent inference folding. A supplied torch model is copied
|
| 68 |
-
during fit, leaving the constructor parameter and its device unchanged.
|
| 69 |
-
"""
|
| 70 |
-
|
| 71 |
-
def __init__(self, model=None, checkpoint=None, device="auto", n_estimators=1, max_context=20000,
|
| 72 |
-
chunk_rows=2048, n_threads=None, seed=0, *, chunk_cells="auto", batch_cells="auto", fold=True,
|
| 73 |
-
random_state=None, repo_id=None, revision=None, cache_dir=None, local_files_only=False):
|
| 74 |
-
self.model = model
|
| 75 |
-
self.checkpoint = checkpoint
|
| 76 |
-
self.device = device
|
| 77 |
-
self.fold = fold
|
| 78 |
-
self.n_estimators = n_estimators
|
| 79 |
-
self.max_context = max_context
|
| 80 |
-
self.chunk_rows = chunk_rows
|
| 81 |
-
self.chunk_cells = chunk_cells
|
| 82 |
-
self.batch_cells = batch_cells
|
| 83 |
-
self.n_threads = n_threads
|
| 84 |
-
self.seed = seed
|
| 85 |
-
self.random_state = random_state
|
| 86 |
-
self.repo_id = repo_id
|
| 87 |
-
self.revision = revision
|
| 88 |
-
self.cache_dir = cache_dir
|
| 89 |
-
self.local_files_only = local_files_only
|
| 90 |
-
|
| 91 |
-
def __sklearn_tags__(self):
|
| 92 |
-
tags = super().__sklearn_tags__()
|
| 93 |
-
tags.input_tags.allow_nan = True
|
| 94 |
-
# A pretrained network is not optimized on sklearn's toy check datasets.
|
| 95 |
-
tags.classifier_tags.poor_score = True
|
| 96 |
-
return tags
|
| 97 |
-
|
| 98 |
-
def _initialize_backend(self):
|
| 99 |
-
key = (id(self.model), self.checkpoint, self.device, self.fold, self.repo_id,
|
| 100 |
-
self.revision, self.cache_dir, self.local_files_only)
|
| 101 |
-
if getattr(self, "_backend_key", None) == key:
|
| 102 |
-
return
|
| 103 |
-
self.device_ = resolve_device(self.device)
|
| 104 |
-
torch_device = "cpu" if kind(self.device_) == "vulkan" else self.device_
|
| 105 |
-
if self.model is not None and self.checkpoint is not None:
|
| 106 |
-
raise ValueError("Provide either model or checkpoint, not both.")
|
| 107 |
-
if self.model is not None:
|
| 108 |
-
self.model_ = copy.deepcopy(self.model) if isinstance(self.model, torch.nn.Module) else self.model
|
| 109 |
-
elif self.checkpoint is not None:
|
| 110 |
-
self.model_ = load_model(self.checkpoint, torch_device)
|
| 111 |
-
else:
|
| 112 |
-
self.model_ = load_pretrained(repo_id=self.repo_id, revision=self.revision, device=torch_device,
|
| 113 |
-
cache_dir=self.cache_dir, local_files_only=self.local_files_only)
|
| 114 |
-
if isinstance(self.model_, torch.nn.Module):
|
| 115 |
-
self.model_ = self.model_.to(torch_device).eval()
|
| 116 |
-
# fold=True predicts with LightPFN.folded(), the same function with fewer memory passes
|
| 117 |
-
self.cpu_net_ = None
|
| 118 |
-
if kind(self.device_) == "vulkan":
|
| 119 |
-
from lightpfn.vulkan import VulkanLightPFN
|
| 120 |
-
|
| 121 |
-
index = self.device_.partition(":")[2]
|
| 122 |
-
self.net_ = VulkanLightPFN(self.model_, adapter=int(index) if index else None)
|
| 123 |
-
else:
|
| 124 |
-
self.net_ = self.model_.folded() if self.fold else self.model_
|
| 125 |
-
self._backend_key = key
|
| 126 |
-
|
| 127 |
-
def _validate_parameters(self):
|
| 128 |
-
for name in ("n_estimators", "max_context", "chunk_rows"):
|
| 129 |
-
value = getattr(self, name)
|
| 130 |
-
if isinstance(value, bool) or not isinstance(value, Integral) or value < 1:
|
| 131 |
-
raise ValueError(f"{name} must be a positive integer.")
|
| 132 |
-
if self.n_threads is not None and (isinstance(self.n_threads, bool) or
|
| 133 |
-
not isinstance(self.n_threads, Integral) or self.n_threads < 1):
|
| 134 |
-
raise ValueError("n_threads must be None or a positive integer.")
|
| 135 |
-
for name in ("chunk_cells", "batch_cells"):
|
| 136 |
-
value = getattr(self, name)
|
| 137 |
-
if value == "auto" or (name == "chunk_cells" and value is None):
|
| 138 |
-
continue
|
| 139 |
-
minimum = 1 if name == "chunk_cells" else 0
|
| 140 |
-
if isinstance(value, bool) or not isinstance(value, Integral) or value < minimum:
|
| 141 |
-
raise ValueError(f"{name} must be 'auto' or an integer >= {minimum}." )
|
| 142 |
-
if not isinstance(self.fold, bool):
|
| 143 |
-
raise ValueError("fold must be a boolean.")
|
| 144 |
-
if self.random_state is not None and self.seed not in (0, None):
|
| 145 |
-
raise ValueError("Use random_state or seed, not both.")
|
| 146 |
-
|
| 147 |
-
def __sklearn_is_fitted__(self):
|
| 148 |
-
return getattr(self, "_is_fitted", False)
|
| 149 |
-
|
| 150 |
-
def _auto(self, value, cpu, gpu):
|
| 151 |
-
if value != "auto":
|
| 152 |
-
return value
|
| 153 |
-
return cpu if kind(self.fit_device_) == "cpu" else gpu
|
| 154 |
-
|
| 155 |
-
def _chunk_cells(self):
|
| 156 |
-
"""Cells per piece of the cell stages: small pieces stay in the CPU cache (about 3x faster on
|
| 157 |
-
wide tables), large pieces bound GPU memory. The predictions do not depend on it."""
|
| 158 |
-
return self._auto(self.chunk_cells, CPU_CHUNK_CELLS_PER_THREAD * torch.get_num_threads(), GPU_CHUNK_CELLS)
|
| 159 |
-
|
| 160 |
-
def _group(self, n, m):
|
| 161 |
-
"""Estimators encoded together as one batch: as many as fit in batch_cells training cells
|
| 162 |
-
(fewer, larger operations; the same predictions as one at a time)."""
|
| 163 |
-
return max(1, min(self.n_estimators, self._auto(self.batch_cells, CPU_BATCH_CELLS, GPU_BATCH_CELLS) // max(1, n * m)))
|
| 164 |
-
|
| 165 |
-
@torch.inference_mode()
|
| 166 |
-
def fit(self, X, y, cat=None):
|
| 167 |
-
self._is_fitted = False
|
| 168 |
-
self._validate_parameters()
|
| 169 |
-
X, y = validate_data(self, X, y, dtype=np.float32, ensure_all_finite="allow-nan")
|
| 170 |
-
check_classification_targets(y)
|
| 171 |
-
if cat is not None:
|
| 172 |
-
cat = np.asarray(cat)
|
| 173 |
-
if cat.shape != (X.shape[1],) or cat.dtype != np.bool_:
|
| 174 |
-
raise ValueError("cat must be a boolean mask with one entry per feature.")
|
| 175 |
-
if cat.any():
|
| 176 |
-
warnings.warn("cat does not enable native categorical handling; columns are treated as numeric "
|
| 177 |
-
"codes. Encode categories before fit. The cat parameter is deprecated.",
|
| 178 |
-
FutureWarning, stacklevel=2)
|
| 179 |
-
if self.n_threads is not None:
|
| 180 |
-
torch.set_num_threads(self.n_threads)
|
| 181 |
-
self.classes_, y = np.unique(np.asarray(y), return_inverse=True)
|
| 182 |
-
if len(self.classes_) > 10:
|
| 183 |
-
raise ValueError("LightPFN supports at most 10 classes (trained on 2-10 classes).")
|
| 184 |
-
if self.max_context < len(self.classes_):
|
| 185 |
-
raise ValueError("max_context must be at least the number of classes.")
|
| 186 |
-
self._initialize_backend()
|
| 187 |
-
if len(self.classes_) > self.net_.cfg.label_slots:
|
| 188 |
-
raise ValueError("Number of classes exceeds the model's label slots.")
|
| 189 |
-
random_state = self.seed if self.random_state is None else self.random_state
|
| 190 |
-
if isinstance(random_state, np.random.RandomState):
|
| 191 |
-
random_state = random_state.randint(2**32)
|
| 192 |
-
rng = np.random.default_rng(random_state)
|
| 193 |
-
S = self.net_.cfg.label_slots
|
| 194 |
-
plans = []
|
| 195 |
-
for e in range(self.n_estimators):
|
| 196 |
-
rows = np.arange(len(y))
|
| 197 |
-
if len(rows) > self.max_context:
|
| 198 |
-
rows = stratified_subsample(rng, y, self.max_context)
|
| 199 |
-
feats = np.arange(X.shape[1]) if e == 0 else rng.permutation(X.shape[1])
|
| 200 |
-
slots = np.arange(S) if e == 0 else rng.permutation(S)
|
| 201 |
-
plans.append((rows, feats, slots))
|
| 202 |
-
self._choose_backend(len(plans[0][0]), X.shape[1])
|
| 203 |
-
try:
|
| 204 |
-
self._encode_members(X, y, plans)
|
| 205 |
-
except MemoryError as exc:
|
| 206 |
-
if kind(self.fit_device_) != "vulkan":
|
| 207 |
-
raise
|
| 208 |
-
# Drop partial contexts and unsubmitted work before a retry or a later fit.
|
| 209 |
-
msg = str(exc)
|
| 210 |
-
exc.__traceback__ = None
|
| 211 |
-
self.members_.clear()
|
| 212 |
-
eng = self.fit_net_.eng
|
| 213 |
-
eng.ops.clear()
|
| 214 |
-
eng.pending = 0.0
|
| 215 |
-
eng.scratch.clear()
|
| 216 |
-
eng.finish()
|
| 217 |
-
self._use_cpu(msg)
|
| 218 |
-
self._encode_members(X, y, plans)
|
| 219 |
-
if kind(self.fit_device_) == "cuda":
|
| 220 |
-
torch.cuda.synchronize(self.fit_device_) # fit time includes the GPU work, not only its launch
|
| 221 |
-
elif kind(self.fit_device_) == "vulkan":
|
| 222 |
-
self.fit_net_.eng.finish()
|
| 223 |
-
self._is_fitted = True
|
| 224 |
-
return self
|
| 225 |
-
|
| 226 |
-
def _encode_members(self, X, y, plans):
|
| 227 |
-
group = self._group(len(plans[0][0]), X.shape[1])
|
| 228 |
-
if kind(self.fit_device_) == "vulkan":
|
| 229 |
-
group = min(group, self.fit_net_.train_capacity(len(plans[0][0]), X.shape[1]))
|
| 230 |
-
tdev = self._tensor_device()
|
| 231 |
-
self.members_ = []
|
| 232 |
-
for g in range(0, len(plans), group):
|
| 233 |
-
part = plans[g : g + group]
|
| 234 |
-
Xt = torch.from_numpy(np.stack([X[rows][:, feats] for rows, feats, _ in part])).to(tdev)
|
| 235 |
-
yt = torch.from_numpy(np.stack([y[rows] for rows, _, _ in part])).long().to(tdev)
|
| 236 |
-
st = torch.from_numpy(np.stack([slots for _, _, slots in part])).long().to(tdev)
|
| 237 |
-
ctx = self.fit_net_.encode(Xt, yt, slots=st, n_classes=len(self.classes_), chunk_cells=self._chunk_cells())
|
| 238 |
-
# (feature order, context): one estimator per entry as before, or (G, m) orders for a batch of G
|
| 239 |
-
feats = part[0][1] if len(part) == 1 else np.stack([feats for _, feats, _ in part])
|
| 240 |
-
self.members_.append((feats, ctx))
|
| 241 |
-
|
| 242 |
-
def _choose_backend(self, n, m):
|
| 243 |
-
"""The network and device of this fit: the classifier's, or the CPU when a Vulkan device chosen by
|
| 244 |
-
"auto" cannot hold one estimator's caches and scratch buffers."""
|
| 245 |
-
self.fit_device_, self.fit_net_ = self.device_, self.net_
|
| 246 |
-
if kind(self.device_) == "vulkan" and self.net_.train_capacity(n, m) < 1:
|
| 247 |
-
self._use_cpu(f"{n} x {m} training cells exceed a Vulkan cache/scratch buffer limit; "
|
| 248 |
-
"lower max_context or use device='cpu'")
|
| 249 |
-
|
| 250 |
-
def _use_cpu(self, msg):
|
| 251 |
-
if not resolve_device_was_auto(self.device):
|
| 252 |
-
raise MemoryError(msg)
|
| 253 |
-
warnings.warn(msg + ": running this fit on the CPU", RuntimeWarning, stacklevel=3)
|
| 254 |
-
if self.cpu_net_ is None:
|
| 255 |
-
self.cpu_net_ = self.model_.folded() if self.fold else self.model_
|
| 256 |
-
self.fit_device_, self.fit_net_ = "cpu", self.cpu_net_
|
| 257 |
-
|
| 258 |
-
def _tensor_device(self):
|
| 259 |
-
return "cpu" if kind(self.fit_device_) == "vulkan" else self.fit_device_
|
| 260 |
-
|
| 261 |
-
@torch.inference_mode()
|
| 262 |
-
def predict_proba(self, X):
|
| 263 |
-
check_is_fitted(self)
|
| 264 |
-
X = validate_data(self, X, reset=False, dtype=np.float32, ensure_all_finite="allow-nan")
|
| 265 |
-
P = np.zeros((len(X), len(self.classes_)))
|
| 266 |
-
for feats, ctx in self.members_:
|
| 267 |
-
for i in range(0, len(X), self.chunk_rows):
|
| 268 |
-
Xt = torch.from_numpy(np.stack([X[i : i + self.chunk_rows][:, f] for f in np.atleast_2d(feats)])).to(self._tensor_device())
|
| 269 |
-
probs = torch.softmax(self.fit_net_.predict_logits(ctx, Xt, self._chunk_cells()).float(), -1).cpu().numpy()
|
| 270 |
-
for p in probs: # one estimator at a time, in float64 as before
|
| 271 |
-
P[i : i + self.chunk_rows] += p
|
| 272 |
-
return P / self.n_estimators
|
| 273 |
-
|
| 274 |
-
def predict(self, X):
|
| 275 |
-
probabilities = self.predict_proba(X)
|
| 276 |
-
return self.classes_[probabilities.argmax(1)]
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
def resolve_device_was_auto(device):
|
| 280 |
-
return device is None or str(device).strip().lower() == "auto"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/vulkan/__init__.py
DELETED
|
@@ -1,40 +0,0 @@
|
|
| 1 |
-
"""Vulkan backend: LightPFN inference on any GPU with a Vulkan driver (AMD, Intel, NVIDIA; Linux and
|
| 2 |
-
Windows), through wgpu (`pip install wgpu`). The WGSL kernels in kernels.py are compiled to SPIR-V when the
|
| 3 |
-
backend starts; no Vulkan SDK or compiler is needed.
|
| 4 |
-
|
| 5 |
-
from lightpfn.vulkan import VulkanLightPFN, is_available
|
| 6 |
-
net = VulkanLightPFN(model) # first GPU with a Vulkan driver
|
| 7 |
-
ctx = net.encode(X_train, y_train, n_classes=C)
|
| 8 |
-
logits = net.predict_logits(ctx, X_test)
|
| 9 |
-
|
| 10 |
-
LIGHTPFN_VULKAN_ADAPTER selects the adapter by index or name ("llvmpipe" runs the kernels on the CPU,
|
| 11 |
-
which is how the tests run without a GPU).
|
| 12 |
-
"""
|
| 13 |
-
|
| 14 |
-
from lightpfn.vulkan.engine import ADAPTER_ENV, GPU_TYPES, pick_adapter, vulkan_adapters
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
def adapters():
|
| 18 |
-
"""The Vulkan adapters wgpu sees, as dicts (index, name, type, vendor, driver)."""
|
| 19 |
-
return [dict(index=i, name=info.get("device"), type=info.get("adapter_type"), vendor=info.get("vendor"),
|
| 20 |
-
driver=info.get("description")) for i, (_, info) in enumerate(vulkan_adapters())]
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
def is_available(adapter=None):
|
| 24 |
-
"""True when wgpu is installed and finds a Vulkan GPU (or the adapter selected by `adapter` or by
|
| 25 |
-
LIGHTPFN_VULKAN_ADAPTER, which may be a CPU driver)."""
|
| 26 |
-
try:
|
| 27 |
-
return pick_adapter(adapter) is not None
|
| 28 |
-
except Exception:
|
| 29 |
-
return False
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
def __getattr__(name): # VulkanLightPFN imports torch and the model: load it on first use
|
| 33 |
-
if name in ("VulkanLightPFN", "VulkanContext"):
|
| 34 |
-
from lightpfn.vulkan import model
|
| 35 |
-
|
| 36 |
-
return getattr(model, name)
|
| 37 |
-
raise AttributeError(name)
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
__all__ = ["ADAPTER_ENV", "GPU_TYPES", "VulkanLightPFN", "VulkanContext", "adapters", "is_available"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/vulkan/engine.py
DELETED
|
@@ -1,210 +0,0 @@
|
|
| 1 |
-
"""Vulkan device, buffers and kernel dispatch through wgpu (WebGPU on the Vulkan backend).
|
| 2 |
-
|
| 3 |
-
Operations are recorded into one compute pass and submitted on flush() or before a read; WebGPU orders the
|
| 4 |
-
dispatches of a pass and inserts the barriers between them, so a later kernel sees what an earlier one
|
| 5 |
-
wrote. Parameters share one storage buffer, at offsets aligned to the device (at least 256 bytes).
|
| 6 |
-
"""
|
| 7 |
-
|
| 8 |
-
import math
|
| 9 |
-
import os
|
| 10 |
-
|
| 11 |
-
import numpy as np
|
| 12 |
-
|
| 13 |
-
from lightpfn.vulkan import kernels
|
| 14 |
-
|
| 15 |
-
try:
|
| 16 |
-
import wgpu
|
| 17 |
-
except ImportError: # optional dependency: pip install wgpu
|
| 18 |
-
wgpu = None
|
| 19 |
-
|
| 20 |
-
ADAPTER_ENV = "LIGHTPFN_VULKAN_ADAPTER"
|
| 21 |
-
GPU_TYPES = ("DiscreteGPU", "IntegratedGPU", "VirtualGPU")
|
| 22 |
-
PARAM_SLOT = 256 # bytes of parameters per operation; u32 63 = slice start
|
| 23 |
-
MAX_GROUPS = 65535
|
| 24 |
-
# GPU work per submission, in floating-point operations: a few ms on a discrete GPU, well under a second on
|
| 25 |
-
# an integrated one. Drivers reset a GPU whose job runs too long (amdgpu's ring timeout, Windows TDR after
|
| 26 |
-
# 2 s), so heavy dispatches are cut into slices of workgroups and submitted in several jobs.
|
| 27 |
-
SUBMIT_FLOPS = 5e10
|
| 28 |
-
SUBMIT_OPS = 4096
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
def vulkan_adapters():
|
| 32 |
-
"""Vulkan adapters seen by wgpu, as (adapter, info) pairs; empty without wgpu or a Vulkan driver."""
|
| 33 |
-
if wgpu is None:
|
| 34 |
-
return []
|
| 35 |
-
try:
|
| 36 |
-
found = wgpu.gpu.enumerate_adapters_sync()
|
| 37 |
-
except Exception: # no loader, no driver
|
| 38 |
-
return []
|
| 39 |
-
return [(a, a.info) for a in found if a.info.get("backend_type") == "Vulkan"]
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
def pick_adapter(adapter=None):
|
| 43 |
-
"""adapter: None (LIGHTPFN_VULKAN_ADAPTER, else the first discrete GPU, else any GPU), an index into
|
| 44 |
-
vulkan_adapters() or a substring of the device name (e.g. "llvmpipe", the CPU driver, for tests)."""
|
| 45 |
-
found = vulkan_adapters()
|
| 46 |
-
if adapter is None:
|
| 47 |
-
adapter = os.environ.get(ADAPTER_ENV) or None
|
| 48 |
-
if adapter is None:
|
| 49 |
-
for kind in GPU_TYPES:
|
| 50 |
-
for a, info in found:
|
| 51 |
-
if info.get("adapter_type") == kind:
|
| 52 |
-
return a
|
| 53 |
-
return None
|
| 54 |
-
if isinstance(adapter, int) or str(adapter).isdigit():
|
| 55 |
-
i = int(adapter)
|
| 56 |
-
return found[i][0] if 0 <= i < len(found) else None
|
| 57 |
-
for a, info in found:
|
| 58 |
-
if str(adapter).lower() in str(info.get("device", "")).lower():
|
| 59 |
-
return a
|
| 60 |
-
return None
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
class View:
|
| 64 |
-
"""Rows of a buffer for the kernels: row r = (z, i) with z = r // L, i = r % L, at float offset
|
| 65 |
-
base + (z // zin) * so + (z % zin) * si + i * ss."""
|
| 66 |
-
|
| 67 |
-
__slots__ = ("buf", "base", "L", "zin", "so", "si", "ss")
|
| 68 |
-
|
| 69 |
-
def __init__(self, buf, base=0, L=1, zin=1, so=0, si=0, ss=0):
|
| 70 |
-
self.buf, self.base, self.L, self.zin, self.so, self.si, self.ss = buf, base, L, zin, so, si, ss
|
| 71 |
-
|
| 72 |
-
def params(self):
|
| 73 |
-
return [self.base, self.L, self.zin, self.so, self.si, self.ss]
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
def rows(buf, width, base=0):
|
| 77 |
-
"""Contiguous rows of `width` floats."""
|
| 78 |
-
return View(buf, base, L=1, zin=1, so=width)
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
class Engine:
|
| 82 |
-
def __init__(self, adapter=None):
|
| 83 |
-
if wgpu is None:
|
| 84 |
-
raise RuntimeError("the Vulkan backend needs wgpu: pip install wgpu")
|
| 85 |
-
ad = pick_adapter(adapter)
|
| 86 |
-
if ad is None:
|
| 87 |
-
raise RuntimeError("no Vulkan adapter found" + (f" matching {adapter!r}" if adapter is not None else ""))
|
| 88 |
-
self.adapter, self.info = ad, ad.info
|
| 89 |
-
lim = ad.limits
|
| 90 |
-
want = ("max-storage-buffer-binding-size", "max-buffer-size", "max-storage-buffers-per-shader-stage",
|
| 91 |
-
"max-compute-workgroup-storage-size", "max-compute-invocations-per-workgroup",
|
| 92 |
-
"max-compute-workgroups-per-dimension")
|
| 93 |
-
self.device = ad.request_device_sync(required_limits={k: lim[k] for k in want if k in lim})
|
| 94 |
-
lim = self.device.limits
|
| 95 |
-
self.max_binding = min(lim["max-storage-buffer-binding-size"], lim["max-buffer-size"], (1 << 34) - 16)
|
| 96 |
-
self.max_groups = min(MAX_GROUPS, lim["max-compute-workgroups-per-dimension"])
|
| 97 |
-
self.param_slot = max(PARAM_SLOT, lim["min-storage-buffer-offset-alignment"])
|
| 98 |
-
self.submit_ops = min(SUBMIT_OPS, self.max_binding // self.param_slot)
|
| 99 |
-
self.usage = wgpu.BufferUsage.STORAGE | wgpu.BufferUsage.COPY_SRC | wgpu.BufferUsage.COPY_DST
|
| 100 |
-
self.pipes = {}
|
| 101 |
-
self.ops = [] # (pipeline, buffers by binding, params u32, workgroups)
|
| 102 |
-
self.pending = 0.0 # estimated flops recorded since the last submission
|
| 103 |
-
self.submit_flops = SUBMIT_FLOPS
|
| 104 |
-
self.scratch = {}
|
| 105 |
-
self.sync = self.empty(4)
|
| 106 |
-
|
| 107 |
-
# buffers -----------------------------------------------------------------------------------
|
| 108 |
-
def upload(self, a, dtype=np.float32):
|
| 109 |
-
a = np.ascontiguousarray(a, dtype=dtype)
|
| 110 |
-
if a.nbytes == 0:
|
| 111 |
-
a = np.zeros(4, dtype)
|
| 112 |
-
if a.nbytes > self.max_binding:
|
| 113 |
-
raise MemoryError(f"{a.nbytes} bytes exceed the device's {self.max_binding}-byte buffer limit")
|
| 114 |
-
return self.device.create_buffer_with_data(data=a, usage=self.usage)
|
| 115 |
-
|
| 116 |
-
def empty(self, n_floats):
|
| 117 |
-
size = (max(4, int(n_floats)) + 3) // 4 * 16
|
| 118 |
-
if size > self.max_binding:
|
| 119 |
-
raise MemoryError(f"{size} bytes exceed the device's {self.max_binding}-byte buffer limit")
|
| 120 |
-
return self.device.create_buffer(size=size, usage=self.usage)
|
| 121 |
-
|
| 122 |
-
def temp(self, role, n_floats):
|
| 123 |
-
"""A scratch buffer reused by role (dispatches are ordered, so reuse is safe)."""
|
| 124 |
-
buf = self.scratch.get(role)
|
| 125 |
-
if buf is None or buf.size < 4 * n_floats:
|
| 126 |
-
grow = min(buf.size // 2, self.max_binding // 4) if buf is not None else 0 # doubling, within the limit
|
| 127 |
-
buf = self.scratch[role] = self.empty(max(n_floats, grow))
|
| 128 |
-
return buf
|
| 129 |
-
|
| 130 |
-
def download(self, buf, shape, offset=0):
|
| 131 |
-
self.flush()
|
| 132 |
-
count = int(np.prod(shape))
|
| 133 |
-
if count == 0:
|
| 134 |
-
return np.empty(shape, np.float32)
|
| 135 |
-
data = self.device.queue.read_buffer(buf, 4 * offset, 4 * count)
|
| 136 |
-
return np.frombuffer(data, dtype=np.float32).reshape(shape).copy()
|
| 137 |
-
|
| 138 |
-
# kernels -----------------------------------------------------------------------------------
|
| 139 |
-
def pipeline(self, name, src, flags=(), **consts):
|
| 140 |
-
key = (name, tuple(sorted(flags)), tuple(sorted(consts.items())))
|
| 141 |
-
p = self.pipes.get(key)
|
| 142 |
-
if p is None:
|
| 143 |
-
code = kernels.render(src, flags, **consts)
|
| 144 |
-
module = self.device.create_shader_module(code=code)
|
| 145 |
-
p = self.pipes[key] = self.device.create_compute_pipeline(
|
| 146 |
-
layout="auto", compute={"module": module, "entry_point": "main"})
|
| 147 |
-
return p
|
| 148 |
-
|
| 149 |
-
def dispatch(self, pipe, buffers, params, groups, flops=0.0):
|
| 150 |
-
"""buffers: {binding: buffer}; params: list of u32 (floats as float32 bits), P[1] and P[63] filled
|
| 151 |
-
here; groups: workgroups (over two grid dims past 65535); flops: estimated cost, which cuts the
|
| 152 |
-
dispatch into slices of consecutive workgroups and the work into submissions of ~submit_flops."""
|
| 153 |
-
if groups <= 0:
|
| 154 |
-
return
|
| 155 |
-
params = [int(v) for v in params]
|
| 156 |
-
if len(params) > 63 or any(v < 0 or v > 0xFFFFFFFF for v in params) or groups > 0xFFFFFFFF:
|
| 157 |
-
raise ValueError("kernel parameters must fit u32 and leave P[63] for the slice start")
|
| 158 |
-
parts = max(1, min(groups, math.ceil(flops / self.submit_flops)))
|
| 159 |
-
step = math.ceil(groups / parts)
|
| 160 |
-
start = 0
|
| 161 |
-
while start < groups:
|
| 162 |
-
count = min(step, groups - start)
|
| 163 |
-
gx = min(count, self.max_groups)
|
| 164 |
-
gy = min(count // gx, self.max_groups)
|
| 165 |
-
if parts == 1 and groups <= self.max_groups ** 2:
|
| 166 |
-
gy = math.ceil(count / gx) # the global kernel bound covers a single dispatch's padding
|
| 167 |
-
else:
|
| 168 |
-
count = gx * gy # no padded workgroups may spill into the next slice
|
| 169 |
-
p = list(params) + [0] * (PARAM_SLOT // 4 - len(params))
|
| 170 |
-
p[1], p[63] = gx, start
|
| 171 |
-
cost = flops * count / groups
|
| 172 |
-
if self.pending + cost > self.submit_flops:
|
| 173 |
-
self.flush()
|
| 174 |
-
self.ops.append((pipe, buffers, p, (gx, gy, 1)))
|
| 175 |
-
self.pending += cost
|
| 176 |
-
start += count
|
| 177 |
-
if self.pending >= self.submit_flops or len(self.ops) >= self.submit_ops:
|
| 178 |
-
self.flush()
|
| 179 |
-
|
| 180 |
-
def flush(self):
|
| 181 |
-
if not self.ops:
|
| 182 |
-
return
|
| 183 |
-
slot = self.param_slot // 4
|
| 184 |
-
P = np.zeros(slot * len(self.ops), np.uint32)
|
| 185 |
-
for k, (_, _, params, _) in enumerate(self.ops):
|
| 186 |
-
P[k * slot : k * slot + len(params)] = params
|
| 187 |
-
pbuf = self.device.create_buffer_with_data(data=P, usage=wgpu.BufferUsage.STORAGE)
|
| 188 |
-
enc = self.device.create_command_encoder()
|
| 189 |
-
cp = enc.begin_compute_pass()
|
| 190 |
-
for k, (pipe, buffers, _, grid) in enumerate(self.ops):
|
| 191 |
-
entries = [{"binding": 0, "resource": {"buffer": pbuf, "offset": k * self.param_slot, "size": PARAM_SLOT}}]
|
| 192 |
-
entries += [{"binding": b, "resource": {"buffer": buf, "offset": 0, "size": buf.size}}
|
| 193 |
-
for b, buf in sorted(buffers.items())]
|
| 194 |
-
bg = self.device.create_bind_group(layout=pipe.get_bind_group_layout(0), entries=entries)
|
| 195 |
-
cp.set_pipeline(pipe)
|
| 196 |
-
cp.set_bind_group(0, bg)
|
| 197 |
-
cp.dispatch_workgroups(*grid)
|
| 198 |
-
cp.end()
|
| 199 |
-
self.device.queue.submit([enc.finish()])
|
| 200 |
-
self.ops = []
|
| 201 |
-
self.pending = 0.0
|
| 202 |
-
|
| 203 |
-
def finish(self):
|
| 204 |
-
"""Waits for the submitted work (timing, synchronization with the host)."""
|
| 205 |
-
self.flush()
|
| 206 |
-
self.device.queue.read_buffer(self.sync, 0, 4)
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
def f32bits(x):
|
| 210 |
-
return int(np.array([x], np.float32).view(np.uint32)[0])
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/vulkan/kernels.py
DELETED
|
@@ -1,457 +0,0 @@
|
|
| 1 |
-
"""WGSL compute shaders of the Vulkan backend (compiled to SPIR-V by wgpu when a pipeline is created).
|
| 2 |
-
|
| 3 |
-
Every kernel reads its parameters from a u32 storage buffer P at offset 0. Tensors are float32 rows
|
| 4 |
-
addressed through views (see engine.View): row r = (z, i) with z = r / L, i = r % L lives at float offset
|
| 5 |
-
base + (z / zin) * so + (z % zin) * si + i * ss. Views let one kernel read a column of cells, a row of
|
| 6 |
-
cells, the first tokens of every row or a broadcast parameter without copies or transposes.
|
| 7 |
-
|
| 8 |
-
Precision: everything is float32. sin/cos use a Cody-Waite range reduction with minimax polynomials and
|
| 9 |
-
erf the Abramowitz-Stegun 7.1.26 formula (float32 error below 5e-7 on [-32,32]), so results do not depend on the precision
|
| 10 |
-
of the driver's transcendental functions.
|
| 11 |
-
"""
|
| 12 |
-
|
| 13 |
-
import re
|
| 14 |
-
|
| 15 |
-
COMMON = """
|
| 16 |
-
@group(0) @binding(0) var<storage, read> P: array<u32>;
|
| 17 |
-
|
| 18 |
-
struct View { base: u32, L: u32, zin: u32, so: u32, si: u32, ss: u32 }
|
| 19 |
-
|
| 20 |
-
fn view(o: u32) -> View { return View(P[o], P[o + 1u], P[o + 2u], P[o + 3u], P[o + 4u], P[o + 5u]); }
|
| 21 |
-
fn voff(v: View, z: u32, i: u32) -> u32 { return v.base + (z / v.zin) * v.so + (z % v.zin) * v.si + i * v.ss; }
|
| 22 |
-
fn roff(v: View, r: u32) -> u32 { return voff(v, r / v.L, r % v.L); }
|
| 23 |
-
fn pf(o: u32) -> f32 { return bitcast<f32>(P[o]); }
|
| 24 |
-
|
| 25 |
-
fn erf_(x: f32) -> f32 {
|
| 26 |
-
let z = abs(x);
|
| 27 |
-
let t = 1.0 / (1.0 + 0.3275911 * z);
|
| 28 |
-
let y = 1.0 - (((((1.061405429 * t - 1.453152027) * t) + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t * exp(-z * z);
|
| 29 |
-
return select(-y, y, x >= 0.0);
|
| 30 |
-
}
|
| 31 |
-
fn gelu(x: f32) -> f32 { return 0.5 * x * (1.0 + erf_(x * 0.7071067811865476)); }
|
| 32 |
-
fn gelu4(v: vec4<f32>) -> vec4<f32> { return vec4<f32>(gelu(v.x), gelu(v.y), gelu(v.z), gelu(v.w)); }
|
| 33 |
-
fn tanh_(x: f32) -> f32 {
|
| 34 |
-
let e = exp(2.0 * min(abs(x), 15.0));
|
| 35 |
-
let t = 1.0 - 2.0 / (e + 1.0);
|
| 36 |
-
return select(-t, t, x >= 0.0);
|
| 37 |
-
}
|
| 38 |
-
fn sincos(x: f32) -> vec2<f32> {
|
| 39 |
-
// x = k * pi/2 + r with |r| <= pi/4 (pi/2 split in three parts, Cody-Waite), cephes sinf/cosf polynomials
|
| 40 |
-
let k = round(x * 0.6366197723675814);
|
| 41 |
-
var r = x - k * 1.5703125;
|
| 42 |
-
r = r - k * 4.837512969970703125e-4;
|
| 43 |
-
r = r - k * 7.54978995489188216e-8;
|
| 44 |
-
let r2 = r * r;
|
| 45 |
-
let s = r + r * r2 * (-1.6666654611e-1 + r2 * (8.3321608736e-3 + r2 * -1.9515295891e-4));
|
| 46 |
-
let c = 1.0 - 0.5 * r2 + r2 * r2 * (4.166664568298827e-2 + r2 * (-1.388731625493765e-3 + r2 * 2.443315711809948e-5));
|
| 47 |
-
let q = i32(k - 4.0 * floor(k * 0.25));
|
| 48 |
-
if (q == 0) { return vec2<f32>(s, c); }
|
| 49 |
-
if (q == 1) { return vec2<f32>(c, -s); }
|
| 50 |
-
if (q == 2) { return vec2<f32>(-s, -c); }
|
| 51 |
-
return vec2<f32>(-c, s);
|
| 52 |
-
}
|
| 53 |
-
// P[1]: workgroups along x (2D grids past 65535); P[63]: first workgroup of this slice of the dispatch
|
| 54 |
-
fn wgid(wg: vec3u) -> u32 { return P[63] + wg.x + wg.y * P[1]; }
|
| 55 |
-
"""
|
| 56 |
-
|
| 57 |
-
# Y[r, :O] = epilogue((X[r, :K] @ W^T)), W (O, K) row-major. 64 x 64 output tiles, 256 threads with 4 x 4
|
| 58 |
-
# outputs each. Options (compile time): NORM scales row r by 1 / rms(X[r]) (a folded RMSNorm), BIAS adds
|
| 59 |
-
# b, GELU applies gelu, RES = "self" adds the old Y, "r" adds rows of a separate view R.
|
| 60 |
-
# P: [0] n_tiles, [1] grid x, [2] N, [3] K, [4] O, [5] eps, [6..11] X view, [12..17] Y view, [18..23] R view
|
| 61 |
-
GEMM = """
|
| 62 |
-
@group(0) @binding(1) var<storage, read> X: array<vec4<f32>>;
|
| 63 |
-
@group(0) @binding(2) var<storage, read> W: array<vec4<f32>>;
|
| 64 |
-
@group(0) @binding(3) var<storage, read_write> Y: array<vec4<f32>>;
|
| 65 |
-
#if BIAS
|
| 66 |
-
@group(0) @binding(4) var<storage, read> Bv: array<vec4<f32>>;
|
| 67 |
-
#endif
|
| 68 |
-
#if RES_R
|
| 69 |
-
@group(0) @binding(5) var<storage, read> R: array<vec4<f32>>;
|
| 70 |
-
#endif
|
| 71 |
-
var<workgroup> xs: array<vec4<f32>, 256>; // [k 16][row 64 / 4]
|
| 72 |
-
var<workgroup> ws: array<vec4<f32>, 256>; // [k 16][col 64 / 4]
|
| 73 |
-
var<workgroup> xo: array<u32, 64>;
|
| 74 |
-
var<workgroup> yo: array<u32, 64>;
|
| 75 |
-
var<workgroup> ro: array<u32, 64>;
|
| 76 |
-
|
| 77 |
-
@compute @workgroup_size(16, 16)
|
| 78 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_id) lid: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 79 |
-
let id = wgid(wg);
|
| 80 |
-
if (id >= P[0]) { return; }
|
| 81 |
-
let N = P[2]; let K = P[3]; let O = P[4];
|
| 82 |
-
let tiles_c = (O + 63u) / 64u;
|
| 83 |
-
let row0 = (id / tiles_c) * 64u;
|
| 84 |
-
let col0 = (id % tiles_c) * 64u;
|
| 85 |
-
if (li < 64u) {
|
| 86 |
-
let r = min(row0 + li, N - 1u);
|
| 87 |
-
xo[li] = roff(view(6u), r) / 4u;
|
| 88 |
-
yo[li] = roff(view(12u), r) / 4u;
|
| 89 |
-
ro[li] = roff(view(18u), r) / 4u;
|
| 90 |
-
}
|
| 91 |
-
workgroupBarrier();
|
| 92 |
-
var acc: array<vec4<f32>, 4>;
|
| 93 |
-
var ssq = vec4<f32>(0.0);
|
| 94 |
-
let lr = li / 4u; // row (and column) loaded by this thread
|
| 95 |
-
let lk = li % 4u; // vec4 of k loaded by this thread
|
| 96 |
-
let K4 = K / 4u;
|
| 97 |
-
for (var k4 = 0u; k4 < K4; k4 += 4u) {
|
| 98 |
-
var xv = vec4<f32>(0.0);
|
| 99 |
-
if (row0 + lr < N && k4 + lk < K4) { xv = X[xo[lr] + k4 + lk]; }
|
| 100 |
-
var wv = vec4<f32>(0.0);
|
| 101 |
-
if (col0 + lr < O && k4 + lk < K4) { wv = W[(col0 + lr) * K4 + k4 + lk]; }
|
| 102 |
-
for (var c = 0u; c < 4u; c++) {
|
| 103 |
-
xs[(lk * 4u + c) * 16u + lr / 4u][lr % 4u] = xv[c];
|
| 104 |
-
ws[(lk * 4u + c) * 16u + lr / 4u][lr % 4u] = wv[c];
|
| 105 |
-
}
|
| 106 |
-
workgroupBarrier();
|
| 107 |
-
for (var kk = 0u; kk < 16u; kk++) {
|
| 108 |
-
let a = xs[kk * 16u + lid.y];
|
| 109 |
-
let b = ws[kk * 16u + lid.x];
|
| 110 |
-
acc[0] += a.x * b;
|
| 111 |
-
acc[1] += a.y * b;
|
| 112 |
-
acc[2] += a.z * b;
|
| 113 |
-
acc[3] += a.w * b;
|
| 114 |
-
#if NORM
|
| 115 |
-
ssq += a * a;
|
| 116 |
-
#endif
|
| 117 |
-
}
|
| 118 |
-
workgroupBarrier();
|
| 119 |
-
}
|
| 120 |
-
let c4 = col0 / 4u + lid.x;
|
| 121 |
-
if (col0 + lid.x * 4u >= O) { return; }
|
| 122 |
-
for (var i = 0u; i < 4u; i++) {
|
| 123 |
-
let rl = lid.y * 4u + i;
|
| 124 |
-
if (row0 + rl >= N) { continue; }
|
| 125 |
-
var v = acc[i];
|
| 126 |
-
#if NORM
|
| 127 |
-
v *= inverseSqrt(ssq[i] / f32(K) + pf(5u));
|
| 128 |
-
#endif
|
| 129 |
-
#if BIAS
|
| 130 |
-
v += Bv[c4];
|
| 131 |
-
#endif
|
| 132 |
-
#if GELU
|
| 133 |
-
v = gelu4(v);
|
| 134 |
-
#endif
|
| 135 |
-
#if RES_SELF
|
| 136 |
-
v += Y[yo[rl] + c4];
|
| 137 |
-
#endif
|
| 138 |
-
#if RES_R
|
| 139 |
-
v += R[ro[rl] + c4];
|
| 140 |
-
#endif
|
| 141 |
-
Y[yo[rl] + c4] = v;
|
| 142 |
-
}
|
| 143 |
-
}
|
| 144 |
-
"""
|
| 145 |
-
|
| 146 |
-
# softmax(q k^T / sqrt(D)) v for every (z, head, query), float32 online softmax (flash attention).
|
| 147 |
-
# Tiled: a workgroup holds 64 queries of one (z, head) and walks the keys in shared-memory tiles of KT.
|
| 148 |
-
# P: [0] n_groups, [1] grid x, [2] Lq, [3] Lk, [4] H, [5] scale, [6..11] q, [12..17] k, [18..23] v,
|
| 149 |
-
# [24..29] o views, [30] q_h, [31] k_h, [32] v_h, [33] o_h head strides (floats)
|
| 150 |
-
ATTN_TILED = """
|
| 151 |
-
const D4: u32 = {D4}u;
|
| 152 |
-
const KT: u32 = {KT}u;
|
| 153 |
-
@group(0) @binding(1) var<storage, read> Q: array<vec4<f32>>;
|
| 154 |
-
@group(0) @binding(2) var<storage, read> K: array<vec4<f32>>;
|
| 155 |
-
@group(0) @binding(3) var<storage, read> V: array<vec4<f32>>;
|
| 156 |
-
@group(0) @binding(4) var<storage, read_write> O: array<vec4<f32>>;
|
| 157 |
-
var<workgroup> ks: array<vec4<f32>, {KT} * {D4}>;
|
| 158 |
-
var<workgroup> vs: array<vec4<f32>, {KT} * {D4}>;
|
| 159 |
-
|
| 160 |
-
@compute @workgroup_size(64)
|
| 161 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 162 |
-
let id = wgid(wg);
|
| 163 |
-
if (id >= P[0]) { return; }
|
| 164 |
-
let Lq = P[2]; let Lk = P[3]; let H = P[4];
|
| 165 |
-
let tiles = (Lq + 63u) / 64u;
|
| 166 |
-
let zh = id / tiles;
|
| 167 |
-
let z = zh / H;
|
| 168 |
-
let h = zh % H;
|
| 169 |
-
let qi = (id % tiles) * 64u + li;
|
| 170 |
-
let qv = view(6u); let kv = view(12u); let vv = view(18u); let ov = view(24u);
|
| 171 |
-
let qb = (voff(qv, z, min(qi, Lq - 1u)) + h * P[30]) / 4u;
|
| 172 |
-
let scale = pf(5u);
|
| 173 |
-
var q: array<vec4<f32>, {D4}>;
|
| 174 |
-
var acc: array<vec4<f32>, {D4}>;
|
| 175 |
-
for (var d = 0u; d < D4; d++) { q[d] = Q[qb + d] * scale; acc[d] = vec4<f32>(0.0); }
|
| 176 |
-
var m = -3.0e38;
|
| 177 |
-
var l = 0.0;
|
| 178 |
-
for (var t0 = 0u; t0 < Lk; t0 += KT) {
|
| 179 |
-
for (var e = li; e < KT * D4; e += 64u) {
|
| 180 |
-
let key = min(t0 + e / D4, Lk - 1u);
|
| 181 |
-
ks[e] = K[(voff(kv, z, key) + h * P[31]) / 4u + e % D4];
|
| 182 |
-
vs[e] = V[(voff(vv, z, key) + h * P[32]) / 4u + e % D4];
|
| 183 |
-
}
|
| 184 |
-
workgroupBarrier();
|
| 185 |
-
let nk = min(KT, Lk - t0);
|
| 186 |
-
var s: array<f32, {KT}>;
|
| 187 |
-
var mt = m;
|
| 188 |
-
for (var j = 0u; j < KT; j++) {
|
| 189 |
-
var dot4 = vec4<f32>(0.0);
|
| 190 |
-
for (var d = 0u; d < D4; d++) { dot4 += q[d] * ks[j * D4 + d]; }
|
| 191 |
-
let sj = select(-3.0e38, dot4.x + dot4.y + dot4.z + dot4.w, j < nk);
|
| 192 |
-
s[j] = sj;
|
| 193 |
-
mt = max(mt, sj);
|
| 194 |
-
}
|
| 195 |
-
let corr = exp(m - mt);
|
| 196 |
-
l *= corr;
|
| 197 |
-
for (var d = 0u; d < D4; d++) { acc[d] *= corr; }
|
| 198 |
-
for (var j = 0u; j < KT; j++) {
|
| 199 |
-
let p = select(0.0, exp(s[j] - mt), j < nk);
|
| 200 |
-
l += p;
|
| 201 |
-
for (var d = 0u; d < D4; d++) { acc[d] += p * vs[j * D4 + d]; }
|
| 202 |
-
}
|
| 203 |
-
m = mt;
|
| 204 |
-
workgroupBarrier();
|
| 205 |
-
}
|
| 206 |
-
if (qi < Lq) {
|
| 207 |
-
let ob = (voff(ov, z, qi) + h * P[33]) / 4u;
|
| 208 |
-
for (var d = 0u; d < D4; d++) { O[ob + d] = acc[d] / l; }
|
| 209 |
-
}
|
| 210 |
-
}
|
| 211 |
-
"""
|
| 212 |
-
|
| 213 |
-
# Same function, one thread per (z, head, query) reading keys from global memory: for few queries per
|
| 214 |
-
# sequence (summary tokens) or few keys. Same parameter layout as ATTN_TILED, [0] = Z * H * Lq.
|
| 215 |
-
ATTN_SMALL = """
|
| 216 |
-
const D4: u32 = {D4}u;
|
| 217 |
-
@group(0) @binding(1) var<storage, read> Q: array<vec4<f32>>;
|
| 218 |
-
@group(0) @binding(2) var<storage, read> K: array<vec4<f32>>;
|
| 219 |
-
@group(0) @binding(3) var<storage, read> V: array<vec4<f32>>;
|
| 220 |
-
@group(0) @binding(4) var<storage, read_write> O: array<vec4<f32>>;
|
| 221 |
-
|
| 222 |
-
@compute @workgroup_size(64)
|
| 223 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 224 |
-
let id = wgid(wg) * 64u + li;
|
| 225 |
-
if (id >= P[0]) { return; }
|
| 226 |
-
let Lq = P[2]; let Lk = P[3]; let H = P[4];
|
| 227 |
-
let qi = id % Lq;
|
| 228 |
-
let zh = id / Lq;
|
| 229 |
-
let z = zh / H;
|
| 230 |
-
let h = zh % H;
|
| 231 |
-
let qv = view(6u); let kv = view(12u); let vv = view(18u); let ov = view(24u);
|
| 232 |
-
let qb = (voff(qv, z, qi) + h * P[30]) / 4u;
|
| 233 |
-
let scale = pf(5u);
|
| 234 |
-
var q: array<vec4<f32>, {D4}>;
|
| 235 |
-
var acc: array<vec4<f32>, {D4}>;
|
| 236 |
-
for (var d = 0u; d < D4; d++) { q[d] = Q[qb + d] * scale; acc[d] = vec4<f32>(0.0); }
|
| 237 |
-
var m = -3.0e38;
|
| 238 |
-
var l = 0.0;
|
| 239 |
-
for (var j = 0u; j < Lk; j++) {
|
| 240 |
-
let kb = (voff(kv, z, j) + h * P[31]) / 4u;
|
| 241 |
-
var dot4 = vec4<f32>(0.0);
|
| 242 |
-
for (var d = 0u; d < D4; d++) { dot4 += q[d] * K[kb + d]; }
|
| 243 |
-
let s = dot4.x + dot4.y + dot4.z + dot4.w;
|
| 244 |
-
let mt = max(m, s);
|
| 245 |
-
let corr = exp(m - mt);
|
| 246 |
-
let p = exp(s - mt);
|
| 247 |
-
l = l * corr + p;
|
| 248 |
-
let vb = (voff(vv, z, j) + h * P[32]) / 4u;
|
| 249 |
-
for (var d = 0u; d < D4; d++) { acc[d] = acc[d] * corr + p * V[vb + d]; }
|
| 250 |
-
m = mt;
|
| 251 |
-
}
|
| 252 |
-
let ob = (voff(ov, z, qi) + h * P[33]) / 4u;
|
| 253 |
-
for (var d = 0u; d < D4; d++) { O[ob + d] = acc[d] / l; }
|
| 254 |
-
}
|
| 255 |
-
"""
|
| 256 |
-
|
| 257 |
-
# dst row r = src row r (W4 vec4 per row). P: [0] rows * W4, [1] grid x, [2] W4, [6..11] src, [12..17] dst
|
| 258 |
-
COPY = """
|
| 259 |
-
@group(0) @binding(1) var<storage, read> S: array<vec4<f32>>;
|
| 260 |
-
@group(0) @binding(2) var<storage, read_write> Dst: array<vec4<f32>>;
|
| 261 |
-
|
| 262 |
-
@compute @workgroup_size(64)
|
| 263 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 264 |
-
let id = wgid(wg) * 64u + li;
|
| 265 |
-
if (id >= P[0]) { return; }
|
| 266 |
-
let r = id / P[2];
|
| 267 |
-
let c = id % P[2];
|
| 268 |
-
Dst[roff(view(12u), r) / 4u + c] = S[roff(view(6u), r) / 4u + c];
|
| 269 |
-
}
|
| 270 |
-
"""
|
| 271 |
-
|
| 272 |
-
# In place on rows of W4 vec4: optional LayerNorm (weight, bias, eps), then optional + emb[slot[z]], where
|
| 273 |
-
# slot index = r / P[4] (the training row of a cell).
|
| 274 |
-
# P: [0] rows, [1] grid x, [2] W4, [3] eps, [4] rows per slot, [6..11] view
|
| 275 |
-
ROWNORM = """
|
| 276 |
-
const W4: u32 = {W4}u;
|
| 277 |
-
@group(0) @binding(1) var<storage, read_write> X: array<vec4<f32>>;
|
| 278 |
-
#if LN
|
| 279 |
-
@group(0) @binding(2) var<storage, read> Lw: array<vec4<f32>>;
|
| 280 |
-
@group(0) @binding(3) var<storage, read> Lb: array<vec4<f32>>;
|
| 281 |
-
#endif
|
| 282 |
-
#if EMB
|
| 283 |
-
@group(0) @binding(4) var<storage, read> Emb: array<vec4<f32>>;
|
| 284 |
-
@group(0) @binding(5) var<storage, read> Slot: array<u32>;
|
| 285 |
-
#endif
|
| 286 |
-
|
| 287 |
-
@compute @workgroup_size(64)
|
| 288 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 289 |
-
let r = wgid(wg) * 64u + li;
|
| 290 |
-
if (r >= P[0]) { return; }
|
| 291 |
-
let v = view(6u);
|
| 292 |
-
let b = roff(v, r) / 4u;
|
| 293 |
-
#if LN
|
| 294 |
-
var x: array<vec4<f32>, {W4}>;
|
| 295 |
-
var s = vec4<f32>(0.0);
|
| 296 |
-
for (var c = 0u; c < W4; c++) { x[c] = X[b + c]; s += x[c]; }
|
| 297 |
-
let mean = (s.x + s.y + s.z + s.w) / f32(W4 * 4u);
|
| 298 |
-
var q = vec4<f32>(0.0);
|
| 299 |
-
for (var c = 0u; c < W4; c++) { let d = x[c] - mean; q += d * d; }
|
| 300 |
-
let rstd = inverseSqrt((q.x + q.y + q.z + q.w) / f32(W4 * 4u) + pf(3u));
|
| 301 |
-
for (var c = 0u; c < W4; c++) { x[c] = (x[c] - mean) * rstd * Lw[c] + Lb[c]; }
|
| 302 |
-
#if EMB
|
| 303 |
-
let e = Slot[r / P[4]] * W4;
|
| 304 |
-
for (var c = 0u; c < W4; c++) { x[c] += Emb[e + c]; }
|
| 305 |
-
#endif
|
| 306 |
-
for (var c = 0u; c < W4; c++) { X[b + c] = x[c]; }
|
| 307 |
-
#else
|
| 308 |
-
let e = Slot[r / P[4]] * W4;
|
| 309 |
-
for (var c = 0u; c < W4; c++) { X[b + c] += Emb[e + c]; }
|
| 310 |
-
#endif
|
| 311 |
-
}
|
| 312 |
-
"""
|
| 313 |
-
|
| 314 |
-
# Cell features for the embedding GEMM, one thread per cell of a block of the (B, n, m) table: cells c in
|
| 315 |
-
# (b, ii, jj) order over B x nr rows (from row i0) x mc columns (from column j0). Inputs z, r, nan are the
|
| 316 |
-
# normalized values of the whole table; neighbors are the group offsets (j + o) % m.
|
| 317 |
-
# Output row c, KP floats: fourier [sum_g sin(z_g f), sum_g cos(z_g f)] (2F) or rbf kernels (64), then
|
| 318 |
-
# [z_g, nan_g, sin(r_g pi 2^e), cos(r_g pi 2^e)], zero padded.
|
| 319 |
-
# P: [0] cells, [1] grid x, [2] n, [3] m, [4] i0, [5] nr, [6] j0, [7] mc
|
| 320 |
-
FEATURES = """
|
| 321 |
-
const G: u32 = {G}u;
|
| 322 |
-
const NF: u32 = {NF}u;
|
| 323 |
-
const NE: u32 = {NE}u;
|
| 324 |
-
const KP: u32 = {KP}u;
|
| 325 |
-
const OFFS = array<u32, {G}>({OFFS});
|
| 326 |
-
@group(0) @binding(1) var<storage, read> Zt: array<f32>;
|
| 327 |
-
@group(0) @binding(2) var<storage, read> Rt: array<f32>;
|
| 328 |
-
@group(0) @binding(3) var<storage, read> Nt: array<f32>;
|
| 329 |
-
@group(0) @binding(4) var<storage, read> Fq: array<f32>;
|
| 330 |
-
@group(0) @binding(5) var<storage, read_write> Out: array<f32>;
|
| 331 |
-
|
| 332 |
-
@compute @workgroup_size(64)
|
| 333 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 334 |
-
let c = wgid(wg) * 64u + li;
|
| 335 |
-
if (c >= P[0]) { return; }
|
| 336 |
-
let n = P[2]; let m = P[3]; let nr = P[5]; let mc = P[7];
|
| 337 |
-
let b = c / (nr * mc);
|
| 338 |
-
let rem = c % (nr * mc);
|
| 339 |
-
let i = P[4] + rem / mc;
|
| 340 |
-
let j = P[6] + rem % mc;
|
| 341 |
-
let rowbase = (b * n + i) * m;
|
| 342 |
-
var z: array<f32, {G}>;
|
| 343 |
-
var r: array<f32, {G}>;
|
| 344 |
-
var nn: array<f32, {G}>;
|
| 345 |
-
for (var g = 0u; g < G; g++) {
|
| 346 |
-
let jg = (j + OFFS[g]) % m;
|
| 347 |
-
z[g] = Zt[rowbase + jg];
|
| 348 |
-
r[g] = Rt[rowbase + jg];
|
| 349 |
-
nn[g] = Nt[rowbase + jg];
|
| 350 |
-
}
|
| 351 |
-
let o = c * KP;
|
| 352 |
-
var k = 0u;
|
| 353 |
-
#if FOURIER
|
| 354 |
-
for (var f = 0u; f < NF; f++) {
|
| 355 |
-
var s = 0.0;
|
| 356 |
-
var co = 0.0;
|
| 357 |
-
for (var g = 0u; g < G; g++) {
|
| 358 |
-
let sc = sincos(z[g] * Fq[g * NF + f]);
|
| 359 |
-
s += sc.x;
|
| 360 |
-
co += sc.y;
|
| 361 |
-
}
|
| 362 |
-
Out[o + f] = s;
|
| 363 |
-
Out[o + NF + f] = co;
|
| 364 |
-
}
|
| 365 |
-
k = 2u * NF;
|
| 366 |
-
#else
|
| 367 |
-
for (var t = 0u; t < 64u; t++) {
|
| 368 |
-
var s = 0.0;
|
| 369 |
-
for (var g = 0u; g < G; g++) { let d = z[g] - Fq[t]; s += exp(-0.5 * d * d); }
|
| 370 |
-
Out[o + t] = s;
|
| 371 |
-
}
|
| 372 |
-
k = 64u;
|
| 373 |
-
#endif
|
| 374 |
-
for (var g = 0u; g < G; g++) { Out[o + k + g] = z[g]; Out[o + k + G + g] = nn[g]; }
|
| 375 |
-
k += 2u * G;
|
| 376 |
-
for (var g = 0u; g < G; g++) {
|
| 377 |
-
for (var e = 0u; e < NE; e++) {
|
| 378 |
-
let sc = sincos(r[g] * 3.141592653589793 * f32(1u << e));
|
| 379 |
-
Out[o + k + g * NE + e] = sc.x;
|
| 380 |
-
Out[o + k + G * NE + g * NE + e] = sc.y;
|
| 381 |
-
}
|
| 382 |
-
}
|
| 383 |
-
k += 2u * G * NE;
|
| 384 |
-
for (; k < KP; k++) { Out[o + k] = 0.0; }
|
| 385 |
-
}
|
| 386 |
-
"""
|
| 387 |
-
|
| 388 |
-
# Rotary embedding in place on adjacent pairs (2p, 2p + 1) of each head (the layout of LightPFN.folded()),
|
| 389 |
-
# for tokens i >= p0 of each sequence at position i - p0; cos/sin from a table T (pos, D / 2, 2).
|
| 390 |
-
# P: [0] rows * H * D4, [1] grid x, [2] H, [3] p0, [4] head stride, [6..11] view
|
| 391 |
-
ROPE = """
|
| 392 |
-
const D4: u32 = {D4}u;
|
| 393 |
-
@group(0) @binding(1) var<storage, read_write> X: array<vec4<f32>>;
|
| 394 |
-
@group(0) @binding(2) var<storage, read> T: array<vec4<f32>>;
|
| 395 |
-
|
| 396 |
-
@compute @workgroup_size(64)
|
| 397 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 398 |
-
let id = wgid(wg) * 64u + li;
|
| 399 |
-
if (id >= P[0]) { return; }
|
| 400 |
-
let v = view(6u);
|
| 401 |
-
let per_row = P[2] * D4;
|
| 402 |
-
let r = id / per_row;
|
| 403 |
-
let i = r % v.L;
|
| 404 |
-
if (i < P[3]) { return; }
|
| 405 |
-
let h = (id % per_row) / D4;
|
| 406 |
-
let d = id % D4;
|
| 407 |
-
let a = (roff(v, r) + h * P[4]) / 4u + d;
|
| 408 |
-
let x = X[a];
|
| 409 |
-
let cs = T[(i - P[3]) * D4 + d]; // (cos, sin) of pairs 2d and 2d + 1
|
| 410 |
-
X[a] = vec4<f32>(x.x * cs.x - x.y * cs.y, x.x * cs.y + x.y * cs.x,
|
| 411 |
-
x.z * cs.z - x.w * cs.w, x.z * cs.w + x.w * cs.z);
|
| 412 |
-
}
|
| 413 |
-
"""
|
| 414 |
-
|
| 415 |
-
# Softmax scaling of attention queries in place: q[k] *= base[k % HD] * (1 + tanh(mod[k])) (contiguous).
|
| 416 |
-
# P: [0] total vec4, [1] grid x, [2] HD / 4
|
| 417 |
-
QSCALE = """
|
| 418 |
-
@group(0) @binding(1) var<storage, read_write> Q: array<vec4<f32>>;
|
| 419 |
-
@group(0) @binding(2) var<storage, read> Md: array<vec4<f32>>;
|
| 420 |
-
@group(0) @binding(3) var<storage, read> Bs: array<vec4<f32>>;
|
| 421 |
-
|
| 422 |
-
@compute @workgroup_size(64)
|
| 423 |
-
fn main(@builtin(workgroup_id) wg: vec3u, @builtin(local_invocation_index) li: u32) {
|
| 424 |
-
let id = wgid(wg) * 64u + li;
|
| 425 |
-
if (id >= P[0]) { return; }
|
| 426 |
-
let mv = Md[id];
|
| 427 |
-
let t = vec4<f32>(tanh_(mv.x), tanh_(mv.y), tanh_(mv.z), tanh_(mv.w));
|
| 428 |
-
Q[id] = Q[id] * Bs[id % P[2]] * (1.0 + t);
|
| 429 |
-
}
|
| 430 |
-
"""
|
| 431 |
-
|
| 432 |
-
|
| 433 |
-
def render(src, flags=(), **consts):
|
| 434 |
-
"""Source with `#if NAME` / `#else` / `#endif` blocks resolved (NAME in flags) and {KEY} replaced."""
|
| 435 |
-
out, stack = [], []
|
| 436 |
-
for line in (COMMON + src).splitlines():
|
| 437 |
-
s = line.strip()
|
| 438 |
-
if s.startswith("#if "):
|
| 439 |
-
stack.append(s[4:].strip() in flags)
|
| 440 |
-
elif s == "#else":
|
| 441 |
-
if not stack:
|
| 442 |
-
raise ValueError("unmatched #else in shader template")
|
| 443 |
-
stack[-1] = not stack[-1]
|
| 444 |
-
elif s == "#endif":
|
| 445 |
-
if not stack:
|
| 446 |
-
raise ValueError("unmatched #endif in shader template")
|
| 447 |
-
stack.pop()
|
| 448 |
-
elif all(stack):
|
| 449 |
-
out.append(line)
|
| 450 |
-
if stack:
|
| 451 |
-
raise ValueError("unclosed #if in shader template")
|
| 452 |
-
code = "\n".join(out)
|
| 453 |
-
for k, v in consts.items():
|
| 454 |
-
code = code.replace("{" + k + "}", str(v))
|
| 455 |
-
if re.search(r"\{[A-Z][A-Z0-9_]*\}", code):
|
| 456 |
-
raise ValueError("missing shader template constant")
|
| 457 |
-
return code
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/lightpfn/vulkan/model.py
DELETED
|
@@ -1,548 +0,0 @@
|
|
| 1 |
-
"""LightPFN inference on a Vulkan GPU: the computation of LightPFN.folded().encode / predict_logits with
|
| 2 |
-
WGSL kernels. Column statistics and value normalization (sort, searchsorted) stay on the host in torch;
|
| 3 |
-
everything after the normalized values runs on the GPU, and the training-set context (column caches, ICL
|
| 4 |
-
key/value cache, decoder keys) stays in GPU memory between fit and predict.
|
| 5 |
-
"""
|
| 6 |
-
|
| 7 |
-
import copy
|
| 8 |
-
import math
|
| 9 |
-
from dataclasses import dataclass
|
| 10 |
-
|
| 11 |
-
import numpy as np
|
| 12 |
-
import torch
|
| 13 |
-
|
| 14 |
-
from lightpfn.model.layers import Block
|
| 15 |
-
from lightpfn.model.lightpfn import CellEmbedder
|
| 16 |
-
from lightpfn.vulkan import kernels as K
|
| 17 |
-
from lightpfn.vulkan.engine import Engine, View, f32bits, rows
|
| 18 |
-
|
| 19 |
-
GPU_CHUNK_CELLS = 1 << 20
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
def _eps(norm):
|
| 23 |
-
return torch.finfo(torch.float32).eps if norm.eps is None else norm.eps
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
def _pad4(n):
|
| 27 |
-
return (n + 3) // 4 * 4
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
@dataclass
|
| 31 |
-
class VulkanContext:
|
| 32 |
-
"""What prediction needs from the training rows; the caches are GPU buffers."""
|
| 33 |
-
|
| 34 |
-
stats: dict
|
| 35 |
-
col_kv: list
|
| 36 |
-
col_refine_kv: list | None
|
| 37 |
-
icl_kv: list
|
| 38 |
-
dec_k: object
|
| 39 |
-
onehot: object
|
| 40 |
-
B: int
|
| 41 |
-
n: int
|
| 42 |
-
m: int
|
| 43 |
-
n_classes: int
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
class _Block:
|
| 47 |
-
"""GPU weights of a folded Block (RMSNorm weights already inside the linear layers)."""
|
| 48 |
-
|
| 49 |
-
def __init__(self, eng, blk):
|
| 50 |
-
a = blk.attn
|
| 51 |
-
self.H, self.hd = a.n_heads, a.head_dim
|
| 52 |
-
self.E = a.q.weight.shape[1]
|
| 53 |
-
self.eps = (_eps(blk.norm_q), _eps(blk.norm_kv), _eps(blk.norm_ff))
|
| 54 |
-
w = lambda t: eng.upload(t.detach().float().cpu().numpy()) # noqa: E731
|
| 55 |
-
self.wq, self.wkv, self.wo = w(a.q.weight), w(a.kv.weight), w(a.out.weight)
|
| 56 |
-
hid = blk.mlp.fc1.weight.shape[0]
|
| 57 |
-
self.hid = _pad4(hid) # zero rows/columns: gelu(0) = 0 adds nothing
|
| 58 |
-
w1 = torch.zeros(self.hid, self.E)
|
| 59 |
-
w1[:hid] = blk.mlp.fc1.weight
|
| 60 |
-
w2 = torch.zeros(self.E, self.hid)
|
| 61 |
-
w2[:, :hid] = blk.mlp.fc2.weight
|
| 62 |
-
self.w1, self.w2 = w(w1), w(w2)
|
| 63 |
-
self.scaling = None
|
| 64 |
-
if a.scaling is not None:
|
| 65 |
-
self.scaling = _Scaling(eng, a.scaling)
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
class _Scaling:
|
| 69 |
-
"""SoftmaxScaling: base(log n) on the host (a vector per key count), mod(q) on the GPU."""
|
| 70 |
-
|
| 71 |
-
def __init__(self, eng, sc):
|
| 72 |
-
self.base_module = copy.deepcopy(sc.base).float().cpu().eval()
|
| 73 |
-
self.H, self.hd = sc.n_heads, sc.head_dim
|
| 74 |
-
w = lambda t: eng.upload(t.detach().float().cpu().numpy()) # noqa: E731
|
| 75 |
-
lin1, lin2 = sc.mod[0], sc.mod[2]
|
| 76 |
-
self.hidden = lin1.weight.shape[0]
|
| 77 |
-
assert self.hidden % 4 == 0
|
| 78 |
-
self.w1, self.b1, self.w2, self.b2 = w(lin1.weight), w(lin1.bias), w(lin2.weight), w(lin2.bias)
|
| 79 |
-
self.eng, self.bases = eng, {}
|
| 80 |
-
|
| 81 |
-
def base(self, n_keys):
|
| 82 |
-
buf = self.bases.get(n_keys)
|
| 83 |
-
if buf is None:
|
| 84 |
-
with torch.no_grad():
|
| 85 |
-
b = self.base_module(torch.full((1, 1), math.log(max(n_keys, 2)))).view(-1)
|
| 86 |
-
buf = self.bases[n_keys] = self.eng.upload(b.numpy())
|
| 87 |
-
return buf
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
class _ColumnStage:
|
| 91 |
-
def __init__(self, eng, st):
|
| 92 |
-
w = lambda t: eng.upload(t.detach().float().cpu().numpy()) # noqa: E731
|
| 93 |
-
self.y_emb = w(st.y_emb.weight)
|
| 94 |
-
self.ind = [w(p) for p in st.inducing]
|
| 95 |
-
self.ind_blocks = [_Block(eng, b) for b in st.ind_blocks]
|
| 96 |
-
self.cell_blocks = [_Block(eng, b) for b in st.cell_blocks]
|
| 97 |
-
with torch.no_grad(): # the inducing queries do not depend on the data
|
| 98 |
-
self.ind_q = [w(b.attn.q(b.norm_q(p))) for b, p in zip(st.ind_blocks, st.inducing)]
|
| 99 |
-
self.n_ind = st.inducing[0].shape[0]
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
class VulkanLightPFN:
|
| 103 |
-
"""LightPFN on a Vulkan GPU with the encode / predict_logits interface of LightPFN (inputs and logits
|
| 104 |
-
are CPU torch tensors). model: a LightPFN, folded or not. adapter: see engine.pick_adapter."""
|
| 105 |
-
|
| 106 |
-
def __init__(self, model, adapter=None):
|
| 107 |
-
cfg = model.cfg
|
| 108 |
-
for dim, heads in ((cfg.col_dim, cfg.col_heads), (cfg.col_dim, cfg.row_heads),
|
| 109 |
-
(cfg.icl_dim, cfg.icl_heads), (cfg.icl_dim, cfg.decoder_heads)):
|
| 110 |
-
if heads <= 0 or dim <= 0 or dim % (4 * heads):
|
| 111 |
-
raise NotImplementedError("Vulkan attention head dimensions must be positive multiples of four")
|
| 112 |
-
if not 0 <= cfg.n_ecdf_freq < 32:
|
| 113 |
-
raise NotImplementedError("Vulkan n_ecdf_freq must be between 0 and 31")
|
| 114 |
-
if cfg.n_inducing <= 0:
|
| 115 |
-
raise NotImplementedError("Vulkan inducing token count must be positive")
|
| 116 |
-
m = model.folded()
|
| 117 |
-
self.cfg = cfg = m.cfg
|
| 118 |
-
self.eng = eng = Engine(adapter)
|
| 119 |
-
self.E, self.D, self.C = cfg.col_dim, cfg.icl_dim, cfg.n_cls
|
| 120 |
-
if cfg.icl_kv_heads_test not in (1, cfg.icl_heads):
|
| 121 |
-
raise NotImplementedError("icl_kv_heads_test must be 1 or icl_heads")
|
| 122 |
-
w = lambda t: eng.upload(t.detach().float().cpu().numpy()) # noqa: E731
|
| 123 |
-
ce = m.cells
|
| 124 |
-
self.offsets, self.n_ecdf = tuple(ce.offsets), ce.n_ecdf_freq
|
| 125 |
-
G = len(self.offsets)
|
| 126 |
-
meta = G * (2 + 2 * self.n_ecdf)
|
| 127 |
-
if cfg.cell_embed == "fourier":
|
| 128 |
-
self.fourier, self.n_freq = True, ce.freq.shape[1]
|
| 129 |
-
value_w, self.fq = ce.fourier.weight, w(ce.freq)
|
| 130 |
-
else:
|
| 131 |
-
self.fourier, self.n_freq = False, 64
|
| 132 |
-
value_w, self.fq = ce.rbf.weight, w(ce.rbf_centers)
|
| 133 |
-
k = value_w.shape[1] + meta
|
| 134 |
-
self.kp = _pad4(k)
|
| 135 |
-
wc = torch.zeros(self.E, self.kp)
|
| 136 |
-
wc[:, :k] = torch.cat([value_w, ce.meta.weight], 1)
|
| 137 |
-
self.w_cells, self.ln_w, self.ln_b, self.ln_eps = w(wc), w(ce.norm.weight), w(ce.norm.bias), ce.norm.eps
|
| 138 |
-
self.col = _ColumnStage(eng, m.col)
|
| 139 |
-
self.row_mode, self.rope_base = m.row.mode, m.row.rope_base
|
| 140 |
-
self.cls = w(m.row.cls)
|
| 141 |
-
self.row_blocks = [_Block(eng, b) for b in m.row.blocks]
|
| 142 |
-
self.refine = None
|
| 143 |
-
if cfg.row_refine:
|
| 144 |
-
self.refine = dict(summary=w(m.refine.summary), n_summary=m.refine.summary.shape[0],
|
| 145 |
-
broadcast=[_Block(eng, b) for b in m.refine.broadcast],
|
| 146 |
-
gather=[_Block(eng, b) for b in m.refine.gather])
|
| 147 |
-
self.col_refine = _ColumnStage(eng, m.col_refine)
|
| 148 |
-
icl = m.icl
|
| 149 |
-
self.icl_y_emb, self.thinking = w(icl.y_emb.weight), w(icl.thinking)
|
| 150 |
-
self.n_thinking = icl.thinking.shape[0]
|
| 151 |
-
self.icl_blocks = [_Block(eng, b) for b in icl.blocks]
|
| 152 |
-
self.kv_heads = icl.kv_heads_test
|
| 153 |
-
dec = m.decoder
|
| 154 |
-
self.dec_H, self.dec_hd = dec.n_heads, dec.head_dim
|
| 155 |
-
self.dec_q, self.dec_k = w(dec.q.weight), w(dec.k.weight)
|
| 156 |
-
self.dec_eps = _eps(icl.norm)
|
| 157 |
-
self.dec_scaling = _Scaling(eng, dec.scaling)
|
| 158 |
-
self.rope_tables = {}
|
| 159 |
-
# Cell limits are upper bounds; train_capacity and the piece costs also bound caches and scratch.
|
| 160 |
-
self.max_cells = eng.max_binding // (4 * 2 * self.E)
|
| 161 |
-
self.max_train_cells = eng.max_binding // (4 * self.E)
|
| 162 |
-
|
| 163 |
-
def _piece_costs(self, n, m):
|
| 164 |
-
"""Floats per estimator in one column / one row piece."""
|
| 165 |
-
stages = [self.col] + ([self.col_refine] if self.refine else [])
|
| 166 |
-
col_width = max(self.kp, 2 * self.E,
|
| 167 |
-
*(b.hid for s in stages for b in s.ind_blocks + s.cell_blocks))
|
| 168 |
-
row_blocks = self.row_blocks + (self.refine["broadcast"] + self.refine["gather"] if self.refine else [])
|
| 169 |
-
row_width = max(2 * self.E, *(b.hid for b in row_blocks))
|
| 170 |
-
S = self.refine["n_summary"] if self.refine else 0
|
| 171 |
-
return max(n, self.col.n_ind) * col_width, max(max(m + self.C, S) * row_width, m * self.kp)
|
| 172 |
-
|
| 173 |
-
def _icl_width(self):
|
| 174 |
-
return max(2 * self.D, self.dec_H * self.dec_scaling.hidden,
|
| 175 |
-
*(b.hid for b in self.icl_blocks),
|
| 176 |
-
*(b.H * b.scaling.hidden for b in self.icl_blocks))
|
| 177 |
-
|
| 178 |
-
def train_capacity(self, n, m):
|
| 179 |
-
"""Estimators whose persistent buffers and smallest pieces fit the binding limit."""
|
| 180 |
-
col_piece, row_piece = self._piece_costs(n, m)
|
| 181 |
-
largest = max(n * m * self.E, m * self.col.n_ind * 2 * self.E,
|
| 182 |
-
(n + self.n_thinking) * self._icl_width(), col_piece, row_piece)
|
| 183 |
-
return min(self.max_train_cells // max(1, n * m), self.eng.max_binding // (4 * max(1, largest)))
|
| 184 |
-
|
| 185 |
-
# kernels -------------------------------------------------------------------------------------
|
| 186 |
-
def gemm(self, x, y, w, N, Kd, O, norm_eps=None, bias=None, gelu=False, res=None):
|
| 187 |
-
"""y rows = epilogue(x rows @ w^T) for N rows; res: None, "self" or a View added to the result."""
|
| 188 |
-
flags = []
|
| 189 |
-
if norm_eps is not None:
|
| 190 |
-
flags.append("NORM")
|
| 191 |
-
if bias is not None:
|
| 192 |
-
flags.append("BIAS")
|
| 193 |
-
if gelu:
|
| 194 |
-
flags.append("GELU")
|
| 195 |
-
if res == "self":
|
| 196 |
-
flags.append("RES_SELF")
|
| 197 |
-
elif res is not None:
|
| 198 |
-
flags.append("RES_R")
|
| 199 |
-
assert Kd % 4 == 0 and O % 4 == 0, (Kd, O)
|
| 200 |
-
tiles = math.ceil(N / 64) * math.ceil(O / 64)
|
| 201 |
-
rv = res if isinstance(res, View) else y
|
| 202 |
-
params = [tiles, 0, N, Kd, O, f32bits(norm_eps or 0.0)] + x.params() + y.params() + rv.params()
|
| 203 |
-
bufs = {1: x.buf, 2: w, 3: y.buf}
|
| 204 |
-
if bias is not None:
|
| 205 |
-
bufs[4] = bias
|
| 206 |
-
if isinstance(res, View):
|
| 207 |
-
bufs[5] = res.buf
|
| 208 |
-
self.eng.dispatch(self.eng.pipeline("gemm", K.GEMM, flags), bufs, params, tiles, 2.0 * N * Kd * O)
|
| 209 |
-
|
| 210 |
-
def attention(self, q, k, v, o, Z, H, Lq, Lk, hd, q_h, k_h, v_h, o_h):
|
| 211 |
-
if Z == 0 or Lq == 0:
|
| 212 |
-
return
|
| 213 |
-
if Lk == 0:
|
| 214 |
-
width = H * hd
|
| 215 |
-
zero = self.eng.upload(np.zeros((Z * Lq, width), np.float32))
|
| 216 |
-
self.copy(rows(zero, width), o, Z * Lq, width)
|
| 217 |
-
return
|
| 218 |
-
small = Lq < 32 or Lk < 16
|
| 219 |
-
D4 = hd // 4
|
| 220 |
-
if small:
|
| 221 |
-
pipe = self.eng.pipeline("attn_small", K.ATTN_SMALL, D4=D4)
|
| 222 |
-
total = Z * H * Lq
|
| 223 |
-
groups = math.ceil(total / 64)
|
| 224 |
-
else:
|
| 225 |
-
kt = min(64, 16384 // (2 * hd * 4))
|
| 226 |
-
pipe = self.eng.pipeline("attn_tiled", K.ATTN_TILED, D4=D4, KT=kt)
|
| 227 |
-
total = math.ceil(Lq / 64) * Z * H
|
| 228 |
-
groups = total
|
| 229 |
-
params = [total, 0, Lq, Lk, H, f32bits(hd ** -0.5)] + q.params() + k.params() + v.params() + o.params()
|
| 230 |
-
params += [q_h, k_h, v_h, o_h]
|
| 231 |
-
self.eng.dispatch(pipe, {1: q.buf, 2: k.buf, 3: v.buf, 4: o.buf}, params, groups, 4.0 * Z * H * Lq * Lk * hd)
|
| 232 |
-
|
| 233 |
-
def copy(self, src, dst, N, width):
|
| 234 |
-
W4 = width // 4
|
| 235 |
-
params = [N * W4, 0, W4, 0, 0, 0] + src.params() + dst.params()
|
| 236 |
-
self.eng.dispatch(self.eng.pipeline("copy", K.COPY), {1: src.buf, 2: dst.buf}, params, math.ceil(N * W4 / 64))
|
| 237 |
-
|
| 238 |
-
def rownorm(self, x, N, width, ln=None, emb=None, slots=None, slot_div=1):
|
| 239 |
-
"""In place on N rows: LayerNorm (ln = (w, b, eps)) then + emb[slots[r // slot_div]]."""
|
| 240 |
-
flags = (["LN"] if ln else []) + (["EMB"] if emb is not None else [])
|
| 241 |
-
params = [N, 0, width // 4, f32bits(ln[2] if ln else 0.0), slot_div, 0] + x.params()
|
| 242 |
-
bufs = {1: x.buf}
|
| 243 |
-
if ln:
|
| 244 |
-
bufs[2], bufs[3] = ln[0], ln[1]
|
| 245 |
-
if emb is not None:
|
| 246 |
-
bufs[4], bufs[5] = emb, slots
|
| 247 |
-
self.eng.dispatch(self.eng.pipeline("rownorm", K.ROWNORM, flags, W4=width // 4), bufs, params, math.ceil(N / 64))
|
| 248 |
-
|
| 249 |
-
def rope(self, x, N, H, hd, head_stride, p0, L):
|
| 250 |
-
tab = self.rope_table(L, hd)
|
| 251 |
-
D4 = hd // 4
|
| 252 |
-
params = [N * H * D4, 0, H, p0, head_stride, 0] + x.params()
|
| 253 |
-
self.eng.dispatch(self.eng.pipeline("rope", K.ROPE, D4=D4), {1: x.buf, 2: tab}, params, math.ceil(N * H * D4 / 64))
|
| 254 |
-
|
| 255 |
-
def rope_table(self, L, hd):
|
| 256 |
-
"""cos/sin of positions 0..L-1 for the adjacent pairs, computed as torch's rope() does."""
|
| 257 |
-
key = (hd, L)
|
| 258 |
-
tab = self.rope_tables.get(key)
|
| 259 |
-
if tab is None:
|
| 260 |
-
for (d, n), t in self.rope_tables.items(): # a longer table of the same head dim serves too
|
| 261 |
-
if d == hd and n >= L:
|
| 262 |
-
return t
|
| 263 |
-
inv_freq = self.rope_base ** (-torch.arange(0, hd, 2, dtype=torch.float32) / hd)
|
| 264 |
-
ang = torch.arange(max(L, 1), dtype=torch.float32)[:, None] * inv_freq[None]
|
| 265 |
-
tab = self.rope_tables[key] = self.eng.upload(torch.stack([ang.cos(), ang.sin()], -1).numpy())
|
| 266 |
-
return tab
|
| 267 |
-
|
| 268 |
-
def qscale(self, q, N, sc, n_keys):
|
| 269 |
-
"""Queries q (N rows of H * hd, contiguous) times base(log n_keys) * (1 + tanh(mod(q)))."""
|
| 270 |
-
rows_h = N * sc.H
|
| 271 |
-
mh = self.eng.temp("mod_h", rows_h * sc.hidden)
|
| 272 |
-
md = self.eng.temp("mod_o", rows_h * sc.hd)
|
| 273 |
-
self.gemm(rows(q, sc.hd), rows(mh, sc.hidden), sc.w1, rows_h, sc.hd, sc.hidden, bias=sc.b1, gelu=True)
|
| 274 |
-
self.gemm(rows(mh, sc.hidden), rows(md, sc.hd), sc.w2, rows_h, sc.hidden, sc.hd, bias=sc.b2)
|
| 275 |
-
total = rows_h * sc.hd // 4
|
| 276 |
-
params = [total, 0, sc.H * sc.hd // 4]
|
| 277 |
-
self.eng.dispatch(self.eng.pipeline("qscale", K.QSCALE), {1: q, 2: md, 3: sc.base(n_keys)}, params,
|
| 278 |
-
math.ceil(total / 64))
|
| 279 |
-
|
| 280 |
-
def mlp(self, blk, x, N):
|
| 281 |
-
hid = self.eng.temp("hid", N * blk.hid)
|
| 282 |
-
self.gemm(x, rows(hid, blk.hid), blk.w1, N, blk.E, blk.hid, norm_eps=blk.eps[2], gelu=True)
|
| 283 |
-
self.gemm(rows(hid, blk.hid), x, blk.w2, N, blk.hid, blk.E, res="self")
|
| 284 |
-
|
| 285 |
-
# stages --------------------------------------------------------------------------------------
|
| 286 |
-
def embed(self, zrn, x, B, n, m, i0, nr, j0, mc):
|
| 287 |
-
"""Cell embeddings of rows i0..i0+nr, columns j0..j0+mc of the (B, n, m) table into view x."""
|
| 288 |
-
N = B * nr * mc
|
| 289 |
-
feat = self.eng.temp("feat", N * self.kp)
|
| 290 |
-
G = len(self.offsets)
|
| 291 |
-
consts = dict(G=G, NF=self.n_freq, NE=self.n_ecdf, KP=self.kp,
|
| 292 |
-
OFFS=", ".join(f"{o % m}u" for o in self.offsets))
|
| 293 |
-
pipe = self.eng.pipeline("features", K.FEATURES, ["FOURIER"] if self.fourier else [], **consts)
|
| 294 |
-
params = [N, 0, n, m, i0, nr, j0, mc]
|
| 295 |
-
self.eng.dispatch(pipe, {1: zrn[0], 2: zrn[1], 3: zrn[2], 4: self.fq, 5: feat}, params, math.ceil(N / 64))
|
| 296 |
-
self.gemm(rows(feat, self.kp), x, self.w_cells, N, self.kp, self.E)
|
| 297 |
-
|
| 298 |
-
def column_context(self, st, main, B, n, m, cols, slots, zrn=None):
|
| 299 |
-
"""ColumnStage.context on pieces of `cols` columns of the (B, n, m, E) buffer `main` (embedded
|
| 300 |
-
first when zrn is given); returns the (B * m, n_ind, 2E) key/value caches of the cell blocks."""
|
| 301 |
-
E, I = self.E, st.n_ind
|
| 302 |
-
caches = [self.eng.empty(B * m * I * 2 * E) for _ in st.cell_blocks]
|
| 303 |
-
for j0 in range(0, m, cols):
|
| 304 |
-
mc = min(cols, m - j0)
|
| 305 |
-
N = B * n * mc
|
| 306 |
-
x = View(main, j0 * E, L=mc, zin=1, so=m * E, si=0, ss=E) # rows (b * n + i, jj)
|
| 307 |
-
ln = None
|
| 308 |
-
if zrn is not None:
|
| 309 |
-
self.embed(zrn, x, B, n, m, 0, n, j0, mc)
|
| 310 |
-
ln = (self.ln_w, self.ln_b, self.ln_eps)
|
| 311 |
-
self.rownorm(x, N, E, ln=ln, emb=st.y_emb, slots=slots, slot_div=mc)
|
| 312 |
-
for ind, ind_q, ib, cb, cache in zip(st.ind, st.ind_q, st.ind_blocks, st.cell_blocks, caches):
|
| 313 |
-
H, hd = ib.H, ib.hd
|
| 314 |
-
# inducing points read the cells of their column: Z = B * mc sequences of n keys
|
| 315 |
-
kv = self.eng.temp("kv", N * 2 * E)
|
| 316 |
-
self.gemm(x, rows(kv, 2 * E), ib.wkv, N, E, 2 * E, norm_eps=ib.eps[1])
|
| 317 |
-
kview = lambda base: View(kv, base, L=n, zin=mc, so=n * mc * 2 * E, si=2 * E, ss=mc * 2 * E) # noqa: E731
|
| 318 |
-
o = self.eng.temp("o_ind", B * mc * I * E)
|
| 319 |
-
self.attention(View(ind_q, 0, L=I, zin=1, so=0, si=0, ss=E), kview(0), kview(E),
|
| 320 |
-
View(o, 0, L=I, zin=1, so=I * E, ss=E), B * mc, H, I, n, hd, hd, hd, hd, hd)
|
| 321 |
-
h = self.eng.temp("h_ind", B * mc * I * E)
|
| 322 |
-
self.gemm(rows(o, E), rows(h, E), ib.wo, B * mc * I, E, E,
|
| 323 |
-
res=View(ind, 0, L=I, zin=1, so=0, si=0, ss=E))
|
| 324 |
-
self.mlp(ib, rows(h, E), B * mc * I)
|
| 325 |
-
cview = lambda base: View(cache, j0 * I * 2 * E + base, L=I, zin=mc, so=m * I * 2 * E, # noqa: E731
|
| 326 |
-
si=I * 2 * E, ss=2 * E)
|
| 327 |
-
self.gemm(rows(h, E), cview(0), cb.wkv, B * mc * I, E, 2 * E, norm_eps=cb.eps[1])
|
| 328 |
-
self.cell_block(cb, x, N, B * mc, n, mc, n * mc, cview(0), cview(E), I)
|
| 329 |
-
return caches
|
| 330 |
-
|
| 331 |
-
def cell_block(self, cb, x, N, Z, L, zin, so_rows, kview, vview, I):
|
| 332 |
-
"""Cells attend to the inducing states of their column (Block.forward with cached keys). x rows
|
| 333 |
-
are in (z_outer * L + i, jj) order with zin columns per sequence group."""
|
| 334 |
-
E = self.E
|
| 335 |
-
q = self.eng.temp("q", N * E)
|
| 336 |
-
self.gemm(x, rows(q, E), cb.wq, N, E, E, norm_eps=cb.eps[0])
|
| 337 |
-
qv = View(q, 0, L=L, zin=zin, so=so_rows * E, si=E, ss=zin * E)
|
| 338 |
-
o = self.eng.temp("o", N * E)
|
| 339 |
-
ov = View(o, 0, L=L, zin=zin, so=so_rows * E, si=E, ss=zin * E)
|
| 340 |
-
self.attention(qv, kview, vview, ov, Z, cb.H, L, I, cb.hd, cb.hd, cb.hd, cb.hd, cb.hd)
|
| 341 |
-
self.gemm(rows(o, E), x, cb.wo, N, E, E, res="self")
|
| 342 |
-
self.mlp(cb, x, N)
|
| 343 |
-
|
| 344 |
-
def column_query(self, st, caches, x, B, nr, m):
|
| 345 |
-
"""ColumnStage.query on test cells x: rows (b * nr + ii, j), contiguous (B, nr, m, E)."""
|
| 346 |
-
E, I = self.E, st.n_ind
|
| 347 |
-
N = B * nr * m
|
| 348 |
-
for cb, cache in zip(st.cell_blocks, caches):
|
| 349 |
-
cview = lambda base: View(cache, base, L=I, zin=1, so=I * 2 * E, si=0, ss=2 * E) # noqa: E731
|
| 350 |
-
self.cell_block(cb, x, N, B * m, nr, m, nr * m, cview(0), cview(E), I)
|
| 351 |
-
|
| 352 |
-
def refine_rows(self, x, R, m):
|
| 353 |
-
"""RowRefinement on R rows of m cells, x a view with L = m (row z, cell i)."""
|
| 354 |
-
E, rf = self.E, self.refine
|
| 355 |
-
S = rf["n_summary"]
|
| 356 |
-
N = R * m
|
| 357 |
-
s = self.eng.temp("s", R * S * E)
|
| 358 |
-
self.copy(View(rf["summary"], 0, L=S, zin=1, so=0, si=0, ss=E), rows(s, E), R * S, E)
|
| 359 |
-
for i, bc in enumerate(rf["broadcast"]):
|
| 360 |
-
H, hd = bc.H, bc.hd
|
| 361 |
-
kvs = self.eng.temp("kvs", R * S * 2 * E)
|
| 362 |
-
self.gemm(rows(s, E), rows(kvs, 2 * E), bc.wkv, R * S, E, 2 * E, norm_eps=bc.eps[1])
|
| 363 |
-
q = self.eng.temp("q", N * E)
|
| 364 |
-
self.gemm(x, rows(q, E), bc.wq, N, E, E, norm_eps=bc.eps[0])
|
| 365 |
-
qv = View(q, 0, L=m, zin=1, so=m * E, ss=E)
|
| 366 |
-
self.rope(qv, N, H, hd, hd, 0, m)
|
| 367 |
-
o = self.eng.temp("o", N * E)
|
| 368 |
-
kv_s = lambda base: View(kvs, base, L=S, zin=1, so=S * 2 * E, ss=2 * E) # noqa: E731
|
| 369 |
-
self.attention(qv, kv_s(0), kv_s(E), View(o, 0, L=m, zin=1, so=m * E, ss=E), R, H, m, S, hd, hd, hd, hd, hd)
|
| 370 |
-
self.gemm(rows(o, E), x, bc.wo, N, E, E, res="self")
|
| 371 |
-
self.mlp(bc, x, N)
|
| 372 |
-
if i < len(rf["gather"]):
|
| 373 |
-
g = rf["gather"][i]
|
| 374 |
-
kvx = self.eng.temp("kv", N * 2 * E)
|
| 375 |
-
self.gemm(x, rows(kvx, 2 * E), g.wkv, N, E, 2 * E, norm_eps=g.eps[1])
|
| 376 |
-
kx = lambda base: View(kvx, base, L=m, zin=1, so=m * 2 * E, ss=2 * E) # noqa: E731
|
| 377 |
-
self.rope(kx(0), N, H, hd, hd, 0, m)
|
| 378 |
-
qs = self.eng.temp("qs", R * S * E)
|
| 379 |
-
self.gemm(rows(s, E), rows(qs, E), g.wq, R * S, E, E, norm_eps=g.eps[0])
|
| 380 |
-
os_ = self.eng.temp("os", R * S * E)
|
| 381 |
-
sv = lambda buf: View(buf, 0, L=S, zin=1, so=S * E, ss=E) # noqa: E731
|
| 382 |
-
self.attention(sv(qs), kx(0), kx(E), sv(os_), R, H, S, m, hd, hd, hd, hd, hd)
|
| 383 |
-
self.gemm(rows(os_, E), rows(s, E), g.wo, R * S, E, E, res="self")
|
| 384 |
-
self.mlp(g, rows(s, E), R * S)
|
| 385 |
-
|
| 386 |
-
def row_stage(self, x, R, m, out):
|
| 387 |
-
"""RowStage on R rows of m cells (x a view with L = m); writes the R rows of C * E floats to view out."""
|
| 388 |
-
E, C = self.E, self.C
|
| 389 |
-
T = C + m
|
| 390 |
-
xr = self.eng.temp("xr", R * T * E)
|
| 391 |
-
self.copy(View(self.cls, 0, L=C, zin=1, so=0, si=0, ss=E), View(xr, 0, L=C, zin=1, so=T * E, ss=E), R * C, E)
|
| 392 |
-
self.copy(x, View(xr, C * E, L=m, zin=1, so=T * E, ss=E), R * m, E)
|
| 393 |
-
tok = lambda buf, w, base=0: View(buf, base, L=T, zin=1, so=T * w, ss=w) # noqa: E731
|
| 394 |
-
for blk in self.row_blocks:
|
| 395 |
-
H, hd = blk.H, blk.hd
|
| 396 |
-
kv = self.eng.temp("kv", R * T * 2 * E)
|
| 397 |
-
self.gemm(rows(xr, E), rows(kv, 2 * E), blk.wkv, R * T, E, 2 * E, norm_eps=blk.eps[1])
|
| 398 |
-
self.rope(tok(kv, 2 * E), R * T, H, hd, hd, C, m)
|
| 399 |
-
if self.row_mode == "summary":
|
| 400 |
-
sq = View(xr, 0, L=C, zin=1, so=T * E, ss=E)
|
| 401 |
-
q = self.eng.temp("q", R * C * E)
|
| 402 |
-
self.gemm(sq, rows(q, E), blk.wq, R * C, E, E, norm_eps=blk.eps[0])
|
| 403 |
-
o = self.eng.temp("o", R * C * E)
|
| 404 |
-
cv = lambda buf: View(buf, 0, L=C, zin=1, so=C * E, ss=E) # noqa: E731
|
| 405 |
-
self.attention(cv(q), tok(kv, 2 * E), tok(kv, 2 * E, E), cv(o), R, H, C, T, hd, hd, hd, hd, hd)
|
| 406 |
-
self.gemm(rows(o, E), sq, blk.wo, R * C, E, E, res="self")
|
| 407 |
-
self.mlp(blk, sq, R * C)
|
| 408 |
-
else:
|
| 409 |
-
q = self.eng.temp("q", R * T * E)
|
| 410 |
-
self.gemm(rows(xr, E), rows(q, E), blk.wq, R * T, E, E, norm_eps=blk.eps[0])
|
| 411 |
-
self.rope(tok(q, E), R * T, H, hd, hd, C, m)
|
| 412 |
-
o = self.eng.temp("o", R * T * E)
|
| 413 |
-
self.attention(tok(q, E), tok(kv, 2 * E), tok(kv, 2 * E, E), tok(o, E), R, H, T, T, hd, hd, hd, hd, hd)
|
| 414 |
-
self.gemm(rows(o, E), rows(xr, E), blk.wo, R * T, E, E, res="self")
|
| 415 |
-
self.mlp(blk, rows(xr, E), R * T)
|
| 416 |
-
self.copy(View(xr, 0, L=1, zin=1, so=T * E), out, R, C * E)
|
| 417 |
-
|
| 418 |
-
def icl_block(self, blk, x, B, Lq, kview, vview, Lk, k_h, v_h, kv_heads_cache=None):
|
| 419 |
-
D = self.D
|
| 420 |
-
N = B * Lq
|
| 421 |
-
q = self.eng.temp("q", N * D)
|
| 422 |
-
self.gemm(rows(x, D), rows(q, D), blk.wq, N, D, D, norm_eps=blk.eps[0])
|
| 423 |
-
self.qscale(q, N, blk.scaling, Lk)
|
| 424 |
-
o = self.eng.temp("o", N * D)
|
| 425 |
-
qv = lambda buf: View(buf, 0, L=Lq, zin=1, so=Lq * D, ss=D) # noqa: E731
|
| 426 |
-
self.attention(qv(q), kview, vview, qv(o), B, blk.H, Lq, Lk, blk.hd, blk.hd, k_h, v_h, blk.hd)
|
| 427 |
-
self.gemm(rows(o, D), rows(x, D), blk.wo, N, D, D, res="self")
|
| 428 |
-
self.mlp(blk, rows(x, D), N)
|
| 429 |
-
|
| 430 |
-
# interface -----------------------------------------------------------------------------------
|
| 431 |
-
def _upload_normalized(self, X, stats):
|
| 432 |
-
z, r, nan = CellEmbedder.normalize(X, stats)
|
| 433 |
-
return tuple(self.eng.upload(t.contiguous().numpy()) for t in (z, r, nan))
|
| 434 |
-
|
| 435 |
-
def _chunk(self, chunk_cells):
|
| 436 |
-
return max(1, min(chunk_cells or GPU_CHUNK_CELLS, self.max_cells))
|
| 437 |
-
|
| 438 |
-
@torch.inference_mode()
|
| 439 |
-
def encode(self, X_train, y_train, d=None, slots=None, n_classes=None, has_padding=None, chunk_cells=None):
|
| 440 |
-
if d is not None:
|
| 441 |
-
raise NotImplementedError("the Vulkan backend does not take zero-padded feature counts (d)")
|
| 442 |
-
X = torch.as_tensor(X_train).float().cpu()
|
| 443 |
-
y = torch.as_tensor(y_train).long().cpu()
|
| 444 |
-
B, n, m = X.shape
|
| 445 |
-
if B == 0 or m == 0:
|
| 446 |
-
raise ValueError("Vulkan encode needs a nonempty batch and at least one feature")
|
| 447 |
-
if n == 0 and n_classes is None:
|
| 448 |
-
raise ValueError("n_classes is required for an empty training context")
|
| 449 |
-
n_classes = int(y.max()) + 1 if n_classes is None else n_classes
|
| 450 |
-
if n_classes > self.cfg.label_slots:
|
| 451 |
-
raise ValueError(f"{n_classes} classes, the model supports {self.cfg.label_slots}")
|
| 452 |
-
slots = torch.arange(self.cfg.label_slots).expand(B, -1) if slots is None else torch.as_tensor(slots).long().cpu()
|
| 453 |
-
y_slots = torch.gather(slots, 1, y)
|
| 454 |
-
if B > self.train_capacity(n, m):
|
| 455 |
-
raise MemoryError(f"{B} x {n} x {m} training cells exceed a Vulkan cache/scratch buffer limit")
|
| 456 |
-
eng, E, D = self.eng, self.E, self.D
|
| 457 |
-
stats = CellEmbedder.stats(X)
|
| 458 |
-
zrn = self._upload_normalized(X, stats)
|
| 459 |
-
slot_buf = eng.upload(y_slots.numpy().reshape(-1), np.uint32)
|
| 460 |
-
chunk = self._chunk(chunk_cells)
|
| 461 |
-
col_cost, row_cost = self._piece_costs(n, m)
|
| 462 |
-
cols = min(max(1, chunk // (B * max(n, 1))), eng.max_binding // (4 * B * col_cost))
|
| 463 |
-
nrows = min(max(1, chunk // (B * m)), eng.max_binding // (4 * B * row_cost))
|
| 464 |
-
main = eng.empty(B * n * m * E)
|
| 465 |
-
col_kv = self.column_context(self.col, main, B, n, m, cols, slot_buf, zrn)
|
| 466 |
-
row_view = lambda i0, nr: View(main, i0 * m * E, L=m, zin=nr, so=n * m * E, si=m * E, ss=E) # noqa: E731
|
| 467 |
-
col_refine_kv = None
|
| 468 |
-
if self.refine is not None:
|
| 469 |
-
for i0 in range(0, n, nrows):
|
| 470 |
-
nr = min(nrows, n - i0)
|
| 471 |
-
self.refine_rows(row_view(i0, nr), B * nr, m)
|
| 472 |
-
col_refine_kv = self.column_context(self.col_refine, main, B, n, m, cols, slot_buf)
|
| 473 |
-
rows_buf = eng.empty(B * n * D)
|
| 474 |
-
for i0 in range(0, n, nrows):
|
| 475 |
-
nr = min(nrows, n - i0)
|
| 476 |
-
self.row_stage(row_view(i0, nr), B * nr, m, View(rows_buf, i0 * D, L=1, zin=nr, so=n * D, si=D))
|
| 477 |
-
del main
|
| 478 |
-
# ICL over thinking rows + training rows
|
| 479 |
-
T = self.n_thinking
|
| 480 |
-
L = T + n
|
| 481 |
-
x = eng.empty(B * L * D)
|
| 482 |
-
self.copy(View(self.thinking, 0, L=T, zin=1, so=0, si=0, ss=D), View(x, 0, L=T, zin=1, so=L * D, ss=D), B * T, D)
|
| 483 |
-
train = View(x, T * D, L=n, zin=1, so=L * D, ss=D)
|
| 484 |
-
self.copy(rows(rows_buf, D), train, B * n, D)
|
| 485 |
-
self.rownorm(train, B * n, D, emb=self.icl_y_emb, slots=slot_buf, slot_div=1)
|
| 486 |
-
del rows_buf
|
| 487 |
-
kvh = self.kv_heads * self.icl_blocks[0].hd
|
| 488 |
-
icl_kv = []
|
| 489 |
-
for blk in self.icl_blocks:
|
| 490 |
-
kv = eng.temp("kv_icl", B * L * 2 * D)
|
| 491 |
-
self.gemm(rows(x, D), rows(kv, 2 * D), blk.wkv, B * L, D, 2 * D, norm_eps=blk.eps[1])
|
| 492 |
-
cache = eng.empty(B * L * 2 * kvh)
|
| 493 |
-
self.copy(rows(kv, 2 * D), rows(cache, 2 * kvh), B * L, kvh)
|
| 494 |
-
self.copy(rows(kv, 2 * D, D), rows(cache, 2 * kvh, kvh), B * L, kvh)
|
| 495 |
-
kview = lambda base: View(kv, base, L=L, zin=1, so=L * 2 * D, ss=2 * D) # noqa: E731
|
| 496 |
-
self.icl_block(blk, x, B, L, kview(0), kview(D), L, blk.hd, blk.hd)
|
| 497 |
-
icl_kv.append(cache)
|
| 498 |
-
dec_k = eng.empty(B * n * D)
|
| 499 |
-
self.gemm(train, rows(dec_k, D), self.dec_k, B * n, D, D, norm_eps=self.dec_eps)
|
| 500 |
-
onehot = torch.nn.functional.one_hot(y, self.dec_hd).float()
|
| 501 |
-
ctx = VulkanContext(stats, col_kv, col_refine_kv, icl_kv, dec_k, eng.upload(onehot.numpy()), B, n, m, n_classes)
|
| 502 |
-
eng.flush()
|
| 503 |
-
return ctx
|
| 504 |
-
|
| 505 |
-
@torch.inference_mode()
|
| 506 |
-
def predict_logits(self, ctx, X_test, chunk_cells=None):
|
| 507 |
-
X = torch.as_tensor(X_test).float().cpu()
|
| 508 |
-
B, M, m = X.shape
|
| 509 |
-
if (B, m) != (ctx.B, ctx.m):
|
| 510 |
-
raise ValueError(f"test batch {(B, m)} does not match the context {(ctx.B, ctx.m)}")
|
| 511 |
-
if M == 0:
|
| 512 |
-
return torch.zeros(B, 0, ctx.n_classes)
|
| 513 |
-
eng, E, D = self.eng, self.E, self.D
|
| 514 |
-
col_cost, row_cost = self._piece_costs(1, m)
|
| 515 |
-
row_cost = max(row_cost, m * col_cost // self.col.n_ind)
|
| 516 |
-
if 4 * B * max(M * m, M * self._icl_width(), row_cost) > eng.max_binding:
|
| 517 |
-
raise MemoryError("query cache/scratch buffer exceeds the Vulkan limit; use fewer query rows")
|
| 518 |
-
zrn = self._upload_normalized(X, ctx.stats)
|
| 519 |
-
step = max(1, self._chunk(chunk_cells) // (B * m))
|
| 520 |
-
step = min(step, eng.max_binding // (4 * B * row_cost))
|
| 521 |
-
rows_buf = eng.empty(B * M * D)
|
| 522 |
-
for i0 in range(0, M, step):
|
| 523 |
-
nr = min(step, M - i0)
|
| 524 |
-
cells = eng.temp("test_cells", B * nr * m * E)
|
| 525 |
-
x = View(cells, 0, L=m, zin=1, so=m * E, ss=E)
|
| 526 |
-
self.embed(zrn, x, B, M, m, i0, nr, 0, m)
|
| 527 |
-
self.rownorm(x, B * nr * m, E, ln=(self.ln_w, self.ln_b, self.ln_eps))
|
| 528 |
-
self.column_query(self.col, ctx.col_kv, x, B, nr, m)
|
| 529 |
-
if self.refine is not None:
|
| 530 |
-
self.refine_rows(x, B * nr, m)
|
| 531 |
-
self.column_query(self.col_refine, ctx.col_refine_kv, x, B, nr, m)
|
| 532 |
-
self.row_stage(x, B * nr, m, View(rows_buf, i0 * D, L=1, zin=nr, so=M * D, si=D))
|
| 533 |
-
Lc = self.n_thinking + ctx.n
|
| 534 |
-
kvh = self.kv_heads * self.icl_blocks[0].hd
|
| 535 |
-
for blk, cache in zip(self.icl_blocks, ctx.icl_kv):
|
| 536 |
-
kview = lambda base: View(cache, base, L=Lc, zin=1, so=Lc * 2 * kvh, ss=2 * kvh) # noqa: E731
|
| 537 |
-
head = 0 if self.kv_heads == 1 else blk.hd
|
| 538 |
-
self.icl_block(blk, rows_buf, B, M, kview(0), kview(kvh), Lc, head, head)
|
| 539 |
-
H, hd = self.dec_H, self.dec_hd
|
| 540 |
-
q = eng.temp("q", B * M * D)
|
| 541 |
-
self.gemm(rows(rows_buf, D), rows(q, D), self.dec_q, B * M, D, D, norm_eps=self.dec_eps)
|
| 542 |
-
self.qscale(q, B * M, self.dec_scaling, ctx.n)
|
| 543 |
-
o = eng.temp("o", B * M * D)
|
| 544 |
-
qv = lambda buf: View(buf, 0, L=M, zin=1, so=M * D, ss=D) # noqa: E731
|
| 545 |
-
self.attention(qv(q), View(ctx.dec_k, 0, L=ctx.n, zin=1, so=ctx.n * D, ss=D),
|
| 546 |
-
View(ctx.onehot, 0, L=ctx.n, zin=1, so=ctx.n * hd, ss=hd), qv(o), B, H, M, ctx.n, hd, hd, hd, 0, hd)
|
| 547 |
-
p = torch.from_numpy(eng.download(o, (B, M, H, hd))).mean(2)[..., : ctx.n_classes]
|
| 548 |
-
return torch.log(p.clamp(min=1e-5) + 3e-5)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/pyproject.toml
DELETED
|
@@ -1,43 +0,0 @@
|
|
| 1 |
-
[build-system]
|
| 2 |
-
requires = ["hatchling>=1.27"]
|
| 3 |
-
build-backend = "hatchling.build"
|
| 4 |
-
|
| 5 |
-
[project]
|
| 6 |
-
name = "lightpfn"
|
| 7 |
-
version = "0.1.0"
|
| 8 |
-
description = "A compact tabular classifier pretrained only on synthetic data"
|
| 9 |
-
readme = "docs/PACKAGE_README.md"
|
| 10 |
-
requires-python = ">=3.10"
|
| 11 |
-
license = "Apache-2.0"
|
| 12 |
-
license-files = ["LICENSE", "NOTICE"]
|
| 13 |
-
authors = [{name = "Giorgio", email = "247403232+GioOtto@users.noreply.github.com"}]
|
| 14 |
-
dependencies = ["numpy>=1.24", "torch>=2.6"]
|
| 15 |
-
classifiers = [
|
| 16 |
-
"Development Status :: 3 - Alpha",
|
| 17 |
-
"Intended Audience :: Science/Research",
|
| 18 |
-
"Programming Language :: Python :: 3",
|
| 19 |
-
"Operating System :: OS Independent",
|
| 20 |
-
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
| 21 |
-
]
|
| 22 |
-
|
| 23 |
-
[project.optional-dependencies]
|
| 24 |
-
sklearn = ["scikit-learn>=1.6"]
|
| 25 |
-
hf = ["huggingface-hub>=0.27", "safetensors>=0.5"]
|
| 26 |
-
vulkan = ["wgpu>=0.32,<0.33"]
|
| 27 |
-
dev = ["pytest>=8", "build>=1.2", "hatchling>=1.27"]
|
| 28 |
-
|
| 29 |
-
[project.urls]
|
| 30 |
-
Repository = "https://github.com/GioOtto/gioPFN"
|
| 31 |
-
Models = "https://huggingface.co/ueuegio/LightPFN"
|
| 32 |
-
|
| 33 |
-
[tool.hatch.build.targets.wheel]
|
| 34 |
-
packages = ["lightpfn"]
|
| 35 |
-
exclude = ["lightpfn/train.py", "lightpfn/prior", "lightpfn/eval", "lightpfn/model/ccmm.py"]
|
| 36 |
-
|
| 37 |
-
[tool.hatch.build.targets.sdist]
|
| 38 |
-
include = [
|
| 39 |
-
"/lightpfn", "/pyproject.toml", "/LICENSE", "/NOTICE",
|
| 40 |
-
"/docs/PACKAGE_README.md", "/tests/test_release.py", "/tests/test_sklearn_api.py",
|
| 41 |
-
"/docs/DEPENDENCIES.md",
|
| 42 |
-
]
|
| 43 |
-
exclude = ["lightpfn/train.py", "lightpfn/prior", "lightpfn/eval", "lightpfn/model/ccmm.py", "**/__pycache__"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/tests/test_release.py
DELETED
|
@@ -1,94 +0,0 @@
|
|
| 1 |
-
"""Release checkpoints must be safe, exact, and usable without research dependencies."""
|
| 2 |
-
|
| 3 |
-
import json
|
| 4 |
-
import pickle
|
| 5 |
-
import subprocess
|
| 6 |
-
import sys
|
| 7 |
-
from dataclasses import asdict
|
| 8 |
-
|
| 9 |
-
import numpy as np
|
| 10 |
-
import pytest
|
| 11 |
-
import torch
|
| 12 |
-
|
| 13 |
-
from lightpfn import Config, LightPFN, load_model, save_model
|
| 14 |
-
from lightpfn.checkpoint import load_pretrained
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
def small_model():
|
| 18 |
-
torch.manual_seed(13)
|
| 19 |
-
model = LightPFN(Config(col_dim=16, col_blocks=1, col_heads=4, n_inducing=4,
|
| 20 |
-
row_blocks=1, row_heads=4, n_cls=2, icl_blocks=2,
|
| 21 |
-
icl_heads=4, decoder_heads=2, n_thinking=2, n_freq=4, n_ecdf_freq=2)).eval()
|
| 22 |
-
for parameter in model.parameters():
|
| 23 |
-
torch.nn.init.normal_(parameter, std=0.05)
|
| 24 |
-
return model
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
def test_legacy_and_safetensors_roundtrip_predictions(tmp_path):
|
| 28 |
-
model = small_model()
|
| 29 |
-
path = tmp_path / "ema.pt"
|
| 30 |
-
torch.save(dict(config=asdict(model.cfg), ema=model.state_dict()), path)
|
| 31 |
-
legacy = load_model(path)
|
| 32 |
-
release = save_model(legacy, tmp_path / "release")
|
| 33 |
-
restored = load_model(release)
|
| 34 |
-
assert restored.cfg == model.cfg
|
| 35 |
-
assert set(json.loads((release / "config.json").read_text())) == set(asdict(model.cfg))
|
| 36 |
-
g = torch.Generator().manual_seed(3)
|
| 37 |
-
X = torch.randn(1, 22, 3, generator=g)
|
| 38 |
-
X[0, 1, 0] = float("nan")
|
| 39 |
-
y = torch.arange(16).remainder(3).unsqueeze(0)
|
| 40 |
-
with torch.inference_mode():
|
| 41 |
-
expected = model(X, y, n_classes=3)
|
| 42 |
-
torch.testing.assert_close(legacy(X, y, n_classes=3), expected, rtol=0, atol=0)
|
| 43 |
-
torch.testing.assert_close(restored(X, y, n_classes=3), expected, rtol=0, atol=0)
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
def test_pickle_payload_is_rejected_without_execution(tmp_path):
|
| 47 |
-
class Payload:
|
| 48 |
-
def __reduce__(self):
|
| 49 |
-
return eval, ("__import__('pathlib').Path(%r).touch()" % str(tmp_path / "executed"),)
|
| 50 |
-
path = tmp_path / "untrusted.pt"
|
| 51 |
-
torch.save(dict(config={}, model=Payload()), path)
|
| 52 |
-
with pytest.raises(pickle.UnpicklingError):
|
| 53 |
-
load_model(path)
|
| 54 |
-
assert not (tmp_path / "executed").exists()
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
def test_reject_malformed_and_folded_checkpoints(tmp_path):
|
| 58 |
-
path = tmp_path / "bad.pt"
|
| 59 |
-
torch.save(dict(config={}, model={"not_a_tensor": 5}), path)
|
| 60 |
-
with pytest.raises(ValueError, match="tensor state"):
|
| 61 |
-
load_model(path)
|
| 62 |
-
with pytest.raises(ValueError, match="folded"):
|
| 63 |
-
save_model(small_model().folded(), tmp_path)
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
def test_hub_files_use_same_resolved_commit(tmp_path, monkeypatch):
|
| 67 |
-
revision = "a" * 40
|
| 68 |
-
folder = save_model(small_model(), tmp_path / "snapshots" / revision)
|
| 69 |
-
calls = []
|
| 70 |
-
def download(filename, **kwargs):
|
| 71 |
-
calls.append((filename, kwargs))
|
| 72 |
-
return str(folder / filename)
|
| 73 |
-
monkeypatch.setattr("huggingface_hub.hf_hub_download", download)
|
| 74 |
-
load_pretrained(repo_id="test/model", revision="main", local_files_only=True)
|
| 75 |
-
assert calls[0][1]["revision"] == "main"
|
| 76 |
-
assert calls[1][1]["revision"] == revision
|
| 77 |
-
assert all(options["local_files_only"] for _, options in calls)
|
| 78 |
-
with pytest.raises(ValueError, match="revision"):
|
| 79 |
-
load_pretrained(repo_id="test/other")
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
def test_core_import_does_not_load_optional_dependencies():
|
| 83 |
-
code = """
|
| 84 |
-
import sys
|
| 85 |
-
from importlib.abc import MetaPathFinder
|
| 86 |
-
class Block(MetaPathFinder):
|
| 87 |
-
def find_spec(self, fullname, path=None, target=None):
|
| 88 |
-
if fullname.split('.')[0] in {'sklearn', 'tabicl', 'pandas', 'scipy', 'wgpu', 'huggingface_hub', 'safetensors'}:
|
| 89 |
-
raise ModuleNotFoundError(fullname, name=fullname)
|
| 90 |
-
sys.meta_path.insert(0, Block())
|
| 91 |
-
from lightpfn import LightPFN, Config, load_model
|
| 92 |
-
assert LightPFN(Config()).cfg.label_slots == 16
|
| 93 |
-
"""
|
| 94 |
-
subprocess.run([sys.executable, "-c", code], check=True, capture_output=True, text=True)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
package/tests/test_sklearn_api.py
DELETED
|
@@ -1,124 +0,0 @@
|
|
| 1 |
-
"""Sklearn integration exercises real inference, cloning and input validation."""
|
| 2 |
-
|
| 3 |
-
import numpy as np
|
| 4 |
-
from dataclasses import asdict
|
| 5 |
-
import pytest
|
| 6 |
-
import torch
|
| 7 |
-
from sklearn.base import clone, is_classifier
|
| 8 |
-
from sklearn.exceptions import NotFittedError
|
| 9 |
-
from sklearn.model_selection import GridSearchCV
|
| 10 |
-
from sklearn.pipeline import Pipeline
|
| 11 |
-
from sklearn.preprocessing import StandardScaler
|
| 12 |
-
from sklearn.utils.estimator_checks import check_estimator
|
| 13 |
-
|
| 14 |
-
from lightpfn import LightPFNClassifier
|
| 15 |
-
from test_release import small_model
|
| 16 |
-
|
| 17 |
-
torch.set_num_threads(4)
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
@pytest.fixture
|
| 21 |
-
def task():
|
| 22 |
-
rng = np.random.default_rng(8)
|
| 23 |
-
X = rng.normal(size=(32, 3)).astype(np.float32)
|
| 24 |
-
y = np.where(X[:, 0] > 0, "yes", "no")
|
| 25 |
-
return X, y
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
def test_construction_and_clone_do_not_load_weights(monkeypatch):
|
| 29 |
-
def fail(*args, **kwargs):
|
| 30 |
-
raise AssertionError("constructor performed IO")
|
| 31 |
-
monkeypatch.setattr("lightpfn.sklearn.load_pretrained", fail)
|
| 32 |
-
monkeypatch.setattr("lightpfn.sklearn.resolve_device", fail)
|
| 33 |
-
clf = LightPFNClassifier(device="auto", n_estimators=3, random_state=12)
|
| 34 |
-
cloned = clone(clf)
|
| 35 |
-
assert cloned.get_params() == clf.get_params()
|
| 36 |
-
assert is_classifier(cloned)
|
| 37 |
-
assert not hasattr(cloned, "net_")
|
| 38 |
-
with pytest.raises(NotFittedError):
|
| 39 |
-
cloned.predict([[0, 1]])
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
def test_pipeline_grid_search_and_set_params(task):
|
| 43 |
-
X, y = task
|
| 44 |
-
clf = LightPFNClassifier(model=small_model(), device="cpu", random_state=3)
|
| 45 |
-
pipe = Pipeline([("scale", StandardScaler()), ("model", clf)])
|
| 46 |
-
search = GridSearchCV(pipe, {"model__n_estimators": [1, 2]}, cv=2).fit(X, y)
|
| 47 |
-
assert search.predict(X[:4]).shape == (4,)
|
| 48 |
-
fitted = search.best_estimator_.named_steps["model"]
|
| 49 |
-
assert fitted.n_features_in_ == 3
|
| 50 |
-
assert 0 <= fitted.score(X, y) <= 1
|
| 51 |
-
changed = clone(clf).set_params(n_estimators=2).fit(X, y)
|
| 52 |
-
assert len(changed.members_) == 2
|
| 53 |
-
assert not hasattr(clone(changed), "members_")
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
def test_input_errors_and_nan_support(task):
|
| 57 |
-
X, y = task
|
| 58 |
-
clf = LightPFNClassifier(model=small_model(), device="cpu")
|
| 59 |
-
with pytest.raises(ValueError, match="inconsistent"):
|
| 60 |
-
clf.fit(X, y[:-1])
|
| 61 |
-
with pytest.raises(ValueError, match="Unknown label|continuous"):
|
| 62 |
-
clf.fit(X, np.linspace(0, 1, len(y)))
|
| 63 |
-
X[0, 1] = np.nan
|
| 64 |
-
clf.fit(X, y)
|
| 65 |
-
P = clf.predict_proba(X)
|
| 66 |
-
assert np.isfinite(P).all()
|
| 67 |
-
np.testing.assert_allclose(P.sum(1), 1, atol=1e-6)
|
| 68 |
-
with pytest.raises(ValueError, match="features"):
|
| 69 |
-
clf.predict(X[:, :2])
|
| 70 |
-
with pytest.raises(ValueError, match="infinity"):
|
| 71 |
-
clf.predict([[0, np.inf, 1]])
|
| 72 |
-
with pytest.raises(ValueError, match="10 classes"):
|
| 73 |
-
clf.fit(X, np.arange(len(y)))
|
| 74 |
-
with pytest.raises(ValueError, match="max_context"):
|
| 75 |
-
clf.set_params(max_context=1).fit(X, y)
|
| 76 |
-
with pytest.raises(NotFittedError):
|
| 77 |
-
clf.predict(X)
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
def test_categorical_mask_warns_and_preserves_numeric_behavior(task):
|
| 81 |
-
X, y = task
|
| 82 |
-
model = small_model()
|
| 83 |
-
plain = LightPFNClassifier(model=model, device="cpu").fit(X, y)
|
| 84 |
-
other = LightPFNClassifier(model=model, device="cpu")
|
| 85 |
-
with pytest.warns(FutureWarning, match="native categorical"):
|
| 86 |
-
other.fit(X, y, cat=np.array([True, False, False]))
|
| 87 |
-
np.testing.assert_array_equal(plain.predict_proba(X), other.predict_proba(X))
|
| 88 |
-
with pytest.raises(ValueError, match="boolean mask"):
|
| 89 |
-
other.fit(X, y, cat=[0])
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
def test_refit_and_single_class(task):
|
| 93 |
-
X, y = task
|
| 94 |
-
model = small_model().train()
|
| 95 |
-
clf = LightPFNClassifier(model=model, device="cpu").fit(X, y)
|
| 96 |
-
assert model.training and not clf.model_.training
|
| 97 |
-
clf.fit(X[:, :2], np.full(len(y), "only"))
|
| 98 |
-
assert clf.n_features_in_ == 2
|
| 99 |
-
np.testing.assert_array_equal(clf.predict_proba(X[:, :2]), np.ones((len(y), 1)))
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
def test_feature_name_validation(task):
|
| 103 |
-
pd = pytest.importorskip("pandas")
|
| 104 |
-
X, y = task
|
| 105 |
-
frame = pd.DataFrame(X, columns=["a", "b", "c"])
|
| 106 |
-
clf = LightPFNClassifier(model=small_model(), device="cpu").fit(frame, y)
|
| 107 |
-
np.testing.assert_array_equal(clf.feature_names_in_, frame.columns)
|
| 108 |
-
with pytest.raises(ValueError, match="Feature names|feature names"):
|
| 109 |
-
clf.predict(frame[["c", "b", "a"]])
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
def test_sklearn_estimator_checks(tmp_path):
|
| 113 |
-
# joblib hashes torch storage identity, so raw nn.Module parameters cannot
|
| 114 |
-
# satisfy its deepcopy/hash comparison. Exercise the distribution's path API.
|
| 115 |
-
model = small_model()
|
| 116 |
-
path = tmp_path / "model.pt"
|
| 117 |
-
torch.save(dict(config=asdict(model.cfg), model=model.state_dict()), path)
|
| 118 |
-
check_estimator(
|
| 119 |
-
LightPFNClassifier(checkpoint=path, device="cpu", n_threads=4),
|
| 120 |
-
expected_failed_checks={
|
| 121 |
-
"check_methods_sample_order_invariance": "FP32 GPU/CPU kernels can differ by ~6e-8 after row permutation.",
|
| 122 |
-
"check_methods_subset_invariance": "FP32 kernels can differ slightly when query batch shapes change.",
|
| 123 |
-
},
|
| 124 |
-
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
provenance.json
CHANGED
|
@@ -1,18 +1,20 @@
|
|
| 1 |
-
{
|
| 2 |
-
"package_version": "
|
| 3 |
-
"
|
| 4 |
-
"
|
| 5 |
-
"
|
| 6 |
-
"
|
| 7 |
-
"
|
| 8 |
-
"
|
| 9 |
-
"
|
| 10 |
-
"
|
| 11 |
-
"
|
| 12 |
-
|
| 13 |
-
"
|
| 14 |
-
"
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
"
|
| 18 |
-
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"package_version": "1.0.0",
|
| 3 |
+
"source_checkpoint": "final_long/ema_step039250.pt",
|
| 4 |
+
"source_checkpoint_sha256": "d2142d467a0bf6935f1a620fc216cb8473ba62815a4e52c0dad6e479a43a2c28",
|
| 5 |
+
"parameters": 4603088,
|
| 6 |
+
"torch_version": "2.14.1+cpu",
|
| 7 |
+
"numpy_version": "2.5.3",
|
| 8 |
+
"license": "Apache-2.0",
|
| 9 |
+
"weights_format": "safetensors",
|
| 10 |
+
"dtype": "float32",
|
| 11 |
+
"hashes": {
|
| 12 |
+
"config.json": "cc68fb632a4c4406fd8e1a4d18ddf144c94879e5ab4a61e8f191c2ff4014e9d3",
|
| 13 |
+
"model.safetensors": "a492572bd1892f8af9a92798fa1316de410c49171eae0ad4821caddbb0be2668",
|
| 14 |
+
"LightPFN_report.pdf": "274be6425d074c984ce49cee56f40a0801421017a2851febe916ac7bbc4718f4"
|
| 15 |
+
},
|
| 16 |
+
"roundtrip": "bitwise identical tensor weights and CPU logits",
|
| 17 |
+
"weights_revision": "bd389ab59a89dd0e05c9ecb7c642c08ee52e9637",
|
| 18 |
+
"source_repository": "https://github.com/GioOtto/LightPFN",
|
| 19 |
+
"package": "https://pypi.org/project/LightPFN/"
|
| 20 |
+
}
|