Repository navigation
Fixes/multitask debugging - #75
Conversation
There was a problem hiding this comment.
🟡 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
Datasetsupport tofit_vae. - Adds
ZipDatasetandChainDataset; 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
assertfor this runtime invariant is unsafe: Python removes assertions under-O, allowing unequal-length datasets (or no datasets, which indexesdatasets[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
assertfor this runtime invariant is unsafe: Python removes assertions under-O, allowing unequal-length datasets (or no datasets, which indexesdatasets[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_inputsfrom this public function breaks the existing in-repository call attests/optimize/test_multitask.py:382withTypeErrorbefore 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.
| elif isinstance(X, Dataset): | ||
| data_loader = DataLoader(X, batch_size=batch_size, shuffle=shuffle) |
There was a problem hiding this comment.
Fixed in b1e529a by normalizing VAE batches in vae_step to accept bare tensor batches and validating tuple/list batch shape.
| if isinstance(inputs, (tuple, list)) and unroll_inputs: | ||
| inputs = inputs[0] |
There was a problem hiding this comment.
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, |
There was a problem hiding this comment.
Fixed in b1e529a. unroll_inputs is preserved on both multitask trainer entry points and combined with per-task task.unroll_inputs.
| elif isinstance(X, Dataset): | ||
| data_loader = DataLoader(X, batch_size=batch_size, shuffle=shuffle) |
There was a problem hiding this comment.
Fixed in b1e529a by updating fit_vae annotations and the TypeError contract to include Dataset.
| "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)}) |
There was a problem hiding this comment.
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>
☂️ Python Coverage
Overall Coverage
New FilesNo new covered files... Modified Files
|
Fixing issues found while testing multitask learning examples.