Skip to content
Merged
Show file tree
Hide file tree
Changes from 6 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 34 additions & 1 deletion src/embkit/datasets/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,4 +77,37 @@ def __len__(self):
return len(self.data)

def __getitem__(self, idx):
return [self.data[idx], self.label]
return [self.data[idx], self.label]

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

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):
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

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)
4 changes: 2 additions & 2 deletions src/embkit/models/vae/vae.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
)

Expand Down
15 changes: 12 additions & 3 deletions src/embkit/optimize/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -304,10 +304,12 @@ 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)
Comment on lines +311 to +312

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in b1e529a by normalizing VAE batches in vae_step to accept bare tensor batches and validating tuple/list batch shape.

Comment on lines +311 to +312

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in b1e529a by updating fit_vae annotations and the TypeError contract to include Dataset.

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)

Expand All @@ -320,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)
Expand Down
63 changes: 49 additions & 14 deletions src/embkit/optimize/multitask.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)):
Expand All @@ -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 len(inputs) == 1:
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:
Expand All @@ -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, name="task", auto_encoder=False, unroll_inputs=False):
"""
Initialize a LearningTask.

Expand All @@ -60,6 +72,18 @@ 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
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):
Expand Down Expand Up @@ -111,9 +135,10 @@ def multi_task_train_weighted_sync(
epochs=5,
lr=0.001,
pairing_mode="truncate",
unroll_inputs=False,
gradient_clip_norm=None,
device=None,
unroll_inputs=False,
lr_gamma=0.5,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in b1e529a. unroll_inputs is preserved on both multitask trainer entry points and combined with per-task task.unroll_inputs.

):
"""
Weighted multitask training over an arbitrary list of LearningTask.
Expand All @@ -128,7 +153,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:
Expand All @@ -153,7 +180,8 @@ def multi_task_train_weighted_sync(
task.model,
batch,
device=device,
unroll_inputs=unroll_inputs,
unroll_inputs=(task.unroll_inputs or unroll_inputs),
auto_encoder=task.auto_encoder
)
task_losses.append(loss)
weighted_loss = task.weight * loss
Expand All @@ -170,10 +198,11 @@ 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({_task_loss_key(tasks, 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(
Expand All @@ -182,9 +211,10 @@ def multi_task_train_interleaved(
lr=0.001,
task_schedule=None,
steps_per_epoch=None,
unroll_inputs=False,
gradient_clip_norm=None,
device=None,
unroll_inputs=False,
lr_gamma=0.5,
):
"""
Interleaved multitask training over an arbitrary list of LearningTask.
Expand All @@ -197,13 +227,16 @@ 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)

schedule_cycle = cycle(normalized_schedule)

current_scores = {}
pbar = tqdm(range(epochs))
for epoch in pbar:
task_iters = [cycle(loader) for loader in loaders]
Expand All @@ -220,7 +253,8 @@ def multi_task_train_interleaved(
task.model,
batch,
device=device,
unroll_inputs=unroll_inputs,
unroll_inputs=(task.unroll_inputs or unroll_inputs),
auto_encoder=task.auto_encoder
)
weighted_loss = task.weight * loss

Expand All @@ -230,12 +264,13 @@ def multi_task_train_interleaved(
torch.nn.utils.clip_grad_norm_(trainable_params, gradient_clip_norm)

optimizer.step()
current_scores[_task_loss_key(tasks, task_idx)] = 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
)

scheduler.step()
if scheduler is not None:
scheduler.step()
20 changes: 19 additions & 1 deletion tests/datasets/test_datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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()
4 changes: 4 additions & 0 deletions tests/optimize/test_multitask.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
6 changes: 6 additions & 0 deletions tests/optimize/test_optimize_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()