Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
7 changes: 4 additions & 3 deletions petab/v2/extensions/sciml_lint.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
"CheckNeuralNetworkModel",
"CheckSciMLConditionTable",
"CheckSciMLParameterTable",
"get_nn_entity_petab_ids",
]

#: Placeholder used in messages when a neural network has no ID.
Expand Down Expand Up @@ -67,7 +68,7 @@ def _nn_ids(problem: core.Problem) -> set[str]:
return ids


def _nn_entity_petab_ids(
def get_nn_entity_petab_ids(
problem: core.Problem,
) -> tuple[dict[str, str], dict[str, str], dict[str, str]]:
"""Classify NN entities referenced in the mapping table.
Expand Down Expand Up @@ -209,7 +210,7 @@ def run(self, problem: core.Problem) -> lint.ValidationIssue | None:
condition_targets = {
c.target_id for ct in problem.conditions for c in ct.changes
}
nn_inputs, nn_outputs, nn_params = _nn_entity_petab_ids(problem)
nn_inputs, nn_outputs, nn_params = get_nn_entity_petab_ids(problem)
array_input_ids = _array_input_ids(problem)
array_param_layers = _array_parameter_layers(problem)
array_param_petab_ids = {
Expand Down Expand Up @@ -333,7 +334,7 @@ class CheckSciMLConditionTable(lint.ValidationTask):
def run(self, problem: core.Problem) -> lint.ValidationIssue | None:
messages = []

nn_inputs, nn_outputs, nn_params = _nn_entity_petab_ids(problem)
nn_inputs, nn_outputs, nn_params = get_nn_entity_petab_ids(problem)
array_input_ids = _array_input_ids(problem)
array_param_layers = _array_parameter_layers(problem)
array_param_petab_ids = {
Expand Down
6 changes: 6 additions & 0 deletions petab/v2/lint.py
Original file line number Diff line number Diff line change
Expand Up @@ -1127,6 +1127,8 @@ def append_overrides(overrides):
parameter_ids -= condition_targets

if problem.extensions.sciml is not None:
from .extensions.sciml_lint import get_nn_entity_petab_ids

hybridization_targets = {
hyb.target_id for hyb in problem.extensions.sciml.hybridizations
}
Expand All @@ -1137,6 +1139,10 @@ def append_overrides(overrides):
}
parameter_ids -= hybridization_target_values

# NN outputs should not appear in the parameters table.
_, nn_outputs, _ = get_nn_entity_petab_ids(problem)
parameter_ids -= set(nn_outputs)

return parameter_ids


Expand Down
51 changes: 51 additions & 0 deletions tests/v2/test_sciml.py
Original file line number Diff line number Diff line change
Expand Up @@ -401,6 +401,57 @@ def test_parameter_posterior_requires_bounds_or_prior():
assert "net1_ps" in issue.message


def _add_observable_consuming_nn_output(problem):
"""Add an observable whose formula references an NN output directly."""
problem.add_mapping("net1_output2", "net1.outputs[0][1]")
problem.add_observable("fitness_obs", "net1_output2", noise_formula="0.05")
problem.add_measurement(
"fitness_obs", time=1, measurement=1, experiment_id="e1"
)
return problem


def test_nn_output_in_observable_formula_not_required_parameter():
"""NN outputs should not appear in the parameter table, and can appear in
observable formulas."""
from petab.v2.lint import get_required_parameters_for_parameter_table

problem = _add_observable_consuming_nn_output(_get_test_problem())

assert "net1_output2" not in get_required_parameters_for_parameter_table(
problem
)
assert problem.validate() == []


def test_nn_output_in_noise_formula_not_required_parameter():
"""Same for noise formulas."""
from petab.v2.lint import get_required_parameters_for_parameter_table

problem = _get_test_problem()
problem.add_mapping("net1_output2", "net1.outputs[0][1]")
problem.observable_tables[0]["B_obs"].noise_formula = "net1_output2"

assert "net1_output2" not in get_required_parameters_for_parameter_table(
problem
)
assert problem.validate() == []


def test_genuinely_missing_output_parameter_still_reported():
"""The NN-output carve-out does not mask real missing parameters."""
problem = _add_observable_consuming_nn_output(_get_test_problem())
# `scale` is not an NN entity and is not in the parameter table.
problem.observable_tables[0][
"fitness_obs"
].formula = "scale * net1_output2"

results = problem.validate()
assert results.has_errors()
assert any("scale" in issue.message for issue in results)
assert not any("net1_output2" in issue.message for issue in results)


# ---------------------------------------------------------------------------
# Full-problem integration
# ---------------------------------------------------------------------------
Expand Down