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.
- symbiotic_learning-0.2.1/.gitignore +218 -0
- symbiotic_learning-0.2.1/LICENSE +21 -0
- symbiotic_learning-0.2.1/PKG-INFO +138 -0
- symbiotic_learning-0.2.1/README.md +122 -0
- symbiotic_learning-0.2.1/pyproject.toml +22 -0
- symbiotic_learning-0.2.1/readout_arch.png +0 -0
- symbiotic_learning-0.2.1/requirements.txt +70 -0
- symbiotic_learning-0.2.1/src/symbiotic_learning/__init__.py +4 -0
- symbiotic_learning-0.2.1/src/symbiotic_learning/classify/__init__.py +0 -0
- symbiotic_learning-0.2.1/src/symbiotic_learning/classify/readout.py +111 -0
- symbiotic_learning-0.2.1/src/symbiotic_learning/classify/utils.py +562 -0
- symbiotic_learning-0.2.1/src/symbiotic_learning/loss.py +36 -0
- symbiotic_learning-0.2.1/symlearn_arch.png +0 -0
|
@@ -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
|
+

|
|
131
|
+
|
|
132
|
+
This figure shows the architecture for a Symbiotic Uplift Network with three pre-Readout models.
|
|
133
|
+
|
|
134
|
+
## Readout Architecture
|
|
135
|
+
|
|
136
|
+

|
|
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
|
+

|
|
115
|
+
|
|
116
|
+
This figure shows the architecture for a Symbiotic Uplift Network with three pre-Readout models.
|
|
117
|
+
|
|
118
|
+
## Readout Architecture
|
|
119
|
+
|
|
120
|
+

|
|
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"
|
|
Binary file
|
|
@@ -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
|
|
File without changes
|
|
@@ -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()
|
|
Binary file
|