From 10f1096f0aeee929463314e0c88d82963de78f50 Mon Sep 17 00:00:00 2001 From: Kyle Ellrott Date: Wed, 19 Aug 2026 23:42:46 -0700 Subject: [PATCH 1/6] Testing multitask learning code against new experiment to better build out interface --- src/embkit/datasets/__init__.py | 29 ++++++++++++++++++++++++++++- src/embkit/optimize/multitask.py | 29 ++++++++++++++++++++++------- 2 files changed, 50 insertions(+), 8 deletions(-) diff --git a/src/embkit/datasets/__init__.py b/src/embkit/datasets/__init__.py index 1f13389..29c13f8 100644 --- a/src/embkit/datasets/__init__.py +++ b/src/embkit/datasets/__init__.py @@ -77,4 +77,31 @@ def __len__(self): return len(self.data) def __getitem__(self, idx): - return [self.data[idx], self.label] \ No newline at end of file + return [self.data[idx], self.label] + +class ZipDataset(Dataset): + def __init__(self, *datasets): + self.datasets = datasets + # Ensure all zipped datasets are of equal length + assert all(len(d) == len(datasets[0]) for d in datasets) + + def __len__(self): + return len(self.datasets[0]) + + def __getitem__(self, idx): + return tuple(d[idx] for d in self.datasets) + +class ChainDataset(Dataset): + def __init__(self, *datasets): + self.datasets = datasets + # Ensure all chained datasets are of equal length + assert all(len(d) == len(datasets[0]) for d in datasets) + + def __len__(self): + return len(self.datasets[0]) + + def __getitem__(self, idx): + out = [] + for d in self.datasets: + out.extend(d[idx]) + return tuple(out) \ No newline at end of file diff --git a/src/embkit/optimize/multitask.py b/src/embkit/optimize/multitask.py index 8c5cb96..977669f 100644 --- a/src/embkit/optimize/multitask.py +++ b/src/embkit/optimize/multitask.py @@ -21,8 +21,12 @@ def _unique_parameters(*models): return unique_params -def _task_loss(criterion, model, batch, device=None, unroll_inputs=False): - inputs, targets = batch +def _task_loss(criterion, model, batch, device=None, unroll_inputs=False, auto_encoder=False): + if auto_encoder: + inputs = batch + targets = batch + else: + inputs, targets = batch if device is not None: if isinstance(inputs, (tuple, list)): @@ -32,6 +36,14 @@ def _task_loss(criterion, model, batch, device=None, unroll_inputs=False): if torch.is_tensor(targets): targets = targets.to(device) + if auto_encoder: + if isinstance(inputs, (tuple, list)) and unroll_inputs: + inputs = inputs[0] + res = model(inputs) + total_loss, recon_loss, kl_loss = criterion(res.recon, inputs, res.mu, res.logvar) + return total_loss, res.recon, inputs + + if isinstance(inputs, (tuple, list)) and unroll_inputs: outputs = model(*inputs) else: @@ -44,7 +56,7 @@ def _task_loss(criterion, model, batch, device=None, unroll_inputs=False): class LearningTask: """Defines a learning task with model, dataset, and training configuration.""" - def __init__(self, model, dataset, batch_size, criterion, weight=1.0): + def __init__(self, model, dataset, batch_size, criterion, weight=1.0, auto_encoder=False, unroll_inputs=False): """ Initialize a LearningTask. @@ -60,6 +72,9 @@ def __init__(self, model, dataset, batch_size, criterion, weight=1.0): self.criterion = criterion self.weight = weight self.batch_size = batch_size + self.auto_encoder = auto_encoder + self.unroll_inputs = unroll_inputs + def _prepare_learning_tasks(tasks): @@ -113,7 +128,6 @@ def multi_task_train_weighted_sync( pairing_mode="truncate", gradient_clip_norm=None, device=None, - unroll_inputs=False, ): """ Weighted multitask training over an arbitrary list of LearningTask. @@ -153,7 +167,8 @@ def multi_task_train_weighted_sync( task.model, batch, device=device, - unroll_inputs=unroll_inputs, + unroll_inputs=task.unroll_inputs, + auto_encoder=task.auto_encoder ) task_losses.append(loss) weighted_loss = task.weight * loss @@ -184,7 +199,6 @@ def multi_task_train_interleaved( steps_per_epoch=None, gradient_clip_norm=None, device=None, - unroll_inputs=False, ): """ Interleaved multitask training over an arbitrary list of LearningTask. @@ -220,7 +234,8 @@ def multi_task_train_interleaved( task.model, batch, device=device, - unroll_inputs=unroll_inputs, + unroll_inputs=task.unroll_inputs, + auto_encoder=task.auto_encoder ) weighted_loss = task.weight * loss From 2e720b8dffb0c5d5904ddc2268704067a948dfa9 Mon Sep 17 00:00:00 2001 From: Kyle Ellrott Date: Wed, 19 Aug 2026 23:50:13 -0700 Subject: [PATCH 2/6] Adding parameter for lr schedule gamma --- src/embkit/optimize/multitask.py | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/src/embkit/optimize/multitask.py b/src/embkit/optimize/multitask.py index 977669f..127cc04 100644 --- a/src/embkit/optimize/multitask.py +++ b/src/embkit/optimize/multitask.py @@ -128,6 +128,7 @@ def multi_task_train_weighted_sync( pairing_mode="truncate", gradient_clip_norm=None, device=None, + lr_gamma=0.5, ): """ Weighted multitask training over an arbitrary list of LearningTask. @@ -142,7 +143,9 @@ def multi_task_train_weighted_sync( loaders, trainable_params = _prepare_learning_tasks(tasks) optimizer = Adam(trainable_params, lr=lr) - scheduler = StepLR(optimizer, step_size=1, gamma=0.5) + scheduler = None + if lr_gamma is not None: + scheduler = StepLR(optimizer, step_size=1, gamma=lr_gamma) pbar = tqdm(range(epochs)) for epoch in pbar: @@ -188,7 +191,8 @@ def multi_task_train_weighted_sync( postfix.update({f"loss_{i}": float(loss.detach().cpu()) for i, loss in enumerate(task_losses)}) pbar.set_postfix(**postfix) - scheduler.step() + if scheduler is not None: + scheduler.step() def multi_task_train_interleaved( @@ -199,6 +203,7 @@ def multi_task_train_interleaved( steps_per_epoch=None, gradient_clip_norm=None, device=None, + lr_gamma=0.5, ): """ Interleaved multitask training over an arbitrary list of LearningTask. @@ -211,7 +216,9 @@ def multi_task_train_interleaved( normalized_schedule = _normalize_task_schedule(task_schedule, len(tasks)) optimizer = Adam(trainable_params, lr=lr) - scheduler = StepLR(optimizer, step_size=1, gamma=0.5) + scheduler = None + if lr_gamma is not None: + scheduler = StepLR(optimizer, step_size=1, gamma=lr_gamma) if steps_per_epoch is None: steps_per_epoch = max(len(loader) for loader in loaders) @@ -252,5 +259,5 @@ def multi_task_train_interleaved( weighted_loss=float(weighted_loss.detach().cpu()), lr=optimizer.param_groups[0]["lr"], ) - - scheduler.step() + if scheduler is not None: + scheduler.step() From f1ba07740b1594585647f060eaacce8971e5968a Mon Sep 17 00:00:00 2001 From: Kyle Ellrott Date: Thu, 20 Aug 2026 14:45:19 -0700 Subject: [PATCH 3/6] Adding name values for multitask learning tasks --- src/embkit/optimize/multitask.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/src/embkit/optimize/multitask.py b/src/embkit/optimize/multitask.py index 127cc04..2b0404c 100644 --- a/src/embkit/optimize/multitask.py +++ b/src/embkit/optimize/multitask.py @@ -56,7 +56,7 @@ def _task_loss(criterion, model, batch, device=None, unroll_inputs=False, auto_e class LearningTask: """Defines a learning task with model, dataset, and training configuration.""" - def __init__(self, model, dataset, batch_size, criterion, weight=1.0, auto_encoder=False, unroll_inputs=False): + def __init__(self, model, dataset, batch_size, criterion, name="task", weight=1.0, auto_encoder=False, unroll_inputs=False): """ Initialize a LearningTask. @@ -74,6 +74,7 @@ def __init__(self, model, dataset, batch_size, criterion, weight=1.0, auto_encod self.batch_size = batch_size self.auto_encoder = auto_encoder self.unroll_inputs = unroll_inputs + self.name = name @@ -225,6 +226,7 @@ def multi_task_train_interleaved( schedule_cycle = cycle(normalized_schedule) + current_scores = {} pbar = tqdm(range(epochs)) for epoch in pbar: task_iters = [cycle(loader) for loader in loaders] @@ -252,12 +254,13 @@ def multi_task_train_interleaved( torch.nn.utils.clip_grad_norm_(trainable_params, gradient_clip_norm) optimizer.step() + current_scores[f"loss_{task.name}"] = float(loss.detach().cpu()) pbar.set_postfix( task=task_idx, - loss=float(loss.detach().cpu()), weighted_loss=float(weighted_loss.detach().cpu()), lr=optimizer.param_groups[0]["lr"], + **current_scores ) if scheduler is not None: scheduler.step() From afdd1e53c62cfb79391a4be736fe4eee7068c185 Mon Sep 17 00:00:00 2001 From: Kyle Ellrott Date: Tue, 15 Sep 2026 16:15:21 -0700 Subject: [PATCH 4/6] Working on output logging --- src/embkit/models/vae/vae.py | 4 ++-- src/embkit/optimize/__init__.py | 2 ++ src/embkit/optimize/multitask.py | 2 +- 3 files changed, 5 insertions(+), 3 deletions(-) diff --git a/src/embkit/models/vae/vae.py b/src/embkit/models/vae/vae.py index 8651f0f..42af10c 100644 --- a/src/embkit/models/vae/vae.py +++ b/src/embkit/models/vae/vae.py @@ -76,8 +76,8 @@ def build_layer_list(raw): sampling=data.get("sampling", True), ) return VAE( - encoder=build( data["encoder"]), - decoder=build( data["decoder"]), + encoder=build(data["encoder"]), + decoder=build(data["decoder"]), **data["extra_args"] ) diff --git a/src/embkit/optimize/__init__.py b/src/embkit/optimize/__init__.py index b599c79..bfd3c80 100644 --- a/src/embkit/optimize/__init__.py +++ b/src/embkit/optimize/__init__.py @@ -304,6 +304,8 @@ def fit_vae(model, data_loader = dataframe_loader(X, batch_size=batch_size, shuffle=shuffle, device=device) elif isinstance(X, torch.Tensor): data_loader = DataLoader(TensorDataset(X), batch_size=batch_size, shuffle=shuffle) + elif isinstance(X, Dataset): + data_loader = DataLoader(X, batch_size=batch_size, shuffle=shuffle) elif isinstance(X, DataLoader): data_loader = X else: diff --git a/src/embkit/optimize/multitask.py b/src/embkit/optimize/multitask.py index 2b0404c..270fb7e 100644 --- a/src/embkit/optimize/multitask.py +++ b/src/embkit/optimize/multitask.py @@ -189,7 +189,7 @@ def multi_task_train_weighted_sync( "total_loss": float(total_loss.detach().cpu()), "lr": optimizer.param_groups[0]["lr"], } - postfix.update({f"loss_{i}": float(loss.detach().cpu()) for i, loss in enumerate(task_losses)}) + postfix.update({f"loss_{tasks[i].name}": float(loss.detach().cpu()) for i, loss in enumerate(task_losses)}) pbar.set_postfix(**postfix) if scheduler is not None: From b55be37236030a7094b205c11d125a9eb9f3648e Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Tue, 15 Sep 2026 23:25:43 +0000 Subject: [PATCH 5/6] Restore unroll_inputs compatibility for multitask trainers Co-authored-by: kellrott <113868+kellrott@users.noreply.github.com> --- src/embkit/optimize/multitask.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/embkit/optimize/multitask.py b/src/embkit/optimize/multitask.py index 270fb7e..c4a892d 100644 --- a/src/embkit/optimize/multitask.py +++ b/src/embkit/optimize/multitask.py @@ -127,6 +127,7 @@ def multi_task_train_weighted_sync( epochs=5, lr=0.001, pairing_mode="truncate", + unroll_inputs=False, gradient_clip_norm=None, device=None, lr_gamma=0.5, @@ -171,7 +172,7 @@ def multi_task_train_weighted_sync( task.model, batch, device=device, - unroll_inputs=task.unroll_inputs, + unroll_inputs=(task.unroll_inputs or unroll_inputs), auto_encoder=task.auto_encoder ) task_losses.append(loss) @@ -202,6 +203,7 @@ def multi_task_train_interleaved( lr=0.001, task_schedule=None, steps_per_epoch=None, + unroll_inputs=False, gradient_clip_norm=None, device=None, lr_gamma=0.5, @@ -243,7 +245,7 @@ def multi_task_train_interleaved( task.model, batch, device=device, - unroll_inputs=task.unroll_inputs, + unroll_inputs=(task.unroll_inputs or unroll_inputs), auto_encoder=task.auto_encoder ) weighted_loss = task.weight * loss From b1e529a91ca0c89db7085cc9d6d652456818f742 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Tue, 15 Sep 2026 23:33:27 +0000 Subject: [PATCH 6/6] Address review feedback for multitask and VAE dataset handling Co-authored-by: kellrott <113868+kellrott@users.noreply.github.com> --- src/embkit/datasets/__init__.py | 14 ++++++++++---- src/embkit/optimize/__init__.py | 13 ++++++++++--- src/embkit/optimize/multitask.py | 16 ++++++++++++---- tests/datasets/test_datasets.py | 20 +++++++++++++++++++- tests/optimize/test_multitask.py | 4 ++++ tests/optimize/test_optimize_helpers.py | 6 ++++++ 6 files changed, 61 insertions(+), 12 deletions(-) diff --git a/src/embkit/datasets/__init__.py b/src/embkit/datasets/__init__.py index 29c13f8..d85ab6e 100644 --- a/src/embkit/datasets/__init__.py +++ b/src/embkit/datasets/__init__.py @@ -81,9 +81,12 @@ def __getitem__(self, idx): class ZipDataset(Dataset): def __init__(self, *datasets): + if len(datasets) == 0: + raise ValueError("ZipDataset requires at least one dataset.") + base_len = len(datasets[0]) + if not all(len(d) == base_len for d in datasets): + raise ValueError("All datasets passed to ZipDataset must have equal length.") self.datasets = datasets - # Ensure all zipped datasets are of equal length - assert all(len(d) == len(datasets[0]) for d in datasets) def __len__(self): return len(self.datasets[0]) @@ -93,9 +96,12 @@ def __getitem__(self, idx): class ChainDataset(Dataset): def __init__(self, *datasets): + if len(datasets) == 0: + raise ValueError("ChainDataset requires at least one dataset.") + base_len = len(datasets[0]) + if not all(len(d) == base_len for d in datasets): + raise ValueError("All datasets passed to ChainDataset must have equal length.") self.datasets = datasets - # Ensure all chained datasets are of equal length - assert all(len(d) == len(datasets[0]) for d in datasets) def __len__(self): return len(self.datasets[0]) diff --git a/src/embkit/optimize/__init__.py b/src/embkit/optimize/__init__.py index bfd3c80..0089162 100644 --- a/src/embkit/optimize/__init__.py +++ b/src/embkit/optimize/__init__.py @@ -249,7 +249,7 @@ def recognizer_step(batch, beta_value: float) -> Dict[str, torch.Tensor]: def fit_vae(model, - X: Union[pd.DataFrame, torch.Tensor, torch.utils.data.DataLoader], + X: Union[pd.DataFrame, torch.Tensor, Dataset, DataLoader], epochs: int = 20, lr: Optional[float] = 1e-3, beta: float = 1.0, @@ -309,7 +309,7 @@ def fit_vae(model, elif isinstance(X, DataLoader): data_loader = X else: - raise TypeError("X must be DataFrame, Tensor, or DataLoader") + raise TypeError("X must be DataFrame, Tensor, Dataset, or DataLoader") opt = _resolve_optimizer(model=model, lr=lr, optimizer=optimizer) @@ -322,7 +322,14 @@ def fit_vae(model, phases = _resolve_phases(epochs=epochs, beta=beta, beta_schedule=beta_schedule) def vae_step(batch, beta_value: float) -> Dict[str, torch.Tensor]: - (x_tensor,) = batch + if torch.is_tensor(batch): + x_tensor = batch + elif isinstance(batch, (tuple, list)): + if len(batch) != 1: + raise ValueError("VAE training batches must contain exactly one tensor input.") + x_tensor = batch[0] + else: + raise TypeError("VAE training batches must be a tensor or a single-item tuple/list.") x_tensor = x_tensor.to(device).float() res = model(x_tensor) diff --git a/src/embkit/optimize/multitask.py b/src/embkit/optimize/multitask.py index c4a892d..384b989 100644 --- a/src/embkit/optimize/multitask.py +++ b/src/embkit/optimize/multitask.py @@ -37,7 +37,7 @@ def _task_loss(criterion, model, batch, device=None, unroll_inputs=False, auto_e targets = targets.to(device) if auto_encoder: - if isinstance(inputs, (tuple, list)) and unroll_inputs: + if isinstance(inputs, (tuple, list)) and len(inputs) == 1: inputs = inputs[0] res = model(inputs) total_loss, recon_loss, kl_loss = criterion(res.recon, inputs, res.mu, res.logvar) @@ -56,7 +56,7 @@ def _task_loss(criterion, model, batch, device=None, unroll_inputs=False, auto_e class LearningTask: """Defines a learning task with model, dataset, and training configuration.""" - def __init__(self, model, dataset, batch_size, criterion, name="task", weight=1.0, auto_encoder=False, unroll_inputs=False): + def __init__(self, model, dataset, batch_size, criterion, weight=1.0, name="task", auto_encoder=False, unroll_inputs=False): """ Initialize a LearningTask. @@ -77,6 +77,14 @@ def __init__(self, model, dataset, batch_size, criterion, name="task", weight=1. self.name = name +def _task_loss_key(tasks, idx): + task = tasks[idx] + name = task.name + duplicate_count = sum(1 for t in tasks if t.name == name) + if duplicate_count > 1: + return f"loss_{name}_{idx}" + return f"loss_{name}" + def _prepare_learning_tasks(tasks): if tasks is None or len(tasks) == 0: @@ -190,7 +198,7 @@ def multi_task_train_weighted_sync( "total_loss": float(total_loss.detach().cpu()), "lr": optimizer.param_groups[0]["lr"], } - postfix.update({f"loss_{tasks[i].name}": float(loss.detach().cpu()) for i, loss in enumerate(task_losses)}) + postfix.update({_task_loss_key(tasks, i): float(loss.detach().cpu()) for i, loss in enumerate(task_losses)}) pbar.set_postfix(**postfix) if scheduler is not None: @@ -256,7 +264,7 @@ def multi_task_train_interleaved( torch.nn.utils.clip_grad_norm_(trainable_params, gradient_clip_norm) optimizer.step() - current_scores[f"loss_{task.name}"] = float(loss.detach().cpu()) + current_scores[_task_loss_key(tasks, task_idx)] = float(loss.detach().cpu()) pbar.set_postfix( task=task_idx, diff --git a/tests/datasets/test_datasets.py b/tests/datasets/test_datasets.py index b4dfde1..55011a0 100644 --- a/tests/datasets/test_datasets.py +++ b/tests/datasets/test_datasets.py @@ -4,7 +4,7 @@ import torch from torch.utils.data import Dataset -from embkit.datasets import BalancedMixer, DatasetMask +from embkit.datasets import BalancedMixer, ChainDataset, DatasetMask, ZipDataset class TinyDataset(Dataset): @@ -47,6 +47,24 @@ def test_dataset_mask_applies_mask_and_device(self): self.assertEqual(x0.device.type, "cpu") self.assertEqual(y0.device.type, "cpu") + def test_zip_dataset_validates_inputs(self): + with self.assertRaises(ValueError): + ZipDataset() + + d1 = TinyDataset([1, 2]) + d2 = TinyDataset([3]) + with self.assertRaises(ValueError): + ZipDataset(d1, d2) + + def test_chain_dataset_validates_inputs(self): + with self.assertRaises(ValueError): + ChainDataset() + + d1 = TinyDataset([(1,), (2,)]) + d2 = TinyDataset([(3,)]) + with self.assertRaises(ValueError): + ChainDataset(d1, d2) + if __name__ == "__main__": unittest.main() diff --git a/tests/optimize/test_multitask.py b/tests/optimize/test_multitask.py index 634f4c4..d47fb05 100644 --- a/tests/optimize/test_multitask.py +++ b/tests/optimize/test_multitask.py @@ -181,6 +181,10 @@ def test_custom_weight(self): task = make_task(weight=2.5) self.assertEqual(task.weight, 2.5) + def test_positional_weight_compatibility(self): + task = LearningTask(TinyModel(), make_dataset(), 2, nn.MSELoss(), 2.0) + self.assertEqual(task.weight, 2.0) + class TestPrepareLearningTasks(unittest.TestCase): def test_returns_loaders_and_params(self): diff --git a/tests/optimize/test_optimize_helpers.py b/tests/optimize/test_optimize_helpers.py index 356e866..8d464f2 100644 --- a/tests/optimize/test_optimize_helpers.py +++ b/tests/optimize/test_optimize_helpers.py @@ -60,6 +60,12 @@ def test_fit_vae_accepts_dataloader(self): out = fit_vae(vae, loader, epochs=1, loss=BCEWithLogitsVAELoss(), progress=False) self.assertIsInstance(out, dict) + def test_fit_vae_accepts_tensor_dataset(self): + vae = BaseVAE(features=["G1", "G2"], latent_dim=1) + x = torch.tensor([[0.1, 0.2], [0.2, 0.3]], dtype=torch.float32) + out = fit_vae(vae, TensorDataset(x), epochs=1, loss=BCEWithLogitsVAELoss(), progress=False) + self.assertIsInstance(out, dict) + if __name__ == "__main__": unittest.main()