-
Notifications
You must be signed in to change notification settings - Fork 12
Expand file tree
/
Copy pathgeneral.py
More file actions
217 lines (165 loc) · 8.25 KB
/
Copy pathgeneral.py
File metadata and controls
217 lines (165 loc) · 8.25 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
"""
docstring needed
:copyright: Copyright 2010-2017 by the NineML Python team, see AUTHORS.
:license: BSD-3, see LICENSE for details.
"""
from __future__ import division
from future.utils import itervalues
from collections import defaultdict
from nineml.exceptions import NineMLUsageError
from nineml.utils import assert_no_duplicates
from ....componentclass.visitors.validators import (
AliasesAreNotRecursiveComponentValidator,
NoUnresolvedSymbolsComponentValidator,
CheckNoLHSAssignmentsToMathsNamespaceComponentValidator,
DimensionalityComponentValidator)
from ..base import BaseDynamicsVisitor
import nineml.units as un
import logging
logger = logging.getLogger('NineML')
class TimeDerivativesAreDeclaredDynamicsValidator(BaseDynamicsVisitor):
""" Check all variables used in TimeDerivative blocks are defined
as StateVariables.
"""
def __init__(self, component_class, **kwargs): # @UnusedVariable
BaseDynamicsVisitor.__init__(self)
self.sv_declared = []
self.time_derivatives_used = []
self.visit(component_class)
for td in self.time_derivatives_used:
if td not in self.sv_declared:
raise NineMLUsageError(
"StateVariable '{}' not declared".format(td))
def action_statevariable(self, state_variable, **kwargs): # @UnusedVariable @IgnorePep8
self.sv_declared.append(state_variable.name)
def action_timederivative(self, timederivative, **kwargs): # @UnusedVariable @IgnorePep8
self.time_derivatives_used.append(
timederivative.variable)
def default_action(self, obj, nineml_cls, **kwargs):
pass
class StateAssignmentsAreOnStateVariablesDynamicsValidator(
BaseDynamicsVisitor):
""" Check that we only attempt to make StateAssignments to state-variables.
"""
def __init__(self, component_class, **kwargs): # @UnusedVariable
BaseDynamicsVisitor.__init__(self)
self.sv_declared = []
self.state_assignments_lhs = []
self.visit(component_class)
for sa in self.state_assignments_lhs:
if sa not in self.sv_declared:
raise NineMLUsageError(
"Not Assigning to state-variable: {}".format(sa))
def action_statevariable(self, state_variable, **kwargs): # @UnusedVariable @IgnorePep8
self.sv_declared.append(state_variable.name)
def action_stateassignment(self, state_assignment, **kwargs): # @UnusedVariable @IgnorePep8
self.state_assignments_lhs.append(state_assignment.lhs)
def default_action(self, obj, nineml_cls, **kwargs):
pass
class AliasesAreNotRecursiveDynamicsValidator(
AliasesAreNotRecursiveComponentValidator,
BaseDynamicsVisitor):
"""Check that aliases are not self-referential"""
def action_dynamics(self, dynamics, **kwargs):
return self.action_componentclass(dynamics, **kwargs)
class NoUnresolvedSymbolsDynamicsValidator(
NoUnresolvedSymbolsComponentValidator,
BaseDynamicsVisitor):
"""
Check that aliases and timederivatives are defined in terms of other
parameters, aliases, statevariables and ports
"""
def action_analogreceiveport(self, port, **kwargs): # @UnusedVariable @IgnorePep8
self.available_symbols.append(port.name)
def action_analogreduceport(self, port, **kwargs): # @UnusedVariable @IgnorePep8
self.available_symbols.append(port.name)
def action_statevariable(self, state_variable, **kwargs): # @UnusedVariable @IgnorePep8
self.add_symbol(symbol=state_variable.name)
def action_timederivative(self, time_derivative, **kwargs): # @UnusedVariable @IgnorePep8
self.time_derivatives.append(time_derivative)
def action_stateassignment(self, state_assignment, **kwargs): # @UnusedVariable @IgnorePep8
self.state_assignments.append(state_assignment)
class RegimeGraphDynamicsValidator(BaseDynamicsVisitor):
def __init__(self, component_class, **kwargs): # @UnusedVariable
BaseDynamicsVisitor.__init__(self)
self.connected_regimes_from_regime = defaultdict(set)
self.component_class = component_class
self.regimes = {}
self.visit(component_class)
self.connected = set()
if self.regimes:
first_regime = next(iter(itervalues(self.regimes)))
# Recursively add all regimes connected to the first regime
self._add_connected_regimes_recursive(first_regime)
if len(self.connected) < len(self.regimes):
logger.warning(
"Transition graph of {} contains islands: {} "
"regimes ('{}') and {} connected ('{}'):\n\n{}"
.format(
component_class,
len(self.regimes),
"', '".join(r.name for r in self.regimes),
len(self.connected),
"', '".join(r.name for r in self.connected),
self.connected_regimes_from_regime))
elif len(self.connected) > len(self.regimes):
assert False
def action_regime(self, regime, **kwargs): # @UnusedVariable
self.regimes[regime.id] = regime
for transition in regime.transitions:
self.connected_regimes_from_regime[regime.id].add(
transition.target_regime.id)
self.connected_regimes_from_regime[
transition.target_regime.id].add(regime.id)
def _add_connected_regimes_recursive(self, regime): # @IgnorePep8
self.connected.add(regime.id)
for id_ in self.connected_regimes_from_regime[regime.id]:
if id_ not in self.connected:
self._add_connected_regimes_recursive(self.regimes[id_])
def default_action(self, obj, nineml_cls, **kwargs):
pass
class RegimeOnlyHasOneHandlerPerEventDynamicsValidator(BaseDynamicsVisitor):
def __init__(self, component_class, **kwargs): # @UnusedVariable
BaseDynamicsVisitor.__init__(self)
self.visit(component_class)
def action_regime(self, regime, **kwargs): # @UnusedVariable
event_triggers = [on_event.src_port_name
for on_event in regime.on_events]
assert_no_duplicates(event_triggers)
def default_action(self, obj, nineml_cls, **kwargs):
pass
class CheckNoLHSAssignmentsToMathsNamespaceDynamicsValidator(
CheckNoLHSAssignmentsToMathsNamespaceComponentValidator,
BaseDynamicsVisitor):
"""
This class checks that there is not a mathematical symbols, (e.g. pi, e)
on the left-hand-side of an equation
"""
def action_statevariable(self, state_variable, **kwargs): # @UnusedVariable @IgnorePep8
self.check_lhssymbol_is_valid(state_variable.name)
def action_stateassignment(self, assignment, **kwargs): # @UnusedVariable
self.check_lhssymbol_is_valid(assignment.lhs)
def action_timederivative(self, time_derivative, **kwargs): # @UnusedVariable @IgnorePep8
self.check_lhssymbol_is_valid(time_derivative.variable)
def default_action(self, obj, nineml_cls, **kwargs):
pass
class DimensionalityDynamicsValidator(DimensionalityComponentValidator,
BaseDynamicsVisitor):
def action_timederivative(self, timederivative, **kwargs): # @UnusedVariable @IgnorePep8
dimension = self._get_dimensions(timederivative)
sv = self.component_class.state_variable(timederivative.variable)
self._compare_dimensionality(
dimension, sv.dimension / un.time, timederivative,
'time derivative of ' + sv.name)
def action_stateassignment(self, stateassignment, **kwargs): # @UnusedVariable @IgnorePep8
dimension = self._get_dimensions(stateassignment)
sv = self.component_class.state_variable(stateassignment.variable)
self._compare_dimensionality(
dimension, sv.dimension, stateassignment, 'state variable ' +
sv.name)
def action_analogsendport(self, port, **kwargs): # @UnusedVariable
self._check_send_port(port)
def action_trigger(self, trigger, **kwargs): # @UnusedVariable
self._flatten_dims(trigger.rhs, trigger)
def default_action(self, obj, nineml_cls, **kwargs):
pass