Skip to content

refactor: simplify data module and condition layer - #837

Draft
ndem0 wants to merge 4 commits into
devfrom
data_module_simplification
Draft

ndem0 wants to merge 4 commits into
devfrom
data_module_simplification

Conversation

@ndem0

@ndem0 ndem0 commented Sep 28, 2026 •

Copy link
Copy Markdown
Member

Description

Remove data managers (Tensor/Graph/Batch), aggregator, condition subset, creator and single-batch loader. Conditions now store data directly in tensor/graph variants and materialize into id-based batches; graph tensor fields are attached per-graph so batching concatenates node-wise. Trainer drops automatic_batching in favour of per-condition dataloader_cls and collate_fn resolution. Adapt DataNormalizer and R3Refinement callbacks and rewrite tests for the new API.

This PR fixes #833

Checklist

  • Code follows the project’s Code Style Guidelines
  • Tests have been added or updated
  • Documentation has been updated if necessary
  • Pull request is linked to an open issue

Remove data managers (Tensor/Graph/Batch), aggregator, condition subset,
creator and single-batch loader. Conditions now store data directly in
tensor/graph variants and materialize into id-based batches; graph tensor
fields are attached per-graph so batching concatenates node-wise. Trainer
drops automatic_batching in favour of per-condition dataloader_cls and
collate_fn resolution. Adapt DataNormalizer and R3Refinement callbacks and
rewrite tests for the new API.

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

There are confirmed runtime-breaking issues in the new condition/data path (notably graph input equation evaluation and slice indexing for graph conditions) that must be fixed before merging.

Review effort: Lite
Findings: 2 Medium severity

Open (2)
What changed in this PR

This PR refactors PINA’s data pipeline by removing the legacy data-manager/creator/aggregator stack and shifting batching/materialization responsibilities into condition implementations, with DataModule now operating on per-condition sample-id tensors and materializing condition batches on-device.

Changes:

  • Replaced _Aggregator + _Creator with Batcher + MultiLoader, and rewired DataModule to split per-condition id tensors and call condition.materialize(...) during device transfer.
  • Introduced TensorCondition / GraphCondition base classes to unify data storage + batching for tensor and graph conditions; updated condition variants to use these.
  • Updated Trainer API to drop automatic_batching and add per-condition dataloader_cls / collate_fn resolution; adapted callbacks and tests accordingly.
File Description
tests/​test_trainer.py Updates Trainer constructor tests and adds coverage for per-condition dataloader/collate option validation & warnings.
tests/​test_data/​test_tensor_data_manager.py Removes tests for deleted tensor data manager abstraction.
tests/​test_data/​test_single_batch_data_loader.py Removes tests for deleted single-batch loader.
tests/​test_data/​test_loader.py Replaces _Aggregator tests with MultiLoader tests.
tests/​test_data/​test_graph_data_manager.py Removes tests for deleted graph data manager abstraction.
tests/​test_data/​test_data_module.py Updates DataModule tests to assert id-tensor subsets instead of _ConditionSubset.
tests/​test_data/​test_creator.py Removes tests for deleted dataloader creator.
tests/​test_data/​test_condition_subset.py Removes tests for deleted condition subset wrapper.
tests/​test_condition/​test_time_series_condition.py Updates tests to use materialize and new tensor-backed condition storage.
tests/​test_condition/​test_input_target_condition.py Updates tests for dict-based batches and materialize semantics across tensor/graph variants.
tests/​test_condition/​test_input_equation_condition.py Updates tests to use materialize for tensor/graph inputs.
tests/​test_condition/​test_graph_time_series_condition.py Updates tests for GraphTimeSeriesCondition and materialize behavior.
tests/​test_condition/​test_data_condition.py Updates tests for dict-based batches and conditional variables under new materialization.
tests/​test_callback/​test_data_normalizer.py Aligns DataNormalizer tests with condition-owned data storage.
pina/​data/​manager.py Removes top-level exports for deleted data manager layer.
pina/​data/​__init__.py Exposes new Batcher/MultiLoader APIs and removes old internal exports.
pina/​condition/​__init__.py Updates public condition exports to include new graph/tensor condition base types.
pina/​_src/​data/​single_batch_data_loader.py Deletes single-batch loader implementation.
pina/​_src/​data/​manager/​tensor_data_manager.py Deletes tensor data manager implementation.
pina/​_src/​data/​manager/​graph_data_manager.py Deletes graph data manager implementation.
pina/​_src/​data/​manager/​data_manager.py Deletes data manager factory.
pina/​_src/​data/​manager/​data_manager_interface.py Deletes data manager interface.
pina/​_src/​data/​manager/​batch_manager.py Deletes legacy batch container.
pina/​_src/​data/​loader.py Adds MultiLoader to aggregate per-condition dataloaders.
pina/​_src/​data/​data_module.py Refactors DataModule to id-based subsets + on-device condition.materialize batching.
pina/​_src/​data/​creator.py Deletes old multi-condition dataloader creator.
pina/​_src/​data/​condition_subset.py Deletes old subset wrapper used for cyclic indexing/auto-batching.
pina/​_src/​data/​batcher.py Adds Batcher to construct per-condition dataloaders over id tensors.
pina/​_src/​data/​aggregator.py Deletes old _Aggregator implementation.
pina/​_src/​core/​trainer.py Drops automatic_batching, adds per-condition dataloader/collate option resolution, and passes into DataModule.
pina/​_src/​condition/​time_series_condition.py Refactors time-series condition to tensor-backed storage with shared unroll helper.
pina/​_src/​condition/​tensor_condition.py Introduces tensor-backed storage/materialization implementation for conditions.
pina/​_src/​condition/​input_target_condition.py Switches to tensor/graph condition variants and dict-based batches.
pina/​_src/​condition/​input_equation_condition.py Switches to tensor/graph condition variants and shared evaluation path.
pina/​_src/​condition/​graph_time_series_condition.py Refactors graph time-series condition to graph-backed storage/materialization.
pina/​_src/​condition/​graph_condition.py Introduces graph-backed storage/materialization implementation for conditions.
pina/​_src/​condition/​domain_equation_condition.py Clarifies non-materializable semantics for domain-sampled conditions.
pina/​_src/​condition/​data_condition.py Switches to tensor/graph condition variants and updates semantics/docs.
pina/​_src/​condition/​condition_interface.py Updates condition contract to center on materialize(...).
pina/​_src/​condition/​base_condition.py Removes legacy dataloader/collate mechanisms from the base class and adds _unwrap_single.
pina/​_src/​callback/​refinement/​r3_refinement.py Avoids autograd graph retention by detaching selected refinement points.
pina/​_src/​callback/​refinement/​base_refinement.py Updates refinement bookkeeping for id-tensor datasets.
pina/​_src/​callback/​processing/​data_normalizer.py Refactors normalizer to operate on condition-held data rather than dataset wrappers.

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

Comment on lines +123 to +125
idx = _normalize_ids(idx)
graphs = [self._graphs[i] for i in idx]
return {self.graph_key: graphs}
Comment on lines +82 to +84
samples = batch["input"].requires_grad_(True)
output = solver.forward(samples)
return self.equation.residual(samples, output, solver._params)
@GiovanniCanali
GiovanniCanali changed the base branch from master to dev October 1, 2026 11:33

@GiovanniCanali GiovanniCanali left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Great job, @ndem0! This significantly simplifies the logic connecting conditions and batching. However, there are still a few points that require attention:

  1. The documentation does not currently reflect the changes introduced in the source code. Please update it accordingly.
  2. Is the Condition factory class affected in any way by the proposed changes? Please verify that its behavior remains consistent.
  3. Tests for Batcher are currently missing.
  4. Each condition currently follows a common structure, with a base class that is then specialized into Tensor and Graph implementations. TimeSeriesCondition, however, does not follow this pattern, as it develops two independent classes. Would it be possible to introduce a common base class with corresponding specializations, consistently with the other condition classes?
  5. The materialize method appears to be identical in both TensorCondition and GraphCondition. Consider moving this implementation to BaseCondition to avoid duplication.

# Select points with residual above the mean
mask = (residuals >= residuals.mean()).flatten()
high_residual_pts = current_points[mask]
high_residual_pts = current_points[mask].detach()

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

I believe this .detach() stops gradient tracking for points with high residuals in subsequent epochs. I would remove it to be safe.

Comment thread pina/_src/data/loader.py

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

PINA's naming convention requires each file to be named after the class it contains. Please rename loader.py to multi_loader.py, or update the class name accordingly. Remember to rename the corresponding test file as well.

Comment thread pina/_src/core/trainer.py
:rtype: tuple[dict[str, type | None], dict[str, Callable | None]]
"""

def _normalize(value, label):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

_normalize is somewhat ambiguous, particularly in the context of datasets. Consider renaming the utility function to make its purpose clearer.

@GiovanniCanali GiovanniCanali added enhancement New feature or request pr-to-fix Label for PR that needs modification labels Oct 1, 2026

@dario-coscia dario-coscia left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Thanks @ndem0 ! I have checked the src files only and left some comments. I did not check the tests, I will do once the PR is ready for review and not draft anymore.

Anyway, I like the approach as it also more readable and reduces the number of lines of code!



class BaseCondition(ConditionInterface):
def _unwrap_single(value):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Maybe it is better to put this in utils.py

:return: The single element when ``value`` is a list of length one,
otherwise ``value``.
"""
if isinstance(value, list) and len(value) == 1:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

To be sure I would say (list, tuple) to check instance

:param dict kwargs: The keyword arguments containing the data to be
stored.
:return: The stored data in a suitable format.
:return: The stored data.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

What format? Can we specify it?

:param batch_fn: Optional callable overriding the default batched
construction. It receives the selected raw data and the ids and
returns the batch. Default is ``None``.
:return: The materialized batch.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Can we specify when None is passed what happens. Also can we specify the dict how it is. Right now it is very hard to understand the data structure composition.

"""
Materialize the data points at the given ids.

:raises NotImplementedError: Always raised since the data points are

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

What do you mean sample at training time? Do you mean when we discretise the domain or do you actually mean at each gradient iteration? Maybe a typo but worth checking. The latter could be very expensive in compute and time


from pina._src.condition.base_condition import BaseCondition
from pina._src.condition.tensor_condition import (
_move_to_device,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

This goes to utils.py

from pina._src.condition.base_condition import BaseCondition
from pina._src.condition.tensor_condition import (
_move_to_device,
_normalize_ids,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

This goes to utils.py

_avail_input_cls = (Data, Graph)

# Name of the graph attribute holding the temporal data
_key = "x"

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

You define this but seems to not use it, e.g. in new can you double check if it is needed ?

check_consistency(key, str)
check_positive_integer(n_windows, strict=True)
check_positive_integer(unroll_length, strict=True)
print(input.x)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Remove

"""
samples = batch["input"].requires_grad_(True)
output = solver.forward(samples)
return self.equation.residual(samples, output, solver._params)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Check this, I agree with copilot

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request pr-to-fix Label for PR that needs modification

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants