Implementation of reachability testing for RSC using SMT. Intitially tested version. No known issues.

This commit is contained in:
Artur Meski
2016-02-17 21:41:43 +01:00
parent 8853589ff8
commit 9e025c2c9a
3 changed files with 135 additions and 80 deletions

View File

@@ -368,6 +368,14 @@ class ReactionSystemWithConcentrations(ReactionSystem):
return False
def get_state_ids(self, state):
"""Returns entities of the given state without levels"""
return [e for e,c in state]
def has_non_zero_concentration(self, elem):
if elem[1] < 1:
raise RuntimeError("Unexpected concentration level in state: " + str(elem))
def add_reaction(self, R, I, P):
"""Adds a reaction"""
@@ -378,18 +386,21 @@ class ReactionSystemWithConcentrations(ReactionSystem):
reactants = []
for r in R:
self.is_valid_entity_with_concentration(r)
self.has_non_zero_concentration(r)
entity,level = r
reactants.append((self.get_entity_id(entity),level))
inhibitors = []
for i in I:
self.is_valid_entity_with_concentration(i)
self.has_non_zero_concentration(i)
entity,level = i
inhibitors.append((self.get_entity_id(entity),level))
products = []
for p in P:
self.is_valid_entity_with_concentration(p)
self.has_non_zero_concentration(p)
entity,level = p
products.append((self.get_entity_id(entity),level))
@@ -435,15 +446,20 @@ class ReactionSystemWithConcentrations(ReactionSystem):
producible_entities = set()
for reaction in self.reactions:
producible_entities = producible_entities.union(set(reaction[2]))
product_entities = [e for e,c in reaction[2]]
producible_entities = producible_entities.union(set(product_entities))
reactions_by_prod = {}
reactions_by_prod = {}
for prod_entity in producible_entities:
reactions_by_prod[prod_entity] = []
for reaction in self.reactions:
if prod_entity in reaction[2]:
reactions_by_prod[prod_entity].append([reaction[0],reaction[1]])
for p_e in producible_entities:
reactions_by_prod[p_e] = []
for r in self.reactions:
product_entities = [e for e,c in r[2]]
if p_e in product_entities:
reactants = r[0]
inhibitors = r[1]
products = [(e,c) for e,c in r[2] if e == p_e]
reactions_by_prod[p_e].append((reactants, inhibitors, products))
# save in cache
self.reactions_by_prod = reactions_by_prod
@@ -460,8 +476,17 @@ class ReactionSystemWithAutomaton(object):
def show(self, soft=False):
self.rs.show(soft)
self.ca.show()
def is_with_concentrations(self):
if not isinstance(self.rs, ReactionSystemWithConcentrations):
return False
if not isinstance(self.ca, ContextAutomatonWithConcentrations):
return False
return True
def sanity_check(self):
pass
#
# class ReactionSystemWithConcentrationWithAutomaton(ReactionSystemWithAutomaton):
#
# def __init__(self, reaction_system, context_automaton):
@@ -471,5 +496,3 @@ class ReactionSystemWithAutomaton(object):
# def show(self, soft=False):
# self.rs.show(soft)
# self.ca.show()

View File

@@ -6,10 +6,11 @@
"""
from rctsys import ReactionSystem,ReactionSystemWithConcentrations,ContextAutomatonWithConcentrations
from rctsys import ReactionSystem,ReactionSystemWithConcentrations,ContextAutomatonWithConcentrations,ReactionSystemWithAutomaton
from smtchecker import SmtChecker
from smtcheckerpgrs import SmtCheckerPGRS
from smtcheckerdistribrs import SmtCheckerDistribRS
from smtcheckerrsc import SmtCheckerRSC
import sys
import rs_examples
@@ -67,15 +68,22 @@ def main():
r = ReactionSystemWithConcentrations()
r.add_bg_set_entities(["a","b","c","d"])
r.add_reaction([("a",1)],[("b",2)],[("c",1)])
r.add_reaction([("c",1)],[("b",2)],[("c",1),("b",1)])
r.show()
# print(r.get_reactions_by_product())
c = ContextAutomatonWithConcentrations(r)
c.add_init_state("1")
c.add_transition("1", [("c",1)], "1")
c.add_transition("1", [("d",2)], "1")
c.show()
rc = ReactionSystemWithAutomaton(r,c)
smt = SmtCheckerRSC(rc)
smt.check_reachability([('c',1)],print_time=True,max_level=20)
if __name__ == "__main__":
try:

View File

@@ -69,47 +69,42 @@ class SmtCheckerRSC(object):
return init_state_enc
def enc_enabledness(self, level, prod_entity):
"""Encodes the enabledness condition for a given level and a given entity"""
def enc_produced_concentration(self, level, prod_entity):
"""Encodes the produced concentrations for the given level and entity"""
rcts_for_prod_entity = self.rs.get_reactions_by_product()[prod_entity]
if rcts_for_prod_entity == []:
return False
#
# TODO:
#
# czy potrzebujemy get_reactions_by_product?
# czy może wystarczy jak zakodujemy reakcje jako trójki?
#
# ... na pewno musi być informacja o produkowanym stężeniu
#
return simplify(self.v[level+1][prod_entity] == 0)
enc_rct_prod = False
for reactants,inhibitors in rcts_for_prod_entity:
for reactants,inhibitors,products in rcts_for_prod_entity:
enc_reactants = True
enc_inhibitors = True
for reactant in reactants:
# enc_products -- below
for reactant,concentration in reactants:
enc_reactants = simplify(And(enc_reactants,
Or(self.v[level][reactant], self.v_ctx[level][reactant])))
for inhibitor in inhibitors:
Or(self.v[level][reactant] >= concentration, self.v_ctx[level][reactant] >= concentration)))
for inhibitor,concentration in inhibitors:
enc_inhibitors = simplify(And(enc_inhibitors,
And(self.v[level][inhibitor], self.v_ctx[level][inhibitor])))
And(self.v[level][inhibitor] < concentration, self.v_ctx[level][inhibitor] < concentration)))
enc_rct_prod = simplify(Or(enc_rct_prod, And(enc_reactants, enc_inhibitors)))
enc_products = self.v[level+1][products[0][0]] == products[0][1]
enc_rct_prod = simplify(Or(enc_rct_prod, And(enc_reactants, enc_inhibitors, enc_products)))
return enc_rct_prod
def enc_entity_production(self, level, prod_entity):
"""Encodes the production of a given entity from a given level at level+1"""
enc_enab_cond = self.enc_enabledness(level, prod_entity)
enc_ent_prod = Or(And(enc_enab_cond, self.v[level+1][prod_entity]),
And(Not(enc_enab_cond), Not(self.v[level+1][prod_entity])))
return simplify(enc_ent_prod)
# def enc_entity_production(self, level, prod_entity):
# """Encodes the production of a given entity from a given level at level+1"""
#
# enc_enab_cond = self.enc_enabledness(level, prod_entity)
#
# enc_ent_prod = Or(And(enc_enab_cond, self.v[level+1][prod_entity]),
# And(Not(enc_enab_cond), Not(self.v[level+1][prod_entity])))
#
# return simplify(enc_ent_prod)
def enc_transition_relation(self, level):
return simplify(And(self.enc_rs_trans(level), self.enc_automaton_trans(level)))
@@ -121,13 +116,15 @@ class SmtCheckerRSC(object):
enc_trans = True
for prod_entity in self.rs.get_reactions_by_product():
reactions = self.rs.get_reactions_by_product()
for prod_entity in reactions:
unused_entities.remove(prod_entity)
enc_trans = simplify(And(enc_trans, self.enc_entity_production(level, prod_entity)))
enc_trans = simplify(And(enc_trans, self.enc_produced_concentration(level, prod_entity)))
for prod_entity in unused_entities:
enc_trans = simplify(And(enc_trans, Not(self.v[level+1][prod_entity])))
enc_trans = simplify(And(enc_trans, self.v[level+1][prod_entity] == 0))
return enc_trans
@@ -141,15 +138,17 @@ class SmtCheckerRSC(object):
dst_enc = self.ca_state[level+1] == dst
all_ent = set(range(len(self.rs.background_set)))
incl_ctx = ctx
incl_ctx = set([e for e,c in ctx])
excl_ctx = all_ent - incl_ctx
ctx_enc = True
for c in incl_ctx:
ctx_enc = simplify(And(ctx_enc, self.v_ctx[level][c]))
for c in excl_ctx:
ctx_enc = simplify(And(ctx_enc, Not(self.v_ctx[level][c])))
for e,c in ctx:
ctx_enc = simplify(And(ctx_enc, self.v_ctx[level][e] == c))
for e in excl_ctx:
ctx_enc = simplify(And(ctx_enc, self.v_ctx[level][e] == 0))
cur_trans = simplify(And(src_enc, ctx_enc, dst_enc))
enc_trans = simplify(Or(enc_trans, cur_trans))
@@ -202,19 +201,27 @@ class SmtCheckerRSC(object):
print(" State: {", end=""),
for var_id in range(len(self.v[level])):
if repr(m[self.v[level][var_id]]) == "True":
print(" " + self.rs.get_entity_name(var_id), end="")
var_rep = repr(m[self.v[level][var_id]])
if not var_rep.isdigit():
raise RuntimeError("unexpected: representation is not a positive integer")
if int(var_rep) > 0:
print(" " + self.rs.get_entity_name(var_id) + "=" + var_rep, end="")
# print(" " + repr(m[self.v[level][var_id]]), end="")
print(" }")
if level != max_level:
print(" Context set: ", end="")
print("{", end="")
for var_id in range(len(self.v[level])):
if repr(m[self.v_ctx[level][var_id]]) == "True":
print(" " + self.rs.get_entity_name(var_id), end="")
var_rep = repr(m[self.v_ctx[level][var_id]])
if not var_rep.isdigit():
raise RuntimeError("unexpected: representation is not a positive integer")
if int(var_rep) > 0:
print(" " + self.rs.get_entity_name(var_id) + "=" + var_rep, end="")
print(" }")
def check_reachability(self, state, print_witness=True, print_time=False):
def check_reachability(self, state, exact_state=False, print_witness=True,
print_time=False, print_mem=False, max_level=100):
"""Main testing function"""
if print_time:
@@ -224,33 +231,50 @@ class SmtCheckerRSC(object):
self.solver.add(self.enc_init_state(0))
current_level = 0
print(self.enc_exact_state(current_level,state))
# print(self.enc_exact_state(current_level,state))
#
# print(self.enc_produced_concentration(current_level, 1))
# while True:
# print("\r[i] Level: " + str(current_level), end="")
# stdout.flush()
#
# self.prepare_all_variables()
#
# # reachability test:
# self.solver.push()
# self.solver.add(self.enc_state(current_level,state))
#
# result = self.solver.check()
# if result == sat:
# print("\n[+] SAT at level=" + str(current_level))
# if print_witness:
# self.decode_witness(current_level)
# break
# else:
# self.solver.pop()
#
# self.solver.add(self.enc_transition_relation(current_level))
#
# current_level += 1
#
# if print_time:
# stop = time()
# print()
# print("[i] Time: " + repr(stop-start))
self.prepare_all_variables()
print(And(self.enc_transition_relation(current_level), self.enc_init_state(0)))
while True:
self.prepare_all_variables()
print("-----[ Working at level=" + str(current_level) + " ]-----")
stdout.flush()
# reachability test:
print("[i] Adding the reachability test...")
self.solver.push()
if exact_state:
self.solver.add(self.enc_exact_state(current_level,state))
else:
self.solver.add(self.enc_min_state(current_level,state))
result = self.solver.check()
if result == sat:
print("\n[+] SAT at level=" + str(current_level))
if print_witness:
self.decode_witness(current_level)
break
else:
self.solver.pop()
print("[i] Unrolling the transition relation")
self.solver.add(self.enc_transition_relation(current_level))
print("-----[ level=" + str(current_level) + " done ]")
current_level += 1
if current_level > max_level:
print("Stopping at level=" + str(max_level))
break
if print_time:
stop = time()
print()
print("[i] Time: " + repr(stop-start))