symbiotic-learning 0.2.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,218 @@
1
+ # Byte-compiled / optimized / DLL files
2
+ __pycache__/
3
+ *.py[codz]
4
+ *$py.class
5
+
6
+ # C extensions
7
+ *.so
8
+
9
+ # Distribution / packaging
10
+ .Python
11
+ build/
12
+ develop-eggs/
13
+ dist/
14
+ downloads/
15
+ eggs/
16
+ .eggs/
17
+ lib/
18
+ lib64/
19
+ parts/
20
+ sdist/
21
+ var/
22
+ wheels/
23
+ share/python-wheels/
24
+ *.egg-info/
25
+ .installed.cfg
26
+ *.egg
27
+ MANIFEST
28
+
29
+ # PyInstaller
30
+ # Usually these files are written by a python script from a template
31
+ # before PyInstaller builds the exe, so as to inject date/other infos into it.
32
+ *.manifest
33
+ *.spec
34
+
35
+ # Installer logs
36
+ pip-log.txt
37
+ pip-delete-this-directory.txt
38
+
39
+ # Unit test / coverage reports
40
+ htmlcov/
41
+ .tox/
42
+ .nox/
43
+ .coverage
44
+ .coverage.*
45
+ .cache
46
+ nosetests.xml
47
+ coverage.xml
48
+ *.cover
49
+ *.py.cover
50
+ .hypothesis/
51
+ .pytest_cache/
52
+ cover/
53
+
54
+ # Translations
55
+ *.mo
56
+ *.pot
57
+
58
+ # Django stuff:
59
+ *.log
60
+ local_settings.py
61
+ db.sqlite3
62
+ db.sqlite3-journal
63
+
64
+ # Flask stuff:
65
+ instance/
66
+ .webassets-cache
67
+
68
+ # Scrapy stuff:
69
+ .scrapy
70
+
71
+ # Sphinx documentation
72
+ docs/_build/
73
+
74
+ # PyBuilder
75
+ .pybuilder/
76
+ target/
77
+
78
+ # Jupyter Notebook
79
+ .ipynb_checkpoints
80
+
81
+ # IPython
82
+ profile_default/
83
+ ipython_config.py
84
+
85
+ # pyenv
86
+ # For a library or package, you might want to ignore these files since the code is
87
+ # intended to run in multiple environments; otherwise, check them in:
88
+ # .python-version
89
+
90
+ # pipenv
91
+ # According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
92
+ # However, in case of collaboration, if having platform-specific dependencies or dependencies
93
+ # having no cross-platform support, pipenv may install dependencies that don't work, or not
94
+ # install all needed dependencies.
95
+ # Pipfile.lock
96
+
97
+ # UV
98
+ # Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control.
99
+ # This is especially recommended for binary packages to ensure reproducibility, and is more
100
+ # commonly ignored for libraries.
101
+ # uv.lock
102
+
103
+ # poetry
104
+ # Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
105
+ # This is especially recommended for binary packages to ensure reproducibility, and is more
106
+ # commonly ignored for libraries.
107
+ # https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
108
+ # poetry.lock
109
+ # poetry.toml
110
+
111
+ # pdm
112
+ # Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
113
+ # pdm recommends including project-wide configuration in pdm.toml, but excluding .pdm-python.
114
+ # https://pdm-project.org/en/latest/usage/project/#working-with-version-control
115
+ # pdm.lock
116
+ # pdm.toml
117
+ .pdm-python
118
+ .pdm-build/
119
+
120
+ # pixi
121
+ # Similar to Pipfile.lock, it is generally recommended to include pixi.lock in version control.
122
+ # pixi.lock
123
+ # Pixi creates a virtual environment in the .pixi directory, just like venv module creates one
124
+ # in the .venv directory. It is recommended not to include this directory in version control.
125
+ .pixi
126
+
127
+ # PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
128
+ __pypackages__/
129
+
130
+ # Celery stuff
131
+ celerybeat-schedule
132
+ celerybeat.pid
133
+
134
+ # Redis
135
+ *.rdb
136
+ *.aof
137
+ *.pid
138
+
139
+ # RabbitMQ
140
+ mnesia/
141
+ rabbitmq/
142
+ rabbitmq-data/
143
+
144
+ # ActiveMQ
145
+ activemq-data/
146
+
147
+ # SageMath parsed files
148
+ *.sage.py
149
+
150
+ # Environments
151
+ .env
152
+ .envrc
153
+ .venv
154
+ env/
155
+ venv/
156
+ ENV/
157
+ env.bak/
158
+ venv.bak/
159
+
160
+ # Spyder project settings
161
+ .spyderproject
162
+ .spyproject
163
+
164
+ # Rope project settings
165
+ .ropeproject
166
+
167
+ # mkdocs documentation
168
+ /site
169
+
170
+ # mypy
171
+ .mypy_cache/
172
+ .dmypy.json
173
+ dmypy.json
174
+
175
+ # Pyre type checker
176
+ .pyre/
177
+
178
+ # pytype static type analyzer
179
+ .pytype/
180
+
181
+ # Cython debug symbols
182
+ cython_debug/
183
+
184
+ # PyCharm
185
+ # JetBrains specific template is maintained in a separate JetBrains.gitignore that can
186
+ # be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
187
+ # and can be added to the global gitignore or merged into this file. For a more nuclear
188
+ # option (not recommended) you can uncomment the following to ignore the entire idea folder.
189
+ # .idea/
190
+
191
+ # Abstra
192
+ # Abstra is an AI-powered process automation framework.
193
+ # Ignore directories containing user credentials, local state, and settings.
194
+ # Learn more at https://abstra.io/docs
195
+ .abstra/
196
+
197
+ # Visual Studio Code
198
+ # Visual Studio Code specific template is maintained in a separate VisualStudioCode.gitignore
199
+ # that can be found at https://github.com/github/gitignore/blob/main/Global/VisualStudioCode.gitignore
200
+ # and can be added to the global gitignore or merged into this file. However, if you prefer,
201
+ # you could uncomment the following to ignore the entire vscode folder
202
+ # .vscode/
203
+ # Temporary file for partial code execution
204
+ tempCodeRunnerFile.py
205
+
206
+ # Ruff stuff:
207
+ .ruff_cache/
208
+
209
+ # PyPI configuration file
210
+ .pypirc
211
+
212
+ # Marimo
213
+ marimo/_static/
214
+ marimo/_lsp/
215
+ __marimo__/
216
+
217
+ # Streamlit
218
+ .streamlit/secrets.toml
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 bsimon717
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,138 @@
1
+ Metadata-Version: 2.5
2
+ Name: symbiotic_learning
3
+ Version: 0.2.1
4
+ Summary: A tool for symbiotically training multiple ML models.
5
+ Project-URL: Repository, https://github.com/bsimon717/symlearn
6
+ Author-email: "Benjamin D. Simon" <bsimon71701@gmail.com>
7
+ License-File: LICENSE
8
+ Keywords: ensemble,machine learning,symbiotic
9
+ Requires-Python: >=3.11
10
+ Requires-Dist: matplotlib
11
+ Requires-Dist: numpy
12
+ Requires-Dist: scikit-learn
13
+ Requires-Dist: torch>=2.9.0
14
+ Requires-Dist: tqdm
15
+ Description-Content-Type: text/markdown
16
+
17
+ # Symbiotic-Learning
18
+
19
+ The *symbiotic_learning* package functions as an implementation of "symbiotic learning": a paradigm for simultaneously training multiple machine-learning models at once, in which collaboration between models is **intrinsic** and **incentivized**. The goal for this method is to effectively leverage the collaboration of relatively small models to achieve performance comparable to that of large, computationally expensive models.
20
+
21
+ A key feature of this method is a minimally-sized attention block (termed the Readout) whose task is to aggregate the perspectives and decisions of the symbiotically trained upstream (pre-Readout) models, giving a final prediction. This attention block applies a linear transformation to the input logits before performing (multi-head) scaled dot-product attention. The output of this is then concatenated with a linear transformation of the input embeddings and passed through a user-specified number of fully connected layers, yielding the final prediction.
22
+
23
+ The appending of this block to the overall system occurs after a pre-determined number of training epochs, this event being termed *uplift*. As such, the framework is separated into two phases: pre-uplift and post-uplift.
24
+
25
+ A system of symbiotically trained ML models with a Readout block is termed a *Symbiotic Uplift Network*.
26
+
27
+ ---
28
+
29
+ ## Usage
30
+
31
+ Currently, *symbiotic_learning* is only implemented for classification tasks.
32
+
33
+ To install this package, run the following command:
34
+
35
+ `pip install symbiotic-learning`
36
+
37
+ To use this package, first include the following imports in your training script:
38
+
39
+ ```
40
+ from symbiotic_learning.classify.readout import Readout
41
+ import symbiotic_learning.classify.utils as classify
42
+ ```
43
+
44
+ Then, include a block with a structure similar to the following:
45
+
46
+ ```
47
+ save_path = ## Path to sym_logs folder ##
48
+ save_end = ## Boolean for saving models at the end of training ##
49
+ save_best = ## Boolean for saving models at epoch of highest validation accuracy ##
50
+
51
+ data_loaders = [train_loader, valid_loader, test_loader]
52
+
53
+ num_classes = ## Task Specific ##
54
+
55
+ preR_dim = ## Dimension of pre-Readout embeddings ##
56
+
57
+ readout_hidden_dim = ## Dimension of fully-connected hidden layers ##
58
+ readout_num_hidden = ## Number of fully-connected hidden layers ##
59
+ num_heads = ## Number of attention heads ##
60
+
61
+ collab_params = [## List of collaboration parameters ##]
62
+ temp = ## Temperature hyperparameter in Readout loss ##
63
+ lamb = ## Responsibility hyperparameter ##
64
+ eps = 1e-7 ## Small value to avoid divide-by-zero errors ##
65
+
66
+ models = []
67
+ opts = []
68
+ scheds = []
69
+ num_preR = 3
70
+
71
+ for _ in range(num_preR):
72
+ models.append( ## Base Model Here ## )
73
+ opts.append( ## Optimizer Here ## )
74
+ scheds.append( ## LR Scheduler ## )
75
+
76
+ readout = Readout(hidden_dim=readout_hidden_dim, num_hidden=readout_num_hidden, num_classes=num_classes, num_heads=num_heads, num_preR=num_preR, preR_dim=preR_dim)
77
+
78
+ classify.train(epochs, models, opts, scheds, data_loaders, collab_params, temp, criterion, uplift=uplift, eps=eps, lamb=lamb, save_path=save_path, save_end=save_end, save_best=save_best)
79
+
80
+ ```
81
+
82
+ ---
83
+
84
+ ## Training
85
+ $N$ pre-Readout models are initialized for the primary task, each having an "embedding block" and a "decision block":
86
+
87
+ - The exact architecture of the embedding block is task-dependent; for an image-classification task, for example, the embedding block could consist of convolutional layers.
88
+
89
+ - The only requirement of the decision block is that it must receive the concatenation of all $N$ embeddings as input to yield a task-specific prediction.
90
+
91
+
92
+ ### Pre-Uplift
93
+ 1. Each pre-Readout model performs its initial assessment of the input data using its embedding block.
94
+ 2. The $N$ embeddings are concatenated and used as input to each of the models' decision blocks, resulting in $N$ predictions.
95
+ 3. A pre-Readout model's total (symbiotic) loss is calculated using its own output as well as the outputs of its peers, with an additional term calculated from their initial embeddings to encourage diversity of perspectives. The weighting of each of these terms is determined by that model's *collaboration parameter*.
96
+
97
+ ### Post-Uplift
98
+ 1. Each pre-Readout model performs its initial assessment of the input data using its embedding block.
99
+ 2. The $N$ embeddings are concatenated and used as input to each of the models' decision blocks, resulting in $N$ predictions.
100
+ 3. The $N$ predictions are concatenated and passed to the Readout's attention layer. Additonally, the vector of pre-Readout embeddings is passed through a single fully-connected layer and concatenated with the attention layer's output. This vector is then passed through fully-connected layers, resulting in the final prediction.
101
+ 4. The Readout is then penalized on how strong its own prediction was compared to the strength of the pre-Readout predictions via a non-linearity.
102
+ 5. Each pre-Readout model's symbiotic loss then has a term added to it capturing that model's culpability for the Readout's mistakes. This term is called the model's "blame loss" and is scaled using a global hyperparameter (termed "responsibility").
103
+
104
+ ---
105
+
106
+ ## Definitions
107
+ - Symbiotic Uplift Network: An aggregate network of machine-learning models trained using symbiotic learning.
108
+ - Symbiotic Loss ($L_{sym,i}$): A pre-Readout model's multi-objective loss function. Collaboration parameters enable coupling of models' loss functions such that 1) an individual model's parameters will also be updated based on the other models' personal losses, and 2) diversity of perspective is encouraged via Embedding Loss.
109
+
110
+ $$ L_{sym,i} = (1-\alpha_i)L_i + \alpha_i(\sum_{j \neq i}{L_j}) + \alpha_{i}^{2}L_{embed,i} $$
111
+
112
+ (Note: The only learnable parameters affected by this coupling are those used in the initial embedding blocks.)
113
+
114
+ - Collaboration Parameters ($\alpha_i$): Coupling constants (hyperparameters) in the symbiotic loss functions of pre-Readout models. Must be in the range $[0,1]$.
115
+ - Personal Loss ($L_i$): A term in a pre-Readout model's symbiotic loss computed using only that model's prediction. Task-specific.
116
+ - Embedding Loss ($L_{embed,i}$): A contrastive term in a pre-Readout model's symbiotic loss which encourages diverse initial assessments. *EmbedSim* is defined to be the cosine similarity function scaled to the range $[0,1]$, and $\delta$ is a temperature hyperparameter shared between all pre-Readout models.
117
+
118
+ $$ L_{embed,i} = \frac{1}{N-1}\sum_{j \neq i}[\exp{(EmbedSim(x_i, x_j)/\delta)-1}] $$
119
+
120
+ - Blame Loss ($L_{blame, i}$): A term added to a pre-Readout model's symbiotic loss after uplift, capturing that model's contribution to the Readout's loss. $\lambda$ is termed a "responsibility" hyperparameter shared between all pre-Readout models
121
+
122
+ $$ L_{blame, i} = \lambda(\frac{L_i}{\sum L_i})*L_F $$
123
+
124
+ - Readout Loss ($L_{Readout}$): A loss function specific to the Readout block which penalizes it the lower the sum of pre-Readout personal losses is, where $L_F$ is its personal loss, and $\tau$ is a temperature hyperparameter.
125
+
126
+ $$ L_{Readout} = L_F (1+\exp[-\tau(\sum L_i)]) $$
127
+
128
+ ## Example Symbiotic Uplift Network Architecture
129
+
130
+ ![Symbiotic Uplift Network Architecture](symlearn_arch.png)
131
+
132
+ This figure shows the architecture for a Symbiotic Uplift Network with three pre-Readout models.
133
+
134
+ ## Readout Architecture
135
+
136
+ ![Readout Architecture](readout_arch.png)
137
+
138
+ This figure shows the architecture of the Readout block. The attention mechanism used is (multi-head) scaled dot-product attention.
@@ -0,0 +1,122 @@
1
+ # Symbiotic-Learning
2
+
3
+ The *symbiotic_learning* package functions as an implementation of "symbiotic learning": a paradigm for simultaneously training multiple machine-learning models at once, in which collaboration between models is **intrinsic** and **incentivized**. The goal for this method is to effectively leverage the collaboration of relatively small models to achieve performance comparable to that of large, computationally expensive models.
4
+
5
+ A key feature of this method is a minimally-sized attention block (termed the Readout) whose task is to aggregate the perspectives and decisions of the symbiotically trained upstream (pre-Readout) models, giving a final prediction. This attention block applies a linear transformation to the input logits before performing (multi-head) scaled dot-product attention. The output of this is then concatenated with a linear transformation of the input embeddings and passed through a user-specified number of fully connected layers, yielding the final prediction.
6
+
7
+ The appending of this block to the overall system occurs after a pre-determined number of training epochs, this event being termed *uplift*. As such, the framework is separated into two phases: pre-uplift and post-uplift.
8
+
9
+ A system of symbiotically trained ML models with a Readout block is termed a *Symbiotic Uplift Network*.
10
+
11
+ ---
12
+
13
+ ## Usage
14
+
15
+ Currently, *symbiotic_learning* is only implemented for classification tasks.
16
+
17
+ To install this package, run the following command:
18
+
19
+ `pip install symbiotic-learning`
20
+
21
+ To use this package, first include the following imports in your training script:
22
+
23
+ ```
24
+ from symbiotic_learning.classify.readout import Readout
25
+ import symbiotic_learning.classify.utils as classify
26
+ ```
27
+
28
+ Then, include a block with a structure similar to the following:
29
+
30
+ ```
31
+ save_path = ## Path to sym_logs folder ##
32
+ save_end = ## Boolean for saving models at the end of training ##
33
+ save_best = ## Boolean for saving models at epoch of highest validation accuracy ##
34
+
35
+ data_loaders = [train_loader, valid_loader, test_loader]
36
+
37
+ num_classes = ## Task Specific ##
38
+
39
+ preR_dim = ## Dimension of pre-Readout embeddings ##
40
+
41
+ readout_hidden_dim = ## Dimension of fully-connected hidden layers ##
42
+ readout_num_hidden = ## Number of fully-connected hidden layers ##
43
+ num_heads = ## Number of attention heads ##
44
+
45
+ collab_params = [## List of collaboration parameters ##]
46
+ temp = ## Temperature hyperparameter in Readout loss ##
47
+ lamb = ## Responsibility hyperparameter ##
48
+ eps = 1e-7 ## Small value to avoid divide-by-zero errors ##
49
+
50
+ models = []
51
+ opts = []
52
+ scheds = []
53
+ num_preR = 3
54
+
55
+ for _ in range(num_preR):
56
+ models.append( ## Base Model Here ## )
57
+ opts.append( ## Optimizer Here ## )
58
+ scheds.append( ## LR Scheduler ## )
59
+
60
+ readout = Readout(hidden_dim=readout_hidden_dim, num_hidden=readout_num_hidden, num_classes=num_classes, num_heads=num_heads, num_preR=num_preR, preR_dim=preR_dim)
61
+
62
+ classify.train(epochs, models, opts, scheds, data_loaders, collab_params, temp, criterion, uplift=uplift, eps=eps, lamb=lamb, save_path=save_path, save_end=save_end, save_best=save_best)
63
+
64
+ ```
65
+
66
+ ---
67
+
68
+ ## Training
69
+ $N$ pre-Readout models are initialized for the primary task, each having an "embedding block" and a "decision block":
70
+
71
+ - The exact architecture of the embedding block is task-dependent; for an image-classification task, for example, the embedding block could consist of convolutional layers.
72
+
73
+ - The only requirement of the decision block is that it must receive the concatenation of all $N$ embeddings as input to yield a task-specific prediction.
74
+
75
+
76
+ ### Pre-Uplift
77
+ 1. Each pre-Readout model performs its initial assessment of the input data using its embedding block.
78
+ 2. The $N$ embeddings are concatenated and used as input to each of the models' decision blocks, resulting in $N$ predictions.
79
+ 3. A pre-Readout model's total (symbiotic) loss is calculated using its own output as well as the outputs of its peers, with an additional term calculated from their initial embeddings to encourage diversity of perspectives. The weighting of each of these terms is determined by that model's *collaboration parameter*.
80
+
81
+ ### Post-Uplift
82
+ 1. Each pre-Readout model performs its initial assessment of the input data using its embedding block.
83
+ 2. The $N$ embeddings are concatenated and used as input to each of the models' decision blocks, resulting in $N$ predictions.
84
+ 3. The $N$ predictions are concatenated and passed to the Readout's attention layer. Additonally, the vector of pre-Readout embeddings is passed through a single fully-connected layer and concatenated with the attention layer's output. This vector is then passed through fully-connected layers, resulting in the final prediction.
85
+ 4. The Readout is then penalized on how strong its own prediction was compared to the strength of the pre-Readout predictions via a non-linearity.
86
+ 5. Each pre-Readout model's symbiotic loss then has a term added to it capturing that model's culpability for the Readout's mistakes. This term is called the model's "blame loss" and is scaled using a global hyperparameter (termed "responsibility").
87
+
88
+ ---
89
+
90
+ ## Definitions
91
+ - Symbiotic Uplift Network: An aggregate network of machine-learning models trained using symbiotic learning.
92
+ - Symbiotic Loss ($L_{sym,i}$): A pre-Readout model's multi-objective loss function. Collaboration parameters enable coupling of models' loss functions such that 1) an individual model's parameters will also be updated based on the other models' personal losses, and 2) diversity of perspective is encouraged via Embedding Loss.
93
+
94
+ $$ L_{sym,i} = (1-\alpha_i)L_i + \alpha_i(\sum_{j \neq i}{L_j}) + \alpha_{i}^{2}L_{embed,i} $$
95
+
96
+ (Note: The only learnable parameters affected by this coupling are those used in the initial embedding blocks.)
97
+
98
+ - Collaboration Parameters ($\alpha_i$): Coupling constants (hyperparameters) in the symbiotic loss functions of pre-Readout models. Must be in the range $[0,1]$.
99
+ - Personal Loss ($L_i$): A term in a pre-Readout model's symbiotic loss computed using only that model's prediction. Task-specific.
100
+ - Embedding Loss ($L_{embed,i}$): A contrastive term in a pre-Readout model's symbiotic loss which encourages diverse initial assessments. *EmbedSim* is defined to be the cosine similarity function scaled to the range $[0,1]$, and $\delta$ is a temperature hyperparameter shared between all pre-Readout models.
101
+
102
+ $$ L_{embed,i} = \frac{1}{N-1}\sum_{j \neq i}[\exp{(EmbedSim(x_i, x_j)/\delta)-1}] $$
103
+
104
+ - Blame Loss ($L_{blame, i}$): A term added to a pre-Readout model's symbiotic loss after uplift, capturing that model's contribution to the Readout's loss. $\lambda$ is termed a "responsibility" hyperparameter shared between all pre-Readout models
105
+
106
+ $$ L_{blame, i} = \lambda(\frac{L_i}{\sum L_i})*L_F $$
107
+
108
+ - Readout Loss ($L_{Readout}$): A loss function specific to the Readout block which penalizes it the lower the sum of pre-Readout personal losses is, where $L_F$ is its personal loss, and $\tau$ is a temperature hyperparameter.
109
+
110
+ $$ L_{Readout} = L_F (1+\exp[-\tau(\sum L_i)]) $$
111
+
112
+ ## Example Symbiotic Uplift Network Architecture
113
+
114
+ ![Symbiotic Uplift Network Architecture](symlearn_arch.png)
115
+
116
+ This figure shows the architecture for a Symbiotic Uplift Network with three pre-Readout models.
117
+
118
+ ## Readout Architecture
119
+
120
+ ![Readout Architecture](readout_arch.png)
121
+
122
+ This figure shows the architecture of the Readout block. The attention mechanism used is (multi-head) scaled dot-product attention.
@@ -0,0 +1,22 @@
1
+ [build-system]
2
+ requires = ["hatchling"]
3
+ build-backend = "hatchling.build"
4
+
5
+ [project]
6
+ name = "symbiotic_learning"
7
+ version = "0.2.1"
8
+ requires-python = ">= 3.11"
9
+ authors = [{name = "Benjamin D. Simon", email="bsimon71701@gmail.com"}]
10
+ description = "A tool for symbiotically training multiple ML models."
11
+ readme = "README.md"
12
+ keywords = ["machine learning", "ensemble", "symbiotic"]
13
+ dependencies = [
14
+ "torch>=2.9.0",
15
+ "numpy",
16
+ "matplotlib",
17
+ "scikit_learn",
18
+ "tqdm"
19
+ ]
20
+
21
+ [project.urls]
22
+ Repository = "https://github.com/bsimon717/symlearn"
@@ -0,0 +1,70 @@
1
+ #
2
+ # This file is autogenerated by pip-compile with Python 3.11
3
+ # by the following command:
4
+ #
5
+ # pip-compile --no-index _symlearn/pyproject.toml
6
+ #
7
+ cloudpickle==3.1.2
8
+ # via joblib
9
+ colorama==0.4.6
10
+ # via tqdm
11
+ contourpy==1.3.3
12
+ # via matplotlib
13
+ cycler==0.12.1
14
+ # via matplotlib
15
+ filelock==4.0.5
16
+ # via torch
17
+ fonttools==4.66.0
18
+ # via matplotlib
19
+ fsspec==2026.9.0
20
+ # via torch
21
+ jinja2==3.1.6
22
+ # via torch
23
+ joblib==1.6.0
24
+ # via scikit-learn
25
+ kiwisolver==1.5.1
26
+ # via matplotlib
27
+ markupsafe==3.0.3
28
+ # via jinja2
29
+ matplotlib==3.11.2
30
+ # via symlearn (_symlearn/pyproject.toml)
31
+ mpmath==1.3.0
32
+ # via sympy
33
+ narwhals==2.26.0
34
+ # via scikit-learn
35
+ networkx==3.6.1
36
+ # via torch
37
+ numpy==2.4.6
38
+ # via
39
+ # contourpy
40
+ # matplotlib
41
+ # scikit-learn
42
+ # scipy
43
+ # symlearn (_symlearn/pyproject.toml)
44
+ packaging==26.3
45
+ # via matplotlib
46
+ pillow==12.3.0
47
+ # via matplotlib
48
+ pyparsing==3.3.3
49
+ # via matplotlib
50
+ python-dateutil==2.9.0.post0
51
+ # via matplotlib
52
+ scikit-learn==1.9.1
53
+ # via symlearn (_symlearn/pyproject.toml)
54
+ scipy==1.17.1
55
+ # via scikit-learn
56
+ six==1.17.0
57
+ # via python-dateutil
58
+ sympy==1.14.0
59
+ # via torch
60
+ threadpoolctl==3.7.0
61
+ # via scikit-learn
62
+ torch==2.14.0
63
+ # via symlearn (_symlearn/pyproject.toml)
64
+ tqdm==4.70.0
65
+ # via symlearn (_symlearn/pyproject.toml)
66
+ typing-extensions==4.16.0
67
+ # via torch
68
+
69
+ # The following packages are considered to be unsafe in a requirements file:
70
+ # setuptools
@@ -0,0 +1,4 @@
1
+ __all__ = ['loss', 'classify']
2
+
3
+ from . import classify
4
+ from . import loss
@@ -0,0 +1,111 @@
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+ from torch import Tensor
5
+
6
+ class Readout(nn.Module):
7
+ """
8
+ Attention block intended to aggregate the decisions of pre-Readout models.
9
+ """
10
+
11
+ def __init__(
12
+ self,
13
+ hidden_dim: int = 32,
14
+ preR_dim: int = 32,
15
+ num_hidden: int = 1,
16
+ num_classes: int = 4,
17
+ num_heads: int = 1,
18
+ num_preR: int = 3,
19
+ attn_dropout: float = 0.0,) -> None:
20
+
21
+ super(Readout, self).__init__()
22
+
23
+ self.hidden_dim = hidden_dim
24
+ self.preR_dim = preR_dim
25
+ self.num_hidden = num_hidden
26
+ self.num_classes = num_classes
27
+ self.num_heads = num_heads
28
+ self.num_preR = num_preR
29
+ self.attn_dropout = attn_dropout
30
+
31
+ if self.num_heads > 1:
32
+ self.multi_head = True
33
+ else:
34
+ self.multi_head = False
35
+
36
+ self.batch_norm = nn.BatchNorm1d(self.preR_dim*self.num_preR, affine=False)
37
+ self.fc_embeds = nn.Linear(self.preR_dim*self.num_preR, self.hidden_dim)
38
+ nn.init.kaiming_normal_(self.fc_embeds.weight, nonlinearity='linear')
39
+ nn.init.zeros_(self.fc_embeds.bias)
40
+
41
+ if self.num_hidden == 1:
42
+ layer = nn.Linear(self.num_classes*self.num_preR + self.hidden_dim, self.hidden_dim)
43
+ nn.init.kaiming_normal_(layer.weight, nonlinearity='leaky_relu')
44
+ nn.init.zeros_(layer.bias)
45
+
46
+ self.linears = nn.ModuleList([layer])
47
+
48
+ else:
49
+ first_layer = nn.Linear(self.num_classes*self.num_preR + self.hidden_dim, self.hidden_dim)
50
+ nn.init.kaiming_normal_(first_layer.weight, nonlinearity='leaky_relu')
51
+ nn.init.zeros_(first_layer.bias)
52
+
53
+ self.linears = nn.ModuleList([first_layer])
54
+
55
+ for _ in range(self.num_hidden-1):
56
+ hidden_layer = nn.Linear(self.hidden_dim, self.hidden_dim)
57
+ nn.init.kaiming_normal_(hidden_layer.weight, nonlinearity='leaky_relu')
58
+ nn.init.zeros_(hidden_layer.bias)
59
+
60
+ self.linears.append(hidden_layer)
61
+
62
+ self.out = nn.Linear(self.hidden_dim, self.num_classes)
63
+ nn.init.xavier_uniform_(self.out.weight)
64
+ nn.init.zeros_(self.out.bias)
65
+
66
+ if self.multi_head:
67
+ self.attn_embed = nn.Linear(self.num_classes*self.num_preR, self.hidden_dim*self.num_heads)
68
+
69
+ self.multihead_attn = nn.MultiheadAttention(
70
+ self.num_heads*self.hidden_dim,
71
+ self.num_heads,
72
+ dropout=self.attn_dropout,
73
+ batch_first=True
74
+ )
75
+
76
+ self.attn_out = nn.Linear(self.num_heads*self.hidden_dim, self.num_classes*self.num_preR)
77
+ else:
78
+ self.attn_embed = nn.Linear(self.num_classes*self.num_preR, self.hidden_dim)
79
+ self.attn_out = nn.Linear(self.hidden_dim, self.num_classes*self.num_preR)
80
+
81
+ def input(self,
82
+ ind_embeds: Tensor) -> Tensor:
83
+
84
+ embed = self.batch_norm(ind_embeds)
85
+ embed = self.fc_embeds(embed)
86
+ return embed
87
+
88
+ def forward(self,
89
+ logits: Tensor,
90
+ ind_embeds: Tensor) -> Tensor:
91
+
92
+ logits = self.attn_embed(logits)
93
+ q = logits
94
+ k = logits
95
+ v = logits
96
+
97
+ if not self.multi_head:
98
+ logits = F.scaled_dot_product_attention(q, k, v, dropout_p=self.attn_dropout)
99
+ else:
100
+ logits, _ = self.multihead_attn(q, k, v, need_weights=False)
101
+
102
+ logits = F.tanh(self.attn_out(logits))
103
+ embed = F.tanh(self.input(ind_embeds))
104
+
105
+ x = torch.cat([logits,embed],dim=1)
106
+
107
+ for layer in self.linears:
108
+ x = layer(x)
109
+ x = F.leaky_relu(x)
110
+
111
+ return self.out(x)
@@ -0,0 +1,562 @@
1
+ import os
2
+ from typing import List
3
+
4
+ import numpy as np
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+ import torch.optim as optim
9
+ from torch.nn import Module
10
+ from torch.optim import Optimizer
11
+ from torch.optim.lr_scheduler import LRScheduler
12
+ from torch.utils.data import DataLoader
13
+
14
+ from tqdm import tqdm
15
+ from sklearn.metrics import accuracy_score
16
+ from datetime import datetime
17
+ import math
18
+ import matplotlib.pyplot as plt
19
+
20
+ from symbiotic_learning.loss import embed_sim, embed_summand, embed_loss
21
+
22
+ def reports_summary(reports: List[dict], epoch: int, tags: List[str]) -> None:
23
+ """
24
+ Prints the contents of a list of training/validation/testing reports with labels.
25
+ """
26
+
27
+ label_lookup = {}
28
+ label_lookup['personal'] = 'Average Personal Loss'
29
+ label_lookup['symbiotic'] = 'Average Symbiotic Loss'
30
+ label_lookup['accuracy'] = 'Accuracy'
31
+ label_lookup['embedding'] = 'Average Embedding Loss'
32
+ label_lookup['blame'] = 'Average Blame Loss'
33
+
34
+ print(f'Summary:')
35
+ for i, report in enumerate(reports):
36
+ tag = tags[i]
37
+
38
+ print(f'\t- {tag}:')
39
+ for label in report.keys():
40
+ if report[label] == None:
41
+ continue
42
+ else:
43
+ if label != 'accuracy':
44
+ print(f'\t\t-- {label_lookup[label]}: {report[label]:.4}')
45
+ else:
46
+ print()
47
+ print(f'\t\t-- {label_lookup[label]}: {report[label]:.4}')
48
+ print()
49
+
50
+ return
51
+
52
+ def fill_lt_reports(lt_reports: List[dict], reports: List[dict], phase: str) -> None:
53
+ """
54
+ Initializes or appends report data to lifetime reports.
55
+ """
56
+
57
+ for lt_report, report in zip(lt_reports, reports):
58
+ keys = list(report.keys())
59
+ for key in keys:
60
+ if key not in lt_report[phase].keys():
61
+ lt_report[phase][key] = [report[key]]
62
+ else:
63
+ lt_report[phase][key].append(report[key])
64
+
65
+ return
66
+
67
+ def train_one_epoch(
68
+ train_loader: DataLoader,
69
+ models: List[Module],
70
+ opts: List[Optimizer],
71
+ scheds: List[LRScheduler],
72
+ collab_params: List[float],
73
+ temp: float,
74
+ epoch: int,
75
+ criterion: Module,
76
+ uplift: int = 10,
77
+ eps: float = 1e-7,
78
+ lamb: float = 1.0) -> List[dict]:
79
+
80
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
81
+
82
+ reports = [{} for model in models]
83
+ personals = [[] for model in models]
84
+ embs = [[] for model in models[:-1]]
85
+ blames = [[] for model in models[:-1]]
86
+ syms = [[] for model in models]
87
+
88
+ for model in models:
89
+ model.train().to(device)
90
+
91
+ ## TRAINING LOOP
92
+ pbar_train = tqdm(train_loader, total=len(train_loader))
93
+ pbar_train.set_description(f'Epoch {epoch}: Training')
94
+ for img, label in pbar_train:
95
+ for opt in opts:
96
+ opt.zero_grad()
97
+
98
+ img = img.to(device)
99
+
100
+ inds = [model.embed(img) for model in models[:-1]]
101
+
102
+ ind_concat = torch.cat(inds, dim=1)
103
+ ind_stack = torch.stack(inds)
104
+
105
+ logits_list = [model(ind_concat.clone()).clone() for model in models[:-1]]
106
+
107
+ L_is = []
108
+ L_emb_is = []
109
+ for i, personal in enumerate(personals[:-1]):
110
+ logits = logits_list[i].clone().cpu()
111
+ L_i = criterion(logits, label)
112
+ L_is.append(L_i)
113
+ personal.append(float(L_i.clone().detach()))
114
+
115
+ src = ind_stack[i].clone()
116
+ auxs = [ind_stack[i].clone() for i in range(len(ind_stack))]
117
+ auxs.pop(i)
118
+ L_emb_i = embed_loss(src, auxs).cpu()
119
+ L_emb_is.append(L_emb_i)
120
+ embs[i].append(float(L_emb_i.clone().detach()))
121
+
122
+ L_i_tensor = torch.stack(L_is)
123
+ L_emb_tensor = torch.stack(L_emb_is)
124
+
125
+ L_syms = []
126
+ for i in range(len(L_is)):
127
+ param = collab_params[i]
128
+ aux_idxs = [j!=i for j in range(len(L_is))]
129
+ aux_L = L_i_tensor.clone()[aux_idxs]
130
+
131
+ L_sym_i = (1-param)*L_i_tensor[i] + param*torch.sum(aux_L) + (param**2)*L_emb_tensor[i]
132
+
133
+ L_syms.append(L_sym_i)
134
+
135
+ if epoch >= uplift:
136
+ upstream_input = torch.cat(logits_list, dim=1).to(device)
137
+
138
+ final_logits = models[-1](upstream_input, ind_concat.clone()).to('cpu')
139
+
140
+ L_F = criterion(final_logits, label)
141
+ personals[-1].append(float(L_F.clone().detach()))
142
+
143
+ L_up_sum = eps + torch.sum(L_i_tensor).detach()
144
+
145
+ L_sym_F = L_F*(1+torch.exp(eps-temp*L_up_sum))
146
+ L_syms.append(L_sym_F)
147
+
148
+ for i, L_i in enumerate(L_i_tensor.clone().detach()):
149
+ L_blame_i = lamb*(L_i/L_up_sum)*L_F.clone()
150
+ blames[i].append(float(L_blame_i.clone().detach()))
151
+ L_syms[i] = L_syms[i] + L_blame_i
152
+
153
+ for i, L_sym_i in enumerate(L_syms):
154
+ syms[i].append(float(L_sym_i.clone().detach()))
155
+ if i != len(L_syms):
156
+ L_sym_i.backward(retain_graph=True)
157
+ else:
158
+ L_sym_i.backward()
159
+
160
+ for i in range(len(syms)-1):
161
+ opts[i].step()
162
+ scheds[i].step()
163
+
164
+ if epoch >= uplift:
165
+ opts[-1].step()
166
+ scheds[-1].step()
167
+
168
+ ## FILL REPORTS
169
+ for i, report in enumerate(reports):
170
+
171
+ if i != len(reports)-1:
172
+ report['personal'] = np.mean(personals[i])
173
+ report['embedding']= np.mean(embs[i])
174
+ if epoch >= uplift:
175
+ report['blame'] = np.mean(blames[i])
176
+ else:
177
+ report['blame'] = None
178
+
179
+ report['symbiotic'] = np.mean(syms[i])
180
+ else:
181
+ if epoch >= uplift:
182
+ report['personal'] = np.mean(personals[i])
183
+ report['symbiotic'] = np.mean(syms[i])
184
+ else:
185
+ report['personal'] = None
186
+ report['symbiotic'] = None
187
+
188
+ return reports
189
+
190
+ def eval_one_epoch(
191
+ eval_loader: DataLoader,
192
+ models: List[Module],
193
+ collab_params: List[float],
194
+ temp: float,
195
+ epoch: int,
196
+ criterion: Module,
197
+ uplift: int = 10,
198
+ eps: float = 1e-7,
199
+ lamb: float = 1.0,
200
+ phase: str = 'Validation') -> List[dict]:
201
+
202
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
203
+
204
+ reports = [{} for model in models]
205
+ personals = [[] for model in models]
206
+ embs = [[] for model in models[:-1]]
207
+ blames = [[] for model in models[:-1]]
208
+ syms = [[] for model in models]
209
+ preds = [[] for model in models]
210
+
211
+ for model in models:
212
+ model.eval()
213
+ model.to(device)
214
+
215
+ all_labels = []
216
+
217
+ ## VALIDATION LOOP
218
+ pbar_eval = tqdm(eval_loader, total=len(eval_loader))
219
+ pbar_eval.set_description(f'Epoch {epoch}: {phase}')
220
+ with torch.no_grad():
221
+ for img, label in pbar_eval:
222
+
223
+ img = img.to(device)
224
+ all_labels += label.tolist()
225
+
226
+ inds = [model.embed(img) for model in models[:-1]]
227
+
228
+ ind_concat = torch.cat(inds, dim=1)
229
+ ind_stack = torch.stack(inds)
230
+
231
+ logits_list = [model(ind_concat.clone()).clone() for model in models[:-1]]
232
+
233
+ probs_list = [F.softmax(logits.clone(),dim=-1).to('cpu') for logits in logits_list]
234
+ preds_list = [torch.argmax(probs.clone(),dim=1).tolist() for probs in probs_list]
235
+
236
+ L_is = []
237
+ L_emb_is = []
238
+ for i, personal in enumerate(personals[:-1]):
239
+ logits = logits_list[i].clone().cpu()
240
+ L_i = criterion(logits, label)
241
+ L_is.append(L_i)
242
+ personal.append(float(L_i.clone()))
243
+
244
+ preds[i] += preds_list[i]
245
+
246
+ src = ind_stack[i].clone()
247
+ auxs = [ind_stack[i].clone() for i in range(len(ind_stack))]
248
+ auxs.pop(i)
249
+ L_emb_i = embed_loss(src, auxs).cpu()
250
+ L_emb_is.append(L_emb_i)
251
+ embs[i].append(float(L_emb_i.clone()))
252
+
253
+ L_i_tensor = torch.stack(L_is)
254
+ L_emb_tensor = torch.stack(L_emb_is)
255
+
256
+ L_syms = []
257
+ for i in range(len(L_is)):
258
+ param = collab_params[i]
259
+ aux_idxs = [j!=i for j in range(len(L_is))]
260
+ aux_L = L_i_tensor.clone()[aux_idxs]
261
+
262
+ L_sym_i = (1-param)*L_i_tensor[i] + param*torch.sum(aux_L) + (param**2)*L_emb_tensor[i]
263
+
264
+ L_syms.append(L_sym_i)
265
+
266
+ if epoch >= uplift:
267
+ upstream_input = torch.cat(logits_list, dim=1).to(device)
268
+
269
+ final_logits = models[-1](upstream_input, ind_concat.clone()).to('cpu')
270
+ final_probs = F.softmax(final_logits.clone(), dim=-1)
271
+ final_preds = torch.argmax(final_probs, dim=1).tolist()
272
+ preds[-1] += final_preds
273
+
274
+ L_F = criterion(final_logits, label)
275
+ personals[-1].append(float(L_F.clone()))
276
+
277
+ L_up_sum = eps + torch.sum(L_i_tensor)
278
+
279
+ L_sym_F = L_F*(1+torch.exp(eps-temp*L_up_sum))
280
+ L_syms.append(L_sym_F)
281
+
282
+ for i, L_i in enumerate(L_i_tensor.clone()):
283
+ L_blame_i = lamb*(L_i/L_up_sum)*L_F.clone()
284
+ blames[i].append(float(L_blame_i.clone()))
285
+ L_syms[i] = L_syms[i] + L_blame_i
286
+
287
+ for i, L_sym_i in enumerate(L_syms):
288
+ syms[i].append(float(L_sym_i.clone()))
289
+
290
+ ## FILL REPORTS
291
+ for i, report in enumerate(reports):
292
+
293
+ if i != len(reports)-1:
294
+ report['personal'] = np.mean(personals[i])
295
+ report['embedding']= np.mean(embs[i])
296
+ if epoch >= uplift:
297
+ report['blame'] = np.mean(blames[i])
298
+ else:
299
+ report['blame'] = None
300
+
301
+ report['symbiotic'] = np.mean(syms[i])
302
+ report['accuracy'] = accuracy_score(all_labels, preds[i])
303
+ else:
304
+ if epoch >= uplift:
305
+ report['personal'] = np.mean(personals[i])
306
+ report['symbiotic'] = np.mean(syms[i])
307
+ report['accuracy'] = accuracy_score(all_labels, preds[i])
308
+ else:
309
+ report['personal'] = None
310
+ report['symbiotic'] = None
311
+ report['accuracy'] = None
312
+
313
+ return reports
314
+
315
+ def train(
316
+ epochs: int,
317
+ models: List[Module],
318
+ opts: List[Optimizer],
319
+ scheds: List[LRScheduler],
320
+ data_loaders: List[DataLoader],
321
+ collab_params: List[float],
322
+ temp: float,
323
+ criterion: Module,
324
+ tags: List[str] = False,
325
+ uplift: int = 10,
326
+ eps: float = 1e-7,
327
+ lamb: float = 1.0,
328
+ save_path: str = False,
329
+ save_best: bool = False,
330
+ save_end: bool = False,
331
+ save_before_uplift: bool = False,
332
+ load_at_uplift: bool = False) -> None:
333
+
334
+ if not os.path.exists(save_path):
335
+ os.mkdir(save_path)
336
+
337
+ if tags==False:
338
+ tags = [f'Model_{i}' for i in range(len(models)-1)] + ['Readout']
339
+
340
+ for tag in tags:
341
+ os.mkdir(f'{save_path}/{tag}')
342
+
343
+ train_loader, val_loader, test_loader = data_loaders
344
+
345
+ len_train = len(train_loader)
346
+ len_valid = len(val_loader)
347
+ len_test = len(test_loader)
348
+
349
+ lt_reports = []
350
+
351
+ for i in range(len(models)):
352
+ lt_report_i = {
353
+ 'Validation': {},
354
+ 'Testing': {}
355
+ }
356
+
357
+ lt_reports.append(lt_report_i)
358
+
359
+ if load_at_uplift:
360
+ epoch_range = range(uplift, epochs)
361
+ else:
362
+ epoch_range = range(0, epochs)
363
+
364
+ best_epoch = 0.0
365
+ best_readout_acc = 0.0
366
+ for epoch in epoch_range:
367
+
368
+ train_reports = train_one_epoch(
369
+ train_loader,
370
+ models,
371
+ opts,
372
+ scheds,
373
+ collab_params,
374
+ temp,
375
+ epoch,
376
+ criterion,
377
+ uplift=uplift,
378
+ eps=eps,
379
+ lamb=lamb
380
+ )
381
+
382
+ reports_summary(train_reports, epoch, tags)
383
+
384
+ valid_reports = eval_one_epoch(
385
+ val_loader,
386
+ models,
387
+ collab_params,
388
+ temp,
389
+ epoch,
390
+ criterion,
391
+ uplift=uplift,
392
+ eps=eps,
393
+ lamb=lamb,
394
+ phase='Validation'
395
+ )
396
+
397
+ if epoch >= uplift and save_best==True:
398
+ readout_acc = valid_reports[-1]['accuracy']
399
+
400
+ if readout_acc > best_readout_acc:
401
+ best_readout_acc = readout_acc
402
+ best_epoch = epoch
403
+
404
+ best_checkpoints = []
405
+ for i, model in enumerate(models):
406
+ tag = tags[i]
407
+ opt = opts[i]
408
+ sched = scheds[i]
409
+
410
+ checkpoint = {
411
+ 'epoch': best_epoch,
412
+ 'model': model.state_dict(),
413
+ 'opt': opt.state_dict(),
414
+ 'sched': sched.state_dict(),
415
+ 'last_step': sched.last_epoch
416
+ }
417
+
418
+ best_checkpoints.append(checkpoint)
419
+
420
+ reports_summary(valid_reports, epoch, tags)
421
+ fill_lt_reports(lt_reports, valid_reports, 'Validation')
422
+
423
+ if (epoch%5 == 0 or epoch == epochs-1) and epoch != 0:
424
+ test_reports = eval_one_epoch(
425
+ test_loader,
426
+ models,
427
+ collab_params,
428
+ temp,
429
+ epoch,
430
+ criterion,
431
+ uplift=uplift,
432
+ eps=eps,
433
+ lamb=lamb,
434
+ phase='Testing'
435
+ )
436
+
437
+ reports_summary(test_reports, epoch, tags)
438
+ fill_lt_reports(lt_reports, test_reports, 'Testing')
439
+
440
+ if epoch == uplift-1 and save_before_uplift == True:
441
+ for i, model in enumerate(models):
442
+ tag = tags[i]
443
+ opt = opts[i]
444
+ sched = scheds[i]
445
+
446
+ checkpoint = {
447
+ 'epoch': epoch,
448
+ 'model': model.state_dict(),
449
+ 'opt': opt.state_dict(),
450
+ 'sched': sched.state_dict(),
451
+ 'last_step': sched.last_epoch
452
+ }
453
+
454
+ torch.save(checkpoint, f'{save_path}/{tag}_pre-uplift.pt')
455
+
456
+
457
+ for tag, lt_report in zip(tags, lt_reports):
458
+ for phase in lt_report.keys():
459
+ plot_phase(
460
+ lt_report,
461
+ epochs,
462
+ uplift,
463
+ save_path,
464
+ load_at_uplift=load_at_uplift,
465
+ phase=phase,
466
+ tag=tag
467
+ )
468
+
469
+ if save_end:
470
+ for i, model in enumerate(models):
471
+ tag = tags[i]
472
+ opt = opts[i]
473
+ sched = scheds[i]
474
+
475
+ checkpoint = {
476
+ 'epoch': epoch,
477
+ 'model': model.state_dict(),
478
+ 'opt': opt.state_dict(),
479
+ 'sched': sched.state_dict(),
480
+ 'last_step': sched.last_epoch
481
+ }
482
+
483
+ torch.save(checkpoint, f'{save_path}/{tag}.pt')
484
+
485
+ if save_best:
486
+ print('Best Epoch: ', best_epoch)
487
+ print('Best Readout Accuracy: ', best_readout_acc)
488
+ for tag, checkpoint in zip(tags, best_checkpoints):
489
+ torch.save(checkpoint, f'{save_path}/{tag}_Best.pt')
490
+
491
+ return
492
+
493
+ def plot_phase(
494
+ lt_report: dict,
495
+ epochs: int,
496
+ uplift: int,
497
+ save_path: str,
498
+ load_at_uplift: bool = False,
499
+ phase: str = 'Validation',
500
+ tag: str = 'Model_0') -> None:
501
+
502
+ label_lookup = {}
503
+ label_lookup['personal'] = 'Average Personal Loss'
504
+ label_lookup['symbiotic'] = 'Average Symbiotic Loss'
505
+ label_lookup['accuracy'] = 'Accuracy'
506
+ label_lookup['embedding'] = 'Average Embedding Loss'
507
+ label_lookup['blame'] = 'Average Blame Loss'
508
+
509
+ if not load_at_uplift:
510
+ if phase == 'Validation':
511
+ full_axis = list(range(0, epochs))
512
+ post_uplift_axis = list(range(uplift, epochs))
513
+
514
+ elif phase == 'Testing':
515
+ full_axis = list(np.arange(5,epochs,5)) + [epochs]
516
+ post_uplift_axis = list(np.arange( max([math.ceil(uplift/5)*5, 5]),epochs,5)) + [epochs]
517
+ else:
518
+ if phase == 'Validation':
519
+ full_axis = list(range(uplift, epochs))
520
+ post_uplift_axis = list(range(uplift, epochs))
521
+
522
+ elif phase == 'Testing':
523
+ full_axis = list(np.arange(uplift,epochs,5)) + [epochs]
524
+ post_uplift_axis = list(np.arange(math.ceil(uplift/5)*5,epochs,5)) + [epochs]
525
+
526
+ full_axis = np.array(full_axis)
527
+ post_uplift_axis = np.array(post_uplift_axis)
528
+
529
+ for key in lt_report[phase].keys():
530
+ if tag == 'Readout' or key == 'blame':
531
+ x_axis = post_uplift_axis
532
+ else:
533
+ x_axis = full_axis
534
+
535
+ data_clean = [float(x) for x in lt_report[phase][key] if x is not None]
536
+
537
+ if key == 'accuracy' and tag == 'Readout':
538
+ idx_best = np.argmax(data_clean)
539
+ best_epoch = x_axis[idx_best]
540
+ best_acc = data_clean[idx_best]
541
+ label = f'{tag}: Best Accuracy=\n{best_acc:.4f} at Epoch {best_epoch}'
542
+ else:
543
+ label = f'{tag}: {label_lookup[key]}'
544
+
545
+ try:
546
+ plt.plot(x_axis, np.array(data_clean), color='black', linestyle='-', label=label)
547
+ except:
548
+ print(f"Error Encountered Plotting {tag}'s {phase} {key.capitalize()} Report")
549
+ print('x axis: ', x_axis)
550
+ print('data: ', data_clean)
551
+ continue
552
+
553
+ plt.title(f'{tag} {phase}: {label_lookup[key]} Per Epoch')
554
+ plt.xlabel('Epoch')
555
+ plt.ylabel(f'{label_lookup[key]}')
556
+ plt.grid()
557
+ if tag != 'Readout' and uplift != 0: plt.axvline(x=uplift, linestyle='dashed', color='red', label=f'Uplift: Epoch {uplift}')
558
+ plt.legend()
559
+ plt.savefig(f'{save_path}/{tag}/{phase}_{key}.png')
560
+ plt.clf()
561
+
562
+ return
@@ -0,0 +1,36 @@
1
+ import torch
2
+ from torch import Tensor
3
+ import torch.nn as nn
4
+ import torch.nn.functional as F
5
+
6
+ def embed_sim(x1: Tensor, x2: Tensor) -> Tensor:
7
+ """
8
+ Computes the cosine similarity between two embeddings scaled to the range [0,1].
9
+ """
10
+
11
+ cosine_sim = F.cosine_similarity(x1, x2, dim=0)
12
+ return (cosine_sim+1)/2
13
+
14
+ def embed_summand(src: Tensor, aux: Tensor, delta: float = 0.5) -> Tensor:
15
+ """
16
+ Computes a sum-term in a pre-Readout model's Embedding Loss.
17
+ """
18
+ return torch.exp( (1/delta)*embed_sim(src, aux) ) - 1
19
+
20
+ def embed_loss(src: Tensor, auxs: Tensor) -> Tensor:
21
+ """
22
+ Computes a pre-Readout model's Embedding Loss.
23
+ """
24
+
25
+ num_aux = len(auxs)
26
+ temp_func = lambda aux: embed_summand(src, aux)
27
+ temp_func = torch.vmap(temp_func)
28
+
29
+ auxs = torch.stack(auxs)
30
+ sum_terms = temp_func(auxs)
31
+ sum_terms = torch.sum(sum_terms, dim=0)
32
+ pre_factor = 1/num_aux
33
+
34
+ L_embed = pre_factor*sum_terms
35
+
36
+ return L_embed.mean()