diff --git a/petab/v2/extensions/sciml_lint.py b/petab/v2/extensions/sciml_lint.py index e9cf27c1..ba98ad32 100644 --- a/petab/v2/extensions/sciml_lint.py +++ b/petab/v2/extensions/sciml_lint.py @@ -27,6 +27,7 @@ "CheckNeuralNetworkModel", "CheckSciMLConditionTable", "CheckSciMLParameterTable", + "get_nn_entity_petab_ids", ] #: Placeholder used in messages when a neural network has no ID. @@ -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. @@ -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 = { @@ -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 = { diff --git a/petab/v2/lint.py b/petab/v2/lint.py index b0a96d94..064bc888 100644 --- a/petab/v2/lint.py +++ b/petab/v2/lint.py @@ -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 } @@ -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 diff --git a/tests/v2/test_sciml.py b/tests/v2/test_sciml.py index 55554665..65157e41 100644 --- a/tests/v2/test_sciml.py +++ b/tests/v2/test_sciml.py @@ -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 # ---------------------------------------------------------------------------