From 6c8815d29fa4a63681052b7e7a6a558b4b287ca5 Mon Sep 17 00:00:00 2001 From: Linus Heck Date: Tue, 25 Aug 2026 14:50:28 +0200 Subject: [PATCH] Migrate dft and pars to experimental new generic handling --- doc/source/doc/parametric_models.ipynb | 2 +- examples/dfts/03-generic-types.py | 37 ++++ .../parametric_models/01-parametric-models.py | 2 +- lib/stormpy/_template.py | 142 +++++++++++++++ lib/stormpy/dft/__init__.py | 49 +----- lib/stormpy/dft/simulator.py | 6 +- lib/stormpy/pars/__init__.py | 30 +--- src/binding_type_index.h | 81 +++++++++ src/dft/analysis.cpp | 34 ++-- src/dft/analysis.h | 2 +- src/dft/dft.cpp | 11 +- src/dft/dft.h | 2 +- src/dft/dft_elements.cpp | 19 +- src/dft/dft_elements.h | 2 +- src/dft/dft_state.cpp | 20 ++- src/dft/dft_state.h | 2 +- src/dft/io.cpp | 6 +- src/dft/simulator.cpp | 11 +- src/dft/simulator.h | 2 +- src/dft/transformations.cpp | 6 +- src/mod_dft.cpp | 20 +-- src/pars/model_instantiator.cpp | 98 ++++------- src/template_binding.h | 165 ++++++++++++++++++ tests/dft/test_analysis.py | 6 +- tests/dft/test_dft.py | 17 +- tests/dft/test_dft_simulator.py | 18 +- tests/dft/test_transformations.py | 5 +- tests/pars/test_model_instantiator.py | 9 +- 28 files changed, 586 insertions(+), 218 deletions(-) create mode 100644 examples/dfts/03-generic-types.py create mode 100644 lib/stormpy/_template.py create mode 100644 src/binding_type_index.h create mode 100644 src/template_binding.h diff --git a/doc/source/doc/parametric_models.ipynb b/doc/source/doc/parametric_models.ipynb index 334fcd6963..0cde72858d 100644 --- a/doc/source/doc/parametric_models.ipynb +++ b/doc/source/doc/parametric_models.ipynb @@ -58,7 +58,7 @@ "source": [ "import stormpy.pars\n", "\n", - "instantiator = stormpy.pars.PDtmcInstantiator(model)" + "instantiator = stormpy.pars.ModelInstantiator[stormpy.ModelType.DTMC, float](model)" ] }, { diff --git a/examples/dfts/03-generic-types.py b/examples/dfts/03-generic-types.py new file mode 100644 index 0000000000..7af8d4b116 --- /dev/null +++ b/examples/dfts/03-generic-types.py @@ -0,0 +1,37 @@ +"""Demonstrate explicit and inferred stormpy template specializations.""" + +import stormpy +import stormpy.dft +import stormpy.examples.files + + +def demonstrate_generic_dft_types(): + # Readable aliases: + assert stormpy.dft._dft._DFT_Double is stormpy.dft.DFT[float] + assert stormpy.dft._dft._DFT_RationalFunction is stormpy.dft.DFT[stormpy.RationalFunction] + + path = stormpy.examples.files.dft_json_and + double_dft = stormpy.dft.load_dft_json_file(path) + assert type(double_dft) is stormpy.dft.DFT[float] + + rational_function_dft = stormpy.dft.load_parametric_dft_json_file(path) + assert type(rational_function_dft) is stormpy.dft.DFT[stormpy.RationalFunction] + + # Construct a specialization explicitly: + explicit_builder = stormpy.dft.ExplicitDFTModelBuilder[float](double_dft) + assert type(explicit_builder) is stormpy.dft.ExplicitDFTModelBuilder[float] + + # Or deduce the specialization from the constructor argument: + inferred_builder = stormpy.dft.ExplicitDFTModelBuilder(rational_function_dft) + assert type(inferred_builder) is stormpy.dft.ExplicitDFTModelBuilder[stormpy.RationalFunction] + + # Overloaded functions call the specialization: + concrete_model = stormpy.dft.build_model(double_dft) + assert not concrete_model.supports_parameters + + parametric_model = stormpy.dft.build_model(rational_function_dft) + assert parametric_model.supports_parameters + + +if __name__ == "__main__": + demonstrate_generic_dft_types() diff --git a/examples/parametric_models/01-parametric-models.py b/examples/parametric_models/01-parametric-models.py index 8218c5fd3b..aba4e07ded 100644 --- a/examples/parametric_models/01-parametric-models.py +++ b/examples/parametric_models/01-parametric-models.py @@ -16,7 +16,7 @@ def example_parametric_models_01(): parameters = model.collect_all_parameters() assert len(parameters) == 2 - instantiator = stormpy.pars.PDtmcInstantiator(model) + instantiator = stormpy.pars.ModelInstantiator[stormpy.ModelType.DTMC, float](model) point = dict() for x in parameters: print(x.name) diff --git a/lib/stormpy/_template.py b/lib/stormpy/_template.py new file mode 100644 index 0000000000..13c2497eb1 --- /dev/null +++ b/lib/stormpy/_template.py @@ -0,0 +1,142 @@ +"""Runtime representation of C++ template families. + +C++ template specializations still need to be compiled and bound separately. +``TemplateClass`` groups those concrete Python classes behind one public, +subscriptable object. +""" + +from __future__ import annotations + +from collections.abc import Iterator, Mapping +from types import MappingProxyType +from typing import Any + + +class TemplateClass: + """Native C++ template family exposed as a subscriptable Python object. + + ``family[parameters]`` returns a registered concrete Python class. + ``family(*args, **kwargs)`` constructs a class selected by a configured + deduction rule. + """ + + def __init__( + self, + name: str, + module: object, + *, + deduce_from: TemplateClass | None = None, + ) -> None: + """Load a native template family. + + :param name: Family name used by the native registration table. + :param module: Native module containing the registered specializations. + :param deduce_from: Optional constructor deduction guide. For an + unsubscripted call, copy the complete parameter tuple of the first + argument from this template family. For example, + ``deduce_from=DFT`` makes ``Builder(dft)`` select + ``Builder[float]`` when ``dft`` is a ``DFT[float]``. Source and + target families must have equivalent parameter tuples; no defaults + or partial deduction are performed. + :raises RuntimeError: If the module has no registrations for ``name``, + has an empty registration table, or mixes parameter arities. + """ + try: + implementations = module._template_instantiations[name] + except (AttributeError, KeyError): + raise RuntimeError(f"Native module has no registrations for {name}") from None + if not implementations: + raise RuntimeError(f"Native module has no registrations for {name}") + + arities = {len(parameters) for parameters in implementations} + if len(arities) != 1: + raise RuntimeError(f"Native module has inconsistent registrations for {name}") + + self.__name__ = name + self._arity = arities.pop() + self._deduction_source = deduce_from + self._instantiations: dict[tuple[object, ...], type] = {} + for parameters, implementation in implementations.items(): + self.register(parameters, implementation) + + def _normalize(self, parameters: object) -> tuple[object, ...]: + """Convert subscription parameters to a validated tuple.""" + key = parameters if isinstance(parameters, tuple) else (parameters,) + if len(key) != self._arity: + raise TypeError(f"{self.__name__} expects {self._arity} template " f"parameter{'s' if self._arity != 1 else ''}, got {len(key)}") + return key + + def register(self, parameters: object, implementation: type) -> None: + """Add a concrete specialization to this family. + + :param parameters: One parameter or a tuple containing the complete + template-parameter list. + :param implementation: Concrete Python class for those parameters. + :raises TypeError: If the parameter count differs from the family's + native arity. + :raises ValueError: If the parameter tuple is already registered. + """ + key = self._normalize(parameters) + if key in self._instantiations: + raise ValueError(f"{self.__name__}{key!r} is already registered") + self._instantiations[key] = implementation + + def __getitem__(self, parameters: object) -> type: + """Return the concrete class registered for ``parameters``. + + A single parameter may be written directly as ``family[T]``; multiple + parameters use normal subscription tuple syntax, ``family[T, U]``. + + :raises TypeError: If the parameter count is wrong or no matching + specialization is registered. + """ + key = self._normalize(parameters) + try: + return self._instantiations[key] + except KeyError: + raise TypeError(f"{self.__name__} has no instantiation for {key!r}") from None + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + """Construct an instance of the specialization selected from the first argument. + + The first positional argument is matched against ``deduce_from``. If no + deduction source is configured, it is matched against this family, + enabling copy-like construction such as ``DFT(existing_dft)``. + + :raises TypeError: If no arguments permit deduction, the first argument + is not a uniquely registered instance, or the resulting + specialization is unavailable. + """ + if args: + parameters = (self._deduction_source or self).parameters_of(args[0]) + else: + raise TypeError(f"{self.__name__} requires explicit template parameters") + return self[parameters](*args, **kwargs) + + def parameters_of(self, instance: object) -> tuple[object, ...]: + """Return the complete parameter tuple of a registered instance. + + :raises TypeError: If the instance matches zero or multiple registered + specializations. + """ + matches = [parameters for parameters, implementation in self._instantiations.items() if isinstance(instance, implementation)] + if len(matches) != 1: + raise TypeError(f"Cannot infer {self.__name__} template parameters from {type(instance)!r}") + return matches[0] + + @property + def instantiations(self) -> Mapping[tuple[object, ...], type]: + """Map registered parameter tuples to classes without allowing mutation.""" + return MappingProxyType(self._instantiations) + + def is_instantiation(self, instance: object) -> bool: + """Return whether ``instance`` belongs to any registered specialization.""" + return isinstance(instance, tuple(self._instantiations.values())) + + def __iter__(self) -> Iterator[tuple[object, ...]]: + """Iterate over registered parameter tuples in registration order.""" + return iter(self._instantiations) + + def __repr__(self) -> str: + """Return a concise template-family representation.""" + return f"