Skip to content

Commit 7a3e0ce

Browse files
committed
refactor sciml specifics into more helpers
1 parent c7d3929 commit 7a3e0ce

3 files changed

Lines changed: 96 additions & 74 deletions

File tree

petab/v2/core.py

Lines changed: 8 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -312,7 +312,6 @@ def __iadd__(self, other: T) -> BaseTable[T]:
312312
# SciML extension classes — imported after BaseTable is defined to avoid
313313
# circular imports (sciml.py does not import from core.py).
314314
from .extensions.sciml import ( # noqa: E402
315-
HybridizationTable,
316315
SciMLConfig,
317316
SciMLExt,
318317
)
@@ -1350,44 +1349,12 @@ def from_yaml(
13501349
else None
13511350
)
13521351

1353-
extensions = ProblemExtensions()
1354-
if config.extensions and config.extensions.get(C.EXT_ID_SCIML):
1355-
from petab_sciml import ArrayDataStandard, NNModel, NNModelStandard
1356-
1357-
# Neural network classes are constructed via pytorch for now to get
1358-
# the proper inputs
1359-
neural_networks = [
1360-
NNModel.from_pytorch_module(
1361-
NNModelStandard.load_data(
1362-
_generate_path(
1363-
file_path=nn_config.location,
1364-
base_path=base_path,
1365-
)
1366-
).to_pytorch_module(),
1367-
nn_model_id=nn_id,
1368-
)
1369-
for nn_id, nn_config in (
1370-
config.extensions[C.EXT_ID_SCIML].neural_networks or {}
1371-
).items()
1372-
]
1373-
1374-
hybridization_tables = [
1375-
HybridizationTable.from_tsv(f, base_path)
1376-
for f in config.extensions[C.EXT_ID_SCIML].hybridization_files
1377-
]
1378-
1379-
array_data_files = [
1380-
ArrayDataStandard.load_data(_generate_path(f, base_path))
1381-
for f in config.extensions[C.EXT_ID_SCIML].array_files
1382-
]
1383-
1384-
extensions = ProblemExtensions(
1385-
sciml=SciMLExt(
1386-
neural_networks=neural_networks,
1387-
hybridization_tables=hybridization_tables,
1388-
array_data_files=array_data_files,
1389-
)
1390-
)
1352+
sciml = (
1353+
SciMLExt.from_config(config, base_path)
1354+
if config.extensions and config.extensions.get(C.EXT_ID_SCIML)
1355+
else None
1356+
)
1357+
extensions = ProblemExtensions(sciml=sciml)
13911358

13921359
return Problem(
13931360
config=config,
@@ -2647,15 +2614,9 @@ def to_yaml(self, filename: str | Path):
26472614
# The schema requires a valid id or no id field at all.
26482615
del data["id"]
26492616

2650-
for ext_id, d_ext in data[C.EXTENSIONS].items():
2617+
for ext_id in list(data[C.EXTENSIONS]):
26512618
if ext_id == C.EXT_ID_SCIML:
2652-
# convert Paths to strings
2653-
for key in ("array_files", "hybridization_files"):
2654-
d_ext[key] = list(map(str, d_ext[key]))
2655-
for nn in d_ext["neural_networks"]:
2656-
d_ext["neural_networks"][nn][C.MODEL_LOCATION] = str(
2657-
d_ext["neural_networks"][nn][C.MODEL_LOCATION]
2658-
)
2619+
data[C.EXTENSIONS][ext_id] = self.extensions[ext_id].to_yaml()
26592620

26602621
write_yaml(data, filename)
26612622

petab/v2/extensions/sciml.py

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -153,6 +153,19 @@ class SciMLConfig(BaseModel):
153153
validate_assignment=True,
154154
)
155155

156+
def to_yaml(self) -> dict:
157+
"""Return a YAML-serializable dict with Paths converted to strings."""
158+
from . import C
159+
160+
d = self.model_dump(by_alias=True)
161+
for key in ("array_files", "hybridization_files"):
162+
d[key] = list(map(str, d[key]))
163+
for nn in d["neural_networks"] or {}:
164+
d["neural_networks"][nn][C.MODEL_LOCATION] = str(
165+
d["neural_networks"][nn][C.MODEL_LOCATION]
166+
)
167+
return d
168+
156169

157170
class SciMLExt:
158171
"""SciML extension runtime state.
@@ -241,3 +254,52 @@ def add_array_data_from_hdf5(
241254
self.array_data_files.append(
242255
ArrayDataStandard.load_data(_generate_path(file_path, base_path))
243256
)
257+
258+
@staticmethod
259+
def from_config(
260+
config,
261+
base_path: str | Path | None = None,
262+
) -> SciMLExt:
263+
"""Construct a SciMLExt from a ProblemConfig.
264+
265+
Arguments:
266+
config: A ProblemConfig whose ``extensions[EXT_ID_SCIML]`` entry
267+
is a :class:`SciMLConfig`.
268+
base_path: Base path used to resolve relative file paths.
269+
"""
270+
from petab_sciml import ArrayDataStandard, NNModel, NNModelStandard
271+
272+
sciml_config: SciMLConfig = config.extensions[C.EXT_ID_SCIML]
273+
274+
# Neural network classes are constructed via pytorch for now to get
275+
# the proper inputs
276+
neural_networks = [
277+
NNModel.from_pytorch_module(
278+
NNModelStandard.load_data(
279+
_generate_path(
280+
file_path=nn_config.location,
281+
base_path=base_path,
282+
)
283+
).to_pytorch_module(),
284+
nn_model_id=nn_id,
285+
)
286+
for nn_id, nn_config in (
287+
sciml_config.neural_networks or {}
288+
).items()
289+
]
290+
291+
hybridization_tables = [
292+
HybridizationTable.from_tsv(f, base_path)
293+
for f in sciml_config.hybridization_files
294+
]
295+
296+
array_data_files = [
297+
ArrayDataStandard.load_data(_generate_path(f, base_path))
298+
for f in sciml_config.array_files
299+
]
300+
301+
return SciMLExt(
302+
neural_networks=neural_networks,
303+
hybridization_tables=hybridization_tables,
304+
array_data_files=array_data_files,
305+
)

tests/v2/test_sciml.py

Lines changed: 26 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -4,37 +4,36 @@
44
from petab.v2.core import *
55
from petab.v2.core import ModelFile
66
from petab.v2.extensions.sciml import NeuralNetConfig, SciMLConfig
7-
from petab.v2.lint import sciml_validation_tasks
87
from petab.v2.models.sbml_model import SbmlModel
98

109

1110
def _get_test_problem():
12-
problem = Problem()
13-
problem.validation_tasks = sciml_validation_tasks
14-
problem.config = ProblemConfig(
15-
format_version="2.0.0",
16-
model_files=ConfigDict(
17-
{"lv": ModelFile(location="lv.xml", language="sbml")}
18-
),
19-
parameter_files=["parameters.tsv"],
20-
measurement_files=["measurements.tsv"],
21-
observable_files=["observables.tsv"],
22-
experiment_files=["experiments.tsv"],
23-
mapping_files=["mappings.tsv"],
24-
extensions={
25-
"sciml": SciMLConfig(
26-
version="0.1.0",
27-
array_files=["net1_ps.hdf5"],
28-
hybridization_files=["hybridizations.tsv"],
29-
neural_networks={
30-
"net1": NeuralNetConfig(
31-
location="net1.yaml",
32-
pre_initialization=False,
33-
format="YAML",
34-
)
35-
},
36-
)
37-
},
11+
problem = Problem(
12+
config=ProblemConfig(
13+
format_version="2.0.0",
14+
model_files=ConfigDict(
15+
{"lv": ModelFile(location="lv.xml", language="sbml")}
16+
),
17+
parameter_files=["parameters.tsv"],
18+
measurement_files=["measurements.tsv"],
19+
observable_files=["observables.tsv"],
20+
experiment_files=["experiments.tsv"],
21+
mapping_files=["mappings.tsv"],
22+
extensions={
23+
"sciml": SciMLConfig(
24+
version="0.1.0",
25+
array_files=["net1_ps.hdf5"],
26+
hybridization_files=["hybridizations.tsv"],
27+
neural_networks={
28+
"net1": NeuralNetConfig(
29+
location="net1.yaml",
30+
pre_initialization=False,
31+
format="YAML",
32+
)
33+
},
34+
)
35+
},
36+
)
3837
)
3938
problem.model = SbmlModel.from_antimony("""
4039
model lv

0 commit comments

Comments
 (0)