Skip to content

Fixes/multitask debugging - #75

Merged
kellrott merged 7 commits into
developfrom
fixes/multitask-debugging
Sep 15, 2026
Merged

kellrott merged 7 commits into
developfrom
fixes/multitask-debugging

Conversation

@kellrott

Copy link
Copy Markdown
Contributor

Fixing issues found while testing multitask learning examples.

Copilot AI left a comment

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.

🟡 Changes recommended

Unresolved compatibility, batching, validation, and task-loss handling issues remain.

Get a fresh assessment by requesting another Copilot review.

Pull request overview

Fixes multitask learning issues and expands VAE dataset support.

Changes:

  • Adds per-task autoencoder, naming, unrolling, and learning-rate controls.
  • Adds Dataset support to fit_vae.
  • Adds ZipDataset and ChainDataset; cleans VAE formatting.
File summaries
File Review findings
src/embkit/optimize/multitask.py Critical (3 votes): preserve the existing positional weight argument. Moderate (3 votes): retain unroll_inputs compatibility and unwrap singleton autoencoder batches. Nit (3 votes): ensure unnamed tasks produce unique loss keys.
src/embkit/optimize/__init__.py Moderate (3 votes): normalize bare-tensor dataset batches or constrain the accepted dataset shape. Nit (2 votes): update the public annotation and error message to include Dataset.
src/embkit/models/vae/vae.py Formatting cleanup; no final review findings.
src/embkit/datasets/__init__.py Moderate (1 vote each, lines 85 and 97): replace runtime assertions with explicit validation, including the empty-dataset case.
Review details

Suppressed comments (4)

src/embkit/datasets/init.py:86

  • Using assert for this runtime invariant is unsafe: Python removes assertions under -O, allowing unequal-length datasets (or no datasets, which indexes datasets[0]) through until a later and less clear failure. Raise an explicit exception after validating that at least one dataset was supplied.
        # Ensure all zipped datasets are of equal length
        assert all(len(d) == len(datasets[0]) for d in datasets)

src/embkit/datasets/init.py:98

  • Using assert for this runtime invariant is unsafe: Python removes assertions under -O, allowing unequal-length datasets (or no datasets, which indexes datasets[0]) through until a later and less clear failure. Raise an explicit exception after validating that at least one dataset was supplied.
        # Ensure all chained datasets are of equal length
        assert all(len(d) == len(datasets[0]) for d in datasets)

src/embkit/optimize/multitask.py:207

  • Removing unroll_inputs from this public function breaks the existing in-repository call at tests/optimize/test_multitask.py:382 with TypeError before training starts. Preserve the old argument for compatibility and define how it combines with the new per-task flag, rather than silently removing it.
    lr_gamma=0.5,

src/embkit/optimize/multitask.py:257

  • The same default-name collision affects interleaved training: multiple unnamed tasks all write current_scores["loss_task"], so the displayed score for earlier tasks is overwritten. Use a unique index-based key or validate task names before training.
            current_scores[f"loss_{task.name}"] = float(loss.detach().cpu())
  • Files reviewed: 3/4 changed files
  • Comments generated: 6
  • Review effort level: Lite

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread src/embkit/optimize/multitask.py Outdated
Comment on lines +307 to +308
elif isinstance(X, Dataset):
data_loader = DataLoader(X, batch_size=batch_size, shuffle=shuffle)

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 thread src/embkit/optimize/multitask.py Outdated
Comment on lines +40 to +41
if isinstance(inputs, (tuple, list)) and unroll_inputs:
inputs = inputs[0]

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 unwrapping singleton autoencoder inputs regardless of unroll_inputs.

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.

Comment on lines +307 to +308
elif isinstance(X, Dataset):
data_loader = DataLoader(X, batch_size=batch_size, shuffle=shuffle)

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.

Comment thread src/embkit/optimize/multitask.py Outdated
"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)})

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 using unique loss keys when task names collide (index-suffixed fallback for duplicates).

Co-authored-by: kellrott <113868+kellrott@users.noreply.github.com>
Co-authored-by: kellrott <113868+kellrott@users.noreply.github.com>
@github-actions

Copy link
Copy Markdown

☂️ Python Coverage

current status: ✅

Overall Coverage

Lines Covered Coverage Threshold Status
3248 2778 86% 0% 🟢

New Files

No new covered files...

Modified Files

File Coverage Status
src/embkit/datasets/_init_.py 71% 🟢
src/embkit/models/vae/vae.py 53% 🟢
src/embkit/optimize/_init_.py 82% 🟢
src/embkit/optimize/multitask.py 95% 🟢
TOTAL 76% 🟢

updated for commit: 3c2c96c by action🐍

@kellrott
kellrott merged commit 43d5472 into develop Sep 15, 2026
1 check passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants