diff --git a/Examples.py b/Examples.py index 8d98dc491c..e662e99298 100644 --- a/Examples.py +++ b/Examples.py @@ -1164,7 +1164,7 @@ def passive_vpa_learning_arithmetics(): def passive_vpa_learning_on_all_benchmark_models(): from aalpy.learning_algs import run_PAPNI from aalpy.utils.BenchmarkVpaModels import vpa_L1, vpa_L12, vpa_for_odd_parentheses - from aalpy.utils import generate_input_output_data_from_vpa, convert_i_o_traces_for_RPNI + from aalpy.utils import generate_input_output_data_from_vpa for gt in [vpa_L1(), vpa_L12(), vpa_for_odd_parentheses()]: vpa_alphabet = gt.input_alphabet @@ -1200,7 +1200,7 @@ def gsm_edsm(): from aalpy import load_automaton_from_file from aalpy.utils.Sampling import get_io_traces, sample_with_length_limits from aalpy.learning_algs.general_passive.GeneralizedStateMerging import run_GSM - from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation + from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import SimpleScoreCalculation from aalpy.learning_algs.general_passive.GsmNode import GsmNode automaton = load_automaton_from_file("DotModels/car_alarm.dot", "moore") @@ -1212,7 +1212,7 @@ def EDSM_score(part: Dict[GsmNode, GsmNode]): nr_merged = len(part) return nr_merged - nr_partitions - score = ScoreCalculation(score_function=EDSM_score) + score = SimpleScoreCalculation(score_function=EDSM_score) learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="deterministic", score_calc=score) learned_model.visualize() @@ -1221,7 +1221,7 @@ def gsm_likelihood_ratio(): from typing import Dict from scipy.stats import chi2 from aalpy.learning_algs.general_passive.GeneralizedStateMerging import run_GSM - from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreFunction, differential_info, ScoreCalculation + from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreFunction, differential_info, SimpleScoreCalculation from aalpy.learning_algs.general_passive.GsmNode import GsmNode from aalpy.utils.Sampling import get_io_traces, sample_with_length_limits from aalpy import load_automaton_from_file @@ -1246,15 +1246,17 @@ def score_fun(part: Dict[GsmNode, GsmNode]): return score_fun - score = ScoreCalculation(score_function=likelihood_ratio_score()) + score = SimpleScoreCalculation(score_function=likelihood_ratio_score()) learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score) learned_model.visualize() def example_Alergia_extension(): + from typing import Any + from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler from aalpy.learning_algs.general_passive.GeneralizedStateMerging import run_GSM - from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import hoeffding_compatibility, ScoreCalculation from aalpy.learning_algs.general_passive.GsmNode import GsmNode + from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import hoeffding_compatibility, SimpleFutureBasedCompatibility, SpecialScores from aalpy.utils.Sampling import get_io_traces, sample_with_length_limits from aalpy import load_automaton_from_file @@ -1262,39 +1264,40 @@ def example_Alergia_extension(): input_traces = sample_with_length_limits(automaton.get_input_alphabet(), 2000, 20, 30) traces = get_io_traces(automaton, input_traces) - # NOTE THAT This example is equivalent to a call to a function run_Alergia_EDSM + # NOTE: a more general version of this is provided in aalpy.learning_algs.general_passive.ScoreFunctionsGSM + class ScoreIOAlergiaWithEDSM(SimpleFutureBasedCompatibility): + def __init__(self, eps: float): + self.compat = hoeffding_compatibility(eps) + SimpleFutureBasedCompatibility.__init__(self, compatibility_on_pta=True) + self.score = None - class IOAlergiaWithEDSM(ScoreCalculation): - def __init__(self, epsilon): - super().__init__() - self.ioa_compatibility = hoeffding_compatibility(epsilon) - self.evidence = 0 - - def reset(self): - self.evidence = 0 + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: + self.score = 0 + verdict = super().early_score(red, blue) + if verdict is SpecialScores.ImmediateReject: + return verdict + return self.score - def local_compatibility(self, a: GsmNode, b: GsmNode): - self.evidence += 1 - return self.ioa_compatibility(a, b) - - def score_function(self, part): - return self.evidence + def local_compatibility(self, red: GsmNode, blue: GsmNode) -> float: + self.score += 1 + return self.compat(red, blue) epsilon = 0.05 scores = { - "IOA": ScoreCalculation(hoeffding_compatibility(epsilon)), - "IOA+EDSM": IOAlergiaWithEDSM(epsilon), + "IOA": SimpleFutureBasedCompatibility(local_compatibility=hoeffding_compatibility(epsilon, True), compatibility_on_pta=True), + "IOA+EDSM": ScoreIOAlergiaWithEDSM(epsilon), } for name, score in scores.items(): learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score, - compatibility_on_pta=True, compatibility_on_futures=True) + data_handler=CountOnPTADataHandler()) learned_model.visualize(name) def gsm_IOAlergia_domain_knowldege(): from aalpy.learning_algs.general_passive.GeneralizedStateMerging import run_GSM - from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import hoeffding_compatibility, ScoreCalculation + from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import hoeffding_compatibility, SimpleFutureBasedCompatibility + from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler from aalpy.learning_algs.general_passive.GsmNode import GsmNode from aalpy.utils.Sampling import get_io_traces, sample_with_length_limits from aalpy import load_automaton_from_file @@ -1320,12 +1323,12 @@ def ioa_compat_domain_knowledge(a: GsmNode, b: GsmNode): return parity and ioa scores = { - "IOA": ScoreCalculation(ioa_compat), - "IOA+DK": ScoreCalculation(ioa_compat_domain_knowledge), + "IOA": SimpleFutureBasedCompatibility(local_compatibility=ioa_compat, compatibility_on_pta=True), + "IOA+DK": SimpleFutureBasedCompatibility(local_compatibility=ioa_compat_domain_knowledge, compatibility_on_pta=True), } for name, score in scores.items(): learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score, - compatibility_on_pta=True, compatibility_on_futures=True) + data_handler=CountOnPTADataHandler()) learned_model.visualize(name) def k_tails_example(): @@ -1336,16 +1339,13 @@ def k_tails_example(): input_alphabet_size=3, output_alphabet_size=3) + # data is a list of sequences in this format [(i1, o1), (i2, o1), (i1, o3)] data = generate_input_output_data_from_automata(model, num_sequences=2000, min_seq_len=1, max_seq_len=12, sequance_type='io_traces') - # k-trails works with prefix-closed input output traces, not labeled sequences like RPNI - # data is a list of sequences in this format [(i1, o1), (i2, o1), (i1, o3)] - # run k_tails with two different k's - k_trails_1 = run_k_tails(data, k=3, automaton_type='moore', print_info=True) - + k_tails_1 = run_k_tails(data, k=3, automaton_type='moore', print_info=True) k_tails_2 = run_k_tails(data, k=8, automaton_type='mealy', print_info=True) diff --git a/aalpy/learning_algs/__init__.py b/aalpy/learning_algs/__init__.py index b58d8cf48e..a4cc189a76 100644 --- a/aalpy/learning_algs/__init__.py +++ b/aalpy/learning_algs/__init__.py @@ -12,6 +12,6 @@ from .deterministic_passive.PAPNI import run_PAPNI from .deterministic_passive.active_RPNI import run_active_RPNI from .general_passive.GeneralizedStateMerging import run_GSM -from .general_passive.GsmAlgorithms import run_EDSM, run_Alergia_EDSM, run_k_tails +from .general_passive.GsmAlgorithms import run_EDSM, run_Alergia_GSM, run_Alergia_EDSM, run_k_tails from .resetless.hW import run_hW from .resetless.resetless_oracles import hWOracle, RandomhWOracle, RandomWphWOracle \ No newline at end of file diff --git a/aalpy/learning_algs/general_passive/AssociatedData.py b/aalpy/learning_algs/general_passive/AssociatedData.py new file mode 100644 index 0000000000..0a1c4c1138 --- /dev/null +++ b/aalpy/learning_algs/general_passive/AssociatedData.py @@ -0,0 +1,63 @@ +import math +from abc import abstractmethod, ABC +from collections import defaultdict +from typing import Any + + +ProbabilityDict = dict[Any, dict[Any, float]] + +class StochasticData(ABC): + """ + Interface class for data used with `transition_behavior` set to "stochastic". + """ + @abstractmethod + def get_probabilities(self) -> ProbabilityDict: + """ + Method for extracting transition probabilities when converting to automaton models. + + :return ProbabilityDict: Nested dictionary of transition probabilities. + """ + pass + +CountDict = dict[Any, dict[Any, int]] + +def int_dict_increment(c_dict, out_sym, cnt): + c_dict[out_sym] = c_dict.get(out_sym, 0) + cnt + +class CountData(StochasticData): + def __init__(self): + # TODO get rid of this indirection + self.transition_count: CountDict = defaultdict(dict) + + def local_log_likelihood_contribution(self): + llc = 0 + for in_sym, trans in self.transition_count.items(): + total_count = 0 + for out_sym, count in trans.items(): + total_count += count + llc += count * math.log(count) + if total_count != 0: + llc -= total_count * math.log(total_count) + return llc + + def count(self): + return sum(sum(trans.values()) for trans in self.transition_count.values()) + + def get_probabilities(self) -> ProbabilityDict: + ret = dict() + for in_sym, trans in self.transition_count.items(): + total_count = sum(trans.values()) + ret[in_sym] = {out_sym: count / total_count for out_sym, count in trans.items()} + return ret + + +ShadowPTA = dict[Any, dict[Any, 'GsmNode']] +class ShadowPTAData: + def __init__(self): + self.shadow_pta: ShadowPTA = defaultdict(dict) + +class CountOnPTAData(ShadowPTAData, CountData): + def __init__(self): + ShadowPTAData.__init__(self) + CountData.__init__(self) + self.pta_count: CountDict = defaultdict(dict) \ No newline at end of file diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py new file mode 100644 index 0000000000..713f4c3939 --- /dev/null +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -0,0 +1,294 @@ +from abc import abstractmethod, ABC +from typing import Generic, TypeVar, Any + +from aalpy.learning_algs.general_passive.AssociatedData import CountData, CountOnPTAData, int_dict_increment +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, IOTrace, IOExample, unknown_output, no_op_input, OutputBehavior + + +T = TypeVar("T") + +DataFormat = str +DataFormatRange = ["io_traces", "labeled_sequences", "traces", "tree"] + +# TODO reuse in RPNI +def detect_data_format(data: Any, check_consistency: bool = False, guess: bool = False) -> DataFormat: + """ + Guess the data format of the provided learning data. + + :param Any data: Input data: a GsmNode (tree), or a sequence of traces/examples. + :param bool check_consistency: Whether to check all data points instead of returning as soon as a unique format is found. + :param bool guess: Whether to allow guessing a single format when multiple formats remain ambiguous. + :return DataFormat: The detected data format string. The different data formats are: + - "tree": a tree-shaped automaton provided as a GsmNode + - "io_traces": either + - Moore traces [[o, (i,o), (i,o), ...], ...] + - Mealy traces [[(i,o), (i,o), ...], ...] + - "labeled_sequences": [([i, i, ...], o), ...] + - "traces": [[o, o, ...], ...] + """ + + if isinstance(data, GsmNode): + return "tree" + + accepted_types = (tuple, list) + + # mapping data formats to compatibility criteria + check_dict = dict( + io_traces=lambda obj: len(obj) <= 1 or all(isinstance(o, accepted_types) and len(o) == 2 for o in obj[1:]), + labeled_sequences=lambda obj: len(obj) == 2 and isinstance(obj[0], accepted_types), + ) + accept_dict = {k: True for k in check_dict} + + if not isinstance(data, accepted_types): + raise ValueError("wrong input format. expected tuple or list.") + if len(data) == 0: + return "io_traces" + + accepted_formats = list(accept_dict.keys()) + for data_point in data: + if not isinstance(data_point, accepted_types): + raise ValueError("wrong input format. expected tuple or list.") + for k, check in check_dict.items(): + accept_dict[k] &= check(data_point) + accepted_formats = [k for k, v in accept_dict.items() if v] + if len(accepted_formats) == 1 and not check_consistency: + return accepted_formats[0] + if len(accepted_formats) == 0: + return "traces" # default to traces + #raise ValueError("invalid or inconsistent data. no options left") + if len(accepted_formats) != 1 and not guess: + raise ValueError("ambiguous data format. data format needs to be specified explicitly.") + return accepted_formats[0] + +class DataHandler(Generic[T], ABC): + def add_trace(self, root_node: GsmNode[T], trace: IOTrace): + """ + Add an IO trace to a given root node, extending it with new nodes as necessary. + + :param GsmNode root_node: GsmNode to which the trace should be added. + :param IOTrace trace: Sequence of (input, output) pairs to add. + """ + curr_node: GsmNode[T] = root_node + for in_value, out_value in trace: + prefix_access_pair = self.abstract(in_value, out_value) + in_sym, out_sym = prefix_access_pair + transitions = curr_node.transitions[in_sym] + node = transitions.get(out_sym) + if node is None: + node = GsmNode(prefix_access_pair, curr_node, self.init_data()) + transitions[out_sym] = node + self.aggregate_data(curr_node, in_value, out_value, node) + curr_node = node + + def add_labeled_sequence(self, root_node: GsmNode[T], example: IOExample): + """ + Add a labeled input sequence (inputs with a single label attached at the end) to the tree. + + :param IOExample example: (inputs, output) pair, where output labels the state reached by inputs. + :param DataHandler[T] self: IOHandler used for abstraction and aggregation of trace data + """ + inputs, output = example + curr_node: GsmNode = root_node + in_sym = None + + if len(inputs) == 0: + in_sym, out_sym = self.abstract(no_op_input, output) + existing_out_sym = root_node.get_prefix_output() + if existing_out_sym is not unknown_output and existing_out_sym != out_sym: + raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") + self.aggregate_data(None, no_op_input, output, root_node) + + # step through inputs and add transitions + for idx, in_value in enumerate(inputs): + out_value = output if idx == len(inputs) - 1 else unknown_output + in_sym, out_sym = self.abstract(in_value, out_value) + transitions = curr_node.transitions[in_sym] + if len(transitions) == 0: + node = GsmNode((in_sym, out_sym), curr_node, self.init_data()) + transitions[out_sym] = node + elif len(transitions) == 1: + existing_out_sym, node = next(iter(transitions.items())) + if existing_out_sym != out_sym and unknown_output not in [existing_out_sym, out_sym]: + raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") + else: + assert False, "failed to construct deterministic PTA or raise an exception where not possible" + + self.aggregate_data(curr_node, in_value, out_value, node) + curr_node = node + + # fix prefix / predecessor + curr_node.resolve_unknown_prefix_output(out_sym) + pred = curr_node.predecessor + if pred: + transitions = pred.transitions[in_sym] + if unknown_output in transitions: + transitions[out_sym] = transitions.pop(unknown_output) + if out_sym not in transitions: + raise ValueError("nondeterminism encountered for GSM with labeled_sequences. not supported") + + def createPTA(self, data: Any, output_behavior: OutputBehavior, data_format: DataFormat = None) -> 'GsmNode[T]': + """ + Build a prefix tree acceptor (PTA) from the given data. + + :param Any data: Learning data, in one of the supported data formats (or already a GsmNode tree). + :param OutputBehavior output_behavior: Either "moore" or "mealy". + :param DataFormat | None data_format: Explicit data format, or None to auto-detect. + :param DataHandler[T] self: IOHandler used for abstraction and aggregation of trace data + :return GsmNode: The root node of the constructed (or passed-through) PTA. + """ + if data_format is None: + data_format = detect_data_format(data) + if data_format not in DataFormatRange: + raise ValueError(f"invalid data format {data_format}. should be in {DataFormatRange}") + + if data_format == "tree": + if not data.is_tree(): + raise ValueError("provided automaton is not a tree") + return data + # TODO extract method for replaying data on dot model + root_node = GsmNode((no_op_input, unknown_output), None, self.init_data()) + if data_format == "labeled_sequences": + for example in data: + self.add_labeled_sequence(root_node, example) + if data_format == "io_traces" or data_format == "traces": + if output_behavior == "moore": + root_node.prefix_access_pair = self.abstract(no_op_input, data[0][0]) + initial_output_symbol = root_node.prefix_access_pair[1] + + for trace in data: + initial_output = trace[0] + _, ios = self.abstract(no_op_input, initial_output) + if ios != initial_output_symbol: + raise ValueError("expect unique initial output symbol for Moore behavior") + self.aggregate_data(None, no_op_input, initial_output, root_node) + + data = (d[1:] for d in data) + for trace in data: + if data_format == "traces": + trace = (("step", t) for t in trace) + self.add_trace(root_node, trace) + return root_node + + def initialize_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): + """ + Callback triggered just before the partitioning for a merge candidate is constructed. + + :param GsmNode red: The red node of the merge candidate. + :param GsmNode blue: The blue node of the merge candidate. + :param bool first_pass: Indicates whether this is the first partitioning pass for score calculation, or the + second pass for finalizing the partitioning. + """ + pass + + def abstract(self, in_val: Any, out_val: Any) -> tuple[Any, Any]: + """ + Method used during PTA construction to abstract from potentially continuous input data. By default, no abstraction is performed. + + :param Any in_val: The input value. + :param Any out_val: The output value. + :return tuple[Any, Any]: The abstracted input and output symbols. + """ + return in_val, out_val + + @abstractmethod + def init_data(self) -> T: + """ + Provides a fresh data object for a new GsmNode instance. + + :return T: The 'empty' data object. + """ + ... + + @abstractmethod + def aggregate_data(self, src_node: 'GsmNode[T]', in_value, out_value, dst_node: 'GsmNode[T]'): + """ + Method for aggregating data in GsmNodes during PTA construction for an individual transition observed in the data. + + :param GsmNode src_node: The source GsmNode instance of the transition. + :param Any in_value: The (concrete) input value of the transition. + :param Any out_value: The (concrete) output value of the transition. + :param GsmNode dst_node: The destination GsmNode instance of the transition. + """ + ... + + @abstractmethod + def merge(self, x: T, y: T) -> T: + """ + Method for merging the data of two GsmNodes during partition construction. + + :param T x: The data of the first GsmNode instance, corresponding to the partition representative. + :param T y: The data of the second GsmNode instance, corresponding to the GsmNode to be merged into the partition. + :return T: The updated data for the partition. + """ + ... + + @abstractmethod + def copy(self, x: T) -> T: + """ + Copy operation for the data attached to GsmNodes used when computing a reversible partitioning. + + :param T x: The data to be copied. + :return T: The copy. + """ + ... + +class NoOpDataHandler(DataHandler[None]): + """ + DataHandler that neither abstracts the traces, nor tracks any other values during PTA construction. + """ + + def init_data(self) -> None: + return None + + def aggregate_data(self, src_node: 'GsmNode[None]', in_sym, out_value, dst_node: 'GsmNode[None]'): + pass + + def merge(self, x: None, y: None) -> None: + return None + + def copy(self, x: None) -> None: + return None + + +class CountDataHandler(DataHandler[CountData]): + def merge(self, x: CountData, y: CountData) -> CountData: + for in_sym, y_count in y.transition_count.items(): + x_count = x.transition_count[in_sym] + for out_sym, count in y_count.items(): + int_dict_increment(x_count, out_sym, count) + return x + + def copy(self, x: CountData) -> CountData: + ret = CountData() + ret.transition_count.update((k, v.copy()) for k, v in x.transition_count.items()) + return ret + + def init_data(self) -> CountData: + return CountData() + + def aggregate_data(self, src_node: 'GsmNode[CountData]', in_value, out_value, dst_node: 'GsmNode[CountData]'): + if out_value is unknown_output: + raise RuntimeError(f"{self.__class__.__name__} does not support non-prefix-closed data") + if src_node is not None: + int_dict_increment(src_node.data.transition_count[in_value], out_value, 1) + + +class CountOnPTADataHandler(CountDataHandler, DataHandler[CountOnPTAData]): + def init_data(self) -> CountOnPTAData: + return CountOnPTAData() + + def copy(self, x: CountOnPTAData) -> CountOnPTAData: + ret = CountOnPTAData() + ret.transition_count.update((k, v.copy()) for k, v in x.transition_count.items()) + ret.pta_count = x.pta_count + ret.shadow_pta = x.shadow_pta + return ret + + def aggregate_data(self, src_node: 'GsmNode[CountOnPTAData]', in_value, out_value, dst_node: 'GsmNode[CountOnPTAData]'): + if out_value is unknown_output: + raise RuntimeError(f"{self.__class__.__name__} does not support non-prefix-closed data") + if src_node is None: + return + int_dict_increment(src_node.data.transition_count[in_value], out_value, 1) + int_dict_increment(src_node.data.pta_count[in_value], out_value, 1) + src_node.data.shadow_pta[in_value][out_value] = dst_node diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index a597fcf413..7a20873201 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -2,14 +2,18 @@ # state-merging framework used to passively learn deterministic, nondeterministic and # stochastic automata from data. import functools +import warnings from collections import deque -from collections.abc import Callable -from typing import Any +from copy import copy +from typing import Callable, Any -from aalpy.base import Automaton -from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, TransitionInfo, \ - OutputBehaviorRange, TransitionBehaviorRange, DataFormat, intersection_iterator, unknown_output, detect_data_format -from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility +from aalpy import Automaton +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, OutputBehaviorRange, \ + TransitionBehaviorRange, unknown_output +from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler, CountDataHandler, detect_data_format, \ + DataHandler, NoOpDataHandler, DataFormat +from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility, \ + SimpleFutureBasedCompatibility, SpecialScores, SimpleScoreCalculation # TODO add option for making checking of futures and partition non mutual exclusive? @@ -27,10 +31,12 @@ def __init__(self, red: GsmNode, blue: GsmNode) -> None: """ self.red: GsmNode = red self.blue: GsmNode = blue - self.score = False + self.score = SpecialScores.NoScore self.red_mapping: dict[GsmNode, GsmNode] = dict() self.full_mapping: dict[GsmNode, GsmNode] = dict() - + self.new_blue = [] + self.remaining_merges = None + self.nr_merged_states = 0 class Instrumentation: """Base class for hooks that observe/report on the progress of GeneralizedStateMerging.run.""" @@ -81,7 +87,6 @@ def learning_done(self, root: GsmNode) -> None: """ pass - class GeneralizedStateMerging: """Implements the red-blue state-merging framework used to passively learn automata from data.""" @@ -89,24 +94,21 @@ def __init__(self, *, output_behavior: OutputBehavior = "moore", transition_behavior: TransitionBehavior = "deterministic", score_calc: ScoreCalculation = None, - pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, - compatibility_on_pta: bool = False, - compatibility_on_futures: bool = False, - node_order: Callable[[GsmNode, GsmNode], bool] = None, - consider_only_min_blue: bool = False, - depth_first: bool = False) -> None: + data_handler: DataHandler = None, + node_order: Callable[[GsmNode], Any] = None, + consider_only_min_blue = False, + depth_first = False, + ): """ Configure a GeneralizedStateMerging instance. :param OutputBehavior output_behavior: Either "moore" or "mealy". :param TransitionBehavior transition_behavior: Either "deterministic", "nondeterministic" or "stochastic". :param ScoreCalculation score_calc: Local compatibility / global score calculation to use. - :param Callable[[GsmNode], GsmNode] pta_preprocessing: Pre-processing function applied to the constructed PTA. :param Callable[[GsmNode], GsmNode] postprocessing: Post-processing function applied to the learned model. - :param bool compatibility_on_pta: Whether compatibility is evaluated on the PTA instead of the current hypothesis. - :param bool compatibility_on_futures: Whether compatibility is evaluated on futures instead of full partitions. - :param Callable[[GsmNode, GsmNode], bool] node_order: Order in which merge candidates are considered. + :param DataHandler data_handler: IOHandler object governing abstraction and aggregation of data + :param Callable[[GsmNode], Any] node_order: Comparison key to determine the order in which merge candidates are considered. :param bool consider_only_min_blue: Whether to only consider the minimal blue node in each round. :param bool depth_first: Whether compatibility is checked depth-first instead of breadth-first. """ @@ -120,41 +122,30 @@ def __init__(self, *, if score_calc is None: if transition_behavior == "deterministic": - score_calc = ScoreCalculation() + score_calc = SimpleScoreCalculation(GsmNode.deterministic_compatible) elif transition_behavior == "nondeterministic" : raise ValueError("Missing score_calc for nondeterministic transition behavior. No default available.") elif transition_behavior == "stochastic" : - score_calc = ScoreCalculation(hoeffding_compatibility(0.005, compatibility_on_pta)) + lc = hoeffding_compatibility(0.005, True) + score_calc = SimpleFutureBasedCompatibility(local_compatibility=lc, compatibility_on_pta=True) + if data_handler is not None: + raise ValueError("Using default algorithm for stochastic systems but a data_handler was provided.") + data_handler = CountOnPTADataHandler() self.score_calc: ScoreCalculation = score_calc - if node_order is None: - self.node_order = GsmNode.default_order - else: - self.node_order = functools.cmp_to_key(lambda a, b: -1 if node_order(a, b) else 1) + if isinstance(node_order, str) and node_order == "short-lex": + node_order = functools.cmp_to_key(lambda a, b: -1 if GsmNode.short_lex_order(a, b) else 1) + self.node_order = node_order - self.pta_preprocessing = pta_preprocessing or (lambda x: x) self.postprocessing = postprocessing or (lambda x: x) - self.compatibility_on_pta = compatibility_on_pta - self.compatibility_on_futures = compatibility_on_futures + if data_handler is None: + data_handler = CountDataHandler() if transition_behavior == "stochastic" else NoOpDataHandler() + self.data_handler = data_handler self.consider_only_min_blue = consider_only_min_blue self.depth_first = depth_first - def compute_local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: - """ - Check whether two nodes are locally compatible, considering output/transition behavior and the score calculation. - - :param GsmNode a: First node. - :param GsmNode b: Second node. - :return bool: True if the nodes are locally compatible. - """ - if self.output_behavior == "moore" and not GsmNode.moore_compatible(a, b): - return False - if self.transition_behavior == "deterministic" and not GsmNode.deterministic_compatible(a, b): - return False - return self.score_calc.local_compatibility(a, b) - # TODO: make more generic by adding the option to use a different algorithm than red blue # for selecting potential merge candidates. Maybe using inheritance with abstract `run`. def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation | None = None, @@ -178,245 +169,314 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation raise ValueError("learning from labeled_sequences is not possible for nondeterministic systems") if data_format == "traces" and self.transition_behavior == "deterministic": print("learning deterministic systems from (output) traces only. this rarely makes sense. is `data_format` set correctly?") - root = GsmNode.createPTA(data, self.output_behavior, data_format) + root = self.data_handler.createPTA(data, self.output_behavior, data_format) - root = self.pta_preprocessing(root) instrumentation.pta_construction_done(root) instrumentation.log_promote(root) if self.transition_behavior == "deterministic": - if not root.is_deterministic(): - raise ValueError("required deterministic automaton but input data is nondeterministic") + deterministic_pta = root.is_deterministic() + if not deterministic_pta: + warnings.warn("required deterministic automaton but input data is nondeterministic") - # sorted list of states already considered + # sorted list of states already considered as distinct red_states = [root] + red_states_backing_set = {root} + blue_states = list(root.child_iterator()) partition_candidates: dict[tuple[GsmNode, GsmNode], Partitioning] = dict() - while True: - # sort states. states are always sorted using default order on original prefix - if self.node_order is not GsmNode.default_order: - red_states.sort(key=self.node_order) - - # get blue states - blue_states = [] - for r in red_states: - for _, _, t in r.transition_iterator(): - c = t.target - if c in red_states: - continue - blue_states.append(c) - if self.consider_only_min_blue and self.node_order is GsmNode.default_order: - break - - # no blue states left -> done - if len(blue_states) == 0: - break + while len(blue_states) != 0: + blue_states_to_consider = blue_states if self.consider_only_min_blue: # does it make sense to check the score function here? - blue_states = [min(blue_states, key=self.node_order)] - if self.node_order is not GsmNode.default_order: - blue_states.sort(key=self.node_order) + if self.node_order is None: + blue_states_to_consider = [blue_states[0]] + else: + blue_states_to_consider = [min(blue_states, key=self.node_order)] + + # could make this sort unconditional, but i think this is closer to the original in any case? + if self.node_order is not None: + # TODO: this could be done using insort as long as the order is static? + blue_states_to_consider.sort(key=self.node_order) + red_states.sort(key=self.node_order) # loop over blue states - promotion = False - for blue_state in blue_states: + best_candidate = None + best_score = None + for blue_state in blue_states_to_consider: # FUTURE: Parallelize # FUTURE: Save partitions? # calculate partitions resulting from merges with red states if necessary - current_candidates: dict[GsmNode, Partitioning] = dict() - perfect_partitioning = None - red_state = None + no_viable_merge_for_blue = True for red_state in red_states: - partition = partition_candidates.get((red_state, blue_state)) - if partition is None: - partition = self._partition_from_merge(red_state, blue_state) - if partition.score is True: - perfect_partitioning = partition + partitioning = partition_candidates.get((red_state, blue_state)) + if partitioning is None: + partitioning = Partitioning(red_state, blue_state) + self._partition_from_merge(partitioning, red_states_backing_set, True) + partition_candidates[(red_state, blue_state)] = partitioning + if best_candidate is None or best_score < partitioning.score: + best_candidate = partitioning + best_score = partitioning.score + no_viable_merge_for_blue &= partitioning.score is SpecialScores.ImmediateReject + if partitioning.score is SpecialScores.ImmediateAccept: break - current_candidates[red_state] = partition - assert red_state is not None # partition with perfect score found: don't consider anything else - if perfect_partitioning: - partition_candidates = {(red_state, blue_state): perfect_partitioning} + if best_score is SpecialScores.ImmediateAccept: + partition_candidates = {(best_candidate.red, best_candidate.blue): best_candidate} break - # no merge candidates for this blue state -> promote - if all(part.score is False for part in current_candidates.values()): - red_states.append(blue_state) - instrumentation.log_promote(blue_state) - promotion = True - break + # no merge candidates for this blue state -> promotion candidate + if no_viable_merge_for_blue: + score = self.score_calc.promotion_score(blue_state) + if best_candidate is None or best_score < score: + best_candidate = blue_state + best_score = score + if score is SpecialScores.ImmediateAccept: + break - # update tracking dict with new candidates - new_candidates = (((red, blue_state), part) for red, part in current_candidates.items() if - part.score is not False) - partition_candidates.update(new_candidates) - - # a state was promoted -> don't clear candidates - if promotion: - continue - - # find best partitioning and clear candidates - best_candidate = max(partition_candidates.values(), key=lambda part: part.score) - for real_node, partition_node in best_candidate.red_mapping.items(): - real_node.transitions = partition_node.transitions - real_node.predecessor = partition_node.predecessor - real_node.prefix_access_pair = partition_node.prefix_access_pair - instrumentation.log_merge(best_candidate) - # FUTURE: optimizations for compatibility tests where merges can be orthogonal - # FUTURE: caching for aggregating compatibility tests - partition_candidates.clear() + # check for state promotion + if isinstance(best_candidate, GsmNode): + # a state was promoted -> only forget scores for this blue node + for red in red_states: + del partition_candidates[(red, best_candidate)] + + # promote best candidate + red_states.append(best_candidate) + red_states_backing_set.add(best_candidate) + blue_states.remove(best_candidate) + blue_states.extend(best_candidate.child_iterator()) + instrumentation.log_promote(best_candidate) + + # promoting a state affects the cached partitioning information (what states are freshly blue) -> invalidate + for key in [key for key, p in partition_candidates.items() if best_candidate in p.full_mapping]: + del partition_candidates[key] + # TODO: update merge candidates instead of invalidating the cache + # # update cached partitions + # for partitioning in partition_candidates.values(): + # updated_promoted_node = partitioning.red_mapping.get(best_candidate) + # if updated_promoted_node is None: + # continue + # for in_sym, out_sym, successor in updated_promoted_node.transition_iterator(): + # trans = best_candidate.transitions.get(in_sym) + # if trans is None or (out_sym not in trans and unknown_output not in trans): + # partitioning.new_blue.append(successor) + elif isinstance(best_candidate, Partitioning): + # apply best merge candidate + for real_node, partition_node in best_candidate.red_mapping.items(): + real_node.transitions = partition_node.transitions + real_node.predecessor = partition_node.predecessor + real_node.data = partition_node.data + real_node.prefix_access_pair = partition_node.prefix_access_pair + self._partition_from_merge(best_candidate, red_states_backing_set, False) + blue_states.extend(best_candidate.new_blue) + blue_states.remove(best_candidate.blue) + instrumentation.log_merge(best_candidate) + + # a merge was performed -> merge scores are invalidated + # FUTURE: optimizations for compatibility tests where merges can be orthogonal + # FUTURE: caching for aggregating compatibility tests + partition_candidates.clear() + else: + assert False and "best candidate is neither a merge nor a promotion" instrumentation.learning_done(root) root = self.postprocessing(root) + if self.transition_behavior == "deterministic" and not root.is_deterministic(): + if deterministic_pta: + msg = "PTA is deterministic -> GSM is misconfigured" + else: + msg = "PTA is nondeterministic -> data is invalid and/or GSM is misconfigured" + raise ValueError(f"requested deterministic automaton but result is nondeterministic. {msg}") if convert: root = root.to_automaton(self.output_behavior, self.transition_behavior) return root - def _check_futures(self, red: GsmNode, blue: GsmNode) -> bool: - """ - Check compatibility of the futures of two nodes, without constructing a full partition. - - :param GsmNode red: Red (already accepted) node. - :param GsmNode blue: Blue (candidate) node. - :return bool: True if all reachable node pairs are locally compatible. - """ - q: deque[tuple[GsmNode, GsmNode]] = deque([(red, blue)]) - pop = q.pop if self.depth_first else q.popleft - - while len(q) != 0: - red, blue = pop() - - if self.compute_local_compatibility(red, blue) is False: - return False - - for in_sym, red_trans, blue_trans in intersection_iterator(red.transitions, blue.transitions): - for out_sym, red_child, blue_child in intersection_iterator(red_trans, blue_trans): - if self.compatibility_on_pta: - if blue_child.original_count == 0 or red_child.original_count == 0: - continue - q.append((red_child.original_target, blue_child.original_target)) - else: - q.append((red_child.target, blue_child.target)) - - return True - - def _partition_from_merge(self, red: GsmNode, blue: GsmNode) -> Partitioning: + def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNode], first_pass): """ Compute the partitioning resulting from merging blue into red, including its score. Assumes that blue is a tree and red is not reachable from blue. + It works in two passes: + - first pass: create partial partitioning sufficient for score calculation + - second pass: merge has been accepted, partitioning needs to be completed - :param GsmNode red: Red (already accepted) node the merge targets. - :param GsmNode blue: Blue (candidate) node being merged. - :return Partitioning: The resulting partitioning (with score False if incompatible). + :param Partitioning partitioning: Partitioning object indicating which states to merge. + :param first_pass: Which pass to perform. """ - partitioning = Partitioning(red, blue) - - self.score_calc.reset() - - if self.compatibility_on_futures: - if self._check_futures(red, blue) is False: - return partitioning - # when compatibility is determined only by future and scores are disabled, we need not create partitions. - if self.compatibility_on_futures and not self.score_calc.has_score_function(): - def update_partition(red_node: GsmNode, blue_node: GsmNode | None) -> GsmNode: - return red_node - else: + red = partitioning.red + blue = partitioning.blue + + # TODO: consider extracting main loop and split preample into two functions + if first_pass: + # for Moore machines the outputs have to match. for prefix-closed data (io-traces) this check is sufficient + # since Moore-ness is preserved for implied merges. + if not self.score_calc.override_default_checks() and self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): + partitioning.score = SpecialScores.ImmediateReject + return + + # check whether there is an early verdict and adapt helper functions accordingly + # TODO maybe split init from early verdict + partitioning.score = self.score_calc.early_score(red, blue) + if partitioning.score is not SpecialScores.NoScore: + return + partitioning.remaining_merges = [] + + # uncertain -> need to construct partitioning + red_partitions: set[GsmNode] = set() def update_partition(red_node: GsmNode, blue_node: GsmNode | None) -> GsmNode: p = partitioning.full_mapping.get(red_node) # could check smaller .red_mapping? if p is None: - p = red_node.shallow_copy() + # there is no partition yet for the 'red' node -> lazily copy + p = copy(red_node) + p.data = self.data_handler.copy(red_node.data) + # TODO: do lazier copies. currently we have "copy on access". could have true "copy on write" + p.transitions = red_node.transitions.copy() + + # add to partition table partitioning.full_mapping[red_node] = p partitioning.red_mapping[red_node] = p + + # check whether the partition is (proper) red + if red_node in red_nodes: + red_partitions.add(p) + assert red_node not in red_nodes or p in red_partitions if blue_node is not None: partitioning.full_mapping[blue_node] = p return p - # rewire the blue node's parent - blue_parent = update_partition(blue.predecessor, None) - blue_in_sym, blue_out_sym = blue.prefix_access_pair - blue_parent.transitions[blue_in_sym][blue_out_sym].target = red + cow_set = set() + def get_partition_trans(part: GsmNode, in_symbol): + trans = part.transitions[in_symbol] + if id(trans) not in cow_set: + trans = trans.copy() + part.transitions[in_symbol] = trans + cow_set.add(id(trans)) + return trans + elif partitioning.remaining_merges is None or len(partitioning.remaining_merges) != 0: + # best scoring merge candidate -> can manipulate nodes directly + red_partitions = red_nodes + def update_partition(red_node: GsmNode, blue_node: GsmNode | None) -> GsmNode: + return red_node - partition = update_partition(red, None) - if self.output_behavior == "moore": - partition.resolve_unknown_prefix_output(blue_out_sym) + def get_partition_trans(part: GsmNode, in_symbol): + return part.transitions[in_symbol] + else: + # first pass already did all the work + return + + self.score_calc.initialize_merge(red, blue, first_pass) + self.data_handler.initialize_merge(red, blue, first_pass) + q: deque[tuple[GsmNode, GsmNode]] = deque() + + if first_pass or partitioning.remaining_merges is None: + # initialize the merge. this should happen only once: + # - in the first pass if there is no early verdict + # - in the second pass if there is an early verdict + assert first_pass == (partitioning.score is SpecialScores.NoScore) + + # rewire the blue node's parent + blue_parent = update_partition(blue.predecessor, None) + blue_in_sym, blue_out_sym = blue.prefix_access_pair + get_partition_trans(blue_parent, blue_in_sym)[blue_out_sym] = red + + # create a partition for the red node and check, whether the new output data is available + partition = update_partition(red, None) + if self.output_behavior == "moore": + partition.resolve_unknown_prefix_output(blue_out_sym) + + # initialize the work queue to the initial merge pair + q.append((red, blue)) + else: + # work on the remaining merges + q.extend(partitioning.remaining_merges) + partitioning.nr_merged_states -= len(partitioning.remaining_merges) # loop over implied merges - q: deque[tuple[GsmNode, GsmNode]] = deque([(red, blue)]) pop = q.pop if self.depth_first else q.popleft while len(q) != 0: red, blue = pop() partition = update_partition(red, blue) - - if not self.compatibility_on_futures: - if self.compute_local_compatibility(partition, blue) is False: - return partitioning + partitioning.nr_merged_states += 1 + + if first_pass: + if not self.score_calc.override_default_checks(): + moore_violated = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(partition, blue) + det_violated = self.transition_behavior == "deterministic" and not GsmNode.deterministic_compatible(partition, blue) + local_compat = not moore_violated and not det_violated and self.score_calc.local_compatibility(partition, blue) + else: + local_compat = self.score_calc.local_compatibility(partition, blue) + if local_compat is False: + partitioning.score = SpecialScores.ImmediateReject + return + if local_compat is None: + partitioning.remaining_merges.append((red, blue)) + continue + + partition.data = self.data_handler.merge(partition.data, blue.data) # create implied merges for all common successors for in_sym, blue_transitions in blue.transitions.items(): - partition_transitions = partition.transitions[in_sym] - for out_sym, blue_transition in blue_transitions.items(): - partition_transition = partition_transitions.get(out_sym) + partition_transitions = get_partition_trans(partition, in_sym) + for out_sym, blue_successor in blue_transitions.items(): + partition_successor = partition_transitions.get(out_sym) # handle unknown output - if partition_transition is None and len(partition_transitions) != 0: + if partition_successor is None and len(partition_transitions) != 0: if out_sym is unknown_output: + # option A: the output is unknown in the added node assert len(partition_transitions) == 1 - partition_transition = list(partition_transitions.values())[0] + partition_successor = list(partition_transitions.values())[0] if unknown_output in partition_transitions: + # option B: the output is unknown in the partition assert len(partition_transitions) == 1 - partition_transition = partition_transitions.pop(unknown_output) - partition_transitions[out_sym] = partition_transition + partition_successor = partition_transitions.pop(unknown_output) + partition_transitions[out_sym] = partition_successor # re-hook access pair - succ_part = update_partition(partition_transition.target, None) + succ_part = update_partition(partition_successor, None) if self.output_behavior == "moore" or succ_part.predecessor is red: succ_part.resolve_unknown_prefix_output(out_sym) # add pairs - if partition_transition is not None: - q.append((partition_transition.target, blue_transition.target)) - partition_transition.count += blue_transition.count + if partition_successor is not None: + q.append((partition_successor, blue_successor)) else: - # blue child is blue after merging if there is a red state in blue's partition - partition_transition = TransitionInfo(blue_transition.target, blue_transition.count, None, 0) - partition_transitions[out_sym] = partition_transition + # blue_successor is blue after merging if the partition is red + if partition in red_partitions: + partitioning.new_blue.append(blue_successor) + # add new transition to partition + partition_transitions[out_sym] = blue_successor # update predecessor of blue child - blue_target_partition = update_partition(blue_transition.target, None) + blue_target_partition = update_partition(blue_successor, None) blue_target_partition.predecessor = red - partitioning.score = self.score_calc.score_function(partitioning.full_mapping) - return partitioning + if first_pass: + partitioning.score = self.score_calc.score_function(partitioning.full_mapping) def run_GSM(data: list, *, output_behavior: OutputBehavior = "moore", transition_behavior: TransitionBehavior = "deterministic", score_calc: ScoreCalculation = None, - pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, - compatibility_on_pta: bool = False, - compatibility_on_futures: bool = False, - node_order: Callable[[GsmNode, GsmNode], bool] = None, - consider_only_min_blue: bool = False, - depth_first: bool = False, - instrumentation: Instrumentation | None = None, - convert: bool = True, - data_format: DataFormat | None = None, - ) -> Automaton | GsmNode: + data_handler: DataHandler = None, + node_order: Callable[[GsmNode], Any] = None, + consider_only_min_blue=False, + depth_first=False, + instrumentation=None, + convert=True, + data_format=None, + ): """ Performs a state merging algorithm in the red-blue framework on provided data. - :param list data: Data used for learning. Recorded behavior of the system. + :param list data: Data used for learning (recorded behavior of the system), or an already-built GsmNode tree. :param OutputBehavior output_behavior: Specifies whether outputs are emitted by states ("moore") or transitions ("mealy"). :param TransitionBehavior transition_behavior: Either "deterministic", "nondeterministic" or "stochastic". :param ScoreCalculation score_calc: A ScoreCalculation object which determines how compatibility and merge scores are calculated. - :param Callable[[GsmNode], GsmNode] pta_preprocessing: A pre-processing function applied to the PTA. :param Callable[[GsmNode], GsmNode] postprocessing: A postprocessing function applied to the learned automaton. - :param bool compatibility_on_pta: Whether compatibility is evaluated on the PTA or the current hypothesis. - :param bool compatibility_on_futures: Whether compatibility is evaluated using the futures of both states or all partition information. - :param Callable[[GsmNode, GsmNode], bool] node_order: Order in which merge candidates are considered. Defaults to short-lex. + :param DataHandler data_handler: IOHandler object governing abstraction and aggregation of data + :param Callable[[GsmNode], Any] node_order: Sorting key which determines the order in which merge candidates are considered. Defaults to insertion order :param bool consider_only_min_blue: Whether to consider merge candidates from all blue nodes or just a single. :param bool depth_first: Whether compatibility is checked depth- or breadth-first. :param Instrumentation | None instrumentation: Instrumentation object for reporting progress or debugging. @@ -429,10 +489,8 @@ def run_GSM(data: list, *, output_behavior=output_behavior, transition_behavior=transition_behavior, score_calc=score_calc, - pta_preprocessing=pta_preprocessing, postprocessing=postprocessing, - compatibility_on_pta=compatibility_on_pta, - compatibility_on_futures=compatibility_on_futures, + data_handler=data_handler, node_order=node_order, consider_only_min_blue=consider_only_min_blue, depth_first=depth_first, diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index d00e08b0bb..f48562f79e 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -1,15 +1,16 @@ # Convenience wrappers around run_GSM implementing well-known passive learning # algorithms: EDSM, k-tails, and Alergia/IoAlergia (with EDSM-style scoring). from collections import defaultdict +from functools import partial from aalpy import DeterministicAutomaton, Onfsm, NDMooreMachine from aalpy.base import Automaton from aalpy.learning_algs.general_passive.GeneralizedStateMerging import run_GSM +from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler from aalpy.learning_algs.general_passive.Instrumentation import ProgressReport -from aalpy.learning_algs.general_passive.GsmNode import GsmNode -from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility, \ - ScoreWithKTail -from aalpy.utils.HelperFunctions import dfa_from_moore +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, unknown_output +from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import SimpleScoreCalculation, ScoreWithKTail, ScoreIOAlergiaWithEDSM +from aalpy.utils.HelperFunctions import dfa_from_moore, mc_format_to_mdp, mc_from_mdp, ensure_input_complete def run_EDSM(data: list, automaton_type: str, input_completeness: str | None = None, @@ -17,7 +18,7 @@ def run_EDSM(data: list, automaton_type: str, input_completeness: str | None = N """ Run Evidence Driven State Merging. - :param list data: sequence of input sequences and corresponding label. Eg. [[(i1,i2,i3, ...), label], ...] + :param list data: sequence of input sequences and corresponding label, e.g. [[(i1,i2,i3, ...), label], ...] :param str automaton_type: either 'dfa', 'mealy', 'moore'. Note that for 'mealy' machine learning, data has to be prefix-closed. :param str | None input_completeness: either None, 'sink_state', or 'self_loop'. If None, learned model could be input incomplete, sink_state will lead all undefined inputs form some state to the sink state, whereas self_loop will simply create @@ -30,21 +31,21 @@ def run_EDSM(data: list, automaton_type: str, input_completeness: str | None = N print_level = ProgressReport(1) if print_info else None - def EDSM_score(part: dict[GsmNode, GsmNode]) -> int: + def _edsm_evidence_score(part: dict[GsmNode, GsmNode]) -> int: reverse_partition = defaultdict(list) for original_node, resulting_node in part.items(): reverse_partition[resulting_node].append(original_node) evidence = 0 for node, contributing_nodes in reverse_partition.items(): - if node.get_prefix_output() is None: + if node.get_prefix_output() is unknown_output: continue # No evidence whatsoever evidence -= 1 # subtract self-comparison for contributing_node in contributing_nodes: - if contributing_node.get_prefix_output() is not None: + if contributing_node.get_prefix_output() is not unknown_output: evidence += 1 return evidence - score = ScoreCalculation(score_function=EDSM_score) + score = SimpleScoreCalculation(local_compatibility=GsmNode.deterministic_compatible, score_function=_edsm_evidence_score) internal_automaton_type = 'moore' if automaton_type != 'mealy' else automaton_type @@ -55,15 +56,7 @@ def EDSM_score(part: dict[GsmNode, GsmNode]) -> int: if automaton_type == 'dfa': learned_model = dfa_from_moore(learned_model) - if not learned_model.is_input_complete(): - if not input_completeness: - if print_info: - print('Warning: Learned Model is not input complete (inputs not defined for all states). ' - 'Consider calling .make_input_complete()') - else: - if print_info: - print(f'Learned model was not input complete. Adapting it with {input_completeness} transitions.') - learned_model.make_input_complete(input_completeness) + ensure_input_complete(learned_model, input_completeness, print_info) return learned_model @@ -91,93 +84,61 @@ def run_k_tails(data: list, automaton_type: str, k: int, input_completeness: str internal_automaton_type = 'moore' if automaton_type != 'mealy' else automaton_type - score = ScoreWithKTail(ScoreCalculation(GsmNode.deterministic_compatible), k) + score = ScoreWithKTail(SimpleScoreCalculation(GsmNode.deterministic_compatible), k) learned_model = run_GSM(data, output_behavior=internal_automaton_type, transition_behavior="nondeterministic", score_calc=score, data_format='io_traces', instrumentation=print_level) - if not learned_model.is_input_complete(): - if not input_completeness: - if print_info: - print('Warning: Learned Model is not input complete (inputs not defined for all states). ' - 'Consider calling .make_input_complete()') - else: - if print_info: - print(f'Learned model was not input complete. Adapting it with {input_completeness} transitions.') - learned_model.make_input_complete(input_completeness) + ensure_input_complete(learned_model, input_completeness, print_info) return learned_model - -def run_Alergia_EDSM(data: list, automaton_type: str, eps: float = 0.05, print_info: bool = False) -> Automaton: +def run_Alergia_GSM(data: list, automaton_type: str, eps: float = 0.05, compat_on_pta_trans: bool = True, compat_on_pta_count: bool = True, edsm: bool = False, print_info: bool = False) -> Automaton: """ - Run IoAlergia with EDSM on provided data. + Run IOAlergia on provided data. Also supports variants that + - use data more extensively than the original + - use EDSM based scoring - :param list data: [[O,(I,O),(I,O)...], [O,(I,O), (I, O)_,...],..,] if learning MDPs, - or [[I,O,I,O...], [I,O_,...],..,] if learning SMMs (I represent input, O output), or [[O, O, O], ...] if + :param list data: [[O,(I,O),(I,O)...], [O,(I,O), (I, O)_,...],...,] if learning MDPs, + or [[I,O,I,O...], [I,O_,...],...,] if learning SMMs (I represent input, O output), or [[O, O, O], ...] if learning Markov chains. Note that when learning MDPs and MCs the first symbol of each entry should be the same (Initial output). - :param float eps: epsilon value if you are using default HoeffdingCompatibility. :param str automaton_type: either 'mdp' if you wish to learn an MDP, or 'smm' if you want to learn stochastic Mealy machine + :param float eps: epsilon value if you are using default HoeffdingCompatibility. + :param compat_on_pta_trans: evaluate compatibility criterion only on transitions present in the respective PTA nodes + :param compat_on_pta_count: evaluate compatibility criterion using counts from the respective PTA nodes + :param bool edsm: enable EDSM based scoring :param bool print_info: default False :return Automaton: A Mc, Mdp or SMM """ - from aalpy.utils.HelperFunctions import mc_format_to_mdp, mc_from_mdp - assert automaton_type in {'mc', 'mdp', 'smm',} + at_types = ['mc', 'mdp', 'smm'] + if automaton_type not in at_types: + raise ValueError(f"automaton_type {automaton_type} not in {at_types}") - print_level = ProgressReport(1) if print_info else None + if not compat_on_pta_trans and compat_on_pta_count: + raise ValueError("compat_on_pta must be set if compat_on_pta_data is") - class IOAlergiaWithEDSM(ScoreCalculation): - """ScoreCalculation combining IoAlergia's Hoeffding compatibility with an EDSM-style evidence score.""" - - def __init__(self, epsilon: float) -> None: - """ - Create an IoAlergia+EDSM score calculation. - - :param float epsilon: Confidence parameter for the Hoeffding compatibility check. - """ - super().__init__() - self.ioa_compatibility = hoeffding_compatibility(epsilon) - self.evidence = 0 - - def reset(self) -> None: - """ - Reset the accumulated evidence counter. - """ - self.evidence = 0 - - def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: - """ - Check local compatibility of two nodes, accumulating evidence for the score function. - - :param GsmNode a: First node. - :param GsmNode b: Second node. - :return bool: True if the nodes are compatible according to the Hoeffding bound. - """ - self.evidence += 1 - return self.ioa_compatibility(a, b) - - def score_function(self, part: dict[GsmNode, GsmNode]) -> int: - """ - Compute the score of a merge partition as the accumulated evidence. - - :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. - :return int: The accumulated evidence count. - """ - return self.evidence + instrumentation = ProgressReport(1) if print_info else None output_behaviour = 'moore' if automaton_type != 'smm' else 'mealy' learning_data = data if automaton_type != 'mc' else mc_format_to_mdp(data) - learned_model = run_GSM(learning_data, output_behavior=output_behaviour, transition_behavior="stochastic", - score_calc=IOAlergiaWithEDSM(eps), - compatibility_on_pta=True, compatibility_on_futures=True, - instrumentation=print_level, data_format='io_traces') + learned_model = run_GSM( + learning_data, + output_behavior=output_behaviour, + transition_behavior="stochastic", + data_handler=CountOnPTADataHandler(), + score_calc=ScoreIOAlergiaWithEDSM(eps, compat_on_pta_trans, compat_on_pta_count, edsm), + instrumentation=instrumentation, + data_format='io_traces', + ) if automaton_type == 'mc': learned_model = mc_from_mdp(learned_model) return learned_model + +run_Alergia_EDSM = partial(run_Alergia_GSM, edsm=True) \ No newline at end of file diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index e2b5a9e372..5ded550722 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -1,21 +1,21 @@ # Generic prefix-tree / observation-tree node structure used by the general passive # (state-merging) learning algorithms, plus conversion to concrete AALpy automaton types. -import functools -import math import pathlib +import warnings from collections import defaultdict from collections.abc import Callable, Iterable, Iterator, Sequence -from functools import total_ordering -from typing import Any, TypeVar +from typing import Any, TypeVar, Generic import pydot -from copy import copy from aalpy.automata import StochasticMealyMachine, StochasticMealyState, MooreState, MooreMachine, NDMooreState, \ NDMooreMachine, Mdp, MdpState, MealyMachine, MealyState, Onfsm, OnfsmState from aalpy.base import Automaton +from aalpy.learning_algs.general_passive.AssociatedData import StochasticData, CountData + Key = TypeVar("Key") Val = TypeVar("Val") +T = TypeVar("T") OutputBehavior = str OutputBehaviorRange = ["moore", "mealy"] @@ -23,9 +23,6 @@ TransitionBehavior = str TransitionBehaviorRange = ["deterministic", "nondeterministic", "stochastic"] -DataFormat = str -DataFormatRange = ["io_traces", "labeled_sequences", "traces", "tree"] - IOPair = tuple[Any, Any] IOTrace = Sequence[IOPair] IOExample = tuple[Sequence[Any], Any] @@ -33,23 +30,31 @@ StateFunction = Callable[['GsmNode'], str] TransitionFunction = Callable[['GsmNode', Any, Any], str] -unknown_output = None # can be set to a special value if required - +unknown_output = object() # can be set to a special value if required +no_op_input = object() +missing = object() -def intersection_iterator(a: dict[Key, Val], b: dict[Key, Val]) -> Iterator[tuple[Key, Val, Val]]: +def intersection_iterator(a: dict[Key, Val], b: dict[Key, Val], sort_by_length: bool = False) -> Iterator[tuple[Key, Val, Val]]: """ Iterate over the key/value pairs that are present in both dictionaries. :param dict[Key, Val] a: First dictionary. :param dict[Key, Val] b: Second dictionary. + :param bool sort_by_length: If set, iterate over the shorter dictionary. :return Iterator[tuple[Key, Val, Val]]: Iterator of (key, value in a, value in b) for keys common to both dicts. """ - missing = object() - for key, a_val in a.items(): - b_val = b.get(key, missing) - if b_val is missing: - continue - yield key, a_val, b_val + if sort_by_length and len(b) < len(a): + for key, b_val in b.items(): + a_val = a.get(key, missing) + if a_val is missing: + continue + yield key, a_val, b_val + else: + for key, a_val in a.items(): + b_val = b.get(key, missing) + if b_val is missing: + continue + yield key, a_val, b_val def union_iterator(a: dict[Key, Val], b: dict[Key, Val], default: Val = None) -> Iterator[tuple[Key, Val, Val]]: @@ -71,83 +76,8 @@ def union_iterator(a: dict[Key, Val], b: dict[Key, Val], default: Val = None) -> yield key, a_val, b_val -# TODO reuse in RPNI -def detect_data_format(data: Any, check_consistency: bool = False, guess: bool = False) -> DataFormat: - """ - Guess the data format of the provided learning data. - - :param Any data: Input data: a GsmNode (tree), or a sequence of traces/examples. - :param bool check_consistency: Whether to check all data points instead of returning as soon as a unique format is found. - :param bool guess: Whether to allow guessing a single format when multiple formats remain ambiguous. - :return DataFormat: The detected data format string (see DataFormatRange). - """ - # The different data formats are - # - "tree": a tree-shaped automaton provided as a GsmNode - # - "io_traces": either - # - Moore traces [[o, (i,o), (i,o), ...], ...] - # - Mealy traces [[(i,o), (i,o), ...], ...] - # - "labeled_sequences": [([i, i, ...], o), ...] - # - "traces": [[o, o, ...], ...] - - if isinstance(data, GsmNode): - return "tree" - - accepted_types = (tuple, list) - - # mapping data formats to compatibility criteria - check_dict = dict( - io_traces=lambda obj: len(obj) <= 1 or all(isinstance(o, accepted_types) and len(o) == 2 for o in obj[1:]), - labeled_sequences=lambda obj: len(obj) == 2 and isinstance(obj[0], accepted_types), - ) - accept_dict = {k: True for k in check_dict} - - if not isinstance(data, accepted_types): - raise ValueError("wrong input format. expected tuple or list.") - if len(data) == 0: - return "io_traces" - - accepted_formats = list(accept_dict.keys()) - for data_point in data: - if not isinstance(data_point, accepted_types): - raise ValueError("wrong input format. expected tuple or list.") - for k, check in check_dict.items(): - accept_dict[k] &= check(data_point) - accepted_formats = [k for k, v in accept_dict.items() if v] - if len(accepted_formats) == 1 and not check_consistency: - return accepted_formats[0] - if len(accepted_formats) == 0: - return "traces" # default to traces - #raise ValueError("invalid or inconsistent data. no options left") - if len(accepted_formats) != 1 and not guess: - raise ValueError("ambiguous data format. data format needs to be specified explicitly.") - return accepted_formats[0] - - -# TODO maybe split this for maintainability (and perfomance?) -class TransitionInfo: - """Stores the current and original (PTA) target node and count for a single transition.""" - - __slots__ = ["target", "count", "original_target", "original_count"] - - def __init__(self, target: 'GsmNode', count: int, original_target: 'GsmNode | None', - original_count: int | None) -> None: - """ - Create a transition info record. - - :param GsmNode target: Current target node of the transition. - :param int count: Current transition count. - :param GsmNode | None original_target: Target node in the original PTA, if any. - :param int | None original_count: Transition count in the original PTA, if any. - """ - self.target: 'GsmNode' = target - self.count: int = count - self.original_target: 'GsmNode' = original_target - self.original_count: int = original_count - - # TODO add custom pickling code that flattens the Node structure in order to circumvent running into recursion issues for large models -@total_ordering -class GsmNode: +class GsmNode(Generic[T]): """ Generic class for observably deterministic automata. @@ -157,39 +87,21 @@ class GsmNode: Transition count is preferred over state count as it allows to easily count transitions for non-tree-shaped automata """ - __slots__ = ['transitions', 'predecessor', 'prefix_access_pair'] + __slots__ = ['transitions', 'predecessor', 'prefix_access_pair', 'data'] - def __init__(self, prefix_access_pair: IOPair, predecessor: 'GsmNode | None' = None) -> None: + def __init__(self, prefix_access_pair: IOPair, predecessor: 'GsmNode[T] | None', data: T): """ Create a node with the given prefix-access pair and predecessor. :param IOPair prefix_access_pair: (input, output) pair leading from the predecessor to this node. :param GsmNode | None predecessor: Predecessor node, or None for the root node. + :param T data: Algorithm-specific data. """ # TODO try single dict - self.transitions: defaultdict[Any, dict[Any, TransitionInfo]] = defaultdict(dict) + self.transitions: defaultdict[Any, dict[Any, GsmNode[T]]] = defaultdict(dict) self.predecessor: GsmNode = predecessor self.prefix_access_pair = prefix_access_pair - - def __lt__(self, other: 'GsmNode', compare_length_only: bool = False) -> bool: - """ - Compare nodes in short-lex order: first by prefix length, then lexicographically by prefix. - - :param GsmNode other: Node to compare against. - :param bool compare_length_only: Whether to only compare based on prefix length. - :return bool: True if self is ordered before other. - """ - own_l, other_l = self.get_prefix_length(), other.get_prefix_length() - if own_l != other_l: - return own_l < other_l - if compare_length_only: - return False - own_p = self.get_prefix() - other_p = other.get_prefix() - try: - return own_p < other_p - except TypeError: - return [str(x) for x in own_p] < [str(x) for x in other_p] + self.data = data # TODO implicit prefixes as currently implemented require O(length) time for prefix calculations (e.g. to determine the minimal blue node) # other options would be to have more efficient explicit prefixes such as shared list representations @@ -222,14 +134,15 @@ def get_prefix_input(self) -> Any: """ return self.prefix_access_pair[0] - def resolve_unknown_prefix_output(self, value: Any) -> None: + def resolve_unknown_prefix_output(self, value): """ Set the prefix output to the given value if it is currently unknown. :param Any value: Output value to assign if the current prefix output is unknown. """ - if self.get_prefix_output() is unknown_output: - self.prefix_access_pair = (self.get_prefix_input(), value) + p_in, p_out = self.prefix_access_pair + if p_out is unknown_output: + self.prefix_access_pair = (p_in, value) def get_prefix(self, include_output: bool = True) -> list[Any]: """ @@ -260,42 +173,24 @@ def get_root(self) -> 'GsmNode': current = current.predecessor return current - def get_or_create_transitions(self, in_sym: Any) -> dict[Any, TransitionInfo]: - """ - Get the transition dictionary for the given input symbol, creating it if necessary. - - :param Any in_sym: Input symbol. - :return dict[Any, TransitionInfo]: Mapping of output symbol to transition info for this input. - """ - t = self.transitions.get(in_sym) - if t is None: - t = dict() - self.transitions[in_sym] = t - return t - - def transition_iterator(self) -> Iterable[tuple[Any, Any, TransitionInfo]]: + def transition_iterator(self) -> Iterable[tuple[Any, Any, 'GsmNode[T]']]: """ Iterate over all outgoing transitions of this node. - :return Iterable[tuple[Any, Any, TransitionInfo]]: Iterable of (input, output, transition info) triples. + :return Iterable[tuple[Any, Any, GsmNode]]: Iterable of (input, output, successor) triples. """ for in_sym, transitions in self.transitions.items(): for out_sym, node in transitions.items(): yield in_sym, out_sym, node - def shallow_copy(self) -> 'GsmNode': + def child_iterator(self) -> Iterable['GsmNode[T]']: """ - Create a shallow copy of this node, duplicating its transition dict but keeping the same targets. + Iterate over all possible successors of this node. - :return GsmNode: The copied node. + :return Iterable[GsmNode[T]]: Iterable of successor nodes. """ - node = GsmNode(self.prefix_access_pair, self.predecessor) - for in_sym, t in self.transitions.items(): - d = dict() # appears to be faster than dict comprehension - for out_sym, ti in t.items(): - d[out_sym] = TransitionInfo(ti.target, ti.count, ti.original_target, ti.original_count) - node.transitions[in_sym] = d - return node + for transitions in self.transitions.values(): + yield from transitions.values() def get_by_prefix(self, seq: IOTrace) -> 'GsmNode | None': """ @@ -306,18 +201,19 @@ def get_by_prefix(self, seq: IOTrace) -> 'GsmNode | None': """ node: GsmNode = self for in_sym, out_sym in seq: - if in_sym is None: # ignore initial transition of Node.get_prefix() + if in_sym is no_op_input: # ignore noops (e.g. initial transition of Node.get_prefix()) continue trans = node.transitions.get(in_sym) if trans is None: return None - t_info = trans.get(out_sym) - if t_info is None: - return None - node = t_info.target + node = trans.get(out_sym) + if node is None: + node = trans.get(unknown_output) + if node is None: + return None return node - def get_all_nodes(self) -> list['GsmNode']: + def get_all_nodes(self) -> list['GsmNode[T]']: """ Collect all nodes reachable from this node (including itself). @@ -326,8 +222,7 @@ def get_all_nodes(self) -> list['GsmNode']: result = [self] backing_set = {self} for state in result: - for _, _, transition in state.transition_iterator(): - child = transition.target + for child in state.child_iterator(): if child not in backing_set: backing_set.add(child) result.append(child) @@ -343,8 +238,7 @@ def is_tree(self) -> bool: backing_set = {self} while len(q) != 0: current = q.pop(0) - for _, _, transition in current.transition_iterator(): - child = transition.target + for child in current.child_iterator(): if child in backing_set: return False q.append(child) @@ -379,16 +273,25 @@ def to_automaton(self, output_behavior: OutputBehavior, transition_behavior: Tra ("mealy", "stochastic"): (StochasticMealyMachine, StochasticMealyState), } - AutomatonClass, StateClass = type_dict[(output_behavior, transition_behavior)] + automaton_class, state_class = type_dict[(output_behavior, transition_behavior)] + + maybe_input_alphabet = self.transitions.keys() + if any(node.transitions.keys() != maybe_input_alphabet for node in nodes): + warnings.warn("Automaton is not input-complete. Consider calling .make_input_complete().") + + for node in nodes: + if node.get_prefix_output() is unknown_output and (node.predecessor is not None or output_behavior == "moore"): + warnings.warn("Automaton has unknown outputs") + break # create states state_map = dict() for i, node in enumerate(nodes): state_id = f's{i}' if output_behavior == "mealy": - state = StateClass(state_id) + state = state_class(state_id) elif output_behavior == "moore": - state = StateClass(state_id, node.get_prefix_output()) + state = state_class(state_id, node.get_prefix_output()) state_map[node] = state if set_prefix: if transition_behavior == "deterministic": @@ -403,26 +306,28 @@ def to_automaton(self, output_behavior: OutputBehavior, transition_behavior: Tra # add transitions for node in nodes: state = state_map[node] + if automaton_class in [Mdp, StochasticMealyMachine]: + if not isinstance(node.data, StochasticData): + raise TypeError(f"No probability in information available for {automaton_class.__name__}") + prob_info = node.data.get_probabilities() for in_sym, transitions in node.transitions.items(): - total = sum(t.count for t in transitions.values()) for out_sym, target_node in transitions.items(): - target_state = state_map[target_node.target] - count = target_node.count - if AutomatonClass is MooreMachine: + target_state = state_map[target_node] + if automaton_class is MooreMachine: state.transitions[in_sym] = target_state - elif AutomatonClass is MealyMachine: + elif automaton_class is MealyMachine: state.transitions[in_sym] = target_state state.output_fun[in_sym] = out_sym - elif AutomatonClass is NDMooreMachine: + elif automaton_class is NDMooreMachine: state.transitions[in_sym].append(target_state) - elif AutomatonClass is Onfsm: + elif automaton_class is Onfsm: state.transitions[in_sym].append((out_sym, target_state)) - elif AutomatonClass is Mdp: - state.transitions[in_sym].append((target_state, count / total)) - elif AutomatonClass is StochasticMealyMachine: - state.transitions[in_sym].append((target_state, out_sym, count / total)) + elif automaton_class is Mdp: + state.transitions[in_sym].append((target_state, prob_info[in_sym][out_sym])) + elif automaton_class is StochasticMealyMachine: + state.transitions[in_sym].append((target_state, out_sym, prob_info[in_sym][out_sym])) - return AutomatonClass(initial_state, list(state_map.values())) + return automaton_class(initial_state, list(state_map.values())) def visualize(self, path: str | pathlib.Path, output_behavior: OutputBehavior = "mealy", format: str = "dot", engine: str = "dot", *, @@ -455,19 +360,23 @@ def visualize(self, path: str | pathlib.Path, output_behavior: OutputBehavior = if trans_props is None: trans_props = dict() if state_label is None: - if output_behavior == "moore": - def state_label(node: GsmNode) -> str: - return f'{node.get_prefix_output()} {node.count()}' - else: - def state_label(node: GsmNode) -> str: - return f'{sum(t.count for _, _, t in node.transition_iterator())}' + def state_label(node: GsmNode) -> str: + label_parts = [] + if output_behavior == "moore": + label_parts.append(str(node.get_prefix_output())) + if isinstance(node.data, CountData): + label_parts.append(str(node.data.count())) + return " ".join(label_parts) if trans_label is None and "label" not in trans_props: - if output_behavior == "moore": - def trans_label(node: GsmNode, in_sym: Any, out_sym: Any) -> str: - return f'{in_sym} [{node.transitions[in_sym][out_sym].count}]' - else: - def trans_label(node: GsmNode, in_sym: Any, out_sym: Any) -> str: - return f'{in_sym} / {out_sym} [{node.transitions[in_sym][out_sym].count}]' + def trans_label(node: GsmNode, in_sym: Any, out_sym: Any) -> str: + label_parts = [] + if output_behavior == "moore": + label_parts.append(str(in_sym)) + else: + label_parts.append(f'{in_sym} / {out_sym}') + if isinstance(node.data, CountData): + label_parts.append(f'[{node.data.transition_count[in_sym][out_sym]}]') + return " ".join(label_parts) if state_color is None: def state_color(x: 'GsmNode') -> str: return "black" if trans_color is None: @@ -498,7 +407,7 @@ def node_naming(node: GsmNode) -> str: for in_sym, options in node.transitions.items(): for out_sym, c in options.items(): arg_dict = {key: fun(node, in_sym, out_sym) for key, fun in trans_props.items()} - graph.add_edge(pydot.Edge(node_naming(node), node_naming(c.target), **arg_dict)) + graph.add_edge(pydot.Edge(node_naming(node), node_naming(c), **arg_dict)) # add initial state # TODO maybe add option to parameterize this @@ -512,12 +421,18 @@ def node_naming(node: GsmNode) -> str: file_ext = 'dot' graph.write(path=str(path) + "." + file_ext, prog=engine, format=format) - def make_input_complete(self) -> list[tuple['GsmNode', Any, Any]]: + def make_input_complete(self, target: 'GsmNode[T] | str' = "self_loop") -> list[tuple['GsmNode', Any, Any]]: """ - Add self-looping transitions for any input undefined at some node, using the node's prefix output. + For all reachable nodes, add transitions for all undefined inputs. The output is set using the targets prefix output. + This function DOES NOT touch the `data` field of affected nodes. Updating this is in the responsibility of the caller. + :param GsmNode[T] | str target: Target node of missing transitions. The special value "self_loop" adds self transitions. :return list[tuple[GsmNode, Any, Any]]: List of (node, input, output) triples for the added transitions. """ + + if isinstance(target, str) and target != "self_loop": + raise ValueError(f"Invalid target {target}. Should be either 'self_loop' or a GsmNode.") + all_nodes = self.get_all_nodes() inputs = {in_sym for node in all_nodes for in_sym in node.transitions} missing_trans = [] @@ -525,103 +440,15 @@ def make_input_complete(self) -> list[tuple['GsmNode', Any, Any]]: for in_sym in inputs: transitions = node.transitions[in_sym] if len(transitions) == 0: - out_sym = node.prefix_access_pair[1] + if target == "self_loop": + successor = node + else: + successor = target + out_sym = successor.prefix_access_pair[1] missing_trans.append((node, in_sym, out_sym)) - t_info = TransitionInfo(node, 1, None, None) - transitions[out_sym] = t_info + transitions[out_sym] = successor return missing_trans - def add_trace(self, trace: IOTrace) -> None: - """ - Add an IO trace to the tree rooted at this node, extending it with new nodes as necessary. - - :param IOTrace trace: Sequence of (input, output) pairs to add. - """ - curr_node: GsmNode = self - for in_sym, out_sym in trace: - transitions = curr_node.transitions[in_sym] - info = transitions.get(out_sym) - if info is None: - node = GsmNode((in_sym, out_sym), curr_node) - transitions[out_sym] = TransitionInfo(node, 1, node, 1) - else: - info.count += 1 - info.original_count += 1 - node = info.target - curr_node = node - - def add_labeled_sequence(self, example: IOExample) -> None: - """ - Add a labeled input sequence (inputs with a single label attached at the end) to the tree. - - :param IOExample example: (inputs, output) pair, where output labels the state reached by inputs. - """ - inputs, output = example - curr_node: GsmNode = self - in_sym = None - - # step through inputs and add transitions - for in_sym in inputs: - transitions = curr_node.transitions[in_sym] - t_infos = list(transitions.values()) - if len(t_infos) == 0: - node = GsmNode((in_sym, unknown_output), curr_node) - t_info = TransitionInfo(node, 1, node, 1) - transitions[unknown_output] = t_info - elif len(t_infos) == 1: - t_info = t_infos[0] - t_info.count += 1 - t_info.original_count += 1 - node = t_info.target - else: - # This should never happen - raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") - curr_node = node - - # set last output - curr_node.resolve_unknown_prefix_output(output) - pred = curr_node.predecessor - if pred: - transitions = pred.transitions[in_sym] - if unknown_output in transitions: - transitions[output] = transitions.pop(unknown_output) - if output not in transitions: - raise ValueError("nondeterminism encountered for GSM with labeled_sequences. not supported") - - @staticmethod - def createPTA(data: Any, output_behavior: OutputBehavior, data_format: DataFormat | None = None) -> 'GsmNode': - """ - Build a prefix tree acceptor (PTA) from the given data. - - :param Any data: Learning data, in one of the supported data formats (or already a GsmNode tree). - :param OutputBehavior output_behavior: Either "moore" or "mealy". - :param DataFormat | None data_format: Explicit data format, or None to auto-detect. - :return GsmNode: The root node of the constructed (or passed-through) PTA. - """ - if data_format is None: - data_format = detect_data_format(data) - if data_format not in DataFormatRange: - raise ValueError(f"invalid data format {data_format}. should be in {DataFormatRange}") - - if data_format == "tree": - if not data.is_tree(): - raise ValueError("provided automaton is not a tree") - return data - root_node = GsmNode((None, unknown_output), None) - if data_format == "labeled_sequences": - for example in data: - root_node.add_labeled_sequence(example) - if data_format == "io_traces" or data_format == "traces": - if output_behavior == "moore": - initial_output = data[0][0] - root_node.prefix_access_pair = (None, initial_output) - data = (d[1:] for d in data) - for trace in data: - if data_format == "traces": - trace = (("step", t) for t in trace) - root_node.add_trace(trace) - return root_node - def is_locally_deterministic(self) -> bool: """ Check whether this node has at most one outgoing transition per input symbol. @@ -659,8 +486,8 @@ def is_moore(self) -> bool: :return bool: True if the structure is Moore-compatible. """ for node in self.get_all_nodes(): - for in_sym, out_sym, transition in node.transition_iterator(): - child_output = transition.target.get_prefix_output() + for in_sym, out_sym, next_node in node.transition_iterator(): + child_output = next_node.get_prefix_output() if out_sym is not unknown_output and child_output != out_sym: return False return True @@ -676,28 +503,21 @@ def moore_compatible(self, other: 'GsmNode') -> bool: oo = other.get_prefix_output() return so == oo or so is unknown_output or oo is unknown_output - def local_log_likelihood_contribution(self) -> float: + def short_lex_order(self, other: 'GsmNode', compare_length_only: bool = False): """ - Compute this node's contribution to the log-likelihood of the data given the model. + Compute the short-lex order of the two nodes. - :return float: The local log-likelihood contribution. + :param GsmNode other: Node to compare with. + :return bool: whether the prefix of `self` is smaller than the prefix of `other` according to short-lex order """ - llc = 0 - for in_sym, trans in self.transitions.items(): - total_count = 0 - for out_sym, info in trans.items(): - total_count += info.count - llc += info.count * math.log(info.count) - if total_count != 0: - llc -= total_count * math.log(total_count) - return llc - - def count(self) -> int: - """ - Compute the total transition count over all outgoing transitions of this node. - - :return int: Sum of transition counts. - """ - return sum(trans.count for _, _, trans in self.transition_iterator()) - - default_order = functools.cmp_to_key(lambda a, b: -1 if a < b else 1) + own_l, other_l = self.get_prefix_length(), other.get_prefix_length() + if own_l != other_l: + return own_l < other_l + if compare_length_only: + return False + own_p = self.get_prefix() + other_p = other.get_prefix() + try: + return own_p < other_p + except TypeError: + return [str(x) for x in own_p] < [str(x) for x in other_p] diff --git a/aalpy/learning_algs/general_passive/Instrumentation.py b/aalpy/learning_algs/general_passive/Instrumentation.py index e854216975..33d08bed0d 100644 --- a/aalpy/learning_algs/general_passive/Instrumentation.py +++ b/aalpy/learning_algs/general_passive/Instrumentation.py @@ -2,11 +2,10 @@ # progress reporting and a debugging helper that checks merges/promotions against ground truth. from time import perf_counter -from aalpy.learning_algs.general_passive.GeneralizedStateMerging import Instrumentation, Partitioning, \ - GeneralizedStateMerging +from aalpy.learning_algs.general_passive.GeneralizedStateMerging import Partitioning, GeneralizedStateMerging, \ + Instrumentation from aalpy.learning_algs.general_passive.GsmNode import GsmNode - class ProgressReport(Instrumentation): """Instrumentation that prints progress information (timing, state/merge counts) during learning.""" @@ -72,8 +71,11 @@ def print_status(self) -> None: """ reset_char = "\33[2K\r" print_str = reset_char + f'Current automaton size: {self.nr_red_states}' - if 0 < self.lvl and not self.gsm.compatibility_on_futures: - print_str += f' Merged: {self.nr_merged_states_total} Remaining: {self.pta_size - self.nr_red_states - self.nr_merged_states_total}' + if 0 < self.lvl: + time_taken = round(perf_counter() - self.previous_time, 2) + mps = round(self.nr_merged_states_total / time_taken, 2) if time_taken != 0 else "inf" + remaining_merges = self.pta_size - self.nr_red_states - self.nr_merged_states_total + print_str += f' Merged: {self.nr_merged_states_total} Remaining: {remaining_merges} ({time_taken} s -> {mps} merges / second)' print(print_str, end="") def log_promote(self, node: GsmNode) -> None: @@ -93,7 +95,7 @@ def log_merge(self, part: Partitioning) -> None: :param Partitioning part: The partitioning describing the performed merge. """ self.log.append(["merge", (part.red.get_prefix(), part.blue.get_prefix())]) - self.nr_merged_states_total += len(part.full_mapping) - len(part.red_mapping) + self.nr_merged_states_total += part.nr_merged_states self.nr_merged_states += 1 self.print_status() @@ -122,7 +124,7 @@ def __init__(self, ground_truth: GsmNode) -> None: :param GsmNode ground_truth: Root node of the ground-truth model. """ super().__init__() - self.root = ground_truth + self.ground_truth_root = ground_truth self.map: dict[GsmNode, GsmNode] = dict() self.log = [] self.gsm: GeneralizedStateMerging | None = None @@ -144,14 +146,16 @@ def log_promote(self, new_red: GsmNode) -> None: :param GsmNode new_red: The promoted node. """ new_red_prefix = new_red.get_prefix() - node = self.root.get_by_prefix(new_red_prefix) - old_red = self.map.get(node) - if old_red is None: - self.map[node] = new_red + gt_node = self.ground_truth_root.get_by_prefix(new_red_prefix) + old_red = self.map.get(gt_node) + if gt_node is None: + self.log.append(("broken promote", new_red_prefix)) + elif old_red is None: + self.map[gt_node] = new_red self.log.append(("promote", new_red_prefix)) elif old_red is not new_red: print(f"Erroneous promotion detected:") - print(f" Ground truth: {node.get_prefix()}") + print(f" Ground truth: {gt_node.get_prefix()}") print(f" Representative (old): {old_red.get_prefix()}") print(f" Representative (new): {new_red_prefix}") self.log.append(("wrong promote", new_red_prefix)) @@ -164,8 +168,8 @@ def log_merge(self, part: Partitioning) -> None: """ red_prefix = part.red.get_prefix() blue_prefix = part.blue.get_prefix() - red_node = self.root.get_by_prefix(red_prefix) - blue_node = self.root.get_by_prefix(blue_prefix) + red_node = self.ground_truth_root.get_by_prefix(red_prefix) + blue_node = self.ground_truth_root.get_by_prefix(blue_prefix) if red_node is None or blue_node is None: self.log.append(("broken merge", red_prefix, blue_prefix)) elif red_node is blue_node: diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 473573c991..c92d5381c7 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -1,78 +1,133 @@ # Score/compatibility function building blocks used to guide the general passive # state-merging algorithm (local compatibility checks and global merge scores). +from abc import ABC +from collections import deque from collections.abc import Callable, Iterable +from functools import total_ordering from math import sqrt, log from typing import Any -from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, TransitionInfo +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, CountData +from aalpy.learning_algs.general_passive.AssociatedData import ShadowPTAData -LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool] +LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool | None] ScoreFunction = Callable[[dict[GsmNode, GsmNode]], Any] AggregationFunction = Callable[[Iterable], Any] -class ScoreCalculation: +class SpecialScores: + @total_ordering + class _SpecialScore: + def __init__(self, ideal: bool): + self.ideal = ideal + + def __lt__(self, other): + return not self.ideal + + def __bool__(self): + return self.ideal + + ImmediateAccept = _SpecialScore(True) + ImmediateReject = _SpecialScore(False) + NoScore = None + +class ScoreCalculation(ABC): """Bundles a local compatibility check and a global score function used during state merging.""" - def __init__(self, local_compatibility: LocalCompatibilityFunction = None, - score_function: ScoreFunction = None) -> None: + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: """ - Create a score calculation, optionally overriding the default (accept-everything) behavior. + Callback at the beginning of the evaluation of a merge candidate. This function can be used to compute a score + of the merge candidate before constructing the partitioning. - :param LocalCompatibilityFunction local_compatibility: Function determining local compatibility of two nodes. - :param ScoreFunction score_function: Function computing the score of a full merge partition. + :param GsmNode red: GsmNode representing the red node of the merge candidate. + :param GsmNode blue: GsmNode representing the blue node of the merge candidate. + :return: Either an early score for the merge candidate or `None`. """ - # This is a hack that gives a simple implementation where we can easily - determine whether the default is - # overridden (for optimization) - override behavior in a functional way by providing the functions as - # arguments (no extra class) - override behavior in a stateful way by implementing a new class that provides - # `local_compatibility` and / or `score_function` methods - if not hasattr(self, "local_compatibility"): - self.local_compatibility: LocalCompatibilityFunction = local_compatibility or self.default_local_compatibility - if not hasattr(self, "score_function"): - self.score_function: ScoreFunction = score_function or self.default_score_function + return None - def reset(self) -> None: + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): """ - Reset any internal state before starting a new learning run. No-op by default. + Callback at the beginning of the evaluation of a merge candidate. It is only called when no early score is present. + + :param GsmNode red: GsmNode representing the red node of the merge candidate. + :param GsmNode blue: GsmNode representing the blue node of the merge candidate. + :param bool first_pass: Whether this is the first pass (in which the partitioning is only partially constructed) + or the second pass (in which the partitioning is completed) + :return: Either an early score for the merge candidate or `None`. """ pass - @staticmethod - def default_local_compatibility(a: GsmNode, b: GsmNode) -> bool: + def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool | None: """ - Default local compatibility check: always compatible. + Computes whether two `GsmNode` are locally compatible. It is called during partition construction. If not overridden, + any two nodes are considered compatible, unless `output_behavior` is set to `"moore"`. Overriding allows rejecting + a merge candidate early, without having to construct the full partitioning. - :param GsmNode a: First node. - :param GsmNode b: Second node. - :return bool: Always True. + :param GsmNode a: The node corresponding to the current (=partial) partition of a node. + :param GsmNode b: The node to be merged into the partition. + :return bool | None: Whether the two `GsmNode` are locally compatible. Returns `None` if no further descendants + need to be considered to assess the final verdict / score. """ return True - @staticmethod - def default_score_function(part: dict[GsmNode, GsmNode]) -> bool: + def score_function(self, part: dict[GsmNode, GsmNode]) -> Any: """ - Default score function: any partition is acceptable. + Computes the score of a merge candidate based on the partitioning resulting from implied merges (determinization) + starting from the original merge candidate. - :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. - :return bool: Always True. + :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. + :return Any: The score of the merge candidate. Special values are: + - SpecialScores.ImmediateAccept: the merge candidate should be merged without considering others. + - SpecialScores.ImmediateReject: the merge candidate should not be considered. + Default is immediate acceptance. """ - return True + return SpecialScores.ImmediateAccept - def has_score_function(self) -> bool: + def promotion_score(self, promotion_candidate: GsmNode) -> Any: """ - Check whether a non-default score function is configured. + Computes the score of a promotion candidate. By default, promotion candidates are immediately accepted. Override + this function to implement promotion scoring. - :return bool: True if score_function was overridden. + :param GsmNode promotion_candidate: GsmNode which is to be promoted. + :return Any: The score of the promotion candidate. Default is ImmediateAccept. """ - return self.score_function is not self.default_score_function + return SpecialScores.ImmediateAccept def has_local_compatibility(self) -> bool: """ - Check whether a non-default local compatibility function is configured. + Check whether a non-default local compatibility is configured. :return bool: True if local_compatibility was overridden. """ - return self.local_compatibility is not self.default_local_compatibility + return self.__class__.local_compatibility is not ScoreCalculation.local_compatibility + + def has_score_function(self) -> bool: + """ + Check whether a non-default score function is configured. + + :return bool: True if score_function was overridden. + """ + return self.__class__.score_function is not ScoreCalculation.score_function + + def override_default_checks(self) -> bool: + """ + Determines whether default compatibility checks for "Moore-ness" and determinism should be performed. + """ + return False + + +class SimpleScoreCalculation(ScoreCalculation): + def __init__(self, local_compatibility: LocalCompatibilityFunction = None, score_function: ScoreFunction = None) -> None: + self.local_compatibility = local_compatibility or self.local_compatibility + self._has_local_compatibility = local_compatibility is not None + self.score_function = score_function or self.score_function + self._has_score_function = score_function is not None + + def has_local_compatibility(self) -> bool: + return self._has_local_compatibility + + def has_score_function(self) -> bool: + return self._has_score_function def hoeffding_compatibility(eps: float, compare_original: bool = True) -> LocalCompatibilityFunction: @@ -84,20 +139,24 @@ def hoeffding_compatibility(eps: float, compare_original: bool = True) -> LocalC :return LocalCompatibilityFunction: Function checking whether two nodes' output distributions are compatible. """ eps_fact = sqrt(0.5 * log(2 / eps)) - count_name = "original_count" if compare_original else "count" - transition_dummy = TransitionInfo(None, 0, None, 0) - def similar(a: GsmNode, b: GsmNode) -> bool: + def similar(a: GsmNode[CountData], b: GsmNode[CountData]) -> bool: # iterate over inputs that are common to both states - for in_sym, a_trans, b_trans in intersection_iterator(a.transitions, b.transitions): + if compare_original: + a_dict = a.data.pta_count + b_dict = b.data.pta_count + else: + a_dict = a.data.transition_count + b_dict = b.data.transition_count + + for in_sym, a_trans, b_trans in intersection_iterator(a_dict, b_dict, True): # could create appropriate dict here - a_total, b_total = (sum(getattr(x, count_name) for x in trans.values()) for trans in (a_trans, b_trans)) + a_total, b_total = (sum(trans.values()) for trans in (a_trans, b_trans)) if a_total == 0 or b_total == 0: continue # parameter combinations require this check threshold = eps_fact * (sqrt(1 / a_total) + sqrt(1 / b_total)) # iterate over outputs that appear in either distribution - for out_sym, a_info, b_info in union_iterator(a_trans, b_trans, transition_dummy): - ac, bc = (getattr(x, count_name) for x in (a_info, b_info)) + for out_sym, ac, bc in union_iterator(a_trans, b_trans, 0): if abs(ac / a_total - bc / b_total) > threshold: return False return True @@ -105,30 +164,131 @@ def similar(a: GsmNode, b: GsmNode) -> bool: return similar -class ScoreWithKTail(ScoreCalculation): +class SimpleFutureBasedCompatibility(ScoreCalculation): + """ + ScoreCalculation without scoring that checks local compatibility only on common futures (as in Alergia) and not + during the construction of the partitioning. This avoids the need to construct the partitioning in a reversible manner, + which results in a significant speedup. + """ + def __init__(self, + compatibility_on_pta = False, + depth_first = False, + local_compatibility: LocalCompatibilityFunction = None, + ): + """ + Create a new CheckFutureScore instance. + + :param bool compatibility_on_pta: Whether compatibility should be checked on the PTA or the partially merged automaton + :param bool depth_first: Whether to traverse the implied merges DFS or BFS. Defaults to True (BFS). + :param LocalCompatibilityFunction local_compatibility: Compatibility criterion used to check futures. + """ + if local_compatibility: + if self.has_local_compatibility(): + raise ValueError("External local compatibility is provided, but the class already defines a local compatibility criterion.") + self.local_compatibility = local_compatibility or self.local_compatibility + self.compatibility_on_pta = compatibility_on_pta + self.depth_first = depth_first + + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: + if self.compatibility_on_pta and not isinstance(red.data, ShadowPTAData): + raise TypeError("compatibility_on_pta is set but no PTA data is available") + + q: deque[tuple[GsmNode, GsmNode]] = deque([(red, blue)]) + pop = q.pop if self.depth_first else q.popleft + + while len(q) != 0: + red, blue = pop() + + local_compatibility = self.local_compatibility(red, blue) + if local_compatibility is False: + return SpecialScores.ImmediateReject + if local_compatibility is None: + continue + + if self.compatibility_on_pta: + red_data: ShadowPTAData = red.data + blue_data: ShadowPTAData = blue.data + for in_sym, red_trans, blue_trans in intersection_iterator(red_data.shadow_pta, blue_data.shadow_pta, True): + for out_sym, red_child, blue_child in intersection_iterator(red_trans, blue_trans): + q.append((red_child,blue_child)) + else: + for in_sym, red_trans, blue_trans in intersection_iterator(red.transitions, blue.transitions, True): + for out_sym, red_child, blue_child in intersection_iterator(red_trans, blue_trans): + q.append((red_child, blue_child)) + + return SpecialScores.ImmediateAccept + + +class ScoreIOAlergiaWithEDSM(SimpleFutureBasedCompatibility): + def __init__(self, eps: float, compat_on_pta: bool, compat_on_pta_data: bool, edsm: bool): + self.compat = hoeffding_compatibility(eps, compat_on_pta_data) + SimpleFutureBasedCompatibility.__init__(self, compatibility_on_pta=compat_on_pta) + self.edsm = edsm + self.score = None + + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: + self.score = 0 + verdict = super().early_score(red, blue) + if self.edsm is False or verdict is SpecialScores.ImmediateReject: + return verdict + return self.score + + def local_compatibility(self, red: GsmNode, blue: GsmNode) -> float: + self.score += 1 + return self.compat(red, blue) + + +class WrappingScore(ScoreCalculation, ABC): + """Baseclass for wrapping `ScoreCalculation` objects with minor changes.""" + + def __init__(self, wrapped: ScoreCalculation): + self.wrapped = wrapped + # TODO could detect overrides and hardlink to methods of wrapped score otherwise. see below. + # if not hasattr(self, "initialized_merge"): + # self.initialized_merge = wrapped.initialize_merge + + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: + return self.wrapped.early_score(red, blue) + + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): + self.wrapped.initialize_merge(red, blue, first_pass) + + def local_compatibility(self, red: GsmNode, blue: GsmNode) -> bool | None: + return self.wrapped.local_compatibility(red, blue) + + def score_function(self, part: dict[GsmNode, GsmNode]) -> Any: + return self.wrapped.score_function(part) + + def promotion_score(self, promotion_candidate: GsmNode) -> Any: + return self.wrapped.promotion_score(promotion_candidate) + + def has_local_compatibility(self) -> bool: + return self.wrapped.has_local_compatibility() + + def has_score_function(self) -> bool: + return self.wrapped.has_score_function() + + +class ScoreWithKTail(WrappingScore): """Applies k-Tails to a compatibility function: Compatibility is only evaluated up to a certain depth k.""" - def __init__(self, other_score: ScoreCalculation, k: int) -> None: + def __init__(self, wrapped: ScoreCalculation, k: int) -> None: """ Wrap another score calculation, limiting local compatibility checks to depth k. - :param ScoreCalculation other_score: Score calculation to delegate to within depth k. + :param ScoreCalculation wrapped: Score calculation to delegate to within depth k. :param int k: Maximum depth (relative to the blue node's initial depth) at which compatibility is checked. """ - super().__init__(None, other_score.score_function) - self.other_score = other_score + super().__init__(wrapped) self.k = k self.depth_offset = None - def reset(self) -> None: - """ - Reset the wrapped score calculation and the depth offset. - """ - self.other_score.reset() - self.depth_offset = None + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): + self.depth_offset = blue.get_prefix_length() + self.wrapped.initialize_merge(red, blue, first_pass) - def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: + def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool | None: """ Check local compatibility, treating nodes beyond depth k as automatically compatible. @@ -137,57 +297,34 @@ def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: :return bool: True if compatible (or beyond depth k), False otherwise. """ # assuming b is tree shaped. - if self.depth_offset is None: - self.depth_offset = b.get_prefix_length() depth = b.get_prefix_length() - self.depth_offset if self.k <= depth: - return True + return None - return self.other_score.local_compatibility(a, b) + return self.wrapped.local_compatibility(a, b) -class ScoreWithSinks(ScoreCalculation): +class ScoreWithSinks(WrappingScore): """This class allows rejecting merge candidates based on additional criteria for the initial merge""" - def __init__(self, other_score: ScoreCalculation, sink_cond: Callable[[GsmNode], bool], + def __init__(self, wrapped: ScoreCalculation, sink_cond: Callable[[GsmNode], bool], allow_sink_merge: bool = True) -> None: """ - Wrap another score calculation, additionally rejecting merges involving "sink" nodes. + Wrapped score calculation, additionally rejecting merges involving "sink" nodes. - :param ScoreCalculation other_score: Score calculation to delegate to. + :param ScoreCalculation wrapped: Score calculation to delegate to. :param Callable[[GsmNode], bool] sink_cond: Predicate identifying sink nodes. :param bool allow_sink_merge: Whether merges between two sink nodes are allowed. """ - super().__init__(None, other_score.score_function) - self.other_score = other_score + super().__init__(wrapped) self.sink_cond = sink_cond self.allow_sink_merge = allow_sink_merge - self.is_first = True - - def reset(self) -> None: - """ - Reset the wrapped score calculation and the first-call flag. - """ - self.other_score.reset() - self.is_first = True - - def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: - """ - Check local compatibility, additionally applying the sink condition on the first call. - - :param GsmNode a: First (red) node. - :param GsmNode b: Second (blue) node. - :return bool: True if compatible according to the sink condition and the wrapped score calculation. - """ - if self.is_first: - self.is_first = False - a_sink, b_sink = self.sink_cond(a), self.sink_cond(b) - if a_sink != b_sink: - return False - if a_sink and b_sink and not self.allow_sink_merge: - return False - return self.other_score.local_compatibility(a, b) + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: + a_sink, b_sink = self.sink_cond(red), self.sink_cond(blue) + if (a_sink or b_sink) and not (a_sink and b_sink and self.allow_sink_merge): + return SpecialScores.ImmediateReject + return self.wrapped.early_score(red, blue) class ScoreCombinator(ScoreCalculation): @@ -205,17 +342,17 @@ def __init__(self, scores: list[ScoreCalculation], aggregate_compatibility: Aggr :param AggregationFunction aggregate_compatibility: Function aggregating the individual compatibility results. :param AggregationFunction aggregate_score: Function aggregating the individual score results. """ - super().__init__() self.scores = scores self.aggregate_compatibility = aggregate_compatibility or self.default_aggregate_compatibility self.aggregate_score = aggregate_score or self.default_aggregate_score - def reset(self) -> None: - """ - Reset all combined score calculations. - """ + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): for score in self.scores: - score.reset() + score.initialize_merge(red, blue, first_pass) + + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: + scores = [score.early_score(red, blue) for score in self.scores] + return self.aggregate_score(scores) def local_compatibility(self, a: GsmNode, b: GsmNode) -> Any: """ @@ -229,36 +366,55 @@ def local_compatibility(self, a: GsmNode, b: GsmNode) -> Any: def score_function(self, part: dict[GsmNode, GsmNode]) -> Any: """ - Compute the aggregated score of a merge partition over all combined score calculations. + Compute the aggregated score of a merge candidate over all combined score calculations. :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. :return Any: Aggregated score result. """ return self.aggregate_score(score.score_function(part) for score in self.scores) + def promotion_score(self, promotion_candidate: GsmNode) -> Any: + """ + Compute the aggregated score of a promotion over all combined score calculations. + + :param GsmNode promotion_candidate: Node to be promoted. + :return Any: Aggregated score result. + """ + return self.aggregate_score(score.promotion_score(promotion_candidate) for score in self.scores) + @staticmethod def default_aggregate_compatibility(compatibility_iterable: Iterable) -> Any: """ - Commits to the first value that is not inconclusive (== None). Accepts if in doubt. + Returns the least permissive among provided values (False < True < None). :param Iterable compatibility_iterable: Iterable of compatibility results. :return Any: The first non-None result, or True if all are None. """ + highest_value = None for compat in compatibility_iterable: - if compat is None: - continue - return compat - return True + if compat is False: + return False + if compat is True: + highest_value = True + return highest_value @staticmethod - def default_aggregate_score(score_iterable: Iterable) -> list: + def default_aggregate_score(score_iterable: Iterable) -> Any: """ - Default score aggregation: collect all scores into a list. + Default score aggregation: collect all scores into a list, unless a special value decides the outcome. + Rejection wins over anything else and a single undecided score leaves the aggregate undecided, whereas + acceptance has to be unanimous, since a list mixing special values is not a meaningful score. :param Iterable score_iterable: Iterable of score results. - :return list: List of the individual scores. + :return Any: The deciding special value, or the list of the individual scores. """ - return list(score_iterable) + scores = list(score_iterable) + for special in (SpecialScores.ImmediateReject, SpecialScores.NoScore): + if any(score is special for score in scores): + return special + if scores and all(score is SpecialScores.ImmediateAccept for score in scores): + return SpecialScores.ImmediateAccept + return scores def local_to_global_compatibility(local_fun: LocalCompatibilityFunction) -> ScoreFunction: @@ -269,64 +425,49 @@ def local_to_global_compatibility(local_fun: LocalCompatibilityFunction) -> Scor partition, original. :param LocalCompatibilityFunction local_fun: Local compatibility function to lift to a global score function. - :return ScoreFunction: Global score function returning False if any local check fails, True otherwise. + :return ScoreFunction: Global score function rejecting if any local check fails and greedily accepting otherwise. """ - def fun(part: dict[GsmNode, GsmNode]) -> bool: + def fun(part: dict[GsmNode, GsmNode]) -> Any: for old_node, new_node in part.items(): if local_fun(new_node, old_node) is False: # Follows local_fun(red, blue) - return False - return True + return SpecialScores.ImmediateReject + return SpecialScores.ImmediateAccept return fun -def differential_info(part: dict[GsmNode, GsmNode]) -> tuple[float, int]: - """ - Compute the change in log-likelihood and number of parameters caused by a merge partition. - - :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. - :return tuple[float, int]: (log-likelihood difference, parameter count difference) between old and new nodes. - """ - relevant_nodes_old = list(part.keys()) - relevant_nodes_new = set(part.values()) - - partial_llh_old = sum(node.local_log_likelihood_contribution() for node in relevant_nodes_old) - partial_llh_new = sum(node.local_log_likelihood_contribution() for node in relevant_nodes_new) - - num_params_old = sum(1 for node in relevant_nodes_old for _ in node.transition_iterator()) - num_params_new = sum(1 for node in relevant_nodes_new for _ in node.transition_iterator()) - - return partial_llh_old - partial_llh_new, num_params_old - num_params_new - - -def transform_score(score: Any, transform: Callable) -> Any: +def score_transformation(transform: Callable) -> Any: """ - Apply a transformation to a score, a score function, or a ScoreCalculation's score function. + Lifts an operation on a score value to score functions and ScoreCalculation objects. Intended as a decorator - :param Any score: A plain value, a callable score function, or a ScoreCalculation instance. :param Callable transform: Function to apply to the (eventual) score value. - :return Any: The transformed score, callable, or ScoreCalculation. + :return Any: Decorated transformation applicable to a score, callable, or ScoreCalculation. """ - if isinstance(score, Callable): - return lambda *args: transform(score(*args)) - if isinstance(score, ScoreCalculation): - original_score_function = score.score_function - score.score_function = lambda *args: transform(original_score_function(*args)) - return score - return transform(score) - - -def make_greedy(score: Any) -> Any: + def score_function(score: Any, *transformation_args, **transformation_kwargs) -> Any: + if isinstance(score, Callable): + return lambda partitioning: transform(score(partitioning), *transformation_args, **transformation_kwargs) + if isinstance(score, ScoreCalculation): + original_score_function = score.score_function + score.score_function = lambda partitioning: transform(original_score_function(partitioning), *transformation_args, **transformation_kwargs) + return score + return transform(score, *transformation_args, **transformation_kwargs) + + return score_function + +@score_transformation +def greedy_score(score: Any) -> Any: """ Transform a score into a greedy (boolean) score: accept anything but a False/reject result. :param Any score: A plain value, callable score function, or ScoreCalculation instance. :return Any: The transformed score, callable, or ScoreCalculation. """ - return transform_score(score, lambda x: x is not False) + should_accept = score is not False and score is not SpecialScores.ImmediateReject + return SpecialScores.ImmediateAccept if should_accept else SpecialScores.ImmediateReject +@score_transformation def lower_threshold(score: Any, thresh: Any) -> Any: """ Transform a score so that it is rejected (False) unless it exceeds a threshold. @@ -335,7 +476,26 @@ def lower_threshold(score: Any, thresh: Any) -> Any: :param Any thresh: Threshold the score must exceed to be accepted. :return Any: The transformed score, callable, or ScoreCalculation. """ - return transform_score(score, lambda x: x if thresh < x else False) + return score if thresh <= score else SpecialScores.ImmediateReject + + +def differential_info(part: dict[GsmNode[CountData], GsmNode[CountData]]) -> tuple[float, int]: + """ + Compute the change in log-likelihood and number of parameters caused by a merge partition. + + :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. + :return tuple[float, int]: (log-likelihood difference, parameter count difference) between old and new nodes. + """ + relevant_nodes_old = list(part.keys()) + relevant_nodes_new = set(part.values()) + + partial_llh_old = sum(node.data.local_log_likelihood_contribution() for node in relevant_nodes_old) + partial_llh_new = sum(node.data.local_log_likelihood_contribution() for node in relevant_nodes_new) + + num_params_old = sum(1 for node in relevant_nodes_old for _ in node.child_iterator()) + num_params_new = sum(1 for node in relevant_nodes_new for _ in node.child_iterator()) + + return partial_llh_old - partial_llh_new, num_params_old - num_params_new def AIC_score(alpha: float = 0) -> ScoreFunction: @@ -352,20 +512,20 @@ def score(part: dict[GsmNode, GsmNode]) -> Any: return score -def EDSM_frequency_score(min_evidence: int = -1) -> ScoreFunction: +def EDSM_frequency_score(min_evidence: int = 0) -> ScoreFunction: """ - Build a score function counting the total evidence (transition count) contradicted by a merge. + Build a score function counting the total evidence (transition count) accumulated by a merge. :param int min_evidence: Minimum evidence required for the merge to be accepted. - :return ScoreFunction: Score function computing the total contradicting evidence of a merge partition. + :return ScoreFunction: Score function computing the total accumulated evidence of a merge partition. """ - def score(part: dict[GsmNode, GsmNode]) -> Any: + def score(part: dict[GsmNode[CountData], GsmNode[CountData]]) -> Any: total_evidence = 0 for old_node, new_node in part.items(): - for in_sym, trans_old, trans_new in intersection_iterator(old_node.transitions, new_node.transitions): - for out_sym, t_info_old, t_info_new in intersection_iterator(trans_old, trans_new): - if t_info_old.count != t_info_new.count: - total_evidence += t_info_old.count + for in_sym, old_trans, new_trans_new in intersection_iterator(old_node.data.transition_count, new_node.data.transition_count): + for out_sym, old_count, new_count in intersection_iterator(old_trans, new_trans_new): + if old_count != new_count: + total_evidence += old_count return lower_threshold(total_evidence, min_evidence) return score diff --git a/aalpy/utils/HelperFunctions.py b/aalpy/utils/HelperFunctions.py index 4196ade91f..7030066652 100644 --- a/aalpy/utils/HelperFunctions.py +++ b/aalpy/utils/HelperFunctions.py @@ -258,6 +258,25 @@ def make_input_complete(automaton: Any, missing_transition_go_to: str = 'self_lo return automaton +def ensure_input_complete(learned_model: Any, input_completeness: str | None, print_info: bool) -> None: + """ + Warns about (or fixes) input incompleteness of a learned model, as commonly needed after passive learning. + + :param Any learned_model: Automaton to check. + :param str | None input_completeness: None to only warn, else 'self_loop' or 'sink_state' to fix it. + :param bool print_info: Whether to print progress/warning information. + """ + if not learned_model.is_input_complete(): + if not input_completeness: + if print_info: + print('Warning: Learned Model is not input complete (inputs not defined for all states). ' + 'Consider calling .make_input_complete()') + else: + if print_info: + print(f'Learned model was not input complete. Adapting it with {input_completeness} transitions.') + learned_model.make_input_complete(input_completeness) + + def convert_i_o_traces_for_RPNI(sequences: list, automaton_type: str = "mealy") -> list[tuple]: """ Converts a list of input-output sequences to RPNI format. diff --git a/tests/learning_algs/deterministic_passive/test_gsm_rpni_exhaustive.py b/tests/learning_algs/deterministic_passive/test_gsm_rpni_exhaustive.py new file mode 100644 index 0000000000..34f1016e94 --- /dev/null +++ b/tests/learning_algs/deterministic_passive/test_gsm_rpni_exhaustive.py @@ -0,0 +1,142 @@ +import pytest + +from aalpy.learning_algs.deterministic_passive.GsmRPNI import GsmRPNI +from aalpy.learning_algs.general_passive.GeneralizedStateMerging import run_GSM +from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import EDSM_score, SimpleScoreCalculation +from aalpy.utils import generate_random_deterministic_automata +from aalpy.utils.HelperFunctions import dfa_from_moore +from aalpy.utils.Sampling import get_complete_sample +from aalpy.utils.ModelChecking import bisimilar + +pytestmark = pytest.mark.exhaustive + +# Mirrors tests/learning_algs/resetless/test_hW_exhaustive.py: sweep many (states, inputs, outputs, seed) +# combinations for each deterministic automaton type and check that GSM-RPNI and EDSM (general passive +# state merging with the classic EDSM score), given a complete (state-cover x characterization-set) +# sample, both reconstruct a model bisimilar to the ground truth. +SEEDS = list(range(30)) +MODEL_SIZES = [ + (2, 2, 3), + (3, 3, 2), + (4, 3, 2), + (5, 3, 3), + (6, 2, 3), + (6, 3, 2), + (10, 2, 3), + (10, 2, 4), + (20, 2, 3), + (30, 2, 4), + (30, 3, 2), + (50, 3, 5), +] + +TEST_CASES = [ + pytest.param( + automaton_type, + seed_val, + num_states, + input_size, + output_size, + id=f"states={num_states}-inputs={input_size}-outputs={output_size}-seed={seed_val}-automaton_type={automaton_type}", + ) + for num_states, input_size, output_size in MODEL_SIZES + for seed_val in SEEDS + for automaton_type in ['dfa', 'moore', 'mealy'] +] + + +def complete_data_set(automaton, automaton_type): + """ + Builds a complete list of (input_sequence, output_label) pairs for GSM-RPNI, labeling every prefix + of every sequence in the state-cover x characterization-set complete sample with its own final-step + output. Labeling only the full sequences (and not each of their prefixes) would, for mealy machines, + leave individual transitions along the way unobserved as a "final" label and thus undetermined. + """ + automaton.compute_prefixes() + + data = [] + if automaton_type in ('dfa', 'moore'): + data.append(((), automaton.initial_state.output)) + + seen = set() + for seq in get_complete_sample(automaton): + if not seq: + continue + automaton.reset_to_initial() + outputs = automaton.execute_sequence(automaton.initial_state, seq) + for k in range(1, len(seq) + 1): + prefix = tuple(seq[:k]) + if prefix in seen: + continue + seen.add(prefix) + data.append((prefix, outputs[k - 1])) + + return data + + +def complete_mealy_io_traces(automaton): + """ + Builds full input/output traces (as expected by run_GSM's 'io_traces' format for mealy behavior) + covering the same state-cover x characterization-set sample as complete_data_set. + """ + automaton.compute_prefixes() + + traces = [] + for seq in get_complete_sample(automaton): + if not seq: + continue + automaton.reset_to_initial() + outputs = automaton.execute_sequence(automaton.initial_state, seq) + traces.append(list(zip(seq, outputs))) + + return traces + + +def run_edsm(automaton, automaton_type): + """ + Runs general passive state merging with the classic EDSM score function on a complete sample of + the given automaton, returning a model of the same automaton_type as the ground truth. + """ + score_calc = SimpleScoreCalculation(score_function=EDSM_score()) + + if automaton_type == 'mealy': + traces = complete_mealy_io_traces(automaton) + return run_GSM(traces, output_behavior='mealy', transition_behavior='deterministic', + score_calc=score_calc, data_format='io_traces') + + data = complete_data_set(automaton, automaton_type) + learned_moore = run_GSM(data, output_behavior='moore', transition_behavior='deterministic', + score_calc=score_calc, data_format='labeled_sequences') + return dfa_from_moore(learned_moore) if automaton_type == 'dfa' else learned_moore + + +@pytest.mark.parametrize("automaton_type,seed_val,num_states,input_size,output_size", TEST_CASES) +@pytest.mark.timeout(30) +def test_gsm_rpni_seed_exhaustive(automaton_type, seed_val, num_states, input_size, output_size): + from random import seed + + seed(seed_val) + + model = generate_random_deterministic_automata( + automaton_type, + num_states=num_states, + input_alphabet_size=input_size, + output_alphabet_size=output_size, + ) + if not model.is_minimal(): + pytest.skip(f"seed {seed_val} does not produce a minimal model") + + data = complete_data_set(model, automaton_type) + + rpni_learned_model = GsmRPNI(data, automaton_type, print_info=False).run_rpni() + + assert rpni_learned_model is not None + assert rpni_learned_model.is_minimal() + assert bisimilar(model, rpni_learned_model) + assert len(rpni_learned_model.states) == len(model.states) + + edsm_learned_model = run_edsm(model, automaton_type) + + assert edsm_learned_model.is_minimal() + assert bisimilar(model, edsm_learned_model) + assert len(edsm_learned_model.states) == len(model.states) diff --git a/tests/learning_algs/general_passive/test_generalized_state_merging.py b/tests/learning_algs/general_passive/test_generalized_state_merging.py index 2b1630fd09..8565cb730d 100644 --- a/tests/learning_algs/general_passive/test_generalized_state_merging.py +++ b/tests/learning_algs/general_passive/test_generalized_state_merging.py @@ -1,12 +1,12 @@ -import random import unittest from itertools import product +from typing import Any from aalpy.automata import Dfa, DfaState, MooreMachine, MooreState, MealyMachine, MealyState -from aalpy.learning_algs.general_passive.GeneralizedStateMerging import ( - GeneralizedStateMerging, Instrumentation, run_GSM, -) -from aalpy.learning_algs.general_passive.GsmNode import GsmNode, unknown_output +from aalpy.learning_algs.general_passive.DataHandler import DataHandler, NoOpDataHandler, DataFormat +from aalpy.learning_algs.general_passive.GeneralizedStateMerging import GeneralizedStateMerging, Instrumentation, run_GSM +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior +from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import SimpleScoreCalculation, ScoreCalculation, SpecialScores from aalpy.utils.HelperFunctions import dfa_from_moore from aalpy.utils.ModelChecking import bisimilar @@ -95,6 +95,20 @@ def test_learns_correct_mealy_machine(self): self.assertEqual(len(learned.states), 2) self.assertTrue(bisimilar(learned, ground_truth)) + def test_override_default_checks_bypasses_root_moore_check(self): + class AlwaysCompatible(ScoreCalculation): + def override_default_checks(self): + return True + def local_compatibility(self, a, b): + return True + def score_function(self, part): + return SpecialScores.ImmediateAccept + + data = [([], True), (('a',), False)] + learned = run_GSM(data, output_behavior='moore', transition_behavior='deterministic', + data_format='labeled_sequences', score_calc=AlwaysCompatible(), convert=False) + self.assertEqual(len(learned.get_all_nodes()), 1) + def test_dfa_via_moore_and_dfa_from_moore_conversion(self): q0 = DfaState('q0', is_accepting=True) q1 = DfaState('q1', is_accepting=False) @@ -115,6 +129,21 @@ def test_convert_false_returns_gsm_node(self): data_format='labeled_sequences', convert=False) self.assertIsInstance(result, GsmNode) + def test_custom_score_calc_cannot_produce_nondeterministic_model(self): + # a score_calc that does not check determinism itself must not be able to merge away determinism + ground_truth = parity_mealy() + alphabet = ground_truth.get_input_alphabet() + traces = [] + for level in range(1, 4): + for seq in product(alphabet, repeat=level): + ground_truth.reset_to_initial() + outputs = ground_truth.execute_sequence(ground_truth.initial_state, seq) + traces.append(list(zip(seq, outputs))) + greedy = SimpleScoreCalculation(score_function=lambda part: len(part) - len(set(part.values()))) + learned = run_GSM(traces, output_behavior='mealy', transition_behavior='deterministic', + score_calc=greedy, data_format='io_traces', convert=False) + self.assertTrue(learned.is_deterministic()) + def test_raises_for_invalid_output_behavior(self): with self.assertRaises(ValueError): GeneralizedStateMerging(output_behavior='invalid') @@ -155,9 +184,12 @@ class TestRunGsmPreprocessingPostprocessing(unittest.TestCase): def test_preprocessing_and_postprocessing_are_applied(self): calls = [] - def pta_preprocessing(root): - calls.append('pre') - return root + class PreProcessingHandler(NoOpDataHandler): + def createPTA(self, data: Any, output_behavior: OutputBehavior, data_format: DataFormat = None) -> GsmNode[None]: + calls.append('pre') + return super().createPTA(data, output_behavior, data_format) + + dh = PreProcessingHandler() def postprocessing(root): calls.append('post') @@ -165,7 +197,7 @@ def postprocessing(root): data = [((), True), (('a',), False)] run_GSM(data, output_behavior='moore', transition_behavior='deterministic', - data_format='labeled_sequences', pta_preprocessing=pta_preprocessing, + data_format='labeled_sequences', data_handler=dh, postprocessing=postprocessing) self.assertEqual(calls, ['pre', 'post']) @@ -183,8 +215,7 @@ def test_depth_first_still_learns_correct_model(self): ground_truth = alternating_moore() data = labeled_sequence_data(ground_truth, depth=3) learned = run_GSM(data, output_behavior='moore', transition_behavior='deterministic', - data_format='labeled_sequences', depth_first=True, - compatibility_on_futures=True) + data_format='labeled_sequences', depth_first=True) self.assertTrue(bisimilar(learned, ground_truth)) diff --git a/tests/learning_algs/general_passive/test_gsm_algorithms.py b/tests/learning_algs/general_passive/test_gsm_algorithms.py index 9e88397a16..9ad75e8201 100644 --- a/tests/learning_algs/general_passive/test_gsm_algorithms.py +++ b/tests/learning_algs/general_passive/test_gsm_algorithms.py @@ -2,10 +2,9 @@ import unittest from itertools import product -from aalpy.automata import ( - Dfa, DfaState, MooreMachine, MooreState, MealyMachine, MealyState, Mdp, MdpState, StochasticMealyMachine, - StochasticMealyState, -) +import pytest + +from aalpy.automata import Dfa, DfaState, MooreMachine, MooreState, MealyMachine, MealyState, Mdp, MdpState from aalpy.SULs import AutomatonSUL from aalpy.learning_algs.general_passive.GsmAlgorithms import run_EDSM, run_Alergia_EDSM, run_k_tails from aalpy.utils.ModelChecking import bisimilar @@ -54,6 +53,14 @@ def labeled_sequence_data(automaton, depth=3): return data +class TestPublicApi(unittest.TestCase): + def test_passive_algorithms_are_exported_from_learning_algs(self): + import aalpy.learning_algs as learning_algs + + for name in ('run_EDSM', 'run_Alergia_GSM', 'run_Alergia_EDSM', 'run_k_tails', 'run_GSM'): + self.assertTrue(hasattr(learning_algs, name), f"{name} is not exported") + + class TestRunEdsm(unittest.TestCase): def test_learns_minimal_dfa(self): ground_truth = even_a_dfa() @@ -76,6 +83,22 @@ def test_learns_minimal_mealy_machine(self): self.assertEqual(len(learned.states), 2) self.assertTrue(bisimilar(learned, ground_truth)) + @pytest.mark.timeout(10) + def test_learns_consistent_model_from_sparse_non_prefix_closed_data(self): + # regression test: a cached (first-pass) merge candidate that had speculatively touched a node + # (e.g. while resolving that node's still-unknown prefix output) was not invalidated once that + # node got independently promoted to red. Applying the stale candidate later overwrote the + # promoted state and re-queued one of its own descendants as blue, which made GSM merge a node + # into itself and loop forever. Reproduces with only labels for full sequences (no intermediate + # prefixes labeled), which is the typical shape of EDSM input data. + data = [((), 'o2'), (('i1', 'i1'), 'o2'), (('i1', 'i2'), 'o1'), (('i1', 'i1', 'i2'), 'o1'), + (('i1', 'i2', 'i1'), 'o2'), (('i2', 'i1', 'i1'), 'o1'), (('i2', 'i1', 'i2'), 'o2')] + learned = run_EDSM(data, automaton_type='moore', print_info=False) + for seq, label in data: + learned.reset_to_initial() + got = learned.initial_state.output if not seq else learned.execute_sequence(learned.initial_state, seq)[-1] + self.assertEqual(got, label) + def test_input_completeness_sink_state(self): data = [((), True), (('a',), False), (('b',), True)] learned = run_EDSM(data, automaton_type='dfa', input_completeness='sink_state', print_info=False) diff --git a/tests/learning_algs/general_passive/test_gsm_data_handler.py b/tests/learning_algs/general_passive/test_gsm_data_handler.py new file mode 100644 index 0000000000..803f940a8b --- /dev/null +++ b/tests/learning_algs/general_passive/test_gsm_data_handler.py @@ -0,0 +1,62 @@ +import unittest + +from aalpy.learning_algs.general_passive.DataHandler import CountDataHandler, CountOnPTADataHandler + + +def counting_pta(handler, data_format, data): + return handler.createPTA(data, 'moore', data_format=data_format) + + +class TestCountDataHandler(unittest.TestCase): + def test_copy_can_be_merged_into(self): + dh = CountDataHandler() + x, y = dh.init_data(), dh.init_data() + dh.aggregate_data(_node_with(x), 'a', 'x', None) + dh.aggregate_data(_node_with(y), 'b', 'y', None) + # merging into a copy must work for inputs the copy has never seen + merged = dh.merge(dh.copy(x), y) + self.assertEqual(dict(merged.transition_count), {'a': {'x': 1}, 'b': {'y': 1}}) + + def test_copy_is_independent_of_original(self): + dh = CountDataHandler() + x = dh.init_data() + dh.aggregate_data(_node_with(x), 'a', 'x', None) + copy = dh.copy(x) + dh.aggregate_data(_node_with(x), 'a', 'x', None) + self.assertEqual(copy.transition_count['a'], {'x': 1}) + + +class TestCountOnPTADataHandler(unittest.TestCase): + def test_copy_keeps_pta_data(self): + dh = CountOnPTADataHandler() + pta = counting_pta(dh, 'io_traces', [[True, ('a',True), ('b', False)]]) + copy = dh.copy(pta.data) + self.assertIsInstance(copy, type(pta.data)) + self.assertEqual(copy.pta_count, pta.data.pta_count) + self.assertEqual(copy.shadow_pta, pta.data.shadow_pta) + + +class TestCreatePTA(unittest.TestCase): + def test_labeled_sequences_initialize_node_data(self): + dh = CountOnPTADataHandler() + pta = counting_pta(dh, 'io_traces', [[True, ('a',True), ('b', False)]]) + for node in pta.get_all_nodes(): + self.assertIsNotNone(node.data) + + def test_conflicting_empty_input_labels_raise(self): + with self.assertRaises(ValueError): + counting_pta(CountDataHandler(), 'labeled_sequences', [([], 'a'), ([], 'b')]) + +def _node_with(data): + """Minimal stand-in for the source node of aggregate_data, which only accesses .data.""" + + class Node: + pass + + node = Node() + node.data = data + return node + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/learning_algs/general_passive/test_gsm_node.py b/tests/learning_algs/general_passive/test_gsm_node.py index 55fc5b8ba6..12d5a84463 100644 --- a/tests/learning_algs/general_passive/test_gsm_node.py +++ b/tests/learning_algs/general_passive/test_gsm_node.py @@ -1,10 +1,22 @@ import unittest +from typing import TypeVar +from aalpy.learning_algs.general_passive.DataHandler import ( + CountOnPTADataHandler, NoOpDataHandler, DataHandler, CountDataHandler, detect_data_format +) from aalpy.learning_algs.general_passive.GsmNode import ( - GsmNode, TransitionInfo, detect_data_format, intersection_iterator, union_iterator, unknown_output, + GsmNode, intersection_iterator, union_iterator, unknown_output, no_op_input, IOTrace ) +T = TypeVar('T') + +def simple_create_PTA(traces: list[IOTrace], data_handler: DataHandler[T]) -> GsmNode[T]: + root = GsmNode((no_op_input, unknown_output), None, data_handler.init_data()) + for trace in traces: + data_handler.add_trace(root, trace) + return root + class TestIterators(unittest.TestCase): def test_intersection_iterator_only_common_keys(self): a = {'x': 1, 'y': 2} @@ -24,7 +36,7 @@ def test_empty_data_defaults_to_io_traces(self): self.assertEqual(detect_data_format([]), 'io_traces') def test_gsm_node_is_tree_format(self): - node = GsmNode((None, unknown_output), None) + node = GsmNode((None, unknown_output), None, None) self.assertEqual(detect_data_format(node), 'tree') def test_labeled_sequences_detected(self): @@ -49,101 +61,84 @@ def test_ambiguous_short_traces_default_without_consistency_check(self): class TestGsmNodeBasics(unittest.TestCase): def test_root_has_no_predecessor_and_zero_prefix_length(self): - root = GsmNode((None, unknown_output), None) + root = GsmNode((None, unknown_output), None, None) self.assertIsNone(root.predecessor) self.assertEqual(root.get_prefix_length(), 0) self.assertEqual(root.get_prefix(), []) def test_add_trace_builds_chain_with_correct_prefix(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x'), ('b', 'y')]) - node_a = root.transitions['a']['x'].target - node_ab = node_a.transitions['b']['y'].target + root = simple_create_PTA([[('a', 'x'), ('b', 'y')]], NoOpDataHandler()) + node_a = root.transitions['a']['x'] + node_ab = node_a.transitions['b']['y'] self.assertEqual(node_a.get_prefix_length(), 1) self.assertEqual(node_ab.get_prefix(), [('a', 'x'), ('b', 'y')]) self.assertIs(node_ab.get_root(), root) def test_add_trace_increments_count_for_repeated_trace(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - root.add_trace([('a', 'x')]) - t_info = root.transitions['a']['x'] - self.assertEqual(t_info.count, 2) - self.assertEqual(t_info.original_count, 2) + dh = CountOnPTADataHandler() + traces = [[('a', 'x')], [('a', 'x')]] + root = simple_create_PTA(traces, dh) + self.assertEqual(root.data.transition_count['a']['x'], 2) + self.assertEqual(root.data.pta_count['a']['x'], 2) def test_get_by_prefix_returns_none_for_undefined_path(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) + root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) self.assertIsNone(root.get_by_prefix([('b', 'y')])) def test_get_by_prefix_ignores_leading_none_input(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - node = root.get_by_prefix([(None, 'initial'), ('a', 'x')]) - self.assertIs(node, root.transitions['a']['x'].target) + root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) + node = root.get_by_prefix([(no_op_input, unknown_output), ('a', 'x')]) + self.assertIs(node, root.transitions['a']['x']) def test_get_all_nodes_includes_root_and_children(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x'), ('b', 'y')]) + root = simple_create_PTA([[('a', 'x'), ('b', 'y')]], NoOpDataHandler()) nodes = root.get_all_nodes() self.assertEqual(len(nodes), 3) def test_is_tree_true_for_pta(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - root.add_trace([('b', 'y')]) + root = simple_create_PTA([[('a', 'x')], [('b', 'y')]], NoOpDataHandler()) self.assertTrue(root.is_tree()) def test_is_tree_false_when_node_shared(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - shared = root.transitions['a']['x'].target + root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) + shared = root.transitions['a']['x'] # manually introduce a shared target to simulate a merged (non-tree) structure - root.transitions['b']['y'] = TransitionInfo(shared, 1, None, None) + root.transitions['b']['y'] = shared self.assertFalse(root.is_tree()) - def test_shallow_copy_shares_targets_but_independent_transitions_dict(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - copy = root.shallow_copy() - self.assertIsNot(copy.transitions, root.transitions) - self.assertIs(copy.transitions['a']['x'].target, root.transitions['a']['x'].target) - copy.transitions['b']['y'] = TransitionInfo(copy, 1, None, None) - self.assertNotIn('b', root.transitions) - def test_make_input_complete_adds_self_loops_for_missing_inputs(self): - root = GsmNode((None, 'root_out'), None) - root.add_trace([('a', 'x')]) + dh = NoOpDataHandler() + root = GsmNode((None, 'root_out'), None, None) + dh.add_trace(root, [('a', 'x')]) # 'b' is used elsewhere in the tree but not from root - node_a = root.transitions['a']['x'].target - node_a.add_trace([('b', 'y')]) + node_a = root.transitions['a']['x'] + dh.add_trace(node_a, [('b', 'y')]) missing = root.make_input_complete() self.assertIn((root, 'b', 'root_out'), missing) - self.assertIs(root.transitions['b']['root_out'].target, root) + self.assertIs(root.transitions['b']['root_out'], root) class TestGsmNodeOrderingAndOutputs(unittest.TestCase): def test_lt_orders_by_prefix_length_then_lexicographically(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 1), ('a', 1)]) - root.add_trace([('b', 1)]) - node_a = root.transitions['a'][1].target - node_b = root.transitions['b'][1].target - node_aa = node_a.transitions['a'][1].target - self.assertTrue(node_a < node_aa) - self.assertTrue(node_a < node_b) # same length, 'a' < 'b' - self.assertFalse(node_b < node_a) + root = simple_create_PTA([[('a', 1), ('a', 1)], [('b', 1)]], NoOpDataHandler()) + node_a = root.transitions['a'][1] + node_b = root.transitions['b'][1] + node_aa = node_a.transitions['a'][1] + self.assertTrue(node_a.short_lex_order(node_aa)) + self.assertTrue(node_a.short_lex_order(node_b)) # same length, 'a' < 'b' + self.assertFalse(node_b.short_lex_order(node_a)) def test_resolve_unknown_prefix_output_only_updates_if_unknown(self): - node = GsmNode(('a', unknown_output), None) + node = GsmNode(('a', unknown_output), None, None) node.resolve_unknown_prefix_output('resolved') self.assertEqual(node.get_prefix_output(), 'resolved') node.resolve_unknown_prefix_output('other') self.assertEqual(node.get_prefix_output(), 'resolved') def test_add_labeled_sequence_sets_prefix_output_on_final_node(self): - root = GsmNode((None, unknown_output), None) - root.add_labeled_sequence((('a', 'b'), 'label1')) + dh = NoOpDataHandler() + root = GsmNode((None, unknown_output), None, None) + dh.add_labeled_sequence(root, (('a', 'b'), 'label1')) # only the final step's transition dict key is resolved from unknown_output to the real label; # intermediate steps remain keyed by unknown_output. node = root.get_by_prefix([('a', unknown_output), ('b', 'label1')]) @@ -151,141 +146,130 @@ def test_add_labeled_sequence_sets_prefix_output_on_final_node(self): self.assertEqual(node.get_prefix_output(), 'label1') def test_add_labeled_sequence_raises_on_conflicting_label_for_same_sequence(self): - root = GsmNode((None, unknown_output), None) - root.add_labeled_sequence((('a',), 'out1')) + dh= NoOpDataHandler() + root = GsmNode((None, unknown_output), None, None) + dh.add_labeled_sequence(root,(('a',), 'out1')) with self.assertRaises(ValueError): - root.add_labeled_sequence((('a',), 'out2')) + dh.add_labeled_sequence(root, (('a',), 'out2')) def test_is_locally_deterministic_true_for_single_output_per_input(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) + root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) self.assertTrue(root.is_locally_deterministic()) def test_is_locally_deterministic_false_for_two_outputs_same_input(self): - root = GsmNode((None, unknown_output), None) - root.transitions['a']['x'] = TransitionInfo(GsmNode(('a', 'x'), root), 1, None, None) - root.transitions['a']['y'] = TransitionInfo(GsmNode(('a', 'y'), root), 1, None, None) + root = simple_create_PTA([[('a', 'x')], [('a', 'y')]], NoOpDataHandler()) self.assertFalse(root.is_locally_deterministic()) self.assertFalse(root.is_deterministic()) def test_deterministic_compatible_true_when_no_shared_inputs(self): - n1 = GsmNode((None, unknown_output), None) - n1.add_trace([('a', 'x')]) - n2 = GsmNode((None, unknown_output), None) - n2.add_trace([('b', 'y')]) + n1 = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) + n2 = simple_create_PTA([[('b', 'y')]], NoOpDataHandler()) self.assertTrue(n1.deterministic_compatible(n2)) def test_deterministic_compatible_false_on_output_mismatch_for_shared_input(self): - n1 = GsmNode((None, unknown_output), None) - n1.add_trace([('a', 'x')]) - n2 = GsmNode((None, unknown_output), None) - n2.add_trace([('a', 'y')]) + n1 = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) + n2 = simple_create_PTA([[('a', 'y')]], NoOpDataHandler()) self.assertFalse(n1.deterministic_compatible(n2)) def test_deterministic_compatible_true_when_unknown_output_present(self): - n1 = GsmNode((None, unknown_output), None) - n1.transitions['a'][unknown_output] = TransitionInfo(GsmNode(('a', unknown_output), n1), 1, None, None) - n2 = GsmNode((None, unknown_output), None) - n2.add_trace([('a', 'x')]) + n1 = simple_create_PTA([[('a', unknown_output)]], NoOpDataHandler()) + n2 = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) self.assertTrue(n1.deterministic_compatible(n2)) def test_is_moore_true_when_child_output_matches_transition_output(self): - root = GsmNode((None, 'root_out'), None) - root.add_trace([('a', 'child_out')]) + dh = NoOpDataHandler() + root = GsmNode((None, 'root_out'), None, None) + dh.add_trace(root,[('a', 'child_out')]) self.assertTrue(root.is_moore()) def test_is_moore_false_when_child_output_mismatches(self): - root = GsmNode((None, 'root_out'), None) - child = GsmNode(('a', 'transition_out'), root) - root.transitions['a']['transition_out'] = TransitionInfo(child, 1, None, None) + root = GsmNode((None, 'root_out'), None, None) + child = GsmNode(('a', 'transition_out'), root, None) + root.transitions['a']['transition_out'] = child child.prefix_access_pair = ('a', 'different_child_out') self.assertFalse(root.is_moore()) def test_moore_compatible_true_for_matching_or_unknown_outputs(self): - n1 = GsmNode(('a', 'x'), None) - n2 = GsmNode(('a', 'x'), None) - n3 = GsmNode(('a', unknown_output), None) + n1 = GsmNode(('a', 'x'), None, None) + n2 = GsmNode(('a', 'x'), None, None) + n3 = GsmNode(('a', unknown_output), None, None) self.assertTrue(n1.moore_compatible(n2)) self.assertTrue(n1.moore_compatible(n3)) def test_moore_compatible_false_for_conflicting_outputs(self): - n1 = GsmNode(('a', 'x'), None) - n2 = GsmNode(('a', 'y'), None) + n1 = GsmNode(('a', 'x'), None, None) + n2 = GsmNode(('a', 'y'), None, None) self.assertFalse(n1.moore_compatible(n2)) def test_count_sums_transition_counts(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - root.add_trace([('a', 'x')]) - root.add_trace([('b', 'y')]) - self.assertEqual(root.count(), 3) + root = simple_create_PTA([[('a', 'x')], [('a', 'x')], [('b', 'y')]], CountDataHandler()) + self.assertEqual(root.data.count(), 3) def test_local_log_likelihood_contribution_zero_for_single_outcome(self): # a deterministic transition (single outcome for its input) contributes 0 to the log-likelihood, # since n*log(n) - n*log(n) == 0 regardless of count. - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - root.add_trace([('a', 'x')]) - self.assertAlmostEqual(root.local_log_likelihood_contribution(), 0.0) + root = simple_create_PTA([[('a', 'x')],[('a', 'x')]], CountDataHandler()) + self.assertAlmostEqual(root.data.local_log_likelihood_contribution(), 0.0) def test_local_log_likelihood_contribution_negative_for_split_outcomes(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - root.add_trace([('a', 'y')]) - self.assertLess(root.local_log_likelihood_contribution(), 0.0) + root = simple_create_PTA([[('a', 'x')],[('a', 'y')]], CountDataHandler()) + self.assertLess(root.data.local_log_likelihood_contribution(), 0.0) class TestGsmNodeCreatePTA(unittest.TestCase): def test_labeled_sequences_format(self): + dh = NoOpDataHandler() data = [(('a', 'b'), 1), (('a', 'c'), 2)] - root = GsmNode.createPTA(data, output_behavior='moore', data_format='labeled_sequences') - node_a = root.transitions['a'][unknown_output].target + root = dh.createPTA(data, output_behavior='moore', data_format='labeled_sequences') + node_a = root.transitions['a'][unknown_output] self.assertEqual(node_a.get_prefix_length(), 1) def test_io_traces_moore_uses_first_output_as_root_output(self): + dh = NoOpDataHandler() data = [[0, ('a', 1)], [0, ('a', 1)]] - root = GsmNode.createPTA(data, output_behavior='moore', data_format='io_traces') + root = dh.createPTA(data, output_behavior='moore', data_format='io_traces') self.assertEqual(root.get_prefix_output(), 0) def test_io_traces_mealy_has_no_root_output(self): + dh = NoOpDataHandler() data = [[('a', 'x')]] - root = GsmNode.createPTA(data, output_behavior='mealy', data_format='io_traces') + root = dh.createPTA(data, output_behavior='mealy', data_format='io_traces') self.assertEqual(root.get_prefix_output(), unknown_output) def test_tree_format_passthrough_requires_tree_structure(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - result = GsmNode.createPTA(root, output_behavior='mealy', data_format='tree') + dh = NoOpDataHandler() + root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) + result = dh.createPTA(root, output_behavior='mealy', data_format='tree') self.assertIs(result, root) def test_tree_format_rejects_non_tree(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) - shared = root.transitions['a']['x'].target - root.transitions['b']['y'] = TransitionInfo(shared, 1, None, None) + dh = NoOpDataHandler() + root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) + root.transitions['b']['y'] = root.transitions['a']['x'] with self.assertRaises(ValueError): - GsmNode.createPTA(root, output_behavior='mealy', data_format='tree') + dh.createPTA(root, output_behavior='mealy', data_format='tree') class TestGsmNodeToAutomaton(unittest.TestCase): def test_to_automaton_deterministic_moore(self): - root = GsmNode((None, 0), None) - root.add_trace([('a', 1)]) + dh = NoOpDataHandler() + root = GsmNode((None, 0), None, None) + dh.add_trace(root, [('a', 1)]) automaton = root.to_automaton('moore', 'deterministic') self.assertEqual(automaton.initial_state.output, 0) self.assertEqual(automaton.initial_state.transitions['a'].output, 1) def test_to_automaton_raises_on_non_moore_structure_when_moore_requested(self): - root = GsmNode((None, 'root_out'), None) - child = GsmNode(('a', 'transition_out'), root) - root.transitions['a']['transition_out'] = TransitionInfo(child, 1, None, None) - child.prefix_access_pair = ('a', 'different_output') + root = GsmNode((None, 'root_out'), None, None) + child = GsmNode(('a', 'different_output'), root, None) + root.transitions['a']['transition_out'] = child with self.assertRaises(ValueError): root.to_automaton('moore', 'deterministic') def test_to_automaton_deterministic_mealy(self): - root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')]) + dh = NoOpDataHandler() + root = GsmNode((None, unknown_output), None, None) + dh.add_trace(root, [('a', 'x')]) automaton = root.to_automaton('mealy', 'deterministic') self.assertEqual(automaton.initial_state.output_fun['a'], 'x') diff --git a/tests/learning_algs/general_passive/test_instrumentation.py b/tests/learning_algs/general_passive/test_instrumentation.py index cd7d88d832..55587f1376 100644 --- a/tests/learning_algs/general_passive/test_instrumentation.py +++ b/tests/learning_algs/general_passive/test_instrumentation.py @@ -1,8 +1,9 @@ import unittest +from aalpy.learning_algs.general_passive.DataHandler import NoOpDataHandler from aalpy.learning_algs.general_passive.GeneralizedStateMerging import run_GSM from aalpy.learning_algs.general_passive.Instrumentation import MergeViolationDebugger, ProgressReport -from aalpy.learning_algs.general_passive.GsmNode import GsmNode, TransitionInfo +from aalpy.learning_algs.general_passive.GsmNode import GsmNode class TestProgressReport(unittest.TestCase): @@ -35,9 +36,9 @@ def build_matching_ground_truth_tree(self): # data [((), True), (('a',), True), (('b',), True)] is only ever consistent with a single-state # automaton that self-loops on 'a' and 'b'; the ground truth tree must reflect that so that the # actual merges GSM performs (root with 'a', root with 'b') are considered correct. - root = GsmNode((None, True), None) - root.transitions['a'][True] = TransitionInfo(root, 1, None, None) - root.transitions['b'][True] = TransitionInfo(root, 1, None, None) + root = GsmNode((None, True), None, None) + root.transitions['a'][True] = root + root.transitions['b'][True] = root return root def test_logs_correct_merges_and_promotions_against_ground_truth(self): @@ -57,9 +58,10 @@ def test_logs_correct_merges_and_promotions_against_ground_truth(self): def test_flags_wrong_merge_against_mismatched_ground_truth(self): # a ground truth tree where 'a' and 'b' are distinct states never merges them; # comparing against it while the actual run does merge them should be flagged as wrong. - mismatched_ground_truth = GsmNode((None, True), None) - mismatched_ground_truth.add_trace([('a', True)]) - mismatched_ground_truth.add_trace([('b', True)]) + dh = NoOpDataHandler() + mismatched_ground_truth = GsmNode((None, True), None, None) + dh.add_trace(mismatched_ground_truth, [('a', True)]) + dh.add_trace(mismatched_ground_truth, [('b', True)]) # sabotage: make root.get_by_prefix for 'b' point to a node distinct from 'a's, but give it a # different (non-tree) identity so the debugger's identity check for a real merge fails debugger = MergeViolationDebugger(mismatched_ground_truth) diff --git a/tests/learning_algs/general_passive/test_score_functions_gsm.py b/tests/learning_algs/general_passive/test_score_functions_gsm.py index 5180f0a256..fd43d546ab 100644 --- a/tests/learning_algs/general_passive/test_score_functions_gsm.py +++ b/tests/learning_algs/general_passive/test_score_functions_gsm.py @@ -1,39 +1,81 @@ import unittest -from aalpy.learning_algs.general_passive.GsmNode import GsmNode, TransitionInfo, unknown_output +from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, unknown_output from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ( - AIC_score, EDSM_frequency_score, EDSM_score, ScoreCalculation, ScoreCombinator, ScoreWithKTail, + ScoreCalculation, + AIC_score, EDSM_frequency_score, EDSM_score, SimpleScoreCalculation, ScoreCombinator, ScoreWithKTail, ScoreWithSinks, differential_info, hoeffding_compatibility, local_to_global_compatibility, lower_threshold, - make_greedy, transform_score, + greedy_score, score_transformation, SpecialScores ) def node_with_counts(counts, prefix_access_pair=(None, unknown_output)): """Builds a node whose single input 'i' has the given {output: count} outgoing transitions.""" - node = GsmNode(prefix_access_pair, None) + dh = CountOnPTADataHandler() + node = GsmNode(prefix_access_pair, None, dh.init_data()) for out_sym, count in counts.items(): - target = GsmNode(('i', out_sym), node) - node.transitions['i'][out_sym] = TransitionInfo(target, count, target, count) + target = GsmNode(('i', out_sym), node, dh.init_data()) + node.transitions['i'][out_sym] = target + node.data.transition_count['i'] = counts + node.data.pta_count['i'] = counts return node class TestScoreCalculationDefaults(unittest.TestCase): def test_default_local_compatibility_always_true(self): - sc = ScoreCalculation() - self.assertTrue(sc.local_compatibility(GsmNode((None, None), None), GsmNode((None, None), None))) - self.assertFalse(sc.has_local_compatibility()) + sc = SimpleScoreCalculation() + self.assertTrue(sc.local_compatibility(GsmNode((None, None), None, None), GsmNode((None, None), None, None))) def test_default_score_function_always_true(self): - sc = ScoreCalculation() + sc = SimpleScoreCalculation() self.assertTrue(sc.score_function({})) self.assertFalse(sc.has_score_function()) def test_custom_functions_are_detected_as_overridden(self): - sc = ScoreCalculation(local_compatibility=lambda a, b: False, score_function=lambda p: 42) - self.assertTrue(sc.has_local_compatibility()) + sc = SimpleScoreCalculation(local_compatibility=lambda a, b: False, score_function=lambda p: 42) self.assertTrue(sc.has_score_function()) +class TestOverrideDetection(unittest.TestCase): + def test_plain_subclass_reports_no_overrides(self): + class Plain(ScoreCalculation): + pass + + self.assertFalse(Plain().has_local_compatibility()) + self.assertFalse(Plain().has_score_function()) + + def test_subclass_overriding_methods_is_detected(self): + class Custom(ScoreCalculation): + def local_compatibility(self, a, b): + return True + + def score_function(self, part): + return 1 + + self.assertTrue(Custom().has_local_compatibility()) + self.assertTrue(Custom().has_score_function()) + + +class TestScoreCombinatorAggregation(unittest.TestCase): + def test_no_early_verdict_when_no_sub_score_has_one(self): + # a combined early verdict must stay NoScore, otherwise the partitioning is never built + comb = ScoreCombinator([SimpleScoreCalculation(), SimpleScoreCalculation()]) + node = node_with_counts({}) + self.assertIs(comb.early_score(node, node), SpecialScores.NoScore) + + def test_single_rejecting_sub_score_rejects(self): + aggregate = ScoreCombinator.default_aggregate_score + self.assertIs(aggregate([1, SpecialScores.ImmediateReject]), SpecialScores.ImmediateReject) + + def test_single_undecided_sub_score_stays_undecided(self): + aggregate = ScoreCombinator.default_aggregate_score + self.assertIs(aggregate([SpecialScores.NoScore, SpecialScores.ImmediateAccept]), SpecialScores.NoScore) + + def test_plain_scores_are_collected_into_a_list(self): + self.assertEqual(ScoreCombinator.default_aggregate_score([1, 2]), [1, 2]) + + class TestHoeffdingCompatibility(unittest.TestCase): def test_identical_distributions_are_compatible(self): a = node_with_counts({'x': 100, 'y': 100}) @@ -54,77 +96,80 @@ def test_zero_total_count_is_ignored(self): self.assertTrue(compat(a, b)) def test_disjoint_inputs_are_compatible(self): - a = GsmNode((None, None), None) - a.transitions['i']['x'] = TransitionInfo(GsmNode(('i', 'x'), a), 100, GsmNode(('i', 'x'), a), 100) - b = GsmNode((None, None), None) - b.transitions['j']['y'] = TransitionInfo(GsmNode(('j', 'y'), b), 100, GsmNode(('j', 'y'), b), 100) + dh = CountOnPTADataHandler() + a = GsmNode((None, None), None, dh.init_data()) + a.data.transition_count = a.data.pta_count = {'i': {'x': 100}} + b = GsmNode((None, None), None, dh.init_data()) + b.data.transition_count = b.data.pta_count = {'j': {'y': 100}} compat = hoeffding_compatibility(0.05) self.assertTrue(compat(a, b)) class TestScoreWithKTail(unittest.TestCase): def test_beyond_depth_k_is_always_compatible(self): - always_false = ScoreCalculation(local_compatibility=lambda a, b: False) + always_false = SimpleScoreCalculation(local_compatibility=lambda a, b: False) wrapped = ScoreWithKTail(always_false, k=1) - root = GsmNode((None, None), None) - blue_shallow = GsmNode(('a', None), root) - blue_shallow_child = GsmNode(('a', None), blue_shallow) + root = GsmNode((None, None), None, None) + blue_shallow = GsmNode(('a', None), root, None) + blue_shallow_child = GsmNode(('a', None), blue_shallow, None) - wrapped.reset() + wrapped.initialize_merge(root, blue_shallow, True) # first call establishes the depth offset at blue_shallow's depth (1) self.assertFalse(wrapped.local_compatibility(root, blue_shallow)) # a node one level deeper than the offset (depth 2) is beyond k=1 -> compatible regardless - self.assertTrue(wrapped.local_compatibility(root, blue_shallow_child)) + self.assertIsNone(wrapped.local_compatibility(root, blue_shallow_child)) def test_within_depth_k_delegates_to_wrapped_score(self): - always_false = ScoreCalculation(local_compatibility=lambda a, b: False) + always_false = SimpleScoreCalculation(local_compatibility=lambda a, b: False) wrapped = ScoreWithKTail(always_false, k=5) - root = GsmNode((None, None), None) - blue = GsmNode(('a', None), root) - wrapped.reset() + root = GsmNode((None, None), None, None) + blue = GsmNode(('a', None), root, None) + wrapped.initialize_merge(root, blue, True) self.assertFalse(wrapped.local_compatibility(root, blue)) class TestScoreWithSinks(unittest.TestCase): def test_rejects_merge_between_sink_and_non_sink(self): - always_true = ScoreCalculation(local_compatibility=lambda a, b: True) + always_true = SimpleScoreCalculation(local_compatibility=lambda a, b: True) is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink) - wrapped.reset() - sink_node = GsmNode((None, 'sink'), None) - normal_node = GsmNode((None, 'normal'), None) - self.assertFalse(wrapped.local_compatibility(sink_node, normal_node)) + sink_node = GsmNode((None, 'sink'), None, None) + normal_node = GsmNode((None, 'normal'), None, None) + # early reject + self.assertFalse(wrapped.initialize_merge(sink_node, normal_node, True)) + # accept if encountered later + self.assertTrue(wrapped.local_compatibility(sink_node, normal_node)) def test_allows_merge_between_two_sinks_by_default(self): - always_true = ScoreCalculation(local_compatibility=lambda a, b: True) + always_true = SimpleScoreCalculation(local_compatibility=lambda a, b: True) is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink) - wrapped.reset() - sink_a = GsmNode((None, 'sink'), None) - sink_b = GsmNode((None, 'sink'), None) + sink_a = GsmNode((None, 'sink'), None, None) + sink_b = GsmNode((None, 'sink'), None, None) + self.assertIsNone(wrapped.initialize_merge(sink_a, sink_b, True)) self.assertTrue(wrapped.local_compatibility(sink_a, sink_b)) def test_rejects_merge_between_two_sinks_when_disallowed(self): - always_true = ScoreCalculation(local_compatibility=lambda a, b: True) + always_true = SimpleScoreCalculation(local_compatibility=lambda a, b: True) is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink, allow_sink_merge=False) - wrapped.reset() - sink_a = GsmNode((None, 'sink'), None) - sink_b = GsmNode((None, 'sink'), None) - self.assertFalse(wrapped.local_compatibility(sink_a, sink_b)) + sink_a = GsmNode((None, 'sink'), None, None) + sink_b = GsmNode((None, 'sink'), None, None) + self.assertFalse(wrapped.initialize_merge(sink_a, sink_b, True)) + self.assertTrue(wrapped.local_compatibility(sink_a, sink_b)) def test_sink_check_only_applies_on_first_call(self): - always_true = ScoreCalculation(local_compatibility=lambda a, b: True) + always_true = SimpleScoreCalculation(local_compatibility=lambda a, b: True) is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink, allow_sink_merge=False) - wrapped.reset() - sink_a = GsmNode((None, 'sink'), None) - normal = GsmNode((None, 'normal'), None) + sink_a = GsmNode((None, 'sink'), None, None) + normal = GsmNode((None, 'normal'), None, None) + self.assertFalse(wrapped.initialize_merge(normal, normal, True)) # consume the "first call" check with a compatible (non-sink) pair self.assertTrue(wrapped.local_compatibility(normal, normal)) # subsequent calls skip the sink check entirely, so this doesn't get rejected @@ -133,32 +178,32 @@ def test_sink_check_only_applies_on_first_call(self): class TestScoreCombinator(unittest.TestCase): def test_default_aggregate_compatibility_commits_to_first_non_none(self): - s1 = ScoreCalculation(local_compatibility=lambda a, b: None) - s2 = ScoreCalculation(local_compatibility=lambda a, b: False) + s1 = SimpleScoreCalculation(local_compatibility=lambda a, b: None) + s2 = SimpleScoreCalculation(local_compatibility=lambda a, b: False) combined = ScoreCombinator([s1, s2]) self.assertFalse(combined.local_compatibility(None, None)) - def test_default_aggregate_compatibility_true_when_all_none(self): - s1 = ScoreCalculation(local_compatibility=lambda a, b: None) + def test_default_aggregate_compatibility_none_when_all_none(self): + s1 = SimpleScoreCalculation(local_compatibility=lambda a, b: None) combined = ScoreCombinator([s1]) - self.assertTrue(combined.local_compatibility(None, None)) + self.assertIsNone(combined.local_compatibility(None, None)) def test_default_aggregate_score_collects_all_scores(self): - s1 = ScoreCalculation(score_function=lambda p: 1) - s2 = ScoreCalculation(score_function=lambda p: 2) + s1 = SimpleScoreCalculation(score_function=lambda p: 1) + s2 = SimpleScoreCalculation(score_function=lambda p: 2) combined = ScoreCombinator([s1, s2]) self.assertEqual(combined.score_function({}), [1, 2]) def test_reset_delegates_to_all_scores(self): calls = [] - class Tracking(ScoreCalculation): - def reset(self): + class Tracking(SimpleScoreCalculation): + def initialize_merge(self, red, blue, first_pass): calls.append(id(self)) s1, s2 = Tracking(), Tracking() combined = ScoreCombinator([s1, s2]) - combined.reset() + combined.initialize_merge(None, None, True) self.assertEqual(len(calls), 2) @@ -187,34 +232,34 @@ def test_merging_identical_nodes_does_not_change_likelihood(self): class TestScoreTransforms(unittest.TestCase): def test_transform_score_on_plain_value(self): - self.assertEqual(transform_score(5, lambda x: x * 2), 10) + self.assertEqual(score_transformation(lambda x: x * 2)(5), 10) def test_transform_score_on_callable(self): - fun = transform_score(lambda part: 5, lambda x: x * 2) + fun = score_transformation(lambda x: x * 2)(lambda part: 5) self.assertEqual(fun({}), 10) def test_transform_score_on_score_calculation(self): # regression test: transform_score used to reassign score.score_function to a lambda that # referenced score.score_function again, causing infinite recursion on the first call. - sc = ScoreCalculation(score_function=lambda part: 5) - transformed = transform_score(sc, lambda x: x * 2) + sc = SimpleScoreCalculation(score_function=lambda part: 5) + transformed = score_transformation(lambda x: x * 2)(sc) self.assertIs(transformed, sc) self.assertEqual(transformed.score_function({}), 10) def test_transform_score_on_score_calculation_can_be_applied_twice(self): - sc = ScoreCalculation(score_function=lambda part: 5) - transform_score(sc, lambda x: x * 2) - transform_score(sc, lambda x: x + 1) + sc = SimpleScoreCalculation(score_function=lambda part: 5) + score_transformation(lambda x: x * 2)(sc) + score_transformation(lambda x: x + 1)(sc) self.assertEqual(sc.score_function({}), 11) def test_make_greedy_rejects_only_false(self): - self.assertTrue(make_greedy(0)) - self.assertTrue(make_greedy('anything')) - self.assertFalse(make_greedy(False)) + self.assertTrue(greedy_score(0)) + self.assertTrue(greedy_score('anything') is SpecialScores.ImmediateAccept) + self.assertTrue(greedy_score(False) is SpecialScores.ImmediateReject) def test_lower_threshold_rejects_values_at_or_below_threshold(self): self.assertEqual(lower_threshold(5, 3), 5) - self.assertFalse(lower_threshold(3, 3)) + self.assertEqual(lower_threshold(3, 3), 3) self.assertFalse(lower_threshold(1, 3)) @@ -226,10 +271,10 @@ def test_aic_score_rejects_partitions_below_threshold(self): result = score_fun({old1: merged}) self.assertFalse(result) - def test_edsm_frequency_score_counts_contradicted_evidence(self): + def test_edsm_frequency_score_counts_accumulated_evidence(self): score_fun = EDSM_frequency_score(min_evidence=-1) old_node = node_with_counts({'x': 5}) - new_node = node_with_counts({'x': 10}) # count changed by the merge -> contradicted evidence + new_node = node_with_counts({'x': 10}) # count changed by the merge -> accumulated evidence result = score_fun({old_node: new_node}) self.assertEqual(result, 5)