2121from devito .finite_differences .tools import coeff_priority , make_shift_x0
2222from devito .logger import warning
2323from devito .tools import (
24- as_tuple , extract_dtype , filter_ordered , flatten , frozendict , infer_dtype , is_integer ,
25- is_number , memoized_func , split
24+ Tag , as_tuple , extract_dtype , filter_ordered , flatten , frozendict , infer_dtype ,
25+ is_integer , is_number , memoized_func , split
2626)
2727from devito .types import Array , DimensionTuple , Evaluable , StencilDimension
2828from devito .types .basic import AbstractFunction , Indexed
3434 'EvalDerivative' ,
3535 'Imag' ,
3636 'IndexDerivative' ,
37+ 'IndexDerivativeProperty' ,
3738 'Real' ,
3839 'Weights' ,
3940]
@@ -1045,14 +1046,23 @@ def value(self, idx):
10451046 return self [idx ]
10461047
10471048
1049+ class IndexDerivativeProperty (Tag ):
1050+
1051+ """A property controlling how an `IndexDerivative` is lowered."""
1052+
1053+
10481054class IndexDerivative (IndexSum ):
10491055
10501056 __rargs__ = ('expr' , 'mapper' )
1051- __rkwargs__ = IndexSum .__rkwargs__ + ('deriv_order' ,)
1057+ __rkwargs__ = IndexSum .__rkwargs__ + ('deriv_order' , 'properties' )
10521058
1053- def __new__ (cls , expr , mapper , deriv_order = None , ** kwargs ):
1059+ def __new__ (cls , expr , mapper , deriv_order = None , properties = (), ** kwargs ):
10541060 dimensions = as_tuple (set (mapper .values ()))
10551061
1062+ properties = frozenset (as_tuple (properties ))
1063+ if not all (isinstance (i , IndexDerivativeProperty ) for i in properties ):
1064+ raise ValueError ("Expected IndexDerivative properties" )
1065+
10561066 # Detect the Weights among the arguments
10571067 weightss = []
10581068 for a in expr .args :
@@ -1073,20 +1083,25 @@ def __new__(cls, expr, mapper, deriv_order=None, **kwargs):
10731083 obj ._mapper = frozendict (mapper )
10741084
10751085 obj ._deriv_order = deriv_order
1086+ obj ._properties = properties
10761087
10771088 return obj
10781089
10791090 def _hashable_content (self ):
1080- return super ()._hashable_content () + (self .mapper ,)
1091+ properties = tuple (sorted (map (str , self .properties )))
1092+ return super ()._hashable_content () + (self .mapper , properties )
10811093
10821094 def compare (self , other ):
10831095 if self is other :
10841096 return 0
10851097 n1 = self .__class__
10861098 n2 = other .__class__
10871099 if n1 .__name__ == n2 .__name__ :
1100+ p1 = tuple (sorted (map (str , self .properties )))
1101+ p2 = tuple (sorted (map (str , other .properties )))
10881102 return (self .weights .compare (other .weights ) or
1089- self .base .compare (other .base ))
1103+ self .base .compare (other .base ) or
1104+ (p1 > p2 ) - (p1 < p2 ))
10901105 else :
10911106 return super ().compare (other )
10921107
@@ -1110,6 +1125,10 @@ def mapper(self):
11101125 def deriv_order (self ):
11111126 return self ._deriv_order
11121127
1128+ @property
1129+ def properties (self ):
1130+ return self ._properties
1131+
11131132 @property
11141133 def depth (self ):
11151134 iderivs = self .expr .find (IndexDerivative )
@@ -1289,7 +1308,8 @@ def _diff2sympy(obj):
12891308 # Handle special objects
12901309 if isinstance (obj , DiffDerivative ):
12911310 return IndexDerivative (* args , obj .mapper ,
1292- deriv_order = obj .deriv_order ), True
1311+ deriv_order = obj .deriv_order ,
1312+ properties = obj .properties ), True
12931313
12941314 # Handle generic objects such as arithmetic operations
12951315 try :
0 commit comments