symbiotic-learning 0.2.1__py3-none-any.whl

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,4 @@
1
+ __all__ = ['loss', 'classify']
2
+
3
+ from . import classify
4
+ from . import loss
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()
@@ -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,9 @@
1
+ symbiotic_learning/__init__.py,sha256=6U2PQ51oUyEp13qANEZhxR0D3Az0PM8cd8C-lBNmi1c,76
2
+ symbiotic_learning/loss.py,sha256=VLDYkWFQx7qBEBaDfYBtBWlfc5gQpt3sfnCxzK83MrQ,1013
3
+ symbiotic_learning/classify/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
4
+ symbiotic_learning/classify/readout.py,sha256=exNUV7QAzz1j6EGd5xTkB3yTzuF5GnrPZkVs561r3qo,4028
5
+ symbiotic_learning/classify/utils.py,sha256=aeWC_W4EdHILK8h05ZLkvwXMLT-Wp5Fz4R39INhJsVg,19083
6
+ symbiotic_learning-0.2.1.dist-info/METADATA,sha256=OCMqHkk8ydT4n42r9E5cL1I7ys7Gqxm6nH3i6Jo0izM,7963
7
+ symbiotic_learning-0.2.1.dist-info/WHEEL,sha256=W3fkpkm7-wf9vBI5Z-7s0eWkeM-spu78I8Neb98DeEg,87
8
+ symbiotic_learning-0.2.1.dist-info/licenses/LICENSE,sha256=-j1XZHmPH_SqWBn-5L1kREgpC2vC-4AIGc1Y2g8zyhY,1087
9
+ symbiotic_learning-0.2.1.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.32.4
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -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.