black formatting

This commit is contained in:
Artur Meski
2023-07-11 19:48:37 +01:00
parent 8be615b293
commit a1fb5836a5
24 changed files with 965 additions and 727 deletions

View File

@@ -24,13 +24,11 @@ class rsLTL_Encoder(object):
self.init_ncalls()
def load_variables(self, var_rs, var_ctx, var_loop_pos):
self.v = var_rs
self.v_ctx = var_ctx
self.loop_position = var_loop_pos
def get_encoding(self, formula, bound):
assert self.v is not None
assert self.v_ctx is not None
assert self.loop_position is not None
@@ -49,8 +47,8 @@ class rsLTL_Encoder(object):
Cache for formulae encodings
"""
self.cache_hits = 0
self.enc_fcache = [{} for level in range(0, bound+1)]
self.enc_fcache_approx = [{} for level in range(0, bound+1)]
self.enc_fcache = [{} for level in range(0, bound + 1)]
self.enc_fcache_approx = [{} for level in range(0, bound + 1)]
def cache_save(self, formula, level, formula_encoding):
self.enc_fcache[level][formula] = formula_encoding
@@ -83,7 +81,6 @@ class rsLTL_Encoder(object):
return (self.ncalls_encode, self.ncalls_encode_approx)
def encode_bag(self, bag_formula, level, context=False):
if not bag_formula:
raise RuntimeError("bag_formula is None")
@@ -100,41 +97,42 @@ class rsLTL_Encoder(object):
if bag_formula.f_type == BagDesc_oper.l_and:
return And(
self.encode_bag(bag_formula.left_operand, level, context),
self.encode_bag(bag_formula.right_operand, level, context))
self.encode_bag(bag_formula.right_operand, level, context),
)
if bag_formula.f_type == BagDesc_oper.l_or:
return Or(
self.encode_bag(bag_formula.left_operand, level, context),
self.encode_bag(bag_formula.right_operand, level, context))
self.encode_bag(bag_formula.right_operand, level, context),
)
if bag_formula.f_type == BagDesc_oper.l_not:
return Not(self.encode_bag(
bag_formula.left_operand, level, context))
return Not(self.encode_bag(bag_formula.left_operand, level, context))
if bag_formula.f_type == BagDesc_oper.lt:
return self.encode_bag(
bag_formula.left_operand, level, context) < int(
bag_formula.right_operand)
return self.encode_bag(bag_formula.left_operand, level, context) < int(
bag_formula.right_operand
)
if bag_formula.f_type == BagDesc_oper.le:
return self.encode_bag(
bag_formula.left_operand, level, context) <= int(
bag_formula.right_operand)
return self.encode_bag(bag_formula.left_operand, level, context) <= int(
bag_formula.right_operand
)
if bag_formula.f_type == BagDesc_oper.eq:
return self.encode_bag(
bag_formula.left_operand, level, context) == int(
bag_formula.right_operand)
return self.encode_bag(bag_formula.left_operand, level, context) == int(
bag_formula.right_operand
)
if bag_formula.f_type == BagDesc_oper.ge:
return self.encode_bag(
bag_formula.left_operand, level, context) >= int(
bag_formula.right_operand)
return self.encode_bag(bag_formula.left_operand, level, context) >= int(
bag_formula.right_operand
)
if bag_formula.f_type == BagDesc_oper.gt:
return self.encode_bag(
bag_formula.left_operand, level, context) > int(
bag_formula.right_operand)
return self.encode_bag(bag_formula.left_operand, level, context) > int(
bag_formula.right_operand
)
assert False, "Unsupported case {:s}".format(bag_formula.f_type)
@@ -149,7 +147,6 @@ class rsLTL_Encoder(object):
return res
def encode(self, formula, level, bound):
self.ncalls_encode += 1
from_cache = self.cache_query(formula, level)
@@ -159,12 +156,12 @@ class rsLTL_Encoder(object):
enc = None
if not isinstance(formula, Formula_rsLTL):
raise NotImplementedError(
"Unsupported formula type: " + str(type(formula)))
raise NotImplementedError("Unsupported formula type: " + str(type(formula)))
if level > bound:
raise RuntimeError(
"level > bound. Unexpected behaviour. The encoding does not support levels higher than a bound.")
"level > bound. Unexpected behaviour. The encoding does not support levels higher than a bound."
)
if formula.f_type == rsLTL_form_type.bag:
enc = self.encode_bag_state(formula.bag_descr, level)
@@ -179,33 +176,40 @@ class rsLTL_Encoder(object):
elif formula.f_type == rsLTL_form_type.l_and:
enc = And(
self.encode(formula.left_operand, level, bound),
self.encode(formula.right_operand, level, bound)
self.encode(formula.right_operand, level, bound),
)
elif formula.f_type == rsLTL_form_type.l_or:
enc = Or(
self.encode(formula.left_operand, level, bound),
self.encode(formula.right_operand, level, bound)
self.encode(formula.right_operand, level, bound),
)
elif formula.f_type == rsLTL_form_type.l_implies:
enc = Implies(
self.encode(formula.left_operand, level, bound),
self.encode(formula.right_operand, level, bound)
self.encode(formula.right_operand, level, bound),
)
elif formula.f_type == rsLTL_form_type.t_next:
if level < bound:
enc = And(
self.encode(formula.left_operand, level + 1, bound),
self.encode_bag_ctx(formula.sub_operand, level)
self.encode_bag_ctx(formula.sub_operand, level),
)
else:
# level == bound
enc = False
for loop_level in range(1, bound+1):
enc = simplify(Or(enc, And(self.loop_position == loop_level, self.encode(
formula.left_operand, loop_level, bound))))
for loop_level in range(1, bound + 1):
enc = simplify(
Or(
enc,
And(
self.loop_position == loop_level,
self.encode(formula.left_operand, loop_level, bound),
),
)
)
enc = And(enc, self.encode_bag_ctx(formula.sub_operand, level))
enc = simplify(enc)
@@ -214,22 +218,25 @@ class rsLTL_Encoder(object):
enc = And(
self.encode(formula.left_operand, level, bound),
self.encode_bag_ctx(formula.sub_operand, level),
self.encode(formula, level + 1, bound)
self.encode(formula, level + 1, bound),
)
else:
# level == bound
enc_loops = False
for loop_level in range(1, bound+1):
enc_loops = simplify(Or(enc_loops,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound),
)
))
for loop_level in range(1, bound + 1):
enc_loops = simplify(
Or(
enc_loops,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound),
),
)
)
enc = And(
self.encode(formula.left_operand, bound, bound),
enc_loops,
self.encode_bag_ctx(formula.sub_operand, level)
self.encode_bag_ctx(formula.sub_operand, level),
)
enc = simplify(enc)
@@ -239,25 +246,26 @@ class rsLTL_Encoder(object):
self.encode(formula.left_operand, level, bound),
And(
self.encode_bag_ctx(formula.sub_operand, level),
self.encode(formula, level + 1, bound)
)
self.encode(formula, level + 1, bound),
),
)
else:
# level == bound
enc_loops = False
for loop_level in range(1, bound+1):
enc_loops = simplify(Or(enc_loops,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound),
)
))
for loop_level in range(1, bound + 1):
enc_loops = simplify(
Or(
enc_loops,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound),
),
)
)
# print(enc)
enc = Or(self.encode(formula.left_operand, bound, bound),
And(
enc_loops,
self.encode_bag_ctx(formula.sub_operand, level)
)
enc = Or(
self.encode(formula.left_operand, bound, bound),
And(enc_loops, self.encode_bag_ctx(formula.sub_operand, level)),
)
enc = simplify(enc)
@@ -268,21 +276,24 @@ class rsLTL_Encoder(object):
else:
# level == bound
inner_enc = False
for loop_level in range(1, bound+1):
inner_enc = simplify(Or(inner_enc,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound)
)
))
for loop_level in range(1, bound + 1):
inner_enc = simplify(
Or(
inner_enc,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound),
),
)
)
enc = Or(
self.encode(formula.right_operand, level, bound),
And(
self.encode(formula.left_operand, level, bound),
inner_enc,
self.encode_bag_ctx(formula.sub_operand, level)
)
self.encode_bag_ctx(formula.sub_operand, level),
),
)
elif formula.f_type == rsLTL_form_type.t_release:
@@ -291,23 +302,23 @@ class rsLTL_Encoder(object):
else:
# level == bound
inner_enc = False
for loop_level in range(1, bound+1):
inner_enc = simplify(Or(inner_enc,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound)
)
))
for loop_level in range(1, bound + 1):
inner_enc = simplify(
Or(
inner_enc,
And(
self.loop_position == loop_level,
self.encode_approx(formula, loop_level, bound),
),
)
)
enc = And(
self.encode(formula.right_operand, level, bound),
Or(
self.encode(formula.left_operand, level, bound),
And(
inner_enc,
self.encode_bag_ctx(formula.sub_operand, level)
)
)
And(inner_enc, self.encode_bag_ctx(formula.sub_operand, level)),
),
)
else:
@@ -341,8 +352,8 @@ class rsLTL_Encoder(object):
And(
self.encode(formula.left_operand, level, bound),
self.encode_approx(formula, level + 1, bound),
self.encode_bag_ctx(formula.sub_operand, level)
)
self.encode_bag_ctx(formula.sub_operand, level),
),
)
else:
# level == bound
@@ -356,9 +367,9 @@ class rsLTL_Encoder(object):
self.encode(formula.left_operand, level, bound),
And(
self.encode_approx(formula, level + 1, bound),
self.encode_bag_ctx(formula.sub_operand, level)
)
)
self.encode_bag_ctx(formula.sub_operand, level),
),
),
)
else:
# level == bound
@@ -369,7 +380,7 @@ class rsLTL_Encoder(object):
enc = And(
self.encode(formula.left_operand, level, bound),
self.encode_bag_ctx(formula.sub_operand, level),
self.encode_approx(formula, level + 1, bound)
self.encode_approx(formula, level + 1, bound),
)
else:
# level == bound
@@ -381,16 +392,15 @@ class rsLTL_Encoder(object):
self.encode(formula.left_operand, level, bound),
And(
self.encode_bag_ctx(formula.sub_operand, level),
self.encode_approx(formula, level + 1, bound)
)
self.encode_approx(formula, level + 1, bound),
),
)
else:
# level == bound
enc = self.encode(formula.left_operand, bound, bound)
else:
raise NotImplementedError(
"Unsupported operator in approximation encoding")
raise NotImplementedError("Unsupported operator in approximation encoding")
if enc is None:
raise RuntimeError("Encoding is NONE. Should never happen")