ueuegio commited on
Commit
ef2fb69
·
verified ·
1 Parent(s): eab0027

LightPFN 1.0.0: public model card, report, provenance; remove the private 0.1.0 package files

Browse files
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:8bea6f04f76fa4d082d8884445169b20c9e55b901543c205de46d9fa7345d7b0
3
- size 460153
 
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 (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.
 
 
 
 
 
 
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
- - pytorch
 
11
  - safetensors
12
  ---
13
 
14
- # LightPFN 0.1.0
15
 
16
- **4,603,088 parameters**, pretrained from scratch **only on synthetic data**.
17
- This repository contains the final `final_long` model, safe inference weights,
18
- and the installable Python package. Both **code and weights are Apache 2.0**.
19
- The initial repository is private; it is not yet a public release or a PyPI upload.
20
 
21
- Preprint: [A Sling Against Giants: LightPFN, a 4.6M-parameter tabular in-context
22
- classifier designed to stay small](./LightPFN_report.pdf), Giorgio Ottoboni,
23
- 5 October 2026.
 
24
 
25
- No distillation was used. No weights or outputs from other tabular foundation
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
- hf auth login
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(device="auto", n_estimators=4, random_state=0)
49
- clf.fit(X_train, y_train)
50
  proba = clf.predict_proba(X_test)
51
- pred = clf.predict(X_test)
52
  ```
53
 
54
- The package pins the weights/config to an immutable commit. The first fit downloads
55
- them using the HF cache. `local_files_only=True` supports an already cached offline
56
- model. Construction and sklearn cloning do no downloads or GPU initialization.
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
- - Base stage: 120,000 steps, 64 tasks/step, approximately 7.7M tasks seen, on two
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
- `provenance.json` records source checkpoint hash, tensor artifact hashes,
81
- architecture size, software versions, source-code revision and preprint hash.
 
 
 
 
 
 
82
 
83
  ## Evaluation
84
 
85
- Variants were selected on held-out synthetic tasks (D1), interaction probes (D2)
86
- and 55 OpenML-CC18 tasks outside TabArena (D3). TabArena was inspected only for
87
- the finalists. These are local experiments, **not an official TabArena submission**.
88
- GBDT comparisons use library defaults, not tuned TabArena configurations.
89
 
90
- | Evaluation | Result |
91
  |---|---|
92
- | D3, 55 tasks, <=1,000 training rows, one estimator | Mean AUC 0.9109 |
93
- | D3 vs default CatBoost | +0.86 AUC points, paired 95% CI [0.42, 1.40] |
94
- | TabArena classification, 38 tasks, official splits, first repeat, four estimators | Mean AUC 0.8581, mean rank 2.50 |
95
- | Same TabArena run vs default CatBoost | Lower task error on 76% of tasks |
96
 
97
- Task error is 1-AUC for binary tasks and log loss for multiclass tasks. On this
98
- full TabArena evaluation, TabICLv2 ranked first (AUC 0.8637, mean rank 1.58).
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 for **2-10 classes**, without regression. A single observed class
105
- returns its constant probability. The default context limit is 20,000 rows;
106
- larger datasets use a stratified subset for each ensemble member. Large contexts
107
- cost substantially more memory/time. Categorical-heavy datasets are a known
108
- weakness: in_vehicle_coupon_recommendation and Amazon_employee_access trailed
109
- default CatBoost by about 4 AUC points in the full evaluation. CPU inference is
110
- sometimes slower than default GBDTs (median 15/54 seconds on medium/large
111
- TabArena splits versus CatBoost's 4/7 seconds on the evaluated server).
112
-
113
- Inference folding, cell blocking and estimator batching preserve predictions
114
- within FP32 tolerances. Row permutation or query batch changes may shift
115
- probabilities around 1e-7; two strict sklearn invariance checks document this
116
- tolerance. GPU adapters have finite storage-buffer limits; device="auto" can
117
- fall back to CPU when a Vulkan context exceeds those limits.
118
-
119
- ## Source and license
120
-
121
- Source: [GioOtto/gioPFN](https://github.com/GioOtto/gioPFN).
122
- The repo's code may require separate access while it is private.
123
- Code and released weights: [Apache License 2.0](./LICENSE), with [NOTICE](./NOTICE).
124
- Third-party dependencies retain their own licenses, listed in
125
- `package/docs/DEPENDENCIES.md`.
 
 
 
 
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": "0.1.0",
3
- "source_commit": "ef5d60de31519d54e361ffb8554c6cfc9bfa1835",
4
- "source_checkpoint": "final_long/ema_step039250.pt",
5
- "source_checkpoint_sha256": "d2142d467a0bf6935f1a620fc216cb8473ba62815a4e52c0dad6e479a43a2c28",
6
- "parameters": 4603088,
7
- "torch_version": "2.14.1+cpu",
8
- "numpy_version": "2.5.3",
9
- "license": "Apache-2.0",
10
- "weights_format": "safetensors",
11
- "dtype": "float32",
12
- "hashes": {
13
- "config.json": "cc68fb632a4c4406fd8e1a4d18ddf144c94879e5ab4a61e8f191c2ff4014e9d3",
14
- "model.safetensors": "a492572bd1892f8af9a92798fa1316de410c49171eae0ad4821caddbb0be2668",
15
- "LightPFN_report.pdf": "8bea6f04f76fa4d082d8884445169b20c9e55b901543c205de46d9fa7345d7b0"
16
- },
17
- "roundtrip": "bitwise identical tensor weights and CPU logits"
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
+ }