Skip to content
Draft
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
2 changes: 1 addition & 1 deletion doc/source/doc/parametric_models.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@
"source": [
"import stormpy.pars\n",
"\n",
"instantiator = stormpy.pars.PDtmcInstantiator(model)"
"instantiator = stormpy.pars.ModelInstantiator[stormpy.ModelType.DTMC, float](model)"
]
},
{
Expand Down
37 changes: 37 additions & 0 deletions examples/dfts/03-generic-types.py
Original file line number Diff line number Diff line change
@@ -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()
2 changes: 1 addition & 1 deletion examples/parametric_models/01-parametric-models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
142 changes: 142 additions & 0 deletions lib/stormpy/_template.py
Original file line number Diff line number Diff line change
@@ -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"<template class {self.__name__}>"
49 changes: 9 additions & 40 deletions lib/stormpy/dft/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,58 +6,27 @@
from . import _dft
from ._dft import *
from .modules import modules_json
from stormpy._template import TemplateClass

_dft._set_up()

DFT = TemplateClass("DFT", _dft)

def analyze_dft(ft, properties, symred=True, allow_modularisation=False, relevant_events=RelevantEvents(), allow_dc_for_relevant=False):
if isinstance(ft, DFT_double):
return _dft._analyze_dft_double(ft, properties, symred, allow_modularisation, relevant_events, allow_dc_for_relevant)
else:
assert isinstance(ft, DFT_ratfunc)
return _dft._analyze_dft_ratfunc(ft, properties, symred, allow_modularisation, relevant_events, allow_dc_for_relevant)
DFTElement = TemplateClass("DFTElement", _dft)

DFTBE = TemplateClass("DFTBE", _dft)

def build_model(ft, symmetries=DftSymmetries(), relevant_events=RelevantEvents(), allow_dc_for_relevant=False):
if isinstance(ft, DFT_double):
return _dft._build_model_double(ft, symmetries, relevant_events, allow_dc_for_relevant)
else:
assert isinstance(ft, DFT_ratfunc)
return _dft._build_model_ratfunc(ft, symmetries, relevant_events, allow_dc_for_relevant)
DFTDependency = TemplateClass("DFTDependency", _dft)

DFTState = TemplateClass("DFTState", _dft)

def transform_dft(ft, unique_constant_be, binary_fdeps, exponential_distributions):
if isinstance(ft, DFT_double):
return _dft._transform_dft_double(ft, unique_constant_be, binary_fdeps, exponential_distributions)
else:
assert isinstance(ft, DFT_ratfunc)
return _dft._transform_dft_ratfunc(ft, unique_constant_be, binary_fdeps, exponential_distributions)
DFTSimulator = TemplateClass("DFTSimulator", _dft, deduce_from=DFT)

ExplicitDFTModelBuilder = TemplateClass("ExplicitDFTModelBuilder", _dft, deduce_from=DFT)

def compute_dependency_conflicts(ft, use_smt=False, solver_timeout=0):
if isinstance(ft, DFT_double):
return _dft._compute_dependency_conflicts_double(ft, use_smt, solver_timeout)
else:
assert isinstance(ft, DFT_ratfunc)
return _dft._compute_dependency_conflicts_ratfunc(ft, use_smt, solver_timeout)
DFTInstantiator = TemplateClass("DFTInstantiator", _dft)


def prepare_for_analysis(ft):
compute_dependency_conflicts(ft, use_smt=False)
return transform_dft(ft, unique_constant_be=True, binary_fdeps=True, exponential_distributions=True)


def is_well_formed(ft, check_valid_for_analysis=True):
if isinstance(ft, DFT_double):
return _dft._is_well_formed_double(ft, check_valid_for_analysis)
else:
assert isinstance(ft, DFT_ratfunc)
return _dft._is_well_formed_ratfunc(ft, check_valid_for_analysis)


def has_potential_modeling_issues(ft):
if isinstance(ft, DFT_double):
return _dft._has_potential_modeling_issues_double(ft)
else:
assert isinstance(ft, DFT_ratfunc)
return _dft._has_potential_modeling_issues_ratfunc(ft)
6 changes: 3 additions & 3 deletions lib/stormpy/dft/simulator.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ def __init__(self, dft, seed=42, relevant=[]):
# Initialize random generator
generator = stormpy.dft.RandomGenerator.create(seed)
# Create simulator
self._simulator = stormpy.dft.DFTSimulator_double(self._dft, info, generator)
self._simulator = stormpy.dft.DFTSimulator(self._dft, info, generator)
# Initialize variables
self._state = None
self._fail_candidates = dict()
Expand Down Expand Up @@ -104,11 +104,11 @@ def _update(self):
for f in self._state.failable():
if f.is_due_dependency():
self._is_failable_dependency = True
fail_dependency = f.as_dependency_double(self._dft)
fail_dependency = f.as_dependency(self._dft)
fail_be = fail_dependency.dependent_events[0]
self._fail_candidates[fail_be.name] = f
else:
fail_be = f.as_be_double(self._dft)
fail_be = f.as_be(self._dft)
self._fail_candidates[fail_be.name] = f

def let_fail(self, be, dependency_successful=True):
Expand Down
30 changes: 3 additions & 27 deletions lib/stormpy/pars/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,38 +7,14 @@
from ._pars import *

from stormpy import ModelType
from stormpy._template import TemplateClass

_pars._set_up()


class ModelInstantiator:
"""
Class for instantiating models.
"""
ModelInstantiator = TemplateClass("ModelInstantiator", _pars)

def __init__(self, model):
"""
Constructor.
:param model: Model.
"""
if model.model_type == ModelType.MDP:
self._instantiator = PMdpInstantiator(model)
elif model.model_type == ModelType.DTMC:
self._instantiator = PDtmcInstantiator(model)
elif model.model_type == ModelType.CTMC:
self._instantiator = PCtmcInstantiator(model)
elif model.model_type == ModelType.MA:
self._instantiator = PMaInstantiator(model)
else:
raise stormpy.exceptions.StormError("Model type {} not supported".format(model.model_type))

def instantiate(self, valuation):
"""
Instantiate model with given valuation.
:param valuation: Valuation from parameter to value.
:return: Instantiated model.
"""
return self._instantiator.instantiate(valuation)
ModelInstantiationChecker = TemplateClass("ModelInstantiationChecker", _pars)


def simplify_model(model, formula):
Expand Down
Loading
Loading