From e41ac82fc1cdc6e610048f0afe5b0d9a108c7927 Mon Sep 17 00:00:00 2001 From: zwergziege Date: Mon, 22 Sep 2025 12:59:41 +0200 Subject: [PATCH 01/70] reimplement data extension for GSM (rebase was too tedious) --- .../GeneralizedStateMerging.py | 17 ++++-- .../learning_algs/general_passive/GsmNode.py | 52 +++++++++++++------ 2 files changed, 49 insertions(+), 20 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index e1a1b1a01c..589ebec0c5 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -3,7 +3,7 @@ from typing import Dict, Tuple, Callable, List, Optional from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, TransitionInfo, \ - OutputBehaviorRange, TransitionBehaviorRange, intersection_iterator, unknown_output, detect_data_format + OutputBehaviorRange, TransitionBehaviorRange, intersection_iterator, unknown_output, detect_data_format, IOHandler from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility @@ -46,11 +46,13 @@ def __init__(self, *, score_calc: ScoreCalculation = None, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, + data_handler: IOHandler = None, compatibility_on_pta: bool = False, compatibility_on_futures: bool = False, node_order: Callable[[GsmNode, GsmNode], bool] = None, consider_only_min_blue=False, - depth_first=False): + depth_first=False, + ): if output_behavior not in OutputBehaviorRange: raise ValueError(f"invalid output behavior {output_behavior}. should be in {OutputBehaviorRange}") @@ -76,6 +78,8 @@ def __init__(self, *, self.pta_preprocessing = pta_preprocessing or (lambda x: x) self.postprocessing = postprocessing or (lambda x: x) + self.data_handler = data_handler or IOHandler() + self.compatibility_on_pta = compatibility_on_pta self.compatibility_on_futures = compatibility_on_futures @@ -102,7 +106,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo 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 = GsmNode.createPTA(data, self.output_behavior, data_format, self.data_handler) root = self.pta_preprocessing(root) instrumentation.pta_construction_done(root) @@ -186,6 +190,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo 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 instrumentation.log_merge(best_candidate) # FUTURE: optimizations for compatibility tests where merges can be orthogonal @@ -240,7 +245,7 @@ def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: p = partitioning.full_mapping.get(red_node) # could check smaller .red_mapping? if p is None: - p = red_node.shallow_copy() + p = red_node.shallow_copy(self.data_handler) partitioning.full_mapping[red_node] = p partitioning.red_mapping[red_node] = p if blue_node is not None: @@ -267,6 +272,8 @@ def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode if self.compute_local_compatibility(partition, blue) is False: return partitioning + 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] @@ -307,6 +314,7 @@ def run_GSM(data: list, *, score_calc: ScoreCalculation = None, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, + data_handler: IOHandler = None, compatibility_on_pta: bool = False, compatibility_on_futures: bool = False, node_order: Callable[[GsmNode, GsmNode], bool] = None, @@ -357,6 +365,7 @@ def run_GSM(data: list, *, score_calc=score_calc, pta_preprocessing=pta_preprocessing, postprocessing=postprocessing, + data_handler=data_handler, compatibility_on_pta=compatibility_on_pta, compatibility_on_futures=compatibility_on_futures, node_order=node_order, diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 3676b351cf..9d90f67c00 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -3,9 +3,8 @@ import pathlib from collections import defaultdict from functools import total_ordering -from typing import Dict, Any, List, Tuple, Iterable, Callable, Union, TypeVar, Iterator, Optional, Sequence +from typing import Dict, Any, List, Tuple, Iterable, Callable, Union, TypeVar, Iterator, Optional, Sequence, Generic import pydot -from copy import copy from aalpy.automata import StochasticMealyMachine, StochasticMealyState, MooreState, MooreMachine, NDMooreState, \ NDMooreMachine, Mdp, MdpState, MealyMachine, MealyState, Onfsm, OnfsmState @@ -13,6 +12,7 @@ Key = TypeVar("Key") Val = TypeVar("Val") +T = TypeVar("T") OutputBehavior = str OutputBehaviorRange = ["moore", "mealy"] @@ -96,6 +96,21 @@ def detect_data_format(data, check_consistency=False, guess=False): raise ValueError("ambiguous data format. data format needs to be specified explicitly.") return accepted_formats[0] +class IOHandler(Generic[T]): + def abstract(self, in_val, out_val): + return in_val, out_val + + def init_data(self) -> T: + return None + + def aggregate_data(self, data: T, in_sym, out_value) -> T: + return None + + def merge(self, x: T, y: T) -> T: + return None + + def copy(self, x: T) -> T: + return None # TODO maybe split this for maintainability (and perfomance?) class TransitionInfo: @@ -110,7 +125,7 @@ def __init__(self, target, count, original_target, 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. @@ -120,13 +135,14 @@ 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, predecessor: 'GsmNode' = None): + def __init__(self, prefix_access_pair, predecessor: 'GsmNode' = None, data: T = None): # TODO (data-ext) check all invocations # TODO try single dict self.transitions: defaultdict[Any, Dict[Any, TransitionInfo]] = defaultdict(dict) self.predecessor: GsmNode = predecessor self.prefix_access_pair = prefix_access_pair + self.data = data def __lt__(self, other, compare_length_only=False): own_l, other_l = self.get_prefix_length(), other.get_prefix_length() @@ -191,8 +207,8 @@ def transition_iterator(self) -> Iterable[Tuple[Any, Any, TransitionInfo]]: for out_sym, node in transitions.items(): yield in_sym, out_sym, node - def shallow_copy(self) -> 'GsmNode': - node = GsmNode(self.prefix_access_pair, self.predecessor) + def shallow_copy(self, data_handler: IOHandler) -> 'GsmNode': #TODO fix + node = GsmNode(self.prefix_access_pair, self.predecessor, data_handler.copy(self.data)) for in_sym, t in self.transitions.items(): d = dict() # appears to be faster than dict comprehension for out_sym, ti in t.items(): @@ -389,13 +405,15 @@ def make_input_complete(self) -> List[Tuple['GsmNode', Any, Any]]: transitions[out_sym] = t_info return missing_trans - def add_trace(self, trace: IOTrace): + def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): curr_node: GsmNode = self - for in_sym, out_sym in trace: + for in_value, out_value in trace: + in_sym, out_sym = data_handler.abstract(in_value, out_value) + curr_node.data = data_handler.aggregate_data(curr_node.data, in_value, out_value) transitions = curr_node.transitions[in_sym] info = transitions.get(out_sym) if info is None: - node = GsmNode((in_sym, out_sym), curr_node) + node = GsmNode((in_sym, out_sym), curr_node, data_handler.init_data()) transitions[out_sym] = TransitionInfo(node, 1, node, 1) else: info.count += 1 @@ -403,13 +421,15 @@ def add_trace(self, trace: IOTrace): node = info.target curr_node = node - def add_labeled_sequence(self, example: IOExample): + def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = None): inputs, output = example curr_node: GsmNode = self in_sym = None # step through inputs and add transitions - for in_sym in inputs: + for in_value in inputs: + in_sym = data_handler.abstract(in_value, None) + curr_node.data = data_handler.aggregate_data(curr_node.data, in_value, None) transitions = curr_node.transitions[in_sym] t_infos = list(transitions.values()) if len(t_infos) == 0: @@ -437,7 +457,7 @@ def add_labeled_sequence(self, example: IOExample): raise ValueError("nondeterminism encountered for GSM with labeled_sequences. not supported") @staticmethod - def createPTA(data, output_behavior, data_format=None) -> 'GsmNode': + def createPTA(data, output_behavior, data_format=None, data_handler: IOHandler[T] = None) -> 'GsmNode': if data_format is None: data_format = detect_data_format(data) if data_format not in DataFormatRange: @@ -447,10 +467,10 @@ def createPTA(data, output_behavior, data_format=None) -> 'GsmNode': if not data.is_tree(): raise ValueError("provided automaton is not a tree") return data - root_node = GsmNode((None, unknown_output), None) + root_node = GsmNode((None, unknown_output), None, data_handler.init_data()) if data_format == "labeled_sequences": for example in data: - root_node.add_labeled_sequence(example) + root_node.add_labeled_sequence(example, data_handler) if data_format == "io_traces" or data_format == "traces": if output_behavior == "moore": initial_output = data[0][0] @@ -459,7 +479,7 @@ def createPTA(data, output_behavior, data_format=None) -> 'GsmNode': for trace in data: if data_format == "traces": trace = (("step", t) for t in trace) - root_node.add_trace(trace) + root_node.add_trace(trace, data_handler) return root_node def is_locally_deterministic(self): From 4ce915bd21e7147b343e8cd54b415e12be85a6b8 Mon Sep 17 00:00:00 2001 From: zwergziege Date: Mon, 13 Oct 2025 10:16:45 +0200 Subject: [PATCH 02/70] Added init function to datahandler --- .../GeneralizedStateMerging.py | 5 ++- .../learning_algs/general_passive/GsmNode.py | 38 ++++++++++++++++++- 2 files changed, 39 insertions(+), 4 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 589ebec0c5..fe7d446809 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -3,7 +3,8 @@ from typing import Dict, Tuple, Callable, List, Optional from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, TransitionInfo, \ - OutputBehaviorRange, TransitionBehaviorRange, intersection_iterator, unknown_output, detect_data_format, IOHandler + OutputBehaviorRange, TransitionBehaviorRange, intersection_iterator, unknown_output, detect_data_format, IOHandler, \ + NoIOHandler from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility @@ -78,7 +79,7 @@ def __init__(self, *, self.pta_preprocessing = pta_preprocessing or (lambda x: x) self.postprocessing = postprocessing or (lambda x: x) - self.data_handler = data_handler or IOHandler() + self.data_handler = data_handler or NoIOHandler() self.compatibility_on_pta = compatibility_on_pta self.compatibility_on_futures = compatibility_on_futures diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 9d90f67c00..130ae26f92 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -1,6 +1,7 @@ import functools import math import pathlib +from abc import abstractmethod from collections import defaultdict from functools import total_ordering from typing import Dict, Any, List, Tuple, Iterable, Callable, Union, TypeVar, Iterator, Optional, Sequence, Generic @@ -97,6 +98,34 @@ def detect_data_format(data, check_consistency=False, guess=False): return accepted_formats[0] class IOHandler(Generic[T]): + @abstractmethod + def init(self, data, output_format, data_format): + ... + + @abstractmethod + def abstract(self, in_val, out_val): + ... + + @abstractmethod + def init_data(self) -> T: + ... + + @abstractmethod + def aggregate_data(self, data: T, in_sym, out_value) -> T: + ... + + @abstractmethod + def merge(self, x: T, y: T) -> T: + ... + + @abstractmethod + def copy(self, x: T) -> T: + ... + +class NoIOHandler(IOHandler): + def init(self, data, output_format, data_format): + pass + def abstract(self, in_val, out_val): return in_val, out_val @@ -207,7 +236,7 @@ def transition_iterator(self) -> Iterable[Tuple[Any, Any, TransitionInfo]]: for out_sym, node in transitions.items(): yield in_sym, out_sym, node - def shallow_copy(self, data_handler: IOHandler) -> 'GsmNode': #TODO fix + def shallow_copy(self, data_handler: IOHandler) -> 'GsmNode': node = GsmNode(self.prefix_access_pair, self.predecessor, data_handler.copy(self.data)) for in_sym, t in self.transitions.items(): d = dict() # appears to be faster than dict comprehension @@ -426,9 +455,12 @@ def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = curr_node: GsmNode = self in_sym = None + if not isinstance(data_handler, NoIOHandler): + raise NotImplementedError("Data handling is not supported for learning from labeled sequences") + # step through inputs and add transitions for in_value in inputs: - in_sym = data_handler.abstract(in_value, None) + in_sym, out_sym = data_handler.abstract(in_value, None) curr_node.data = data_handler.aggregate_data(curr_node.data, in_value, None) transitions = curr_node.transitions[in_sym] t_infos = list(transitions.values()) @@ -463,6 +495,8 @@ def createPTA(data, output_behavior, data_format=None, data_handler: IOHandler[T if data_format not in DataFormatRange: raise ValueError(f"invalid data format {data_format}. should be in {DataFormatRange}") + data_handler.init(data, output_behavior, data_format) + if data_format == "tree": if not data.is_tree(): raise ValueError("provided automaton is not a tree") From 6bd7f76380985b5856d7b4fbfaaf5bc8758edbca Mon Sep 17 00:00:00 2001 From: zwergziege Date: Thu, 23 Oct 2025 09:47:39 +0200 Subject: [PATCH 03/70] FIX: also use abstraction on initial output in GSM + Data --- aalpy/learning_algs/general_passive/GsmNode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 130ae26f92..dac34a1959 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -508,7 +508,7 @@ def createPTA(data, output_behavior, data_format=None, data_handler: IOHandler[T 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) + root_node.prefix_access_pair = data_handler.abstract(None, initial_output) data = (d[1:] for d in data) for trace in data: if data_format == "traces": From 5270972bf1170537d484f2a735337d3034010ad7 Mon Sep 17 00:00:00 2001 From: zwergziege Date: Thu, 23 Oct 2025 10:06:40 +0200 Subject: [PATCH 04/70] Allow IOHandler.aggregate to modify src and dst data. Also aggregate initial output for Moore --- .../learning_algs/general_passive/GsmNode.py | 22 +++++++++++++------ 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index dac34a1959..dc988819ee 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -111,7 +111,7 @@ def init_data(self) -> T: ... @abstractmethod - def aggregate_data(self, data: T, in_sym, out_value) -> T: + def aggregate_data(self, src_node: 'GsmNode[T]', in_sym, out_value, dst_node: 'GsmNode[T]'): ... @abstractmethod @@ -132,8 +132,8 @@ def abstract(self, in_val, out_val): def init_data(self) -> T: return None - def aggregate_data(self, data: T, in_sym, out_value) -> T: - return None + def aggregate_data(self, src_node: 'GsmNode[T]', in_sym, out_value, dst_node: 'GsmNode[T]'): + pass def merge(self, x: T, y: T) -> T: return None @@ -438,7 +438,6 @@ def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): curr_node: GsmNode = self for in_value, out_value in trace: in_sym, out_sym = data_handler.abstract(in_value, out_value) - curr_node.data = data_handler.aggregate_data(curr_node.data, in_value, out_value) transitions = curr_node.transitions[in_sym] info = transitions.get(out_sym) if info is None: @@ -448,6 +447,7 @@ def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): info.count += 1 info.original_count += 1 node = info.target + data_handler.aggregate_data(curr_node, in_value, out_value, node) curr_node = node def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = None): @@ -461,7 +461,6 @@ def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = # step through inputs and add transitions for in_value in inputs: in_sym, out_sym = data_handler.abstract(in_value, None) - curr_node.data = data_handler.aggregate_data(curr_node.data, in_value, None) transitions = curr_node.transitions[in_sym] t_infos = list(transitions.values()) if len(t_infos) == 0: @@ -476,6 +475,7 @@ def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = else: # This should never happen raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") + data_handler.aggregate_data(curr_node, in_value, None, node) curr_node = node # set last output @@ -507,8 +507,16 @@ def createPTA(data, output_behavior, data_format=None, data_handler: IOHandler[T root_node.add_labeled_sequence(example, data_handler) if data_format == "io_traces" or data_format == "traces": if output_behavior == "moore": - initial_output = data[0][0] - root_node.prefix_access_pair = data_handler.abstract(None, initial_output) + root_node.prefix_access_pair = data_handler.abstract(None, data[0][0]) + initial_output_symbol = root_node.prefix_access_pair[1] + + for trace in data: + initial_output = trace[0] + _, ios = data_handler.abstract(None, initial_output) + if ios != initial_output_symbol: + raise ValueError("expect unique initial output symbol for Moore behavior") + data_handler.aggregate_data(None, None, initial_output, root_node) + data = (d[1:] for d in data) for trace in data: if data_format == "traces": From b705ac481ac11959f6540ba484d4b441446dbb44 Mon Sep 17 00:00:00 2001 From: zwergziege Date: Fri, 24 Oct 2025 15:56:36 +0200 Subject: [PATCH 05/70] added tiny todo --- aalpy/learning_algs/general_passive/GsmNode.py | 1 + 1 file changed, 1 insertion(+) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index dc988819ee..ccb744fc4f 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -501,6 +501,7 @@ def createPTA(data, output_behavior, data_format=None, data_handler: IOHandler[T 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((None, unknown_output), None, data_handler.init_data()) if data_format == "labeled_sequences": for example in data: From 18c4733a4df33cdf42320328bbb47b22636b9238 Mon Sep 17 00:00:00 2001 From: zwergziege Date: Mon, 3 Nov 2025 10:13:55 +0100 Subject: [PATCH 06/70] minor change in data aggregation --- aalpy/learning_algs/general_passive/GsmNode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index ccb744fc4f..fa0685be08 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -111,7 +111,7 @@ def init_data(self) -> T: ... @abstractmethod - def aggregate_data(self, src_node: 'GsmNode[T]', in_sym, out_value, dst_node: 'GsmNode[T]'): + def aggregate_data(self, src_node: 'GsmNode[T]', in_value, out_value, dst_node: 'GsmNode[T]'): ... @abstractmethod From 03c0c35aef473d1fbcc16bb06788444911c91595 Mon Sep 17 00:00:00 2001 From: zwergziege Date: Thu, 25 Jun 2026 17:11:05 +0200 Subject: [PATCH 07/70] type hints for data --- aalpy/learning_algs/general_passive/GsmNode.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index fa0685be08..54ab2d5b7a 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -142,13 +142,13 @@ def copy(self, x: T) -> T: return None # TODO maybe split this for maintainability (and perfomance?) -class TransitionInfo: +class TransitionInfo(Generic[T]): __slots__ = ["target", "count", "original_target", "original_count"] def __init__(self, target, count, original_target, original_count): - self.target: 'GsmNode' = target + self.target: 'GsmNode[T]' = target self.count: int = count - self.original_target: 'GsmNode' = original_target + self.original_target: 'GsmNode[T]' = original_target self.original_count: int = original_count @@ -166,9 +166,9 @@ class GsmNode(Generic[T]): """ __slots__ = ['transitions', 'predecessor', 'prefix_access_pair', 'data'] - def __init__(self, prefix_access_pair, predecessor: 'GsmNode' = None, data: T = None): # TODO (data-ext) check all invocations + def __init__(self, prefix_access_pair, predecessor: 'GsmNode[T]' = None, data: T = None): # TODO (data-ext) check all invocations # TODO try single dict - self.transitions: defaultdict[Any, Dict[Any, TransitionInfo]] = defaultdict(dict) + self.transitions: defaultdict[Any, Dict[Any, TransitionInfo[T]]] = defaultdict(dict) self.predecessor: GsmNode = predecessor self.prefix_access_pair = prefix_access_pair self.data = data From 2fa02c27e21fb0f82a04a65847fc1628a905f09e Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 3 Jul 2026 15:23:46 +0200 Subject: [PATCH 08/70] intersection iterator: add option to iterate over smaller --- .../learning_algs/general_passive/GsmNode.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 54ab2d5b7a..7008da6639 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -34,13 +34,20 @@ unknown_output = None # can be set to a special value if required -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=False) -> Iterator[Tuple[Key, Val, Val]]: 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]]: From ab5f8da823790176e35d6a8eb20392452f9b428a Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 3 Jul 2026 15:26:11 +0200 Subject: [PATCH 09/70] minor cleanup --- .../learning_algs/general_passive/GsmNode.py | 25 ++++++++++--------- 1 file changed, 13 insertions(+), 12 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 7008da6639..423bf5fa8c 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -210,8 +210,9 @@ def get_prefix_input(self): return self.prefix_access_pair[0] def resolve_unknown_prefix_output(self, value): - 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=True): node = self @@ -309,16 +310,16 @@ 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)] # 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": @@ -338,21 +339,21 @@ def to_automaton(self, output_behavior: OutputBehavior, transition_behavior: Tra for out_sym, target_node in transitions.items(): target_state = state_map[target_node.target] count = target_node.count - if AutomatonClass is MooreMachine: + 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: + elif automaton_class is Mdp: state.transitions[in_sym].append((target_state, count / total)) - elif AutomatonClass is StochasticMealyMachine: + elif automaton_class is StochasticMealyMachine: state.transitions[in_sym].append((target_state, out_sym, count / total)) - return AutomatonClass(initial_state, list(state_map.values())) + return automaton_class(initial_state, list(state_map.values())) def visualize(self, path: Union[str, pathlib.Path], output_behavior: OutputBehavior = "mealy", format: str = "dot", engine="dot", *, From 5876e3042286734eea977560d32892cc63ebdb82 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 3 Jul 2026 15:28:43 +0200 Subject: [PATCH 10/70] add options for input completion --- aalpy/learning_algs/general_passive/GsmNode.py | 16 ++++++++++++++-- 1 file changed, 14 insertions(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 423bf5fa8c..687f47d006 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -428,7 +428,13 @@ def node_naming(node: GsmNode): 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, ic_mode=None) -> List[Tuple['GsmNode', Any, Any]]: + ic_modes = ["self-loop", "sink-state", "root"] + if ic_mode is None: + ic_mode = ic_modes[0] + if ic_mode not in ic_modes: + raise ValueError(f"Invalid ic_mode {ic_mode}. Should be one of {ic_modes}") + all_nodes = self.get_all_nodes() inputs = {in_sym for node in all_nodes for in_sym in node.transitions} missing_trans = [] @@ -438,7 +444,13 @@ def make_input_complete(self) -> List[Tuple['GsmNode', Any, Any]]: if len(transitions) == 0: out_sym = node.prefix_access_pair[1] missing_trans.append((node, in_sym, out_sym)) - t_info = TransitionInfo(node, 1, None, None) + if ic_mode == "self-loop": + successor = node + elif ic_mode == "sink-state": + raise NotImplementedError() + elif ic_mode == "root": + successor = self + t_info = TransitionInfo(successor, 1, None, None) transitions[out_sym] = t_info return missing_trans From 20db16c66bf336932f0701b837880d75609396cf Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 3 Jul 2026 15:30:29 +0200 Subject: [PATCH 11/70] extract IOHandlers in own file --- .../learning_algs/general_passive/GsmNode.py | 46 +--------------- .../general_passive/IOHandler.py | 52 +++++++++++++++++++ 2 files changed, 53 insertions(+), 45 deletions(-) create mode 100644 aalpy/learning_algs/general_passive/IOHandler.py diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 687f47d006..75c9c381eb 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -1,7 +1,6 @@ import functools import math import pathlib -from abc import abstractmethod from collections import defaultdict from functools import total_ordering from typing import Dict, Any, List, Tuple, Iterable, Callable, Union, TypeVar, Iterator, Optional, Sequence, Generic @@ -10,6 +9,7 @@ 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.IOHandler import IOHandler, NoIOHandler Key = TypeVar("Key") Val = TypeVar("Val") @@ -104,50 +104,6 @@ def detect_data_format(data, check_consistency=False, guess=False): raise ValueError("ambiguous data format. data format needs to be specified explicitly.") return accepted_formats[0] -class IOHandler(Generic[T]): - @abstractmethod - def init(self, data, output_format, data_format): - ... - - @abstractmethod - def abstract(self, in_val, out_val): - ... - - @abstractmethod - def init_data(self) -> T: - ... - - @abstractmethod - def aggregate_data(self, src_node: 'GsmNode[T]', in_value, out_value, dst_node: 'GsmNode[T]'): - ... - - @abstractmethod - def merge(self, x: T, y: T) -> T: - ... - - @abstractmethod - def copy(self, x: T) -> T: - ... - -class NoIOHandler(IOHandler): - def init(self, data, output_format, data_format): - pass - - def abstract(self, in_val, out_val): - return in_val, out_val - - def init_data(self) -> T: - return None - - def aggregate_data(self, src_node: 'GsmNode[T]', in_sym, out_value, dst_node: 'GsmNode[T]'): - pass - - def merge(self, x: T, y: T) -> T: - return None - - def copy(self, x: T) -> T: - return None - # TODO maybe split this for maintainability (and perfomance?) class TransitionInfo(Generic[T]): __slots__ = ["target", "count", "original_target", "original_count"] diff --git a/aalpy/learning_algs/general_passive/IOHandler.py b/aalpy/learning_algs/general_passive/IOHandler.py new file mode 100644 index 0000000000..da0ff2c04c --- /dev/null +++ b/aalpy/learning_algs/general_passive/IOHandler.py @@ -0,0 +1,52 @@ +from abc import abstractmethod, ABC +from typing import Generic, TypeVar + +T = TypeVar("T") + +class IOHandler(Generic[T]): + @abstractmethod + def init(self, data, output_format, data_format): + ... + + def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]'): + pass + + @abstractmethod + def abstract(self, in_val, out_val): + ... + + @abstractmethod + def init_data(self) -> T: + ... + + @abstractmethod + def aggregate_data(self, src_node: 'GsmNode[T]', in_value, out_value, dst_node: 'GsmNode[T]'): + ... + + @abstractmethod + def merge(self, x: T, y: T) -> T: + ... + + @abstractmethod + def copy(self, x: T) -> T: + ... + +class NoAbstractionIOHandler(IOHandler[T], ABC): + def init(self, data, output_format, data_format): + pass + + def abstract(self, in_val, out_val): + return in_val, out_val + +class NoIOHandler(NoAbstractionIOHandler[T]): + def init_data(self) -> T: + return None + + def aggregate_data(self, src_node: 'GsmNode[T]', in_sym, out_value, dst_node: 'GsmNode[T]'): + pass + + def merge(self, x: T, y: T) -> T: + return None + + def copy(self, x: T) -> T: + return None From 042028301d174cd087f8af2a146a2c6b64de8dfe Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 3 Jul 2026 15:33:19 +0200 Subject: [PATCH 12/70] add IOHandler callback for new merge computation --- aalpy/learning_algs/general_passive/GeneralizedStateMerging.py | 1 + 1 file changed, 1 insertion(+) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index fe7d446809..e69bf17096 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -252,6 +252,7 @@ def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode if blue_node is not None: partitioning.full_mapping[blue_node] = p return p + self.data_handler.init_merge(red, blue) # rewire the blue node's parent blue_parent = update_partition(blue.predecessor, None) From 282c60fe45c1125690d1314e19e34f99790060f0 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 3 Jul 2026 15:34:12 +0200 Subject: [PATCH 13/70] add copy on write IOHandler --- .../general_passive/IOHandler.py | 25 +++++++++++++++++++ 1 file changed, 25 insertions(+) diff --git a/aalpy/learning_algs/general_passive/IOHandler.py b/aalpy/learning_algs/general_passive/IOHandler.py index da0ff2c04c..6fd39645fe 100644 --- a/aalpy/learning_algs/general_passive/IOHandler.py +++ b/aalpy/learning_algs/general_passive/IOHandler.py @@ -50,3 +50,28 @@ def merge(self, x: T, y: T) -> T: def copy(self, x: T) -> T: return None + +class CopyOnWriteIOHandler(IOHandler[T], ABC): + def __init__(self): + self.copied_on_write = set() + + def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]'): + self.copied_on_write.clear() + + def merge(self, x: T, y: T) -> T: + if id(x) not in self.copied_on_write: + x = self.copy_on_write(x) + self.copied_on_write.add(id(x)) + self.merge_into_x(x, y) + return x + + @abstractmethod + def merge_into_x(self, x: T, y: T): + pass + + @abstractmethod + def copy_on_write(self, x: T) -> T: + pass + + def copy(self, x: T) -> T: + return x From 0cdaaf0321dcecaae7c9b8a70b40c1d00c1ab170 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Mon, 6 Jul 2026 15:25:53 +0200 Subject: [PATCH 14/70] added merges per second to ProgressReport --- aalpy/learning_algs/general_passive/Instrumentation.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/Instrumentation.py b/aalpy/learning_algs/general_passive/Instrumentation.py index bd20bfb36b..36fa99efca 100644 --- a/aalpy/learning_algs/general_passive/Instrumentation.py +++ b/aalpy/learning_algs/general_passive/Instrumentation.py @@ -52,7 +52,10 @@ def print_status(self): 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}' + time_taken = round(perf_counter() - self.previous_time, 2) + mps = round(self.nr_merged_states_total / time_taken, 2) + 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): From 08eab981711f8412b3ad94f26f8b528796d8624a Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 3 Jul 2026 12:46:19 +0200 Subject: [PATCH 15/70] wip commit --- .../GeneralizedStateMerging.py | 94 +++++------- .../general_passive/GsmAlgorithms.py | 2 +- .../learning_algs/general_passive/GsmNode.py | 134 +++++++----------- .../general_passive/IOHandler.py | 91 +++++++++++- .../general_passive/Instrumentation.py | 2 +- .../general_passive/ScoreFunctionsGSM.py | 106 ++++++++++---- 6 files changed, 257 insertions(+), 172 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index e69bf17096..8793013170 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -2,7 +2,7 @@ from collections import deque from typing import Dict, Tuple, Callable, List, Optional -from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, TransitionInfo, \ +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, \ OutputBehaviorRange, TransitionBehaviorRange, intersection_iterator, unknown_output, detect_data_format, IOHandler, \ NoIOHandler from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility @@ -48,8 +48,7 @@ def __init__(self, *, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, data_handler: IOHandler = None, - compatibility_on_pta: bool = False, - compatibility_on_futures: bool = False, + use_early_verdicts: bool = False, node_order: Callable[[GsmNode, GsmNode], bool] = None, consider_only_min_blue=False, depth_first=False, @@ -68,7 +67,7 @@ def __init__(self, *, 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)) + score_calc = ScoreCalculation(hoeffding_compatibility(0.005, True)) self.score_calc: ScoreCalculation = score_calc if node_order is None: @@ -81,8 +80,7 @@ def __init__(self, *, self.data_handler = data_handler or NoIOHandler() - self.compatibility_on_pta = compatibility_on_pta - self.compatibility_on_futures = compatibility_on_futures + self.use_early_verdicts = use_early_verdicts self.consider_only_min_blue = consider_only_min_blue self.depth_first = depth_first @@ -129,8 +127,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo # get blue states blue_states = [] for r in red_states: - for _, _, t in r.transition_iterator(): - c = t.target + for _, _, c in r.transition_iterator(): if c in red_states: continue blue_states.append(c) @@ -205,48 +202,26 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo root = root.to_automaton(self.output_behavior, self.transition_behavior) return root - def _check_futures(self, red: GsmNode, blue: GsmNode) -> bool: - 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: # Compatibility check based on partitions. # assumes that blue is a tree and red is not reachable from blue 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(): + early_verdict = self.score_calc.initialize_merge(red, blue) if self.use_early_verdicts else None + if early_verdict is False: + # early reject -> return failing partitioning + return partitioning + elif early_verdict is True: + # early accept -> can manipulate nodes directly def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: return red_node else: + # uncertain -> need to construct partitioning def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: p = partitioning.full_mapping.get(red_node) # could check smaller .red_mapping? if p is None: - p = red_node.shallow_copy(self.data_handler) + p = red_node.shallow_copy(self.data_handler) # TODO maybe don't do a full copy here partitioning.full_mapping[red_node] = p partitioning.red_mapping[red_node] = p if blue_node is not None: @@ -257,7 +232,7 @@ def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode # 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 + blue_parent.transitions[blue_in_sym][blue_out_sym] = red partition = update_partition(red, None) if self.output_behavior == "moore": @@ -270,40 +245,39 @@ def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode 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 + if early_verdict is None and self.compute_local_compatibility(partition, blue) is False: + return partitioning 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) + 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 + 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) @@ -317,8 +291,7 @@ def run_GSM(data: list, *, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, data_handler: IOHandler = None, - compatibility_on_pta: bool = False, - compatibility_on_futures: bool = False, + use_early_verdicts: bool = False, node_order: Callable[[GsmNode, GsmNode], bool] = None, consider_only_min_blue=False, depth_first=False, @@ -342,9 +315,7 @@ def run_GSM(data: list, *, postprocessing: A postprocessing function applied to the learned automaton. - compatibility_on_pta: Whether compatibility is evaluated on the PTA or the current hypothesis. - - compatibility_on_futures: Whether compatibility is evaluated using the futures of both states or all partition information. + use_early_verdicts: Whether to use potential early verdicts from the score object. Defaults to False. node_order: Order in which merge candidates are considered. Defaults to short-lex. @@ -368,8 +339,7 @@ def run_GSM(data: list, *, pta_preprocessing=pta_preprocessing, postprocessing=postprocessing, data_handler=data_handler, - compatibility_on_pta=compatibility_on_pta, - compatibility_on_futures=compatibility_on_futures, + use_early_verdicts=use_early_verdicts, 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 02d85d4f74..65f701596f 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -146,7 +146,7 @@ def __init__(self, epsilon): self.ioa_compatibility = hoeffding_compatibility(epsilon) self.evidence = 0 - def reset(self): + def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: self.evidence = 0 def local_compatibility(self, a: GsmNode, b: GsmNode): diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 75c9c381eb..e4481c0fe1 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -1,5 +1,4 @@ import functools -import math import pathlib from collections import defaultdict from functools import total_ordering @@ -9,7 +8,12 @@ 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.IOHandler import IOHandler, NoIOHandler +from aalpy.learning_algs.general_passive.IOHandler import ( + IOHandler, + NoIOHandler, + StochasticData, + CountData, +) Key = TypeVar("Key") Val = TypeVar("Val") @@ -104,17 +108,6 @@ def detect_data_format(data, check_consistency=False, guess=False): 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(Generic[T]): - __slots__ = ["target", "count", "original_target", "original_count"] - - def __init__(self, target, count, original_target, original_count): - self.target: 'GsmNode[T]' = target - self.count: int = count - self.original_target: 'GsmNode[T]' = 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(Generic[T]): @@ -131,7 +124,7 @@ class GsmNode(Generic[T]): def __init__(self, prefix_access_pair, predecessor: 'GsmNode[T]' = None, data: T = None): # TODO (data-ext) check all invocations # TODO try single dict - self.transitions: defaultdict[Any, Dict[Any, TransitionInfo[T]]] = defaultdict(dict) + self.transitions: defaultdict[Any, Dict[Any, GsmNode[T]]] = defaultdict(dict) self.predecessor: GsmNode = predecessor self.prefix_access_pair = prefix_access_pair self.data = data @@ -188,28 +181,25 @@ def get_root(self): current = current.predecessor return current - def get_or_create_transitions(self, in_sym) -> Dict[Any, TransitionInfo]: + def get_or_create_transitions(self, in_sym) -> Dict[Any, 'GsmNode[T]']: 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]']]: 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, data_handler: IOHandler) -> 'GsmNode': + def shallow_copy(self, data_handler: IOHandler) -> 'GsmNode[T]': node = GsmNode(self.prefix_access_pair, self.predecessor, data_handler.copy(self.data)) 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 + node.transitions[in_sym] = t.copy() return node - def get_by_prefix(self, seq: IOTrace) -> Optional['GsmNode']: + def get_by_prefix(self, seq: IOTrace) -> Optional['GsmNode[T]']: node: GsmNode = self for in_sym, out_sym in seq: if in_sym is None: # ignore initial transition of Node.get_prefix() @@ -217,18 +207,16 @@ def get_by_prefix(self, seq: IOTrace) -> Optional['GsmNode']: trans = node.transitions.get(in_sym) if trans is None: return None - t_info = trans.get(out_sym) - if t_info is None: + node = trans.get(out_sym) + if node is None: return None - node = t_info.target return node - def get_all_nodes(self) -> List['GsmNode']: + def get_all_nodes(self) -> List['GsmNode[T]']: result = [self] backing_set = {self} for state in result: - for _, _, transition in state.transition_iterator(): - child = transition.target + for _, _, child in state.transition_iterator(): if child not in backing_set: backing_set.add(child) result.append(child) @@ -239,8 +227,7 @@ def is_tree(self): backing_set = {self} while len(q) != 0: current = q.pop(0) - for _, _, transition in current.transition_iterator(): - child = transition.target + for _, _, child in current.transition_iterator(): if child in backing_set: return False q.append(child) @@ -290,11 +277,13 @@ 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 + target_state = state_map[target_node] if automaton_class is MooreMachine: state.transitions[in_sym] = target_state elif automaton_class is MealyMachine: @@ -305,9 +294,9 @@ def to_automaton(self, output_behavior: OutputBehavior, transition_behavior: Tra elif automaton_class is Onfsm: state.transitions[in_sym].append((out_sym, target_state)) elif automaton_class is Mdp: - state.transitions[in_sym].append((target_state, count / total)) + 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, count / total)) + state.transitions[in_sym].append((target_state, out_sym, prob_info[in_sym][out_sym])) return automaton_class(initial_state, list(state_map.values())) @@ -327,19 +316,23 @@ def visualize(self, path: Union[str, pathlib.Path], output_behavior: OutputBehav if trans_props is None: trans_props = dict() if state_label is None: - if output_behavior == "moore": - def state_label(node: GsmNode): - return f'{node.get_prefix_output()} {node.count()}' - else: - def state_label(node: GsmNode): - return f'{sum(t.count for _, _, t in node.transition_iterator())}' + def state_label(node: GsmNode): + 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, out_sym): - return f'{in_sym} [{node.transitions[in_sym][out_sym].count}]' - else: - def trans_label(node: GsmNode, in_sym, out_sym): - return f'{in_sym} / {out_sym} [{node.transitions[in_sym][out_sym].count}]' + def trans_label(node: GsmNode, in_sym, out_sym): + 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): return "black" if trans_color is None: @@ -370,7 +363,7 @@ def node_naming(node: GsmNode): 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 @@ -406,8 +399,7 @@ def make_input_complete(self, ic_mode=None) -> List[Tuple['GsmNode', Any, Any]]: raise NotImplementedError() elif ic_mode == "root": successor = self - t_info = TransitionInfo(successor, 1, None, None) - transitions[out_sym] = t_info + transitions[out_sym] = successor return missing_trans def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): @@ -415,14 +407,10 @@ def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): for in_value, out_value in trace: in_sym, out_sym = data_handler.abstract(in_value, out_value) transitions = curr_node.transitions[in_sym] - info = transitions.get(out_sym) - if info is None: + node = transitions.get(out_sym) + if node is None: node = GsmNode((in_sym, out_sym), curr_node, data_handler.init_data()) - transitions[out_sym] = TransitionInfo(node, 1, node, 1) - else: - info.count += 1 - info.original_count += 1 - node = info.target + transitions[out_sym] = node data_handler.aggregate_data(curr_node, in_value, out_value, node) curr_node = node @@ -438,16 +426,12 @@ def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = for in_value in inputs: in_sym, out_sym = data_handler.abstract(in_value, None) transitions = curr_node.transitions[in_sym] - t_infos = list(transitions.values()) - if len(t_infos) == 0: + successors = list(transitions.values()) + if len(successors) == 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 + transitions[unknown_output] = node + elif len(successors) == 1: + node = successors[0] else: # This should never happen raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") @@ -517,8 +501,8 @@ def deterministic_compatible(self, other: 'GsmNode'): def is_moore(self): 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 @@ -528,18 +512,4 @@ def moore_compatible(self, other: 'GsmNode'): oo = other.get_prefix_output() return so == oo or so is unknown_output or oo is unknown_output - def local_log_likelihood_contribution(self): - 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): - 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) diff --git a/aalpy/learning_algs/general_passive/IOHandler.py b/aalpy/learning_algs/general_passive/IOHandler.py index 6fd39645fe..be21cdacba 100644 --- a/aalpy/learning_algs/general_passive/IOHandler.py +++ b/aalpy/learning_algs/general_passive/IOHandler.py @@ -1,5 +1,8 @@ +import math from abc import abstractmethod, ABC -from typing import Generic, TypeVar +from collections import defaultdict +from copy import copy +from typing import Generic, TypeVar, Any T = TypeVar("T") @@ -75,3 +78,89 @@ def copy_on_write(self, x: T) -> T: def copy(self, x: T) -> T: return x + +ProbabilityDict = dict[Any, dict[Any, float]] + +class StochasticData(ABC): + @abstractmethod + def get_probabilities(self) -> ProbabilityDict: pass + +CountDict = dict[Any, dict[Any, int]] + +class CountData(StochasticData): + def __init__(self): + # TODO get rid of this indirection + self.transition_count: CountDict = CountData.default_ctor() + + 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()) + + @staticmethod + def default_ctor(): + return defaultdict(lambda: defaultdict(int)) + + @staticmethod + def merge(x: CountDict, y: CountDict) -> CountDict: + for in_sym, y_o_dict in y.items(): + x_o_dict = x.get(in_sym, None) + if x_o_dict is None: + x[in_sym] = y_o_dict + continue + for out_sym, count in y_o_dict.items(): + x_o_dict[out_sym] += count + return x + + 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 + +class CountHandler(NoAbstractionIOHandler[CountData], CopyOnWriteIOHandler): + def init_data(self) -> CountData: + return CountData() + + def aggregate_data(self, src_node: 'GsmNode[CountData]', in_value, out_value, dst_node: 'GsmNode[CountData]'): + src_node.data.transition_count[in_value][out_value] += 1 + + def merge_into_x(self, x: CountData, y: CountData): + CountData.merge(x.transition_count, y.transition_count) + + def copy_on_write(self, x: CountData) -> CountData: + new_x = copy(x) + new_x.transition_count = CountData.default_ctor() + for in_sym, trans in x.transition_count.items(): + new_x.transition_count[in_sym] = trans.copy() + return x + +class ShadowPTAData: + def __init__(self): + self.shadow_pta: ShadowPTA = defaultdict(dict) + +ShadowPTA = dict[Any, dict[Any, 'GsmNode']] +class CountOnPTAData(ShadowPTAData, CountData): + def __init__(self): + ShadowPTAData.__init__(self) + CountData.__init__(self) + self.pta_count: CountDict = CountData.default_ctor() + +class CountOnPTAHandler(CountHandler): + def init_data(self) -> CountOnPTAData: + return CountOnPTAData() + + def aggregate_data(self, src_node: 'GsmNode[CountOnPTAData]', in_value, out_value, dst_node: 'GsmNode[CountOnPTAData]'): + src_node.data.transition_count[in_value][out_value] += 1 + 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/Instrumentation.py b/aalpy/learning_algs/general_passive/Instrumentation.py index 36fa99efca..9741789b81 100644 --- a/aalpy/learning_algs/general_passive/Instrumentation.py +++ b/aalpy/learning_algs/general_passive/Instrumentation.py @@ -51,7 +51,7 @@ def pta_construction_done(self, root): def print_status(self): 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: + if 0 < self.lvl and not self.gsm.use_early_verdicts: time_taken = round(perf_counter() - self.previous_time, 2) mps = round(self.nr_merged_states_total / time_taken, 2) remaining_merges = self.pta_size - self.nr_red_states - self.nr_merged_states_total diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 0f1cd977ae..f4640601f2 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -1,7 +1,10 @@ +from collections import deque from math import sqrt, log -from typing import Callable, Dict, List, Iterable, Any +from typing import Callable, Dict, List, Iterable, Any, Tuple -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.IOHandler import ShadowPTAData LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool] ScoreFunction = Callable[[Dict[GsmNode, GsmNode]], Any] @@ -19,8 +22,8 @@ def __init__(self, local_compatibility: LocalCompatibilityFunction = None, score if not hasattr(self, "score_function"): self.score_function: ScoreFunction = score_function or self.default_score_function - def reset(self): - pass + def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: + return None @staticmethod def default_local_compatibility(a: GsmNode, b: GsmNode): @@ -39,26 +42,74 @@ def has_local_compatibility(self): def hoeffding_compatibility(eps, compare_original=True) -> LocalCompatibilityFunction: eps_fact = sqrt(0.5 * log(2 / eps)) - count_name = "original_count" if compare_original else "count" - transition_dummy = TransitionInfo(None, 0, None, 0) + count_dict_name = "pta_count" if compare_original else "transition_count" - def similar(a: GsmNode, b: GsmNode): + def similar(a: GsmNode[CountData], b: GsmNode[CountData]): # iterate over inputs that are common to both states - for in_sym, a_trans, b_trans in intersection_iterator(a.transitions, b.transitions): + for in_sym, a_trans, b_trans in intersection_iterator(getattr(a.data, count_dict_name), getattr(b.data, count_dict_name)): # 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): if abs(ac / a_total - bc / b_total) > threshold: return False return True return similar +class CheckFutureScore(ScoreCalculation): + def __init__(self, + local_compatibility: LocalCompatibilityFunction = None, + score_function: ScoreFunction = None, + output_behavior = "moore", + transition_behavior = "deterministic", + compatibility_on_pta = False, + depth_first = False + ): + super().__init__(local_compatibility, score_function) + # TODO: auto init from GSM + self.output_behavior = output_behavior + self.transition_behavior = transition_behavior + self.compatibility_on_pta = compatibility_on_pta + self.depth_first = depth_first + + def compute_local_compatibility(self, a: GsmNode, b: GsmNode): + 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.local_compatibility(a, b) + + def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: + 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() + + if self.compute_local_compatibility(red, blue) is False: + return False + + if not self.compatibility_on_pta: + 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): + q.append((red_child, blue_child)) + else: + 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): + for out_sym, red_child, blue_child in intersection_iterator(red_trans, blue_trans): + q.append((red_child,blue_child)) + + if self.has_score_function(): + return None + return True class ScoreWithKTail(ScoreCalculation): """Applies k-Tails to a compatibility function: Compatibility is only evaluated up to a certain depth k.""" @@ -70,9 +121,9 @@ def __init__(self, other_score: ScoreCalculation, k: int): self.depth_offset = None - def reset(self): - self.other_score.reset() + def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool: self.depth_offset = None + return self.other_score.initialize_merge(red, blue) def local_compatibility(self, a: GsmNode, b: GsmNode): # assuming b is tree shaped. @@ -96,9 +147,9 @@ def __init__(self, other_score: ScoreCalculation, sink_cond: Callable[[GsmNode], self.is_first = True - def reset(self): - self.other_score.reset() + def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: self.is_first = True + return self.other_score.initialize_merge(red, blue) def local_compatibility(self, a: GsmNode, b: GsmNode): if self.is_first: @@ -124,9 +175,14 @@ def __init__(self, scores: List[ScoreCalculation], aggregate_compatibility: Aggr self.aggregate_compatibility = aggregate_compatibility or self.default_aggregate_compatibility self.aggregate_score = aggregate_score or self.default_aggregate_score - def reset(self): + def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: + all_true = True for score in self.scores: - score.reset() + verdict = score.initialize_merge(red, blue) + if verdict is False: + return False + all_true = all_true and verdict + return all_true or None def local_compatibility(self, a: GsmNode, b: GsmNode): return self.aggregate_compatibility(score.local_compatibility(a, b) for score in self.scores) @@ -165,12 +221,12 @@ def fun(part: Dict[GsmNode, GsmNode]): return fun -def differential_info(part: Dict[GsmNode, GsmNode]): +def differential_info(part: Dict[GsmNode[CountData], GsmNode[CountData]]): 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) + 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.transition_iterator()) num_params_new = sum(1 for node in relevant_nodes_new for _ in node.transition_iterator()) @@ -204,13 +260,13 @@ def score(part: Dict[GsmNode, GsmNode]): def EDSM_frequency_score(min_evidence=-1) -> ScoreFunction: - def score(part: Dict[GsmNode, GsmNode]): + def score(part: Dict[GsmNode[CountData], GsmNode[CountData]]): 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 From ca4bfef2c119e16fa21827a3c0ca48161ffa9dc8 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Thu, 16 Jul 2026 11:26:01 +0200 Subject: [PATCH 16/70] minor fixes --- aalpy/learning_algs/general_passive/GsmNode.py | 6 ++++-- aalpy/learning_algs/general_passive/IOHandler.py | 2 ++ aalpy/learning_algs/general_passive/Instrumentation.py | 2 ++ aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py | 2 +- 4 files changed, 9 insertions(+), 3 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index e4481c0fe1..484c3924fa 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -35,7 +35,7 @@ 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 def intersection_iterator(a: Dict[Key, Val], b: Dict[Key, Val], sort_by_length=False) -> Iterator[Tuple[Key, Val, Val]]: @@ -209,7 +209,9 @@ def get_by_prefix(self, seq: IOTrace) -> Optional['GsmNode[T]']: return None node = trans.get(out_sym) if node is None: - return None + node = trans.get(unknown_output) + if node is None: + return None return node def get_all_nodes(self) -> List['GsmNode[T]']: diff --git a/aalpy/learning_algs/general_passive/IOHandler.py b/aalpy/learning_algs/general_passive/IOHandler.py index be21cdacba..4e1dc3dfac 100644 --- a/aalpy/learning_algs/general_passive/IOHandler.py +++ b/aalpy/learning_algs/general_passive/IOHandler.py @@ -161,6 +161,8 @@ def init_data(self) -> CountOnPTAData: return CountOnPTAData() def aggregate_data(self, src_node: 'GsmNode[CountOnPTAData]', in_value, out_value, dst_node: 'GsmNode[CountOnPTAData]'): + if src_node is None: + return src_node.data.transition_count[in_value][out_value] += 1 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/Instrumentation.py b/aalpy/learning_algs/general_passive/Instrumentation.py index 9741789b81..d544e8949a 100644 --- a/aalpy/learning_algs/general_passive/Instrumentation.py +++ b/aalpy/learning_algs/general_passive/Instrumentation.py @@ -99,6 +99,8 @@ def log_promote(self, new_red: GsmNode): if old_red is None: self.map[node] = new_red self.log.append(("promote", new_red_prefix)) + elif node is None: + self.log.append(("broken promote", new_red_prefix)) elif old_red is not new_red: print(f"Erroneous promotion detected:") print(f" Ground truth: {node.get_prefix()}") diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index f4640601f2..0b2d5378f4 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -53,7 +53,7 @@ def similar(a: GsmNode[CountData], b: GsmNode[CountData]): 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, ac, bc in union_iterator(a_trans, b_trans): + 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 From eb8cd64194e4f47c97fa52c2499dd97b11685eb0 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 8 Jul 2026 16:59:56 +0200 Subject: [PATCH 17/70] add simplified version of GsmNode.transition_iterator --- .../general_passive/GeneralizedStateMerging.py | 2 +- aalpy/learning_algs/general_passive/GsmNode.py | 8 ++++++-- aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py | 4 ++-- 3 files changed, 9 insertions(+), 5 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 8793013170..eab3659017 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -127,7 +127,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo # get blue states blue_states = [] for r in red_states: - for _, _, c in r.transition_iterator(): + for c in r.child_iterator(): if c in red_states: continue blue_states.append(c) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 484c3924fa..e2a9b3a6e2 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -193,6 +193,10 @@ def transition_iterator(self) -> Iterable[Tuple[Any, Any, 'GsmNode[T]']]: for out_sym, node in transitions.items(): yield in_sym, out_sym, node + def child_iterator(self) -> Iterable['GsmNode[T]']: + for transitions in self.transitions.values(): + yield from transitions.values() + def shallow_copy(self, data_handler: IOHandler) -> 'GsmNode[T]': node = GsmNode(self.prefix_access_pair, self.predecessor, data_handler.copy(self.data)) for in_sym, t in self.transitions.items(): @@ -218,7 +222,7 @@ def get_all_nodes(self) -> List['GsmNode[T]']: result = [self] backing_set = {self} for state in result: - for _, _, child in state.transition_iterator(): + for child in state.child_iterator(): if child not in backing_set: backing_set.add(child) result.append(child) @@ -229,7 +233,7 @@ def is_tree(self): backing_set = {self} while len(q) != 0: current = q.pop(0) - for _, _, child in current.transition_iterator(): + for child in current.child_iterator(): if child in backing_set: return False q.append(child) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 0b2d5378f4..fbf07bdfef 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -228,8 +228,8 @@ def differential_info(part: Dict[GsmNode[CountData], GsmNode[CountData]]): 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.transition_iterator()) - num_params_new = sum(1 for node in relevant_nodes_new for _ in node.transition_iterator()) + 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 From ab4deac4d29e9c320dbe6598636b116c6b8fc027 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 8 Jul 2026 17:05:02 +0200 Subject: [PATCH 18/70] slight improvement for blue state computation --- .../general_passive/GeneralizedStateMerging.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index eab3659017..26c231dcf8 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -117,6 +117,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo # sorted list of states already considered red_states = [root] + red_states_backing_set = {root} partition_candidates: Dict[Tuple[GsmNode, GsmNode], Partitioning] = dict() while True: @@ -125,10 +126,12 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo red_states.sort(key=self.node_order) # get blue states + # TODO: eliminate explicit blue set construction. + # should be constructed from merge/promotion info blue_states = [] for r in red_states: for c in r.child_iterator(): - if c in red_states: + if c in red_states_backing_set: continue blue_states.append(c) if self.consider_only_min_blue and self.node_order is GsmNode.default_order: @@ -170,6 +173,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo # 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) + red_states_backing_set.add(blue_state) instrumentation.log_promote(blue_state) promotion = True break From 9096343cc259db33c20efae2c81f837bfeee5e6c Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Thu, 16 Jul 2026 12:04:05 +0200 Subject: [PATCH 19/70] lazy copy transition info during merge creation --- .../GeneralizedStateMerging.py | 23 ++++++++++++++++--- .../learning_algs/general_passive/GsmNode.py | 6 ----- 2 files changed, 20 insertions(+), 9 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 26c231dcf8..033eb5db9b 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -1,5 +1,6 @@ import functools from collections import deque +from copy import copy from typing import Dict, Tuple, Callable, List, Optional from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, \ @@ -220,23 +221,39 @@ def _partition_from_merge(self, red: GsmNode, blue: GsmNode) -> Partitioning: # early accept -> can manipulate nodes directly def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: return red_node + + def get_partition_trans(part: GsmNode, in_symbol): + return part.transitions[in_symbol] else: # uncertain -> need to construct partitioning def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: p = partitioning.full_mapping.get(red_node) # could check smaller .red_mapping? if p is None: - p = red_node.shallow_copy(self.data_handler) # TODO maybe don't do a full copy here + p = copy(red_node) + p.data = self.data_handler.copy(red_node.data) + p.transitions = red_node.transitions.copy() + partitioning.full_mapping[red_node] = p partitioning.red_mapping[red_node] = p if blue_node is not None: partitioning.full_mapping[blue_node] = p return p + + 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 + self.data_handler.init_merge(red, blue) # 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] = red + get_partition_trans(blue_parent, blue_in_sym)[blue_out_sym] = red partition = update_partition(red, None) if self.output_behavior == "moore": @@ -256,7 +273,7 @@ def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode # create implied merges for all common successors for in_sym, blue_transitions in blue.transitions.items(): - partition_transitions = partition.transitions[in_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 diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index e2a9b3a6e2..514eb22e09 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -197,12 +197,6 @@ def child_iterator(self) -> Iterable['GsmNode[T]']: for transitions in self.transitions.values(): yield from transitions.values() - def shallow_copy(self, data_handler: IOHandler) -> 'GsmNode[T]': - node = GsmNode(self.prefix_access_pair, self.predecessor, data_handler.copy(self.data)) - for in_sym, t in self.transitions.items(): - node.transitions[in_sym] = t.copy() - return node - def get_by_prefix(self, seq: IOTrace) -> Optional['GsmNode[T]']: node: GsmNode = self for in_sym, out_sym in seq: From e4759e44efcaca3a8b082370a85bb32f27065cf0 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Thu, 16 Jul 2026 11:25:29 +0200 Subject: [PATCH 20/70] eliminated explicit blue state construction --- .../GeneralizedStateMerging.py | 54 +++++++++++-------- 1 file changed, 32 insertions(+), 22 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 033eb5db9b..d342700e63 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -19,6 +19,7 @@ def __init__(self, red: GsmNode, blue: GsmNode): self.score = False self.red_mapping: Dict[GsmNode, GsmNode] = dict() self.full_mapping: Dict[GsmNode, GsmNode] = dict() + self.new_blue = [] class Instrumentation: @@ -119,36 +120,29 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo # sorted list of states already considered 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 - # TODO: eliminate explicit blue set construction. - # should be constructed from merge/promotion info - blue_states = [] - for r in red_states: - for c in r.child_iterator(): - if c in red_states_backing_set: - 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 + + 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)] + 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 GsmNode.default_order: - blue_states.sort(key=self.node_order) + blue_states_to_consider.sort(key=self.node_order) + + # sort red 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) # loop over blue states promotion = False - for blue_state in blue_states: + for blue_state in blue_states_to_consider: # FUTURE: Parallelize # FUTURE: Save partitions? @@ -159,7 +153,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo 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) + partition = self._partition_from_merge(red_state, blue_state, red_states_backing_set) if partition.score is True: perfect_partitioning = partition break @@ -175,6 +169,8 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo if all(part.score is False for part in current_candidates.values()): red_states.append(blue_state) red_states_backing_set.add(blue_state) + blue_states.remove(blue_state) + blue_states.extend(blue_state.child_iterator()) instrumentation.log_promote(blue_state) promotion = True break @@ -195,6 +191,8 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo real_node.predecessor = partition_node.predecessor real_node.data = partition_node.data real_node.prefix_access_pair = partition_node.prefix_access_pair + blue_states.extend(best_candidate.new_blue) + blue_states.remove(best_candidate.blue) instrumentation.log_merge(best_candidate) # FUTURE: optimizations for compatibility tests where merges can be orthogonal # FUTURE: caching for aggregating compatibility tests @@ -207,7 +205,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo root = root.to_automaton(self.output_behavior, self.transition_behavior) return root - def _partition_from_merge(self, red: GsmNode, blue: GsmNode) -> Partitioning: + def _partition_from_merge(self, red: GsmNode, blue: GsmNode, red_nodes: set[GsmNode]) -> Partitioning: # Compatibility check based on partitions. # assumes that blue is a tree and red is not reachable from blue @@ -219,6 +217,7 @@ def _partition_from_merge(self, red: GsmNode, blue: GsmNode) -> Partitioning: return partitioning elif early_verdict is True: # early accept -> can manipulate nodes directly + red_partitions = red_nodes def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: return red_node @@ -226,15 +225,23 @@ def get_partition_trans(part: GsmNode, in_symbol): return part.transitions[in_symbol] else: # uncertain -> need to construct partitioning + red_partitions: set[GsmNode] = set() def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: p = partitioning.full_mapping.get(red_node) # could check smaller .red_mapping? if p is None: + # there is no partition yet for the 'red' node -> lazily copy p = copy(red_node) p.data = self.data_handler.copy(red_node.data) 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 @@ -295,7 +302,10 @@ def get_partition_trans(part: GsmNode, in_symbol): 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 + # 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_successor, None) From ebe9cb8850cb745863429e72ade751d050f16f01 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Thu, 16 Jul 2026 15:48:08 +0200 Subject: [PATCH 21/70] less comparisons for moore machines --- .../general_passive/GeneralizedStateMerging.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index d342700e63..c0d9ec7f61 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -88,8 +88,6 @@ def __init__(self, *, self.depth_first = depth_first def compute_local_compatibility(self, a: GsmNode, b: GsmNode): - 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) @@ -211,6 +209,11 @@ def _partition_from_merge(self, red: GsmNode, blue: GsmNode, red_nodes: set[GsmN partitioning = Partitioning(red, blue) + # for Moore machines the outputs have to match. but only once since Moore-ness is preserved for implied merges + if self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): + return partitioning + + # check whether there is an early verdict and adapt helper functions accordingly early_verdict = self.score_calc.initialize_merge(red, blue) if self.use_early_verdicts else None if early_verdict is False: # early reject -> return failing partitioning From 79de36b1b1fbda4c61feeb17b6a666ef116e5f5d Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Mon, 20 Jul 2026 10:50:10 +0200 Subject: [PATCH 22/70] minor fixes --- aalpy/learning_algs/general_passive/GsmNode.py | 4 ++-- aalpy/learning_algs/general_passive/Instrumentation.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 514eb22e09..e5b36d979c 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -391,14 +391,14 @@ def make_input_complete(self, ic_mode=None) -> 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] - missing_trans.append((node, in_sym, out_sym)) if ic_mode == "self-loop": successor = node elif ic_mode == "sink-state": raise NotImplementedError() elif ic_mode == "root": successor = self + out_sym = successor.prefix_access_pair[1] + missing_trans.append((node, in_sym, out_sym)) transitions[out_sym] = successor return missing_trans diff --git a/aalpy/learning_algs/general_passive/Instrumentation.py b/aalpy/learning_algs/general_passive/Instrumentation.py index d544e8949a..3d3794daf2 100644 --- a/aalpy/learning_algs/general_passive/Instrumentation.py +++ b/aalpy/learning_algs/general_passive/Instrumentation.py @@ -53,7 +53,7 @@ def print_status(self): print_str = reset_char + f'Current automaton size: {self.nr_red_states}' if 0 < self.lvl and not self.gsm.use_early_verdicts: time_taken = round(perf_counter() - self.previous_time, 2) - mps = round(self.nr_merged_states_total / time_taken, 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="") From 5f210af26e1e42e0090554816be53e2eb2c3495a Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 16:37:40 +0200 Subject: [PATCH 23/70] introduced two phase merging to GSM --- .../GeneralizedStateMerging.py | 114 +++++++++++------- .../general_passive/IOHandler.py | 10 +- .../general_passive/ScoreFunctionsGSM.py | 2 +- 3 files changed, 80 insertions(+), 46 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index c0d9ec7f61..54f16908a1 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -9,6 +9,8 @@ from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility +defer_merge = object() + # TODO add option for making checking of futures and partition non mutual exclusive? # Easiest done by adding a new method / field to ScoreCalculation @@ -16,10 +18,11 @@ class Partitioning: def __init__(self, red: GsmNode, blue: GsmNode): self.red: GsmNode = red self.blue: GsmNode = blue - self.score = False + self.score = None self.red_mapping: Dict[GsmNode, GsmNode] = dict() self.full_mapping: Dict[GsmNode, GsmNode] = dict() self.new_blue = [] + self.remaining_merges = None class Instrumentation: @@ -149,13 +152,14 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo perfect_partitioning = None red_state = None 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, red_states_backing_set) - 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) + if partitioning.score is True: + perfect_partitioning = partitioning break - current_candidates[red_state] = partition + current_candidates[red_state] = partitioning assert red_state is not None # partition with perfect score found: don't consider anything else @@ -189,6 +193,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo 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) @@ -203,30 +208,28 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo root = root.to_automaton(self.output_behavior, self.transition_behavior) return root - def _partition_from_merge(self, red: GsmNode, blue: GsmNode, red_nodes: set[GsmNode]) -> Partitioning: + def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNode], first_pass): # Compatibility check based on partitions. # assumes that blue is a tree and red is not reachable from blue + # works in two passes: + # - first pass: create partial partitioning sufficient for score calculation + # - second pass: merge has been accepted, partitioning needs to be completed + + red = partitioning.red + blue = partitioning.blue + + if first_pass: + # for Moore machines the outputs have to match. but only once since Moore-ness is preserved for implied merges + if self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): + partitioning.score = False + return + # check whether there is an early verdict and adapt helper functions accordingly + # TODO maybe split init from early verdict and also call init (maybe with first_pass as an argument) in both cases + partitioning.score = self.score_calc.initialize_merge(red, blue) if self.use_early_verdicts else None + if partitioning.score is not None: + return + partitioning.remaining_merges = [] - partitioning = Partitioning(red, blue) - - # for Moore machines the outputs have to match. but only once since Moore-ness is preserved for implied merges - if self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): - return partitioning - - # check whether there is an early verdict and adapt helper functions accordingly - early_verdict = self.score_calc.initialize_merge(red, blue) if self.use_early_verdicts else None - if early_verdict is False: - # early reject -> return failing partitioning - return partitioning - elif early_verdict is True: - # early accept -> can manipulate nodes directly - red_partitions = red_nodes - def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: - return red_node - - def get_partition_trans(part: GsmNode, in_symbol): - return part.transitions[in_symbol] - else: # uncertain -> need to construct partitioning red_partitions: set[GsmNode] = set() def update_partition(red_node: GsmNode, blue_node: Optional[GsmNode]) -> GsmNode: @@ -257,27 +260,56 @@ def get_partition_trans(part: GsmNode, in_symbol): 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: Optional[GsmNode]) -> GsmNode: + return red_node + + def get_partition_trans(part: GsmNode, in_symbol): + return part.transitions[in_symbol] + else: + return + + self.data_handler.init_merge(red, blue, first_pass) + q: deque[Tuple[GsmNode, GsmNode]] = deque() - self.data_handler.init_merge(red, blue) + 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 and partitioning.score is None) or (not first_pass and partitioning.score is not None) - # 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 + # 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 - partition = update_partition(red, None) - if self.output_behavior == "moore": - partition.resolve_unknown_prefix_output(blue_out_sym) + # 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) # 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 early_verdict is None and self.compute_local_compatibility(partition, blue) is False: - return partitioning + if first_pass: + local_compat = self.compute_local_compatibility(partition, blue) + if local_compat is False: + partitioning.score = False + return + if local_compat is defer_merge: + partitioning.remaining_merges.append((red, blue)) + continue partition.data = self.data_handler.merge(partition.data, blue.data) @@ -314,8 +346,8 @@ def get_partition_trans(part: GsmNode, in_symbol): 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, *, diff --git a/aalpy/learning_algs/general_passive/IOHandler.py b/aalpy/learning_algs/general_passive/IOHandler.py index 4e1dc3dfac..c34f37e0d7 100644 --- a/aalpy/learning_algs/general_passive/IOHandler.py +++ b/aalpy/learning_algs/general_passive/IOHandler.py @@ -11,7 +11,7 @@ class IOHandler(Generic[T]): def init(self, data, output_format, data_format): ... - def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]'): + def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): pass @abstractmethod @@ -58,11 +58,13 @@ class CopyOnWriteIOHandler(IOHandler[T], ABC): def __init__(self): self.copied_on_write = set() - def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]'): - self.copied_on_write.clear() + def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): + self.first_pass = first_pass + if first_pass: + self.copied_on_write.clear() def merge(self, x: T, y: T) -> T: - if id(x) not in self.copied_on_write: + if self.first_pass and id(x) not in self.copied_on_write: x = self.copy_on_write(x) self.copied_on_write.add(id(x)) self.merge_into_x(x, y) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index fbf07bdfef..2af3799947 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -22,7 +22,7 @@ def __init__(self, local_compatibility: LocalCompatibilityFunction = None, score if not hasattr(self, "score_function"): self.score_function: ScoreFunction = score_function or self.default_score_function - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: + def initialize_merge(self, red: GsmNode, blue: GsmNode) -> Any: return None @staticmethod From f1f55dd04c57331614ed33ec08cdbd132b35987e Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 11:44:37 +0200 Subject: [PATCH 24/70] shift responsibility to ensure determinism to user (with sensible default) --- .../general_passive/GeneralizedStateMerging.py | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 54f16908a1..54f6c07bf6 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -1,4 +1,5 @@ import functools +import warnings from collections import deque from copy import copy from typing import Dict, Tuple, Callable, List, Optional @@ -68,7 +69,7 @@ def __init__(self, *, if score_calc is None: if transition_behavior == "deterministic": - score_calc = ScoreCalculation() + score_calc = ScoreCalculation(GsmNode.deterministic_compatible) elif transition_behavior == "nondeterministic" : raise ValueError("Missing score_calc for nondeterministic transition behavior. No default available.") elif transition_behavior == "stochastic" : @@ -90,11 +91,6 @@ def __init__(self, *, self.consider_only_min_blue = consider_only_min_blue self.depth_first = depth_first - def compute_local_compatibility(self, a: GsmNode, b: GsmNode): - 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, convert=True, instrumentation: Instrumentation=None, data_format=None): @@ -116,7 +112,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo if self.transition_behavior == "deterministic": if not root.is_deterministic(): - raise ValueError("required deterministic automaton but input data is nondeterministic") + warnings.warn("required deterministic automaton but input data is nondeterministic") # sorted list of states already considered red_states = [root] @@ -303,7 +299,7 @@ def get_partition_trans(part: GsmNode, in_symbol): partition = update_partition(red, blue) if first_pass: - local_compat = self.compute_local_compatibility(partition, blue) + local_compat = self.score_calc.local_compatibility(partition, blue) if local_compat is False: partitioning.score = False return From c4dd5970993e17d1fe1ec4ca45d2ca077e886de3 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 13:14:10 +0200 Subject: [PATCH 25/70] add option to rank candidates for promotion in GSM --- .../GeneralizedStateMerging.py | 77 +++++++++++-------- .../general_passive/ScoreFunctionsGSM.py | 3 + 2 files changed, 46 insertions(+), 34 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 54f6c07bf6..3a0e2250d9 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -2,7 +2,7 @@ import warnings from collections import deque from copy import copy -from typing import Dict, Tuple, Callable, List, Optional +from typing import Dict, Tuple, Callable, List, Optional, Any from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, \ OutputBehaviorRange, TransitionBehaviorRange, intersection_iterator, unknown_output, detect_data_format, IOHandler, \ @@ -114,7 +114,7 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo if not root.is_deterministic(): 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()) @@ -138,7 +138,8 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo red_states.sort(key=self.node_order) # loop over blue states - promotion = False + best_promotion_candidate = None + best_promotion_score = None for blue_state in blue_states_to_consider: # FUTURE: Parallelize # FUTURE: Save partitions? @@ -156,46 +157,54 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo perfect_partitioning = partitioning break current_candidates[red_state] = partitioning - 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} 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) - red_states_backing_set.add(blue_state) - blue_states.remove(blue_state) - blue_states.extend(blue_state.child_iterator()) - instrumentation.log_promote(blue_state) - promotion = True - 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) + new_candidates = (((red, blue_state), part) for red, part in current_candidates.items()) 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.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) - # FUTURE: optimizations for compatibility tests where merges can be orthogonal - # FUTURE: caching for aggregating compatibility tests - partition_candidates.clear() + # no merge candidates for this blue state -> promotion candidate + if all(part.score is False for part in current_candidates.values()): + score = self.score_calc.promotion_score(blue_state) + if best_promotion_candidate is None or score is True or best_promotion_score < score: + best_promotion_candidate = blue_state + best_promotion_score = score + if score is True: + break + + # check for state promotion + if best_promotion_candidate is not None: + # a state was promoted -> only forget scores for this blue node + for red in red_states: + del partition_candidates[(red, best_promotion_candidate)] + + # promote best candidate + red_states.append(best_promotion_candidate) + red_states_backing_set.add(best_promotion_candidate) + blue_states.remove(best_promotion_candidate) + blue_states.extend(best_promotion_candidate.child_iterator()) + instrumentation.log_promote(best_promotion_candidate) + else: + # find best partitioning and apply + 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.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() instrumentation.learning_done(root) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 2af3799947..890e8348e4 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -25,6 +25,9 @@ def __init__(self, local_compatibility: LocalCompatibilityFunction = None, score def initialize_merge(self, red: GsmNode, blue: GsmNode) -> Any: return None + def promotion_score(self, promotion_candidate: GsmNode) -> Any: + return True + @staticmethod def default_local_compatibility(a: GsmNode, b: GsmNode): return True From ebb6c17f21000dab6f9400a8e4a9ef8b79b397d6 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 13:03:32 +0200 Subject: [PATCH 26/70] eliminate determinism check from score calculation for comparing futures only --- .../general_passive/ScoreFunctionsGSM.py | 14 +------------- 1 file changed, 1 insertion(+), 13 deletions(-) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 890e8348e4..6294be3f09 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -67,25 +67,13 @@ class CheckFutureScore(ScoreCalculation): def __init__(self, local_compatibility: LocalCompatibilityFunction = None, score_function: ScoreFunction = None, - output_behavior = "moore", - transition_behavior = "deterministic", compatibility_on_pta = False, depth_first = False ): super().__init__(local_compatibility, score_function) - # TODO: auto init from GSM - self.output_behavior = output_behavior - self.transition_behavior = transition_behavior self.compatibility_on_pta = compatibility_on_pta self.depth_first = depth_first - def compute_local_compatibility(self, a: GsmNode, b: GsmNode): - 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.local_compatibility(a, b) - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: if self.compatibility_on_pta and not isinstance(red.data, ShadowPTAData): raise TypeError("compatibility_on_pta is set but no PTA data is available") @@ -96,7 +84,7 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: while len(q) != 0: red, blue = pop() - if self.compute_local_compatibility(red, blue) is False: + if not self.local_compatibility(red, blue): return False if not self.compatibility_on_pta: From e43e23671c1a16badd457ea640d6017be621a9d9 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 16:38:34 +0200 Subject: [PATCH 27/70] fix learning moore machines from examples --- aalpy/learning_algs/general_passive/GeneralizedStateMerging.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 3a0e2250d9..38e4d54bdb 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -309,7 +309,8 @@ def get_partition_trans(part: GsmNode, in_symbol): if first_pass: local_compat = self.score_calc.local_compatibility(partition, blue) - if local_compat is False: + moore_check = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(red, blue) + if local_compat is False or moore_check: partitioning.score = False return if local_compat is defer_merge: From fe820719a92637cb8472b4a0f7543883c04ed341 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 16:43:51 +0200 Subject: [PATCH 28/70] switch from special object to none for deferring merges --- .../general_passive/GeneralizedStateMerging.py | 4 +--- aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py | 7 +++---- 2 files changed, 4 insertions(+), 7 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 38e4d54bdb..a09a2128dd 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -10,8 +10,6 @@ from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility -defer_merge = object() - # TODO add option for making checking of futures and partition non mutual exclusive? # Easiest done by adding a new method / field to ScoreCalculation @@ -313,7 +311,7 @@ def get_partition_trans(part: GsmNode, in_symbol): if local_compat is False or moore_check: partitioning.score = False return - if local_compat is defer_merge: + if local_compat is None: partitioning.remaining_merges.append((red, blue)) continue diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 6294be3f09..8bb489b7d4 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -1,12 +1,11 @@ from collections import deque from math import sqrt, log -from typing import Callable, Dict, List, Iterable, Any, Tuple +from typing import Callable, Dict, List, Iterable, Any, Tuple, Optional -from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, \ - CountData +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, CountData from aalpy.learning_algs.general_passive.IOHandler import ShadowPTAData -LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool] +LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], Optional[bool]] ScoreFunction = Callable[[Dict[GsmNode, GsmNode]], Any] AggregationFunction = Callable[[Iterable], Any] From 797102fc775376d6962fb21de8a1daf7d2520473 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 16:57:22 +0200 Subject: [PATCH 29/70] keep track of number of merges in all cases --- .../learning_algs/general_passive/GeneralizedStateMerging.py | 2 ++ aalpy/learning_algs/general_passive/Instrumentation.py | 4 ++-- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index a09a2128dd..f3817cc013 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -22,6 +22,7 @@ def __init__(self, red: GsmNode, blue: GsmNode): self.full_mapping: Dict[GsmNode, GsmNode] = dict() self.new_blue = [] self.remaining_merges = None + self.nr_merged_states = 0 class Instrumentation: @@ -304,6 +305,7 @@ def get_partition_trans(part: GsmNode, in_symbol): while len(q) != 0: red, blue = pop() partition = update_partition(red, blue) + partitioning.nr_merged_states += 1 if first_pass: local_compat = self.score_calc.local_compatibility(partition, blue) diff --git a/aalpy/learning_algs/general_passive/Instrumentation.py b/aalpy/learning_algs/general_passive/Instrumentation.py index 3d3794daf2..21cf32e33f 100644 --- a/aalpy/learning_algs/general_passive/Instrumentation.py +++ b/aalpy/learning_algs/general_passive/Instrumentation.py @@ -51,7 +51,7 @@ def pta_construction_done(self, root): def print_status(self): 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.use_early_verdicts: + 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 @@ -65,7 +65,7 @@ def log_promote(self, node: GsmNode): def log_merge(self, part: Partitioning): 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() From 63608e4ab1b47d865358d9d401f98a6c9b5946b2 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 21 Jul 2026 17:00:51 +0200 Subject: [PATCH 30/70] eliminate use_early_verdicts argument --- .../general_passive/GeneralizedStateMerging.py | 9 +-------- 1 file changed, 1 insertion(+), 8 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index f3817cc013..e6ff1007dc 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -53,7 +53,6 @@ def __init__(self, *, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, data_handler: IOHandler = None, - use_early_verdicts: bool = False, node_order: Callable[[GsmNode, GsmNode], bool] = None, consider_only_min_blue=False, depth_first=False, @@ -85,8 +84,6 @@ def __init__(self, *, self.data_handler = data_handler or NoIOHandler() - self.use_early_verdicts = use_early_verdicts - self.consider_only_min_blue = consider_only_min_blue self.depth_first = depth_first @@ -229,7 +226,7 @@ def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNo return # check whether there is an early verdict and adapt helper functions accordingly # TODO maybe split init from early verdict and also call init (maybe with first_pass as an argument) in both cases - partitioning.score = self.score_calc.initialize_merge(red, blue) if self.use_early_verdicts else None + partitioning.score = self.score_calc.initialize_merge(red, blue) if partitioning.score is not None: return partitioning.remaining_merges = [] @@ -363,7 +360,6 @@ def run_GSM(data: list, *, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, data_handler: IOHandler = None, - use_early_verdicts: bool = False, node_order: Callable[[GsmNode, GsmNode], bool] = None, consider_only_min_blue=False, depth_first=False, @@ -387,8 +383,6 @@ def run_GSM(data: list, *, postprocessing: A postprocessing function applied to the learned automaton. - use_early_verdicts: Whether to use potential early verdicts from the score object. Defaults to False. - node_order: Order in which merge candidates are considered. Defaults to short-lex. consider_only_min_blue: Whether to consider merge candidates from all blue nodes or just a single. @@ -411,7 +405,6 @@ def run_GSM(data: list, *, pta_preprocessing=pta_preprocessing, postprocessing=postprocessing, data_handler=data_handler, - use_early_verdicts=use_early_verdicts, node_order=node_order, consider_only_min_blue=consider_only_min_blue, depth_first=depth_first, From e941175fc2755e6fe951a2110a71ec2637aa9a1e Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 22 Jul 2026 16:31:17 +0200 Subject: [PATCH 31/70] speed up count computation --- .../general_passive/IOHandler.py | 22 +++++++++---------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/aalpy/learning_algs/general_passive/IOHandler.py b/aalpy/learning_algs/general_passive/IOHandler.py index c34f37e0d7..f76700633f 100644 --- a/aalpy/learning_algs/general_passive/IOHandler.py +++ b/aalpy/learning_algs/general_passive/IOHandler.py @@ -89,10 +89,13 @@ def get_probabilities(self) -> ProbabilityDict: 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 = CountData.default_ctor() + self.transition_count: CountDict = defaultdict(dict) def local_log_likelihood_contribution(self): llc = 0 @@ -108,10 +111,6 @@ def local_log_likelihood_contribution(self): def count(self): return sum(sum(trans.values()) for trans in self.transition_count.values()) - @staticmethod - def default_ctor(): - return defaultdict(lambda: defaultdict(int)) - @staticmethod def merge(x: CountDict, y: CountDict) -> CountDict: for in_sym, y_o_dict in y.items(): @@ -120,7 +119,7 @@ def merge(x: CountDict, y: CountDict) -> CountDict: x[in_sym] = y_o_dict continue for out_sym, count in y_o_dict.items(): - x_o_dict[out_sym] += count + int_dict_increment(x_o_dict, out_sym, count) return x def get_probabilities(self) -> ProbabilityDict: @@ -135,14 +134,15 @@ def init_data(self) -> CountData: return CountData() def aggregate_data(self, src_node: 'GsmNode[CountData]', in_value, out_value, dst_node: 'GsmNode[CountData]'): - src_node.data.transition_count[in_value][out_value] += 1 + if src_node is not None: + int_dict_increment(src_node.data.transition_count[in_value], out_value, 1) def merge_into_x(self, x: CountData, y: CountData): CountData.merge(x.transition_count, y.transition_count) def copy_on_write(self, x: CountData) -> CountData: new_x = copy(x) - new_x.transition_count = CountData.default_ctor() + new_x.transition_count = defaultdict(dict) for in_sym, trans in x.transition_count.items(): new_x.transition_count[in_sym] = trans.copy() return x @@ -156,7 +156,7 @@ class CountOnPTAData(ShadowPTAData, CountData): def __init__(self): ShadowPTAData.__init__(self) CountData.__init__(self) - self.pta_count: CountDict = CountData.default_ctor() + self.pta_count: CountDict = defaultdict(dict) class CountOnPTAHandler(CountHandler): def init_data(self) -> CountOnPTAData: @@ -165,6 +165,6 @@ def init_data(self) -> CountOnPTAData: def aggregate_data(self, src_node: 'GsmNode[CountOnPTAData]', in_value, out_value, dst_node: 'GsmNode[CountOnPTAData]'): if src_node is None: return - src_node.data.transition_count[in_value][out_value] += 1 - src_node.data.pta_count[in_value][out_value] += 1 + 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 From 3d5b51678fce2c6a15563827069e0ee14b8b98ea Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 22 Jul 2026 16:31:59 +0200 Subject: [PATCH 32/70] marginally simplify hoeffding compatibility --- .../learning_algs/general_passive/ScoreFunctionsGSM.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 8bb489b7d4..b67128b149 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -44,11 +44,17 @@ def has_local_compatibility(self): def hoeffding_compatibility(eps, compare_original=True) -> LocalCompatibilityFunction: eps_fact = sqrt(0.5 * log(2 / eps)) - count_dict_name = "pta_count" if compare_original else "transition_count" def similar(a: GsmNode[CountData], b: GsmNode[CountData]): # iterate over inputs that are common to both states - for in_sym, a_trans, b_trans in intersection_iterator(getattr(a.data, count_dict_name), getattr(b.data, count_dict_name)): + 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(trans.values()) for trans in (a_trans, b_trans)) if a_total == 0 or b_total == 0: From 7ce5a9ddb7235e0b9c496e57bbce9ea348c700b4 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 22 Jul 2026 16:32:54 +0200 Subject: [PATCH 33/70] add simple edsm mode to check based on futures --- .../general_passive/ScoreFunctionsGSM.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index b67128b149..62be0013ec 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -73,38 +73,46 @@ def __init__(self, local_compatibility: LocalCompatibilityFunction = None, score_function: ScoreFunction = None, compatibility_on_pta = False, - depth_first = False + depth_first = False, + edsm = False, ): super().__init__(local_compatibility, score_function) self.compatibility_on_pta = compatibility_on_pta self.depth_first = depth_first + self.edsm = edsm - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: + def initialize_merge(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 + evidence = 0 + while len(q) != 0: red, blue = pop() if not self.local_compatibility(red, blue): return False + evidence += 1 + if not self.compatibility_on_pta: - for in_sym, red_trans, blue_trans in intersection_iterator(red.transitions, blue.transitions): + 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)) else: 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): + 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)) if self.has_score_function(): return None + if self.edsm: + return evidence return True class ScoreWithKTail(ScoreCalculation): From 89c717ed2d45c144101325661b32f5fabfeb7765 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 22 Jul 2026 16:34:13 +0200 Subject: [PATCH 34/70] avoid repeated creation of sentinel object --- aalpy/learning_algs/general_passive/GsmNode.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index e5b36d979c..80bea03f53 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -36,10 +36,9 @@ TransitionFunction = Callable[['GsmNode', Any, Any], str] unknown_output = object() # can be set to a special value if required - +missing = object() def intersection_iterator(a: Dict[Key, Val], b: Dict[Key, Val], sort_by_length=False) -> Iterator[Tuple[Key, Val, Val]]: - missing = object() if sort_by_length and len(b) < len(a): for key, b_val in b.items(): a_val = a.get(key, missing) From 9262acb73d140a2d032df583faf9f4343673ad02 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 22 Jul 2026 16:34:48 +0200 Subject: [PATCH 35/70] avoid spurious tuple construction --- aalpy/learning_algs/general_passive/GsmNode.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 80bea03f53..5c41e35cd7 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -404,11 +404,12 @@ def make_input_complete(self, ic_mode=None) -> List[Tuple['GsmNode', Any, Any]]: def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): curr_node: GsmNode = self for in_value, out_value in trace: - in_sym, out_sym = data_handler.abstract(in_value, out_value) + prefix_access_pair = data_handler.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((in_sym, out_sym), curr_node, data_handler.init_data()) + node = GsmNode(prefix_access_pair, curr_node, data_handler.init_data()) transitions[out_sym] = node data_handler.aggregate_data(curr_node, in_value, out_value, node) curr_node = node From 5c90a62b379ba20d71ab89c35ba7426c13c2fc68 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 22 Jul 2026 16:35:44 +0200 Subject: [PATCH 36/70] restore sensible gsm default for learning stochastic systems --- .../general_passive/GeneralizedStateMerging.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index e6ff1007dc..f4d6fa5776 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -7,7 +7,9 @@ from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, \ OutputBehaviorRange, TransitionBehaviorRange, intersection_iterator, unknown_output, detect_data_format, IOHandler, \ NoIOHandler -from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility +from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler +from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility, \ + CheckFutureScore # TODO add option for making checking of futures and partition non mutual exclusive? @@ -71,7 +73,10 @@ def __init__(self, *, 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, True)) + score_calc = CheckFutureScore(hoeffding_compatibility(0.005, True), 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 = CountOnPTAHandler() self.score_calc: ScoreCalculation = score_calc if node_order is None: From 92773704ea7a1d63a976906838b95492e2a77f47 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Thu, 23 Jul 2026 12:37:07 +0200 Subject: [PATCH 37/70] fix and clarify default node order. also changed node order interface and added string value for short-lex --- .../GeneralizedStateMerging.py | 27 +++++++++------- .../learning_algs/general_passive/GsmNode.py | 32 +++++++++---------- 2 files changed, 30 insertions(+), 29 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index f4d6fa5776..72549be3d7 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -55,9 +55,9 @@ def __init__(self, *, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, data_handler: IOHandler = None, - node_order: Callable[[GsmNode, GsmNode], bool] = None, - consider_only_min_blue=False, - depth_first=False, + node_order: Callable[[GsmNode], Any] = None, + consider_only_min_blue = False, + depth_first = False, ): if output_behavior not in OutputBehaviorRange: @@ -79,10 +79,9 @@ def __init__(self, *, data_handler = CountOnPTAHandler() 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 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) @@ -128,14 +127,17 @@ def run(self, data, convert=True, instrumentation: Instrumentation=None, data_fo blue_states_to_consider = blue_states if self.consider_only_min_blue: # does it make sense to check the score function here? - blue_states_to_consider = [min(blue_states, 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 GsmNode.default_order: + if self.node_order is not None: blue_states_to_consider.sort(key=self.node_order) # sort red states. states are always sorted using default order on original prefix - if self.node_order is not GsmNode.default_order: + if self.node_order is not None: red_states.sort(key=self.node_order) # loop over blue states @@ -365,7 +367,7 @@ def run_GSM(data: list, *, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, data_handler: IOHandler = None, - node_order: Callable[[GsmNode, GsmNode], bool] = None, + node_order: Callable[[GsmNode], Any] = None, consider_only_min_blue=False, depth_first=False, instrumentation=None, @@ -388,7 +390,8 @@ def run_GSM(data: list, *, postprocessing: A postprocessing function applied to the learned automaton. - node_order: Order in which merge candidates are considered. Defaults to short-lex. + node_order: Sorting key which determines the order in which merge candidates are considered. + Defaults to insertion order consider_only_min_blue: Whether to consider merge candidates from all blue nodes or just a single. diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 5c41e35cd7..aef2e9728c 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -1,7 +1,6 @@ import functools import pathlib from collections import defaultdict -from functools import total_ordering from typing import Dict, Any, List, Tuple, Iterable, Callable, Union, TypeVar, Iterator, Optional, Sequence, Generic import pydot @@ -63,7 +62,6 @@ def union_iterator(a: Dict[Key, Val], b: Dict[Key, Val], default: Val = None) -> a_val = a.get(key, default) yield key, a_val, b_val - # TODO reuse in RPNI def detect_data_format(data, check_consistency=False, guess=False): # The different data formats are @@ -108,7 +106,6 @@ def detect_data_format(data, check_consistency=False, guess=False): return accepted_formats[0] # 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(Generic[T]): """ Generic class for observably deterministic automata. @@ -128,19 +125,6 @@ def __init__(self, prefix_access_pair, predecessor: 'GsmNode[T]' = None, data: T self.prefix_access_pair = prefix_access_pair self.data = data - def __lt__(self, other, compare_length_only=False): - 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] - # 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 def get_prefix_length(self): @@ -512,4 +496,18 @@ def moore_compatible(self, other: 'GsmNode'): oo = other.get_prefix_output() return so == oo or so is unknown_output or oo is unknown_output - default_order = functools.cmp_to_key(lambda a, b: -1 if a < b else 1) + + insertion_order = object() + + def short_lex_order(self, other, compare_length_only=False): + 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] From 6db97e105930e97d184b13bf33e1626329fd9260 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Mon, 3 Aug 2026 15:53:20 +0200 Subject: [PATCH 38/70] fix corner case in local compat (probably not relevant yet) --- aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 62be0013ec..87373fa4e3 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -93,7 +93,7 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode) -> Any: while len(q) != 0: red, blue = pop() - if not self.local_compatibility(red, blue): + if self.local_compatibility(red, blue) is False: return False evidence += 1 From 327b5f010bb4f63c43ad2d0aca65d7f90d0590c2 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 1 Sep 2026 16:50:07 +0200 Subject: [PATCH 39/70] minor: fix doc-strings and simplify logic --- .../GeneralizedStateMerging.py | 19 ++++++++----------- .../general_passive/ScoreFunctionsGSM.py | 8 ++++---- 2 files changed, 12 insertions(+), 15 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index fe66f71482..62764dff27 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -108,7 +108,7 @@ def __init__(self, *, :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 Callable[[GsmNode, GsmNode], bool] node_order: Order in which merge candidates are considered. + :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. """ @@ -183,11 +183,7 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation blue_states = list(root.child_iterator()) partition_candidates: dict[tuple[GsmNode, GsmNode], Partitioning] = dict() - while True: - # 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? if self.node_order is None: @@ -197,10 +193,8 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation # 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) - - # sort red states. states are always sorted using default order on original prefix - if self.node_order is not None: red_states.sort(key=self.node_order) # loop over blue states @@ -296,10 +290,12 @@ def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNo blue = partitioning.blue if first_pass: - # for Moore machines the outputs have to match. but only once since Moore-ness is preserved for implied merges + # 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 self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): partitioning.score = False return + # check whether there is an early verdict and adapt helper functions accordingly # TODO maybe split init from early verdict and also call init (maybe with first_pass as an argument) in both cases partitioning.score = self.score_calc.initialize_merge(red, blue) @@ -346,6 +342,7 @@ def update_partition(red_node: GsmNode, blue_node: GsmNode | None) -> GsmNode: def get_partition_trans(part: GsmNode, in_symbol): return part.transitions[in_symbol] else: + # first pass already did all the work return self.data_handler.init_merge(red, blue, first_pass) @@ -453,7 +450,7 @@ def run_GSM(data: list, *, :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 IOHandler data_handler: IOHandler object governing abstraction and aggregation of data - :param Callable[[GsmNode, GsmNode], bool] node_order: Sorting key which determines the order in which merge candidates are considered. Defaults to insertion order + :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. diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 24f532e7f5..a13fa89e62 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -24,10 +24,10 @@ def __init__(self, local_compatibility: LocalCompatibilityFunction = None, :param LocalCompatibilityFunction local_compatibility: Function determining local compatibility of two nodes. :param ScoreFunction score_function: Function computing the score of a full merge partition. """ - # 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 + # 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"): From 23a2d3fc82b2fcd69a7ffb12b244f4c4a13f7a7d Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 1 Sep 2026 17:55:19 +0200 Subject: [PATCH 40/70] added explicit values for reject and greedy accept scores (instead of Booleans) --- .../GeneralizedStateMerging.py | 68 +++++++++--------- .../general_passive/GsmAlgorithms.py | 2 +- .../general_passive/ScoreFunctionsGSM.py | 70 ++++++++++++------- 3 files changed, 79 insertions(+), 61 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 62764dff27..473f7ccdde 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -12,7 +12,7 @@ TransitionBehaviorRange, unknown_output, detect_data_format, IOHandler, NoIOHandler, DataFormat from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility, \ - CheckFutureScore + CheckFutureScore, SpecialScores # TODO add option for making checking of futures and partition non mutual exclusive? @@ -198,59 +198,55 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation red_states.sort(key=self.node_order) # loop over blue states - best_promotion_candidate = None - best_promotion_score = None + 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: 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) - if partitioning.score is True: - perfect_partitioning = partitioning + 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.InstantReject + if partitioning.score is SpecialScores.InstantAccept: break - current_candidates[red_state] = partitioning # 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.InstantAccept: + partition_candidates = {(best_candidate.red, best_candidate.blue): best_candidate} break - # update tracking dict with new candidates - new_candidates = (((red, blue_state), part) for red, part in current_candidates.items()) - partition_candidates.update(new_candidates) - # no merge candidates for this blue state -> promotion candidate - if all(part.score is False for part in current_candidates.values()): + if no_viable_merge_for_blue: score = self.score_calc.promotion_score(blue_state) - if best_promotion_candidate is None or score is True or best_promotion_score < score: - best_promotion_candidate = blue_state - best_promotion_score = score - if score is True: + if best_candidate is None or best_score < score: + best_candidate = blue_state + best_score = score + if score is SpecialScores.InstantAccept: break # check for state promotion - if best_promotion_candidate is not None: + 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_promotion_candidate)] + del partition_candidates[(red, best_candidate)] # promote best candidate - red_states.append(best_promotion_candidate) - red_states_backing_set.add(best_promotion_candidate) - blue_states.remove(best_promotion_candidate) - blue_states.extend(best_promotion_candidate.child_iterator()) - instrumentation.log_promote(best_promotion_candidate) - else: - # find best partitioning and apply - best_candidate = max(partition_candidates.values(), key=lambda part: part.score) + 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) + 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 @@ -265,6 +261,8 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation # 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) @@ -293,12 +291,12 @@ def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNo # 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 self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): - partitioning.score = False + partitioning.score = SpecialScores.InstantReject return # check whether there is an early verdict and adapt helper functions accordingly - # TODO maybe split init from early verdict and also call init (maybe with first_pass as an argument) in both cases - partitioning.score = self.score_calc.initialize_merge(red, blue) + # TODO maybe split init from early verdict + partitioning.score = self.score_calc.initialize_merge(red, blue, first_pass) if partitioning.score is not None: return partitioning.remaining_merges = [] @@ -334,6 +332,8 @@ def get_partition_trans(part: GsmNode, in_symbol): cow_set.add(id(trans)) return trans elif partitioning.remaining_merges is None or len(partitioning.remaining_merges) != 0: + self.score_calc.initialize_merge(red, blue, first_pass) + # best scoring merge candidate -> can manipulate nodes directly red_partitions = red_nodes def update_partition(red_node: GsmNode, blue_node: GsmNode | None) -> GsmNode: @@ -381,7 +381,7 @@ def get_partition_trans(part: GsmNode, in_symbol): local_compat = self.score_calc.local_compatibility(partition, blue) moore_check = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(red, blue) if local_compat is False or moore_check: - partitioning.score = False + partitioning.score = SpecialScores.InstantReject return if local_compat is None: partitioning.remaining_merges.append((red, blue)) diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index cb63b46253..5cfec4eab9 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -145,7 +145,7 @@ def __init__(self, epsilon: float) -> None: self.ioa_compatibility = hoeffding_compatibility(epsilon) self.evidence = 0 - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): """ Reset the accumulated evidence counter. """ diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index a13fa89e62..8ccd4a5591 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -2,6 +2,7 @@ # state-merging algorithm (local compatibility checks and global merge scores). from collections import deque from collections.abc import Callable, Iterable +from functools import total_ordering from math import sqrt, log from typing import Any @@ -13,6 +14,19 @@ AggregationFunction = Callable[[Iterable], Any] +class SpecialScores: + @total_ordering + class _SpecialScore: + def __init__(self, ideal: bool): + self.ideal = ideal + + def __lt__(self, other): + return not self.ideal + + InstantAccept = _SpecialScore(True) + InstantReject = _SpecialScore(False) + NoScore = None + class ScoreCalculation: """Bundles a local compatibility check and a global score function used during state merging.""" @@ -33,11 +47,20 @@ def __init__(self, local_compatibility: LocalCompatibilityFunction = None, if not hasattr(self, "score_function"): self.score_function: ScoreFunction = score_function or self.default_score_function - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> Any: + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + """ + Callback at the beginning of the evaluation of a merge candidate. + + :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`. + """ return None def promotion_score(self, promotion_candidate: GsmNode) -> Any: - return True + return SpecialScores.InstantAccept @staticmethod def default_local_compatibility(a: GsmNode, b: GsmNode) -> bool: @@ -51,14 +74,14 @@ def default_local_compatibility(a: GsmNode, b: GsmNode) -> bool: return True @staticmethod - def default_score_function(part: dict[GsmNode, GsmNode]) -> bool: + def default_score_function(part: dict[GsmNode, GsmNode]) -> Any: """ Default score function: any partition is acceptable. :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. - :return bool: Always True. + :return Any: Always accept. """ - return True + return SpecialScores.InstantAccept def has_score_function(self) -> bool: """ @@ -123,7 +146,7 @@ def __init__(self, self.depth_first = depth_first self.edsm = edsm - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> Any: + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> 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") @@ -136,7 +159,7 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode) -> Any: red, blue = pop() if self.local_compatibility(red, blue) is False: - return False + return SpecialScores.InstantReject evidence += 1 @@ -155,7 +178,7 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode) -> Any: return None if self.edsm: return evidence - return True + return SpecialScores.InstantAccept class ScoreWithKTail(ScoreCalculation): """Applies k-Tails to a compatibility function: Compatibility is only evaluated up to a certain depth k.""" @@ -173,9 +196,9 @@ def __init__(self, other_score: ScoreCalculation, k: int) -> None: self.depth_offset = None - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool: + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: self.depth_offset = None - return self.other_score.initialize_merge(red, blue) + return self.other_score.initialize_merge(red, blue, first_pass) def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: """ @@ -214,9 +237,9 @@ def __init__(self, other_score: ScoreCalculation, sink_cond: Callable[[GsmNode], self.is_first = True - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: self.is_first = True - return self.other_score.initialize_merge(red, blue) + return self.other_score.initialize_merge(red, blue, first_pass) def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: """ @@ -256,14 +279,9 @@ def __init__(self, scores: list[ScoreCalculation], aggregate_compatibility: Aggr self.aggregate_compatibility = aggregate_compatibility or self.default_aggregate_compatibility self.aggregate_score = aggregate_score or self.default_aggregate_score - def initialize_merge(self, red: GsmNode, blue: GsmNode) -> bool | None: - all_true = True - for score in self.scores: - verdict = score.initialize_merge(red, blue) - if verdict is False: - return False - all_true = all_true and verdict - return all_true or None + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + scores = [score.initialize_merge(red, blue, first_pass) for score in self.scores] + return self.aggregate_score(scores) def local_compatibility(self, a: GsmNode, b: GsmNode) -> Any: """ @@ -317,14 +335,14 @@ 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.InstantReject + return SpecialScores.InstantAccept return fun @@ -372,7 +390,7 @@ def make_greedy(score: Any) -> Any: :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) + return transform_score(score, lambda x: x is not False and x is not SpecialScores.InstantReject) def lower_threshold(score: Any, thresh: Any) -> Any: @@ -383,7 +401,7 @@ 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 transform_score(score, lambda x: x if thresh < x else SpecialScores.InstantReject) def AIC_score(alpha: float = 0) -> ScoreFunction: From e05fe3db7229ec554170950ac2a1509b36650bc6 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 2 Sep 2026 11:05:29 +0200 Subject: [PATCH 41/70] eliminate edsm from CheckFutureScore --- .../general_passive/ScoreFunctionsGSM.py | 18 +++++------------- 1 file changed, 5 insertions(+), 13 deletions(-) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 8ccd4a5591..efbd19585b 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -139,12 +139,10 @@ def __init__(self, score_function: ScoreFunction = None, compatibility_on_pta = False, depth_first = False, - edsm = False, ): super().__init__(local_compatibility, score_function) self.compatibility_on_pta = compatibility_on_pta self.depth_first = depth_first - self.edsm = edsm def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: if self.compatibility_on_pta and not isinstance(red.data, ShadowPTAData): @@ -153,31 +151,25 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any q: deque[tuple[GsmNode, GsmNode]] = deque([(red, blue)]) pop = q.pop if self.depth_first else q.popleft - evidence = 0 - while len(q) != 0: red, blue = pop() if self.local_compatibility(red, blue) is False: return SpecialScores.InstantReject - evidence += 1 - - if not self.compatibility_on_pta: - 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)) - else: + 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)) if self.has_score_function(): return None - if self.edsm: - return evidence return SpecialScores.InstantAccept class ScoreWithKTail(ScoreCalculation): From ecd63e7873bb881f083b8ed422df365e744dca7d Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 4 Sep 2026 16:10:30 +0200 Subject: [PATCH 42/70] eliminate unused code --- .../learning_algs/general_passive/GsmNode.py | 20 ++----------------- 1 file changed, 2 insertions(+), 18 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index abdfc5e134..8af48ab5fa 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -223,24 +223,11 @@ def get_root(self) -> 'GsmNode': current = current.predecessor return current - def get_or_create_transitions(self, in_sym: Any) -> dict[Any, 'GsmNode[T]']: - """ - 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, '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(): @@ -248,7 +235,7 @@ def transition_iterator(self) -> Iterable[tuple[Any, Any, 'GsmNode[T]']]: def child_iterator(self) -> Iterable['GsmNode[T]']: """ - Iterate over all possible succesors of this node. + Iterate over all possible successors of this node. :return Iterable[GsmNode[T]]: Iterable of successor nodes. """ @@ -663,9 +650,6 @@ 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 - # TODO maybe group orders? - insertion_order = object() - def short_lex_order(self, other: 'GsmNode', compare_length_only: bool = False): """ Compute the short-lex order of the two nodes. From 3a099cf1f43a89dbab31fbb98647912012542466 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 4 Sep 2026 17:39:18 +0200 Subject: [PATCH 43/70] docstrings, fixing examples and some refactoring --- Examples.py | 60 +++---- aalpy/learning_algs/__init__.py | 2 +- .../GeneralizedStateMerging.py | 30 ++-- .../general_passive/GsmAlgorithms.py | 84 +++------ .../general_passive/ScoreFunctionsGSM.py | 167 +++++++++++------- 5 files changed, 175 insertions(+), 168 deletions(-) diff --git a/Examples.py b/Examples.py index 8d98dc491c..e93d829d17 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 @@ -1252,9 +1252,11 @@ def score_fun(part: Dict[GsmNode, GsmNode]): def example_Alergia_extension(): + from typing import Any + from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler 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, SimpleFutureBasedScore, SpecialScores from aalpy.utils.Sampling import get_io_traces, sample_with_length_limits from aalpy import load_automaton_from_file @@ -1262,39 +1264,39 @@ 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(SimpleFutureBasedScore): + def __init__(self, eps: float): + self.compat = hoeffding_compatibility(eps) + SimpleFutureBasedScore.__init__(self, None, 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 initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + self.score = 0 + verdict = super().initialize_merge(red, blue, first_pass) + 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": SimpleFutureBasedScore(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) + learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score, data_handler=CountOnPTAHandler()) 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, SimpleFutureBasedScore + from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler 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 +1322,11 @@ 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": SimpleFutureBasedScore(ioa_compat, compatibility_on_pta=True), + "IOA+DK": SimpleFutureBasedScore(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) + learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score, data_handler=CountOnPTAHandler()) learned_model.visualize(name) def k_tails_example(): @@ -1336,16 +1337,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..0dbbd4848a 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_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/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 473f7ccdde..edd8b2295c 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -10,9 +10,9 @@ from aalpy import Automaton from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, OutputBehaviorRange, \ TransitionBehaviorRange, unknown_output, detect_data_format, IOHandler, NoIOHandler, DataFormat -from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler +from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler, CountHandler from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility, \ - CheckFutureScore, SpecialScores + SimpleFutureBasedScore, SpecialScores # TODO add option for making checking of futures and partition non mutual exclusive? @@ -30,7 +30,7 @@ def __init__(self, red: GsmNode, blue: GsmNode) -> None: """ self.red: GsmNode = red self.blue: GsmNode = blue - self.score = None + self.score = SpecialScores.NoScore self.red_mapping: dict[GsmNode, GsmNode] = dict() self.full_mapping: dict[GsmNode, GsmNode] = dict() self.new_blue = [] @@ -108,6 +108,7 @@ def __init__(self, *, :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 IOHandler 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. @@ -126,7 +127,7 @@ def __init__(self, *, elif transition_behavior == "nondeterministic" : raise ValueError("Missing score_calc for nondeterministic transition behavior. No default available.") elif transition_behavior == "stochastic" : - score_calc = CheckFutureScore(hoeffding_compatibility(0.005, True), compatibility_on_pta=True) + score_calc = SimpleFutureBasedScore(hoeffding_compatibility(0.005, True), 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 = CountOnPTAHandler() @@ -139,7 +140,9 @@ def __init__(self, *, self.pta_preprocessing = pta_preprocessing or (lambda x: x) self.postprocessing = postprocessing or (lambda x: x) - self.data_handler = data_handler or NoIOHandler() + if data_handler is None: + data_handler = CountHandler() if transition_behavior == "stochastic" else NoIOHandler() + self.data_handler = data_handler self.consider_only_min_blue = consider_only_min_blue self.depth_first = depth_first @@ -215,12 +218,12 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation 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.InstantReject - if partitioning.score is SpecialScores.InstantAccept: + no_viable_merge_for_blue &= partitioning.score is SpecialScores.ImmediateReject + if partitioning.score is SpecialScores.ImmediateAccept: break # partition with perfect score found: don't consider anything else - if best_score is SpecialScores.InstantAccept: + if best_score is SpecialScores.ImmediateAccept: partition_candidates = {(best_candidate.red, best_candidate.blue): best_candidate} break @@ -230,7 +233,7 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation if best_candidate is None or best_score < score: best_candidate = blue_state best_score = score - if score is SpecialScores.InstantAccept: + if score is SpecialScores.ImmediateAccept: break # check for state promotion @@ -287,17 +290,18 @@ def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNo 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 self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): - partitioning.score = SpecialScores.InstantReject + 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.initialize_merge(red, blue, first_pass) - if partitioning.score is not None: + if partitioning.score is not SpecialScores.NoScore: return partitioning.remaining_merges = [] @@ -352,7 +356,7 @@ def get_partition_trans(part: GsmNode, in_symbol): # 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 and partitioning.score is None) or (not first_pass and partitioning.score is not None) + assert first_pass == (partitioning.score is SpecialScores.NoScore) # rewire the blue node's parent blue_parent = update_partition(blue.predecessor, None) @@ -381,7 +385,7 @@ def get_partition_trans(part: GsmNode, in_symbol): local_compat = self.score_calc.local_compatibility(partition, blue) moore_check = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(red, blue) if local_compat is False or moore_check: - partitioning.score = SpecialScores.InstantReject + partitioning.score = SpecialScores.ImmediateReject return if local_compat is None: partitioning.remaining_merges.append((red, blue)) diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index 5cfec4eab9..f1d28d6434 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.IOHandler import CountOnPTAHandler 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.ScoreFunctionsGSM import ScoreCalculation, ScoreWithKTail, ScoreIOAlergiaWithEDSM +from aalpy.utils.HelperFunctions import dfa_from_moore, mc_format_to_mdp, mc_from_mdp def run_EDSM(data: list, automaton_type: str, input_completeness: str | None = None, @@ -109,78 +110,51 @@ def run_k_tails(data: list, automaton_type: str, k: int, input_completeness: str 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 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 - # TODO: This needs to be reworked... - raise NotImplementedError + at_types = ['mc', 'mdp', 'smm'] + if automaton_type not in at_types: + raise ValueError(f"automaton_type {automaton_type} not in {at_types}") - assert automaton_type in {'mc', 'mdp', 'smm',} + 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") - print_level = ProgressReport(1) if print_info else None - - 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 initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): - """ - 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=CountOnPTAHandler(), + 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/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index efbd19585b..440ca40a85 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -7,7 +7,7 @@ from typing import Any from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, CountData -from aalpy.learning_algs.general_passive.IOHandler import ShadowPTAData +from aalpy.learning_algs.general_passive.IOHandler import ShadowPTAData, CountOnPTAData LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool | None] ScoreFunction = Callable[[dict[GsmNode, GsmNode]], Any] @@ -23,8 +23,11 @@ def __init__(self, ideal: bool): def __lt__(self, other): return not self.ideal - InstantAccept = _SpecialScore(True) - InstantReject = _SpecialScore(False) + def __bool__(self): + return self.ideal + + ImmediateAccept = _SpecialScore(True) + ImmediateReject = _SpecialScore(False) NoScore = None class ScoreCalculation: @@ -60,7 +63,14 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any return None def promotion_score(self, promotion_candidate: GsmNode) -> Any: - return SpecialScores.InstantAccept + """ + Computes the score of a promotion candidate. By default, promotion candidates are immediately accepted. Override + this function to implement promotion scoring. + + :param GsmNode promotion_candidate: GsmNode which is to be promoted. + :return Any: The score of the promotion candidate. Default is ImmediateAccept. + """ + return SpecialScores.ImmediateAccept @staticmethod def default_local_compatibility(a: GsmNode, b: GsmNode) -> bool: @@ -81,7 +91,7 @@ def default_score_function(part: dict[GsmNode, GsmNode]) -> Any: :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. :return Any: Always accept. """ - return SpecialScores.InstantAccept + return SpecialScores.ImmediateAccept def has_score_function(self) -> bool: """ @@ -91,14 +101,6 @@ def has_score_function(self) -> bool: """ return self.score_function is not self.default_score_function - def has_local_compatibility(self) -> bool: - """ - Check whether a non-default local compatibility function is configured. - - :return bool: True if local_compatibility was overridden. - """ - return self.local_compatibility is not self.default_local_compatibility - def hoeffding_compatibility(eps: float, compare_original: bool = True) -> LocalCompatibilityFunction: """ @@ -133,13 +135,27 @@ def similar(a: GsmNode[CountData], b: GsmNode[CountData]) -> bool: return similar -class CheckFutureScore(ScoreCalculation): +class SimpleFutureBasedScore(ScoreCalculation): + """ + ScoreCalculation that checks local compatibility only on common futures (as in Alergia) and not during the + construction of the partitioning. As long as no scoring (based on the partitioning) is used, this results in a + significant speedup. + """ def __init__(self, local_compatibility: LocalCompatibilityFunction = None, score_function: ScoreFunction = None, compatibility_on_pta = False, depth_first = False, ): + """ + Create a new CheckFutureScore instance. + + :param LocalCompatibilityFunction local_compatibility: Compatibility criterion used to check futures. + :param ScoreFunction score_function: The score function to rank merge candidates. Values other than None (default) + negate the speedup. + :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). + """ super().__init__(local_compatibility, score_function) self.compatibility_on_pta = compatibility_on_pta self.depth_first = depth_first @@ -155,7 +171,7 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any red, blue = pop() if self.local_compatibility(red, blue) is False: - return SpecialScores.InstantReject + return SpecialScores.ImmediateReject if self.compatibility_on_pta: red_data: ShadowPTAData = red.data @@ -169,8 +185,28 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any q.append((red_child, blue_child)) if self.has_score_function(): - return None - return SpecialScores.InstantAccept + return SpecialScores.NoScore + return SpecialScores.ImmediateAccept + + +class ScoreIOAlergiaWithEDSM(SimpleFutureBasedScore): + 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) + SimpleFutureBasedScore.__init__(self, None, compatibility_on_pta=compat_on_pta) + self.edsm = edsm + self.score = None + + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + self.score = 0 + verdict = super().initialize_merge(red, blue, first_pass) + 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 ScoreWithKTail(ScoreCalculation): """Applies k-Tails to a compatibility function: Compatibility is only evaluated up to a certain depth k.""" @@ -189,10 +225,10 @@ def __init__(self, other_score: ScoreCalculation, k: int) -> None: self.depth_offset = None def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: - self.depth_offset = None + self.depth_offset = blue.get_prefix_length() return self.other_score.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. @@ -201,11 +237,9 @@ 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) @@ -227,10 +261,10 @@ def __init__(self, other_score: ScoreCalculation, sink_cond: Callable[[GsmNode], self.sink_cond = sink_cond self.allow_sink_merge = allow_sink_merge - self.is_first = True - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: - self.is_first = True + 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.other_score.initialize_merge(red, blue, first_pass) def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: @@ -241,13 +275,6 @@ def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool: :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) @@ -333,58 +360,43 @@ def local_to_global_compatibility(local_fun: LocalCompatibilityFunction) -> Scor 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 SpecialScores.InstantReject - return SpecialScores.InstantAccept + return SpecialScores.ImmediateReject + return SpecialScores.ImmediateAccept return fun -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 transform_score(score: Any, transform: Callable) -> Any: +def transform_score(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 fun(score: Any, *args) -> Any: + if isinstance(score, Callable): + return lambda partitioning: transform(score(partitioning), *args) + if isinstance(score, ScoreCalculation): + original_score_function = score.score_function + score.score_function = lambda partitioning: transform(original_score_function(partitioning), *args) + return score + return transform(score, *args) + return fun -def make_greedy(score: Any) -> Any: +@transform_score +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 and x is not SpecialScores.InstantReject) + should_accept = score is not False and score is not SpecialScores.ImmediateReject + return SpecialScores.ImmediateAccept if should_accept else SpecialScores.ImmediateReject +@transform_score def lower_threshold(score: Any, thresh: Any) -> Any: """ Transform a score so that it is rejected (False) unless it exceeds a threshold. @@ -393,7 +405,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 SpecialScores.InstantReject) + 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: @@ -410,7 +441,7 @@ 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. From 2460dde45e14ba38d6a94e362a629f102f01f821 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 4 Sep 2026 17:43:12 +0200 Subject: [PATCH 44/70] Rename IOHandler -> DataHandler --- Examples.py | 10 ++++++---- .../{IOHandler.py => DataHandler.py} | 12 ++++++------ .../general_passive/GeneralizedStateMerging.py | 16 ++++++++-------- .../general_passive/GsmAlgorithms.py | 4 ++-- aalpy/learning_algs/general_passive/GsmNode.py | 16 ++++++++-------- .../general_passive/ScoreFunctionsGSM.py | 2 +- 6 files changed, 31 insertions(+), 29 deletions(-) rename aalpy/learning_algs/general_passive/{IOHandler.py => DataHandler.py} (93%) diff --git a/Examples.py b/Examples.py index e93d829d17..7c9e4f2694 100644 --- a/Examples.py +++ b/Examples.py @@ -1253,7 +1253,7 @@ def score_fun(part: Dict[GsmNode, GsmNode]): def example_Alergia_extension(): from typing import Any - from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler + 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.GsmNode import GsmNode from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import hoeffding_compatibility, SimpleFutureBasedScore, SpecialScores @@ -1289,14 +1289,15 @@ def local_compatibility(self, red: GsmNode, blue: GsmNode) -> float: } for name, score in scores.items(): - learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score, data_handler=CountOnPTAHandler()) + learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score, + 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, SimpleFutureBasedScore - from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler + 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 @@ -1326,7 +1327,8 @@ def ioa_compat_domain_knowledge(a: GsmNode, b: GsmNode): "IOA+DK": SimpleFutureBasedScore(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, data_handler=CountOnPTAHandler()) + learned_model = run_GSM(traces, output_behavior="moore", transition_behavior="stochastic", score_calc=score, + data_handler=CountOnPTADataHandler()) learned_model.visualize(name) def k_tails_example(): diff --git a/aalpy/learning_algs/general_passive/IOHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py similarity index 93% rename from aalpy/learning_algs/general_passive/IOHandler.py rename to aalpy/learning_algs/general_passive/DataHandler.py index f76700633f..124079e344 100644 --- a/aalpy/learning_algs/general_passive/IOHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -6,7 +6,7 @@ T = TypeVar("T") -class IOHandler(Generic[T]): +class DataHandler(Generic[T]): @abstractmethod def init(self, data, output_format, data_format): ... @@ -34,14 +34,14 @@ def merge(self, x: T, y: T) -> T: def copy(self, x: T) -> T: ... -class NoAbstractionIOHandler(IOHandler[T], ABC): +class NoAbstractionDataHandler(DataHandler[T], ABC): def init(self, data, output_format, data_format): pass def abstract(self, in_val, out_val): return in_val, out_val -class NoIOHandler(NoAbstractionIOHandler[T]): +class NoOpDataHandler(NoAbstractionDataHandler[T]): def init_data(self) -> T: return None @@ -54,7 +54,7 @@ def merge(self, x: T, y: T) -> T: def copy(self, x: T) -> T: return None -class CopyOnWriteIOHandler(IOHandler[T], ABC): +class CopyOnWriteDataHandler(DataHandler[T], ABC): def __init__(self): self.copied_on_write = set() @@ -129,7 +129,7 @@ def get_probabilities(self) -> ProbabilityDict: ret[in_sym] = {out_sym: count / total_count for out_sym, count in trans.items()} return ret -class CountHandler(NoAbstractionIOHandler[CountData], CopyOnWriteIOHandler): +class CountDataHandler(NoAbstractionDataHandler[CountData], CopyOnWriteDataHandler): def init_data(self) -> CountData: return CountData() @@ -158,7 +158,7 @@ def __init__(self): CountData.__init__(self) self.pta_count: CountDict = defaultdict(dict) -class CountOnPTAHandler(CountHandler): +class CountOnPTADataHandler(CountDataHandler): def init_data(self) -> CountOnPTAData: return CountOnPTAData() diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index edd8b2295c..9f2e73d863 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -9,8 +9,8 @@ from aalpy import Automaton from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, OutputBehaviorRange, \ - TransitionBehaviorRange, unknown_output, detect_data_format, IOHandler, NoIOHandler, DataFormat -from aalpy.learning_algs.general_passive.IOHandler import CountOnPTAHandler, CountHandler + TransitionBehaviorRange, unknown_output, detect_data_format, DataHandler, NoOpDataHandler, DataFormat +from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler, CountDataHandler from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility, \ SimpleFutureBasedScore, SpecialScores @@ -95,7 +95,7 @@ def __init__(self, *, score_calc: ScoreCalculation = None, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, - data_handler: IOHandler = None, + data_handler: DataHandler = None, node_order: Callable[[GsmNode], Any] = None, consider_only_min_blue = False, depth_first = False, @@ -108,7 +108,7 @@ def __init__(self, *, :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 IOHandler data_handler: IOHandler object governing abstraction and aggregation of data + :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. @@ -130,7 +130,7 @@ def __init__(self, *, score_calc = SimpleFutureBasedScore(hoeffding_compatibility(0.005, True), 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 = CountOnPTAHandler() + data_handler = CountOnPTADataHandler() self.score_calc: ScoreCalculation = score_calc if isinstance(node_order, str) and node_order == "short-lex": @@ -141,7 +141,7 @@ def __init__(self, *, self.postprocessing = postprocessing or (lambda x: x) if data_handler is None: - data_handler = CountHandler() if transition_behavior == "stochastic" else NoIOHandler() + data_handler = CountDataHandler() if transition_behavior == "stochastic" else NoOpDataHandler() self.data_handler = data_handler self.consider_only_min_blue = consider_only_min_blue @@ -436,7 +436,7 @@ def run_GSM(data: list, *, score_calc: ScoreCalculation = None, pta_preprocessing: Callable[[GsmNode], GsmNode] = None, postprocessing: Callable[[GsmNode], GsmNode] = None, - data_handler: IOHandler = None, + data_handler: DataHandler = None, node_order: Callable[[GsmNode], Any] = None, consider_only_min_blue=False, depth_first=False, @@ -453,7 +453,7 @@ def run_GSM(data: list, *, :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 IOHandler data_handler: IOHandler object governing abstraction and aggregation of data + :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. diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index f1d28d6434..4ccdf9357a 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -6,7 +6,7 @@ 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.IOHandler import CountOnPTAHandler +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, ScoreWithKTail, ScoreIOAlergiaWithEDSM @@ -146,7 +146,7 @@ def run_Alergia_GSM(data: list, automaton_type: str, eps: float = 0.05, compat_o learning_data, output_behavior=output_behaviour, transition_behavior="stochastic", - data_handler=CountOnPTAHandler(), + data_handler=CountOnPTADataHandler(), score_calc=ScoreIOAlergiaWithEDSM(eps, compat_on_pta_trans, compat_on_pta_count, edsm), instrumentation=instrumentation, data_format='io_traces', diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 8af48ab5fa..361541976d 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -9,7 +9,7 @@ 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.IOHandler import IOHandler, NoIOHandler, StochasticData, CountData +from aalpy.learning_algs.general_passive.DataHandler import DataHandler, NoOpDataHandler, StochasticData, CountData Key = TypeVar("Key") @@ -491,12 +491,12 @@ def make_input_complete(self, ic_mode: str = "self-loop") -> list[tuple['GsmNode transitions[out_sym] = successor return missing_trans - def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): + def add_trace(self, trace: IOTrace, data_handler: DataHandler[T]): """ 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. - :param IOHandler[T] data_handler: IOHandler used for abstraction and aggregation of trace data + :param DataHandler[T] data_handler: IOHandler used for abstraction and aggregation of trace data """ curr_node: GsmNode = self for in_value, out_value in trace: @@ -510,18 +510,18 @@ def add_trace(self, trace: IOTrace, data_handler: IOHandler[T]): data_handler.aggregate_data(curr_node, in_value, out_value, node) curr_node = node - def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = None): + def add_labeled_sequence(self, example: IOExample, data_handler: DataHandler[T] = 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. - :param IOHandler[T] data_handler: IOHandler used for abstraction and aggregation of trace data + :param DataHandler[T] data_handler: IOHandler used for abstraction and aggregation of trace data """ inputs, output = example curr_node: GsmNode = self in_sym = None - if not isinstance(data_handler, NoIOHandler): + if not isinstance(data_handler, NoOpDataHandler): raise NotImplementedError("Data handling is not supported for learning from labeled sequences") # step through inputs and add transitions @@ -551,14 +551,14 @@ def add_labeled_sequence(self, example: IOExample, data_handler: IOHandler[T] = raise ValueError("nondeterminism encountered for GSM with labeled_sequences. not supported") @staticmethod - def createPTA(data: Any, output_behavior: OutputBehavior, data_format: DataFormat = None, data_handler: IOHandler[T] = None) -> 'GsmNode': + def createPTA(data: Any, output_behavior: OutputBehavior, data_format: DataFormat = None, data_handler: DataHandler[T] = 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. - :param IOHandler[T] data_handler: IOHandler used for abstraction and aggregation of trace data + :param DataHandler[T] data_handler: 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: diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 440ca40a85..4e97ab9998 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -7,7 +7,7 @@ from typing import Any from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, CountData -from aalpy.learning_algs.general_passive.IOHandler import ShadowPTAData, CountOnPTAData +from aalpy.learning_algs.general_passive.DataHandler import ShadowPTAData, CountOnPTAData LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool | None] ScoreFunction = Callable[[dict[GsmNode, GsmNode]], Any] From 5d7511d56eed3aaf6ec165c143e5259c174c88b9 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Mon, 7 Sep 2026 13:51:40 +0200 Subject: [PATCH 45/70] in SimpleFutureBasedScore, added option to not check compatibility of children at any point. --- aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 4e97ab9998..05cd7a504f 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -170,8 +170,11 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any while len(q) != 0: red, blue = pop() - if self.local_compatibility(red, blue) is False: + 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 From f8ddfa10e980da8b24aac86d1a169aa98a33048b Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Mon, 7 Sep 2026 18:02:51 +0200 Subject: [PATCH 46/70] Added dedicated `SimpleScoreCalculation` to eliminate the hack in `ScoreCalculation.__init__`. also some other score calc stuff and docstrings --- Examples.py | 22 +- .../general_passive/DataHandler.py | 2 +- .../GeneralizedStateMerging.py | 7 +- .../general_passive/GsmAlgorithms.py | 6 +- .../general_passive/ScoreFunctionsGSM.py | 214 ++++++++++-------- 5 files changed, 144 insertions(+), 107 deletions(-) diff --git a/Examples.py b/Examples.py index 7c9e4f2694..0a6b02eee6 100644 --- a/Examples.py +++ b/Examples.py @@ -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,7 +1246,7 @@ 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() @@ -1256,7 +1256,7 @@ def example_Alergia_extension(): 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.GsmNode import GsmNode - from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import hoeffding_compatibility, SimpleFutureBasedScore, SpecialScores + 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 @@ -1265,10 +1265,10 @@ def example_Alergia_extension(): traces = get_io_traces(automaton, input_traces) # NOTE: a more general version of this is provided in aalpy.learning_algs.general_passive.ScoreFunctionsGSM - class ScoreIOAlergiaWithEDSM(SimpleFutureBasedScore): + class ScoreIOAlergiaWithEDSM(SimpleFutureBasedCompatibility): def __init__(self, eps: float): self.compat = hoeffding_compatibility(eps) - SimpleFutureBasedScore.__init__(self, None, compatibility_on_pta=True) + SimpleFutureBasedCompatibility.__init__(self, compatibility_on_pta=True) self.score = None def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: @@ -1284,7 +1284,7 @@ def local_compatibility(self, red: GsmNode, blue: GsmNode) -> float: epsilon = 0.05 scores = { - "IOA": SimpleFutureBasedScore(hoeffding_compatibility(epsilon, True), compatibility_on_pta=True), + "IOA": SimpleFutureBasedCompatibility(local_compatibility=hoeffding_compatibility(epsilon, True), compatibility_on_pta=True), "IOA+EDSM": ScoreIOAlergiaWithEDSM(epsilon), } @@ -1296,7 +1296,7 @@ def local_compatibility(self, red: GsmNode, blue: GsmNode) -> float: 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, SimpleFutureBasedScore + 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 @@ -1323,8 +1323,8 @@ def ioa_compat_domain_knowledge(a: GsmNode, b: GsmNode): return parity and ioa scores = { - "IOA": SimpleFutureBasedScore(ioa_compat, compatibility_on_pta=True), - "IOA+DK": SimpleFutureBasedScore(ioa_compat_domain_knowledge, compatibility_on_pta=True), + "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, diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index 124079e344..5abbc0f9d1 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -147,11 +147,11 @@ def copy_on_write(self, x: CountData) -> CountData: new_x.transition_count[in_sym] = trans.copy() return x +ShadowPTA = dict[Any, dict[Any, 'GsmNode']] class ShadowPTAData: def __init__(self): self.shadow_pta: ShadowPTA = defaultdict(dict) -ShadowPTA = dict[Any, dict[Any, 'GsmNode']] class CountOnPTAData(ShadowPTAData, CountData): def __init__(self): ShadowPTAData.__init__(self) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 9f2e73d863..11020059f3 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -12,7 +12,7 @@ TransitionBehaviorRange, unknown_output, detect_data_format, DataHandler, NoOpDataHandler, DataFormat from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler, CountDataHandler from aalpy.learning_algs.general_passive.ScoreFunctionsGSM import ScoreCalculation, hoeffding_compatibility, \ - SimpleFutureBasedScore, SpecialScores + SimpleFutureBasedCompatibility, SpecialScores, SimpleScoreCalculation # TODO add option for making checking of futures and partition non mutual exclusive? @@ -123,11 +123,12 @@ def __init__(self, *, if score_calc is None: if transition_behavior == "deterministic": - score_calc = ScoreCalculation(GsmNode.deterministic_compatible) + 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 = SimpleFutureBasedScore(hoeffding_compatibility(0.005, True), compatibility_on_pta=True) + 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() diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index 4ccdf9357a..f10e8f6e46 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -9,7 +9,7 @@ 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, ScoreWithKTail, ScoreIOAlergiaWithEDSM +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 @@ -45,7 +45,7 @@ def EDSM_score(part: dict[GsmNode, GsmNode]) -> int: evidence += 1 return evidence - score = ScoreCalculation(score_function=EDSM_score) + score = SimpleScoreCalculation(score_function=EDSM_score) internal_automaton_type = 'moore' if automaton_type != 'mealy' else automaton_type @@ -92,7 +92,7 @@ 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", diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 05cd7a504f..e5dcf39b69 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -1,5 +1,6 @@ # 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 @@ -7,7 +8,7 @@ from typing import Any from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, CountData -from aalpy.learning_algs.general_passive.DataHandler import ShadowPTAData, CountOnPTAData +from aalpy.learning_algs.general_passive.DataHandler import ShadowPTAData LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool | None] ScoreFunction = Callable[[dict[GsmNode, GsmNode]], Any] @@ -30,26 +31,9 @@ def __bool__(self): ImmediateReject = _SpecialScore(False) NoScore = None -class ScoreCalculation: +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: - """ - Create a score calculation, optionally overriding the default (accept-everything) behavior. - - :param LocalCompatibilityFunction local_compatibility: Function determining local compatibility of two nodes. - :param ScoreFunction score_function: Function computing the score of a full merge partition. - """ - # 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 - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: """ Callback at the beginning of the evaluation of a merge candidate. @@ -62,6 +46,32 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any """ return None + def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool | None: + """ + 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: 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 + + def score_function(self, part: dict[GsmNode, GsmNode]) -> Any: + """ + 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 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 SpecialScores.ImmediateAccept + def promotion_score(self, promotion_candidate: GsmNode) -> Any: """ Computes the score of a promotion candidate. By default, promotion candidates are immediately accepted. Override @@ -72,26 +82,13 @@ def promotion_score(self, promotion_candidate: GsmNode) -> Any: """ return SpecialScores.ImmediateAccept - @staticmethod - def default_local_compatibility(a: GsmNode, b: GsmNode) -> bool: - """ - Default local compatibility check: always compatible. - - :param GsmNode a: First node. - :param GsmNode b: Second node. - :return bool: Always True. - """ - return True - - @staticmethod - def default_score_function(part: dict[GsmNode, GsmNode]) -> Any: + def has_local_compatibility(self) -> bool: """ - Default score function: any partition is acceptable. + Check whether a non-default local compatibility is configured. - :param dict[GsmNode, GsmNode] part: Mapping of original nodes to their merged partition representative. - :return Any: Always accept. + :return bool: True if local_compatibility was overridden. """ - return SpecialScores.ImmediateAccept + return self.__class__.local_compatibility is ScoreCalculation.local_compatibility def has_score_function(self) -> bool: """ @@ -99,7 +96,21 @@ def has_score_function(self) -> bool: :return bool: True if score_function was overridden. """ - return self.score_function is not self.default_score_function + return self.__class__.score_function is ScoreCalculation.score_function + + +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: @@ -135,28 +146,29 @@ def similar(a: GsmNode[CountData], b: GsmNode[CountData]) -> bool: return similar -class SimpleFutureBasedScore(ScoreCalculation): + +class SimpleFutureBasedCompatibility(ScoreCalculation): """ - ScoreCalculation that checks local compatibility only on common futures (as in Alergia) and not during the - construction of the partitioning. As long as no scoring (based on the partitioning) is used, this results in a - significant speedup. + 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, - local_compatibility: LocalCompatibilityFunction = None, - score_function: ScoreFunction = None, compatibility_on_pta = False, depth_first = False, + local_compatibility: LocalCompatibilityFunction = None, ): """ Create a new CheckFutureScore instance. - :param LocalCompatibilityFunction local_compatibility: Compatibility criterion used to check futures. - :param ScoreFunction score_function: The score function to rank merge candidates. Values other than None (default) - negate the speedup. :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. """ - super().__init__(local_compatibility, score_function) + if local_compatibility: + if self.has_local_compatibility(): + raise ValueError("Exernal 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 @@ -187,15 +199,13 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any for out_sym, red_child, blue_child in intersection_iterator(red_trans, blue_trans): q.append((red_child, blue_child)) - if self.has_score_function(): - return SpecialScores.NoScore return SpecialScores.ImmediateAccept -class ScoreIOAlergiaWithEDSM(SimpleFutureBasedScore): +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) - SimpleFutureBasedScore.__init__(self, None, compatibility_on_pta=compat_on_pta) + SimpleFutureBasedCompatibility.__init__(self, compatibility_on_pta=compat_on_pta) self.edsm = edsm self.score = None @@ -211,25 +221,52 @@ def local_compatibility(self, red: GsmNode, blue: GsmNode) -> float: return self.compat(red, blue) -class ScoreWithKTail(ScoreCalculation): +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 initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + return 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 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 initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: self.depth_offset = blue.get_prefix_length() - return self.other_score.initialize_merge(red, blue, first_pass) + return self.wrapped.initialize_merge(red, blue, first_pass) def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool | None: """ @@ -244,23 +281,22 @@ def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool | None: if self.k <= depth: 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 @@ -268,17 +304,7 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> 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.other_score.initialize_merge(red, blue, first_pass) - - 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. - """ - return self.other_score.local_compatibility(a, b) + return self.wrapped.initialize_merge(red, blue, first_pass) class ScoreCombinator(ScoreCalculation): @@ -296,7 +322,6 @@ 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 @@ -317,26 +342,37 @@ 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: @@ -369,25 +405,25 @@ def fun(part: dict[GsmNode, GsmNode]) -> Any: return fun -def transform_score(transform: Callable) -> Any: +def score_transformation(transform: Callable) -> Any: """ Lifts an operation on a score value to score functions and ScoreCalculation objects. Intended as a decorator :param Callable transform: Function to apply to the (eventual) score value. :return Any: Decorated transformation applicable to a score, callable, or ScoreCalculation. """ - def fun(score: Any, *args) -> Any: + def score_function(score: Any, *transformation_args, **transformation_kwargs) -> Any: if isinstance(score, Callable): - return lambda partitioning: transform(score(partitioning), *args) + 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), *args) + score.score_function = lambda partitioning: transform(original_score_function(partitioning), *transformation_args, **transformation_kwargs) return score - return transform(score, *args) + return transform(score, *transformation_args, **transformation_kwargs) - return fun + return score_function -@transform_score +@score_transformation def greedy_score(score: Any) -> Any: """ Transform a score into a greedy (boolean) score: accept anything but a False/reject result. @@ -399,7 +435,7 @@ def greedy_score(score: Any) -> Any: return SpecialScores.ImmediateAccept if should_accept else SpecialScores.ImmediateReject -@transform_score +@score_transformation def lower_threshold(score: Any, thresh: Any) -> Any: """ Transform a score so that it is rejected (False) unless it exceeds a threshold. From 117212d06d6e9870e559de0fe86a53998f469602 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 8 Sep 2026 11:07:10 +0200 Subject: [PATCH 47/70] added special value for input of root GsmNode objects --- aalpy/learning_algs/general_passive/GsmNode.py | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 361541976d..8c61c6e340 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -33,6 +33,7 @@ TransitionFunction = Callable[['GsmNode', Any, Any], str] 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], sort_by_length: bool = False) -> Iterator[tuple[Key, Val, Val]]: @@ -251,7 +252,7 @@ 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: @@ -551,7 +552,7 @@ def add_labeled_sequence(self, example: IOExample, data_handler: DataHandler[T] raise ValueError("nondeterminism encountered for GSM with labeled_sequences. not supported") @staticmethod - def createPTA(data: Any, output_behavior: OutputBehavior, data_format: DataFormat = None, data_handler: DataHandler[T] = None) -> 'GsmNode': + def createPTA(data: Any, output_behavior: OutputBehavior, data_format: DataFormat = None, data_handler: DataHandler[T] = None) -> 'GsmNode[T]': """ Build a prefix tree acceptor (PTA) from the given data. @@ -573,21 +574,21 @@ def createPTA(data: Any, output_behavior: OutputBehavior, data_format: DataForma raise ValueError("provided automaton is not a tree") return data # TODO extract method for replaying data on dot model - root_node = GsmNode((None, unknown_output), None, data_handler.init_data()) + root_node = GsmNode((no_op_input, unknown_output), None, data_handler.init_data()) if data_format == "labeled_sequences": for example in data: root_node.add_labeled_sequence(example, data_handler) if data_format == "io_traces" or data_format == "traces": if output_behavior == "moore": - root_node.prefix_access_pair = data_handler.abstract(None, data[0][0]) + root_node.prefix_access_pair = data_handler.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 = data_handler.abstract(None, initial_output) + _, ios = data_handler.abstract(no_op_input, initial_output) if ios != initial_output_symbol: raise ValueError("expect unique initial output symbol for Moore behavior") - data_handler.aggregate_data(None, None, initial_output, root_node) + data_handler.aggregate_data(None, no_op_input, initial_output, root_node) data = (d[1:] for d in data) for trace in data: From 5659764c2621f0bc7b394f195bd14490d377df0e Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 8 Sep 2026 11:07:46 +0200 Subject: [PATCH 48/70] eliminate unimplemented input completion method for GSM --- .../learning_algs/general_passive/GsmNode.py | 21 +++++++++---------- 1 file changed, 10 insertions(+), 11 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 8c61c6e340..95a14652ed 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -463,16 +463,17 @@ 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, ic_mode: str = "self-loop") -> 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 str ic_mode: determines how input completenes is achieved ("self-loop", "sink-state" or "root". + :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. """ - ic_modes = ["self-loop", "sink-state", "root"] - if ic_mode not in ic_modes: - raise ValueError(f"Invalid ic_mode {ic_mode}. Should be one of {ic_modes}") + + 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} @@ -481,12 +482,10 @@ def make_input_complete(self, ic_mode: str = "self-loop") -> list[tuple['GsmNode for in_sym in inputs: transitions = node.transitions[in_sym] if len(transitions) == 0: - if ic_mode == "self-loop": + if target == "self-loop": successor = node - elif ic_mode == "sink-state": - raise NotImplementedError() - elif ic_mode == "root": - successor = self + else: + successor = target out_sym = successor.prefix_access_pair[1] missing_trans.append((node, in_sym, out_sym)) transitions[out_sym] = successor From 8d8b4a4e54e7029ff69d04652cbe9068d9ca618b Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 8 Sep 2026 16:39:34 +0200 Subject: [PATCH 49/70] fix edsm --- aalpy/learning_algs/general_passive/GsmAlgorithms.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index f10e8f6e46..2dee5d1832 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -45,7 +45,7 @@ def EDSM_score(part: dict[GsmNode, GsmNode]) -> int: evidence += 1 return evidence - score = SimpleScoreCalculation(score_function=EDSM_score) + score = SimpleScoreCalculation(local_compatibility=GsmNode.deterministic_compatible, score_function=EDSM_score) internal_automaton_type = 'moore' if automaton_type != 'mealy' else automaton_type From 0aa095bf0ac22ea752c90f9e634e4bab6fa5a635 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 8 Sep 2026 16:43:22 +0200 Subject: [PATCH 50/70] minor refactoring / cleanup and lots of docstrings --- .../general_passive/DataHandler.py | 152 ++++++++++++------ .../GeneralizedStateMerging.py | 3 +- .../learning_algs/general_passive/GsmNode.py | 15 +- 3 files changed, 112 insertions(+), 58 deletions(-) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index 5abbc0f9d1..59ff462dd3 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -1,91 +1,139 @@ import math from abc import abstractmethod, ABC from collections import defaultdict -from copy import copy from typing import Generic, TypeVar, Any T = TypeVar("T") -class DataHandler(Generic[T]): +class DataHandler(Generic[T], ABC): + # TODO: consider merging `init` with `GsmNode.createPTA`. Could then eliminate PTA post-processing. @abstractmethod - def init(self, data, output_format, data_format): + def init(self, data: Any, output_behavior: 'OutputBehavior', data_format: 'DataFormat'): + """ + Initializes the data handler on the data from which the PTA is constructed. + + :param 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 data_format: Indicates the format of the provided data. Options are: + - "io_traces": describes prefix-closed data. either + - Moore traces [[o, (i,o), (i,o), ...], ...] + - Mealy traces [[(i,o), (i,o), ...], ...] + - "labeled_sequences": [([i, i, ...], o), ...] + - "traces": [[o, o, ...], ...] + - "tree": a tree-shaped automaton provided as a GsmNode + """ ... + @abstractmethod def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): - pass + """ + 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. + """ + ... @abstractmethod - def abstract(self, in_val, out_val): + def abstract(self, in_val: Any, out_val: Any) -> tuple[Any, Any]: + """ + Method used during PTA construction to abstract from potentially continuous input data. + + :param Any in_val: The input value. + :param Any out_val: The output value. + :return tuple[Any, Any]: The abstract output value and the input symbols. + """ ... @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 NoAbstractionDataHandler(DataHandler[T], ABC): - def init(self, data, output_format, data_format): + """ + DataHandler using input and output symbols "as is". Might still keep track of other information. + """ + + def init(self, data, output_behavior, data_format): pass def abstract(self, in_val, out_val): return in_val, out_val -class NoOpDataHandler(NoAbstractionDataHandler[T]): - def init_data(self) -> T: +class NoOpDataHandler(NoAbstractionDataHandler[None]): + """ + DataHandler that does nothing. + """ + + def init_data(self) -> None: return None - def aggregate_data(self, src_node: 'GsmNode[T]', in_sym, out_value, dst_node: 'GsmNode[T]'): + def aggregate_data(self, src_node: 'GsmNode[None]', in_sym, out_value, dst_node: 'GsmNode[None]'): pass - def merge(self, x: T, y: T) -> T: + def init_merge(self, red: 'GsmNode[None]', blue: 'GsmNode[None]', first_pass: bool): return None - def copy(self, x: T) -> T: + def merge(self, x: None, y: None) -> None: return None -class CopyOnWriteDataHandler(DataHandler[T], ABC): - def __init__(self): - self.copied_on_write = set() - - def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): - self.first_pass = first_pass - if first_pass: - self.copied_on_write.clear() - - def merge(self, x: T, y: T) -> T: - if self.first_pass and id(x) not in self.copied_on_write: - x = self.copy_on_write(x) - self.copied_on_write.add(id(x)) - self.merge_into_x(x, y) - return x - - @abstractmethod - def merge_into_x(self, x: T, y: T): - pass - - @abstractmethod - def copy_on_write(self, x: T) -> T: - pass - - def copy(self, x: T) -> T: - return x + def copy(self, x: None) -> None: + return None 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: pass + 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]] @@ -129,7 +177,22 @@ def get_probabilities(self) -> ProbabilityDict: ret[in_sym] = {out_sym: count / total_count for out_sym, count in trans.items()} return ret -class CountDataHandler(NoAbstractionDataHandler[CountData], CopyOnWriteDataHandler): +class CountDataHandler(NoAbstractionDataHandler[CountData]): + def init_merge(self, red: 'GsmNode[CountData]', blue: 'GsmNode[CountData]', first_pass: bool): + pass + + 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 = {k: v.copy() for k, v in x.transition_count.items()} + return ret + def init_data(self) -> CountData: return CountData() @@ -137,15 +200,6 @@ def aggregate_data(self, src_node: 'GsmNode[CountData]', in_value, out_value, ds if src_node is not None: int_dict_increment(src_node.data.transition_count[in_value], out_value, 1) - def merge_into_x(self, x: CountData, y: CountData): - CountData.merge(x.transition_count, y.transition_count) - - def copy_on_write(self, x: CountData) -> CountData: - new_x = copy(x) - new_x.transition_count = defaultdict(dict) - for in_sym, trans in x.transition_count.items(): - new_x.transition_count[in_sym] = trans.copy() - return x ShadowPTA = dict[Any, dict[Any, 'GsmNode']] class ShadowPTAData: @@ -158,7 +212,7 @@ def __init__(self): CountData.__init__(self) self.pta_count: CountDict = defaultdict(dict) -class CountOnPTADataHandler(CountDataHandler): +class CountOnPTADataHandler(CountDataHandler, DataHandler[CountOnPTAData]): def init_data(self) -> CountOnPTAData: return CountOnPTAData() diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 11020059f3..70c9ec1565 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -171,7 +171,7 @@ 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, self.data_handler) + root = GsmNode.createPTA(data, self.data_handler, self.output_behavior, data_format) root = self.pta_preprocessing(root) instrumentation.pta_construction_done(root) @@ -314,6 +314,7 @@ def update_partition(red_node: GsmNode, blue_node: GsmNode | None) -> GsmNode: # 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 diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 95a14652ed..ab5e52b92e 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -510,7 +510,7 @@ def add_trace(self, trace: IOTrace, data_handler: DataHandler[T]): data_handler.aggregate_data(curr_node, in_value, out_value, node) curr_node = node - def add_labeled_sequence(self, example: IOExample, data_handler: DataHandler[T] = None): + def add_labeled_sequence(self, example: IOExample, data_handler: DataHandler[T]): """ Add a labeled input sequence (inputs with a single label attached at the end) to the tree. @@ -526,18 +526,17 @@ def add_labeled_sequence(self, example: IOExample, data_handler: DataHandler[T] # step through inputs and add transitions for in_value in inputs: - in_sym, out_sym = data_handler.abstract(in_value, None) + in_sym, out_sym = data_handler.abstract(in_value, unknown_output) transitions = curr_node.transitions[in_sym] - successors = list(transitions.values()) - if len(successors) == 0: + if len(transitions) == 0: node = GsmNode((in_sym, unknown_output), curr_node) transitions[unknown_output] = node - elif len(successors) == 1: - node = successors[0] + elif len(transitions) == 1: + node = next(iter(transitions.values())) else: # This should never happen raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") - data_handler.aggregate_data(curr_node, in_value, None, node) + data_handler.aggregate_data(curr_node, in_value, unknown_output, node) curr_node = node # set last output @@ -551,7 +550,7 @@ def add_labeled_sequence(self, example: IOExample, data_handler: DataHandler[T] raise ValueError("nondeterminism encountered for GSM with labeled_sequences. not supported") @staticmethod - def createPTA(data: Any, output_behavior: OutputBehavior, data_format: DataFormat = None, data_handler: DataHandler[T] = None) -> 'GsmNode[T]': + def createPTA(data: Any, data_handler: DataHandler[T], output_behavior: OutputBehavior, data_format: DataFormat = None) -> 'GsmNode[T]': """ Build a prefix tree acceptor (PTA) from the given data. From fe4c18e2cc33fb03b06df67c7a598a5abc6c532c Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 8 Sep 2026 17:12:35 +0200 Subject: [PATCH 51/70] fix testsuite --- .../test_generalized_state_merging.py | 3 +- .../general_passive/test_gsm_node.py | 183 ++++++++---------- .../general_passive/test_instrumentation.py | 12 +- .../test_score_functions_gsm.py | 107 +++++----- 4 files changed, 147 insertions(+), 158 deletions(-) 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..e4f4d7244d 100644 --- a/tests/learning_algs/general_passive/test_generalized_state_merging.py +++ b/tests/learning_algs/general_passive/test_generalized_state_merging.py @@ -183,8 +183,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_node.py b/tests/learning_algs/general_passive/test_gsm_node.py index 55fc5b8ba6..4ff9126fe2 100644 --- a/tests/learning_algs/general_passive/test_gsm_node.py +++ b/tests/learning_algs/general_passive/test_gsm_node.py @@ -1,10 +1,30 @@ import unittest +from typing import TypeVar +from aalpy.learning_algs.general_passive.DataHandler import ( + CountOnPTADataHandler, + NoOpDataHandler, + DataHandler, + CountDataHandler, +) from aalpy.learning_algs.general_passive.GsmNode import ( - GsmNode, TransitionInfo, detect_data_format, intersection_iterator, union_iterator, unknown_output, + GsmNode, + detect_data_format, + 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: + root.add_trace(trace, data_handler) + return root + class TestIterators(unittest.TestCase): def test_intersection_iterator_only_common_keys(self): a = {'x': 1, 'y': 2} @@ -55,84 +75,66 @@ def test_root_has_no_predecessor_and_zero_prefix_length(self): 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): + dh = NoOpDataHandler() root = GsmNode((None, 'root_out'), None) - root.add_trace([('a', 'x')]) + root.add_trace([('a', 'x')], dh) # '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'] + node_a.add_trace([('b', 'y')], dh) 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) @@ -143,7 +145,7 @@ def test_resolve_unknown_prefix_output_only_updates_if_unknown(self): 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')) + root.add_labeled_sequence((('a', 'b'), 'label1'), NoOpDataHandler()) # 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')]) @@ -152,52 +154,43 @@ def test_add_labeled_sequence_sets_prefix_output_on_final_node(self): 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')) + root.add_labeled_sequence((('a',), 'out1'), NoOpDataHandler()) with self.assertRaises(ValueError): - root.add_labeled_sequence((('a',), 'out2')) + root.add_labeled_sequence((('a',), 'out2'), NoOpDataHandler()) 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')]) + root.add_trace([('a', 'child_out')], NoOpDataHandler()) 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.transitions['a']['transition_out'] = child child.prefix_access_pair = ('a', 'different_child_out') self.assertFalse(root.is_moore()) @@ -214,78 +207,68 @@ def test_moore_compatible_false_for_conflicting_outputs(self): 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): 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 = GsmNode.createPTA(data, data_handler=NoOpDataHandler(), 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): data = [[0, ('a', 1)], [0, ('a', 1)]] - root = GsmNode.createPTA(data, output_behavior='moore', data_format='io_traces') + root = GsmNode.createPTA(data, data_handler=NoOpDataHandler(), output_behavior='moore', data_format='io_traces') self.assertEqual(root.get_prefix_output(), 0) def test_io_traces_mealy_has_no_root_output(self): data = [[('a', 'x')]] - root = GsmNode.createPTA(data, output_behavior='mealy', data_format='io_traces') + root = GsmNode.createPTA(data, data_handler=NoOpDataHandler(), 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') + root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) + result = GsmNode.createPTA(root, data_handler=NoOpDataHandler(), 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) + 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') + GsmNode.createPTA(root, data_handler=NoOpDataHandler(), 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)]) + root.add_trace([('a', 1)], NoOpDataHandler()) 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') + child = GsmNode(('a', 'different_output'), root) + 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')]) + root.add_trace([('a', 'x')], NoOpDataHandler()) 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..3342d62e58 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): @@ -36,8 +37,8 @@ def build_matching_ground_truth_tree(self): # 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.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. + dh = NoOpDataHandler() mismatched_ground_truth = GsmNode((None, True), None) - mismatched_ground_truth.add_trace([('a', True)]) - mismatched_ground_truth.add_trace([('b', True)]) + mismatched_ground_truth.add_trace([('a', True)], dh) + mismatched_ground_truth.add_trace([('b', True)], dh) # 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..4e893ffb51 100644 --- a/tests/learning_algs/general_passive/test_score_functions_gsm.py +++ b/tests/learning_algs/general_passive/test_score_functions_gsm.py @@ -1,36 +1,38 @@ 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, + 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() + sc = SimpleScoreCalculation() self.assertTrue(sc.local_compatibility(GsmNode((None, None), None), GsmNode((None, None), None))) - self.assertFalse(sc.has_local_compatibility()) 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()) @@ -54,77 +56,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) - 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() + 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)) + # 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) + 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)) + 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) + 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 +138,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 +192,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)) From 50b55d6122234dc5ae9f0080e927f36e85a592c5 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 9 Sep 2026 11:00:19 +0200 Subject: [PATCH 52/70] refactor: createPTA is now a method of DataHandler + minor stuff --- .../general_passive/AssociatedData.py | 74 +++++ .../general_passive/DataHandler.py | 259 +++++++++++------- .../GeneralizedStateMerging.py | 7 +- .../learning_algs/general_passive/GsmNode.py | 159 +---------- .../general_passive/ScoreFunctionsGSM.py | 4 +- .../test_generalized_state_merging.py | 3 +- .../general_passive/test_gsm_algorithms.py | 5 +- .../general_passive/test_gsm_node.py | 41 +-- .../general_passive/test_instrumentation.py | 4 +- 9 files changed, 270 insertions(+), 286 deletions(-) create mode 100644 aalpy/learning_algs/general_passive/AssociatedData.py diff --git a/aalpy/learning_algs/general_passive/AssociatedData.py b/aalpy/learning_algs/general_passive/AssociatedData.py new file mode 100644 index 0000000000..1b50143ee4 --- /dev/null +++ b/aalpy/learning_algs/general_passive/AssociatedData.py @@ -0,0 +1,74 @@ +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()) + + @staticmethod + def merge(x: CountDict, y: CountDict) -> CountDict: + for in_sym, y_o_dict in y.items(): + x_o_dict = x.get(in_sym, None) + if x_o_dict is None: + x[in_sym] = y_o_dict + continue + for out_sym, count in y_o_dict.items(): + int_dict_increment(x_o_dict, out_sym, count) + return x + + 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 index 59ff462dd3..d0b2cd6430 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -1,28 +1,168 @@ -import math from abc import abstractmethod, ABC -from collections import defaultdict 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 (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] + class DataHandler(Generic[T], ABC): - # TODO: consider merging `init` with `GsmNode.createPTA`. Could then eliminate PTA post-processing. - @abstractmethod - def init(self, data: Any, output_behavior: 'OutputBehavior', data_format: 'DataFormat'): + 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. """ - Initializes the data handler on the data from which the PTA is constructed. + 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 data: Learning data, in one of the supported data formats (or already a GsmNode 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 + + # TODO check implementation and eliminate + if not isinstance(self, NoOpDataHandler): + raise NotImplementedError("Data handling is not supported for learning from labeled sequences") + + # step through inputs and add transitions + for in_value in inputs: + in_sym, out_sym = self.abstract(in_value, unknown_output) + transitions = curr_node.transitions[in_sym] + if len(transitions) == 0: + node = GsmNode((in_sym, unknown_output), curr_node) + transitions[unknown_output] = node + elif len(transitions) == 1: + node = next(iter(transitions.values())) + else: + # This should never happen + raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") + self.aggregate_data(curr_node, in_value, unknown_output, node) + 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") + + 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 data_format: Indicates the format of the provided data. Options are: - - "io_traces": describes prefix-closed data. either - - Moore traces [[o, (i,o), (i,o), ...], ...] - - Mealy traces [[(i,o), (i,o), ...], ...] - - "labeled_sequences": [([i, i, ...], o), ...] - - "traces": [[o, o, ...], ...] - - "tree": a tree-shaped automaton provided as a GsmNode + :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 @abstractmethod def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): @@ -36,16 +176,15 @@ def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): """ ... - @abstractmethod def abstract(self, in_val: Any, out_val: Any) -> tuple[Any, Any]: """ - Method used during PTA construction to abstract from potentially continuous input data. + 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 abstract output value and the input symbols. """ - ... + return in_val, out_val @abstractmethod def init_data(self) -> T: @@ -89,20 +228,9 @@ def copy(self, x: T) -> T: """ ... -class NoAbstractionDataHandler(DataHandler[T], ABC): +class NoOpDataHandler(DataHandler[None]): """ - DataHandler using input and output symbols "as is". Might still keep track of other information. - """ - - def init(self, data, output_behavior, data_format): - pass - - def abstract(self, in_val, out_val): - return in_val, out_val - -class NoOpDataHandler(NoAbstractionDataHandler[None]): - """ - DataHandler that does nothing. + DataHandler that neither abstracts the traces, nor tracks any other values during PTA construction. """ def init_data(self) -> None: @@ -120,64 +248,8 @@ def merge(self, x: None, y: None) -> None: def copy(self, x: None) -> None: return None -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()) - - @staticmethod - def merge(x: CountDict, y: CountDict) -> CountDict: - for in_sym, y_o_dict in y.items(): - x_o_dict = x.get(in_sym, None) - if x_o_dict is None: - x[in_sym] = y_o_dict - continue - for out_sym, count in y_o_dict.items(): - int_dict_increment(x_o_dict, out_sym, count) - return x - - 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 - -class CountDataHandler(NoAbstractionDataHandler[CountData]): +class CountDataHandler(DataHandler[CountData]): def init_merge(self, red: 'GsmNode[CountData]', blue: 'GsmNode[CountData]', first_pass: bool): pass @@ -201,17 +273,6 @@ def aggregate_data(self, src_node: 'GsmNode[CountData]', in_value, out_value, ds int_dict_increment(src_node.data.transition_count[in_value], out_value, 1) -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) - class CountOnPTADataHandler(CountDataHandler, DataHandler[CountOnPTAData]): def init_data(self) -> CountOnPTAData: return CountOnPTAData() diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 70c9ec1565..327bc907a8 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -9,8 +9,9 @@ from aalpy import Automaton from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior, TransitionBehavior, OutputBehaviorRange, \ - TransitionBehaviorRange, unknown_output, detect_data_format, DataHandler, NoOpDataHandler, DataFormat -from aalpy.learning_algs.general_passive.DataHandler import CountOnPTADataHandler, CountDataHandler + 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 @@ -171,7 +172,7 @@ 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.data_handler, 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) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index ab5e52b92e..3dfbda7420 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -9,7 +9,7 @@ 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.DataHandler import DataHandler, NoOpDataHandler, StochasticData, CountData +from aalpy.learning_algs.general_passive.AssociatedData import StochasticData, CountData Key = TypeVar("Key") @@ -22,9 +22,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] @@ -77,56 +74,6 @@ def union_iterator(a: dict[Key, Val], b: dict[Key, Val], default: Val = None) -> a_val = a.get(key, default) 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 add custom pickling code that flattens the Node structure in order to circumvent running into recursion issues for large models class GsmNode(Generic[T]): @@ -491,110 +438,6 @@ def make_input_complete(self, target: 'GsmNode[T] | str' = "self-loop") -> list[ transitions[out_sym] = successor return missing_trans - def add_trace(self, trace: IOTrace, data_handler: DataHandler[T]): - """ - 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. - :param DataHandler[T] data_handler: IOHandler used for abstraction and aggregation of trace data - """ - curr_node: GsmNode = self - for in_value, out_value in trace: - prefix_access_pair = data_handler.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, data_handler.init_data()) - transitions[out_sym] = node - data_handler.aggregate_data(curr_node, in_value, out_value, node) - curr_node = node - - def add_labeled_sequence(self, example: IOExample, data_handler: DataHandler[T]): - """ - 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] data_handler: IOHandler used for abstraction and aggregation of trace data - """ - inputs, output = example - curr_node: GsmNode = self - in_sym = None - - if not isinstance(data_handler, NoOpDataHandler): - raise NotImplementedError("Data handling is not supported for learning from labeled sequences") - - # step through inputs and add transitions - for in_value in inputs: - in_sym, out_sym = data_handler.abstract(in_value, unknown_output) - transitions = curr_node.transitions[in_sym] - if len(transitions) == 0: - node = GsmNode((in_sym, unknown_output), curr_node) - transitions[unknown_output] = node - elif len(transitions) == 1: - node = next(iter(transitions.values())) - else: - # This should never happen - raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") - data_handler.aggregate_data(curr_node, in_value, unknown_output, node) - 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, data_handler: DataHandler[T], 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] data_handler: 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}") - - data_handler.init(data, output_behavior, data_format) - - 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, data_handler.init_data()) - if data_format == "labeled_sequences": - for example in data: - root_node.add_labeled_sequence(example, data_handler) - if data_format == "io_traces" or data_format == "traces": - if output_behavior == "moore": - root_node.prefix_access_pair = data_handler.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 = data_handler.abstract(no_op_input, initial_output) - if ios != initial_output_symbol: - raise ValueError("expect unique initial output symbol for Moore behavior") - data_handler.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) - root_node.add_trace(trace, data_handler) - return root_node - def is_locally_deterministic(self) -> bool: """ Check whether this node has at most one outgoing transition per input symbol. diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index e5dcf39b69..a80a580fba 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -8,7 +8,7 @@ from typing import Any from aalpy.learning_algs.general_passive.GsmNode import GsmNode, intersection_iterator, union_iterator, CountData -from aalpy.learning_algs.general_passive.DataHandler import ShadowPTAData +from aalpy.learning_algs.general_passive.AssociatedData import ShadowPTAData LocalCompatibilityFunction = Callable[[GsmNode, GsmNode], bool | None] ScoreFunction = Callable[[dict[GsmNode, GsmNode]], Any] @@ -256,7 +256,7 @@ 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__(wrapped) 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 e4f4d7244d..d263909d0b 100644 --- a/tests/learning_algs/general_passive/test_generalized_state_merging.py +++ b/tests/learning_algs/general_passive/test_generalized_state_merging.py @@ -1,4 +1,3 @@ -import random import unittest from itertools import product @@ -6,7 +5,7 @@ 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.GsmNode import GsmNode from aalpy.utils.HelperFunctions import dfa_from_moore from aalpy.utils.ModelChecking import bisimilar diff --git a/tests/learning_algs/general_passive/test_gsm_algorithms.py b/tests/learning_algs/general_passive/test_gsm_algorithms.py index 9e88397a16..01aa7ca57e 100644 --- a/tests/learning_algs/general_passive/test_gsm_algorithms.py +++ b/tests/learning_algs/general_passive/test_gsm_algorithms.py @@ -2,10 +2,7 @@ import unittest from itertools import product -from aalpy.automata import ( - Dfa, DfaState, MooreMachine, MooreState, MealyMachine, MealyState, Mdp, MdpState, StochasticMealyMachine, - StochasticMealyState, -) +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 diff --git a/tests/learning_algs/general_passive/test_gsm_node.py b/tests/learning_algs/general_passive/test_gsm_node.py index 4ff9126fe2..bc56172a75 100644 --- a/tests/learning_algs/general_passive/test_gsm_node.py +++ b/tests/learning_algs/general_passive/test_gsm_node.py @@ -6,10 +6,10 @@ NoOpDataHandler, DataHandler, CountDataHandler, + detect_data_format, ) from aalpy.learning_algs.general_passive.GsmNode import ( GsmNode, - detect_data_format, intersection_iterator, union_iterator, unknown_output, @@ -22,7 +22,7 @@ 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: - root.add_trace(trace, data_handler) + data_handler.add_trace(root, trace) return root class TestIterators(unittest.TestCase): @@ -117,10 +117,10 @@ def test_is_tree_false_when_node_shared(self): def test_make_input_complete_adds_self_loops_for_missing_inputs(self): dh = NoOpDataHandler() root = GsmNode((None, 'root_out'), None) - root.add_trace([('a', 'x')], dh) + dh.add_trace(root, [('a', 'x')]) # 'b' is used elsewhere in the tree but not from root node_a = root.transitions['a']['x'] - node_a.add_trace([('b', 'y')], dh) + 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'], root) @@ -144,8 +144,9 @@ def test_resolve_unknown_prefix_output_only_updates_if_unknown(self): self.assertEqual(node.get_prefix_output(), 'resolved') def test_add_labeled_sequence_sets_prefix_output_on_final_node(self): + dh = NoOpDataHandler() root = GsmNode((None, unknown_output), None) - root.add_labeled_sequence((('a', 'b'), 'label1'), NoOpDataHandler()) + 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')]) @@ -153,10 +154,11 @@ 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): + dh= NoOpDataHandler() root = GsmNode((None, unknown_output), None) - root.add_labeled_sequence((('a',), 'out1'), NoOpDataHandler()) + dh.add_labeled_sequence(root,(('a',), 'out1')) with self.assertRaises(ValueError): - root.add_labeled_sequence((('a',), 'out2'), NoOpDataHandler()) + dh.add_labeled_sequence(root, (('a',), 'out2')) def test_is_locally_deterministic_true_for_single_output_per_input(self): root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) @@ -183,8 +185,9 @@ def test_deterministic_compatible_true_when_unknown_output_present(self): self.assertTrue(n1.deterministic_compatible(n2)) def test_is_moore_true_when_child_output_matches_transition_output(self): + dh = NoOpDataHandler() root = GsmNode((None, 'root_out'), None) - root.add_trace([('a', 'child_out')], NoOpDataHandler()) + dh.add_trace(root,[('a', 'child_out')]) self.assertTrue(root.is_moore()) def test_is_moore_false_when_child_output_mismatches(self): @@ -223,38 +226,43 @@ def test_local_log_likelihood_contribution_negative_for_split_outcomes(self): class TestGsmNodeCreatePTA(unittest.TestCase): def test_labeled_sequences_format(self): + dh = NoOpDataHandler() data = [(('a', 'b'), 1), (('a', 'c'), 2)] - root = GsmNode.createPTA(data, data_handler=NoOpDataHandler(), output_behavior='moore', - data_format='labeled_sequences') + 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, data_handler=NoOpDataHandler(), 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, data_handler=NoOpDataHandler(), 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): + dh = NoOpDataHandler() root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) - result = GsmNode.createPTA(root, data_handler=NoOpDataHandler(), output_behavior='mealy', data_format='tree') + result = dh.createPTA(root, output_behavior='mealy', data_format='tree') self.assertIs(result, root) def test_tree_format_rejects_non_tree(self): + dh = NoOpDataHandler() root = simple_create_PTA([[('a', 'x')]], NoOpDataHandler()) root.transitions['b']['y'] = root.transitions['a']['x'] with self.assertRaises(ValueError): - GsmNode.createPTA(root, data_handler=NoOpDataHandler(), 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): + dh = NoOpDataHandler() root = GsmNode((None, 0), None) - root.add_trace([('a', 1)], NoOpDataHandler()) + 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) @@ -267,8 +275,9 @@ def test_to_automaton_raises_on_non_moore_structure_when_moore_requested(self): root.to_automaton('moore', 'deterministic') def test_to_automaton_deterministic_mealy(self): + dh = NoOpDataHandler() root = GsmNode((None, unknown_output), None) - root.add_trace([('a', 'x')], NoOpDataHandler()) + 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 3342d62e58..1b145a9404 100644 --- a/tests/learning_algs/general_passive/test_instrumentation.py +++ b/tests/learning_algs/general_passive/test_instrumentation.py @@ -60,8 +60,8 @@ def test_flags_wrong_merge_against_mismatched_ground_truth(self): # comparing against it while the actual run does merge them should be flagged as wrong. dh = NoOpDataHandler() mismatched_ground_truth = GsmNode((None, True), None) - mismatched_ground_truth.add_trace([('a', True)], dh) - mismatched_ground_truth.add_trace([('b', True)], dh) + 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) From 92b8a7e5c07343341c43fd5eb2b93935474eb0f2 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 9 Sep 2026 11:26:15 +0200 Subject: [PATCH 53/70] implement datahandlers for labeled sequences --- .../general_passive/DataHandler.py | 26 +++++++++---------- 1 file changed, 13 insertions(+), 13 deletions(-) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index d0b2cd6430..9bba312861 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -92,33 +92,33 @@ def add_labeled_sequence(self, root_node: GsmNode[T], example: IOExample): curr_node: GsmNode = root_node in_sym = None - # TODO check implementation and eliminate - if not isinstance(self, NoOpDataHandler): - raise NotImplementedError("Data handling is not supported for learning from labeled sequences") + if len(inputs) == 0: + self.aggregate_data(None, no_op_input, output, root_node) + in_sym, out_sym = self.abstract(no_op_input, output) # step through inputs and add transitions - for in_value in inputs: - in_sym, out_sym = self.abstract(in_value, unknown_output) + 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, unknown_output), curr_node) - transitions[unknown_output] = node + node = GsmNode((in_sym, out_sym), curr_node) + transitions[out_sym] = node elif len(transitions) == 1: node = next(iter(transitions.values())) else: - # This should never happen raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") - self.aggregate_data(curr_node, in_value, unknown_output, node) + self.aggregate_data(curr_node, in_value, out_value, node) curr_node = node - # set last output - curr_node.resolve_unknown_prefix_output(output) + # 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[output] = transitions.pop(unknown_output) - if output not 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]': From d32e71512c0d8ac26fcee977a8c8bf3bfb4a12bf Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 9 Sep 2026 11:35:41 +0200 Subject: [PATCH 54/70] slimmer imports --- .../test_generalized_state_merging.py | 4 +--- tests/learning_algs/general_passive/test_gsm_node.py | 12 ++---------- 2 files changed, 3 insertions(+), 13 deletions(-) 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 d263909d0b..5b05d2f0bc 100644 --- a/tests/learning_algs/general_passive/test_generalized_state_merging.py +++ b/tests/learning_algs/general_passive/test_generalized_state_merging.py @@ -2,9 +2,7 @@ from itertools import product 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.GeneralizedStateMerging import GeneralizedStateMerging, Instrumentation, run_GSM from aalpy.learning_algs.general_passive.GsmNode import GsmNode from aalpy.utils.HelperFunctions import dfa_from_moore from aalpy.utils.ModelChecking import bisimilar diff --git a/tests/learning_algs/general_passive/test_gsm_node.py b/tests/learning_algs/general_passive/test_gsm_node.py index bc56172a75..e325e36c29 100644 --- a/tests/learning_algs/general_passive/test_gsm_node.py +++ b/tests/learning_algs/general_passive/test_gsm_node.py @@ -2,18 +2,10 @@ from typing import TypeVar from aalpy.learning_algs.general_passive.DataHandler import ( - CountOnPTADataHandler, - NoOpDataHandler, - DataHandler, - CountDataHandler, - detect_data_format, + CountOnPTADataHandler, NoOpDataHandler, DataHandler, CountDataHandler, detect_data_format ) from aalpy.learning_algs.general_passive.GsmNode import ( - GsmNode, - intersection_iterator, - union_iterator, - unknown_output, - no_op_input, IOTrace, + GsmNode, intersection_iterator, union_iterator, unknown_output, no_op_input, IOTrace ) From 2d05160c1b002439a6478cf2f29a6309b7cd3c92 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 9 Sep 2026 11:56:47 +0200 Subject: [PATCH 55/70] better docstring --- .../learning_algs/general_passive/DataHandler.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index 9bba312861..0f7d0f0f25 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -18,15 +18,14 @@ def detect_data_format(data: Any, check_consistency: bool = False, guess: bool = :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). + :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, ...], ...] """ - # 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" From ad72b9c4b8551649ade272ac6f995470e5b04df5 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 9 Sep 2026 12:02:15 +0200 Subject: [PATCH 56/70] eliminate pta_preprocessing in favor for DH override --- .../general_passive/GeneralizedStateMerging.py | 7 ------- .../test_generalized_state_merging.py | 15 ++++++++++----- 2 files changed, 10 insertions(+), 12 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 327bc907a8..f0d981595a 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -94,7 +94,6 @@ 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, data_handler: DataHandler = None, node_order: Callable[[GsmNode], Any] = None, @@ -107,7 +106,6 @@ def __init__(self, *, :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 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. @@ -139,7 +137,6 @@ def __init__(self, *, 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) if data_handler is None: @@ -174,7 +171,6 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation print("learning deterministic systems from (output) traces only. this rarely makes sense. is `data_format` set correctly?") 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) @@ -437,7 +433,6 @@ 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, data_handler: DataHandler = None, node_order: Callable[[GsmNode], Any] = None, @@ -454,7 +449,6 @@ def run_GSM(data: list, *, :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 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 @@ -470,7 +464,6 @@ def run_GSM(data: list, *, output_behavior=output_behavior, transition_behavior=transition_behavior, score_calc=score_calc, - pta_preprocessing=pta_preprocessing, postprocessing=postprocessing, data_handler=data_handler, node_order=node_order, 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 5b05d2f0bc..4e70d43121 100644 --- a/tests/learning_algs/general_passive/test_generalized_state_merging.py +++ b/tests/learning_algs/general_passive/test_generalized_state_merging.py @@ -1,9 +1,11 @@ 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.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 +from aalpy.learning_algs.general_passive.GsmNode import GsmNode, OutputBehavior from aalpy.utils.HelperFunctions import dfa_from_moore from aalpy.utils.ModelChecking import bisimilar @@ -152,9 +154,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') @@ -162,7 +167,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']) From 49b96cd2cdae7153e79fda4ca53ce8af036d83ed Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Thu, 10 Sep 2026 12:54:45 +0200 Subject: [PATCH 57/70] gsm v2 bughunt --- .../general_passive/DataHandler.py | 11 ++++-- .../GeneralizedStateMerging.py | 28 ++++++++++++--- .../general_passive/GsmAlgorithms.py | 6 ++-- .../learning_algs/general_passive/GsmNode.py | 3 +- .../general_passive/Instrumentation.py | 20 +++++------ .../general_passive/ScoreFunctionsGSM.py | 15 +++++--- .../general_passive/test_gsm_node.py | 36 +++++++++---------- .../general_passive/test_instrumentation.py | 4 +-- .../test_score_functions_gsm.py | 28 +++++++-------- 9 files changed, 92 insertions(+), 59 deletions(-) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index 0f7d0f0f25..0cc598523a 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -101,7 +101,7 @@ def add_labeled_sequence(self, root_node: GsmNode[T], example: IOExample): 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) + node = GsmNode((in_sym, out_sym), curr_node, self.init_data()) transitions[out_sym] = node elif len(transitions) == 1: node = next(iter(transitions.values())) @@ -261,7 +261,7 @@ def merge(self, x: CountData, y: CountData) -> CountData: def copy(self, x: CountData) -> CountData: ret = CountData() - ret.transition_count = {k: v.copy() for k, v in x.transition_count.items()} + ret.transition_count.update((k, v.copy()) for k, v in x.transition_count.items()) return ret def init_data(self) -> CountData: @@ -276,6 +276,13 @@ 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 src_node is None: return diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index f0d981595a..787a414088 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -175,7 +175,8 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation instrumentation.log_promote(root) if self.transition_behavior == "deterministic": - if not root.is_deterministic(): + 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 as distinct @@ -213,9 +214,9 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation 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 + 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 @@ -246,6 +247,16 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation blue_states.remove(best_candidate) blue_states.extend(best_candidate.child_iterator()) instrumentation.log_promote(best_candidate) + + # check cached partitions + for partitioning in partition_candidates.values(): + updated_promoted_node = partitioning.full_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: + 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(): @@ -268,6 +279,12 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation 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 @@ -372,6 +389,7 @@ def get_partition_trans(part: GsmNode, in_symbol): else: # work on the remaining merges q.extend(partitioning.remaining_merges) + partitioning.nr_merged_states -= len(partitioning.remaining_merges) # loop over implied merges pop = q.pop if self.depth_first else q.popleft @@ -382,7 +400,7 @@ def get_partition_trans(part: GsmNode, in_symbol): if first_pass: local_compat = self.score_calc.local_compatibility(partition, blue) - moore_check = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(red, blue) + moore_check = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(partition, blue) if local_compat is False or moore_check: partitioning.score = SpecialScores.ImmediateReject return diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index 2dee5d1832..b07c455f97 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -8,7 +8,7 @@ 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.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 @@ -37,11 +37,11 @@ def EDSM_score(part: dict[GsmNode, GsmNode]) -> int: 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 diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 3dfbda7420..95e7ab69e4 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -88,12 +88,13 @@ class GsmNode(Generic[T]): """ __slots__ = ['transitions', 'predecessor', 'prefix_access_pair', 'data'] - def __init__(self, prefix_access_pair: IOPair, predecessor: 'GsmNode[T]' = None, data: T = None): # TODO (data-ext) check all invocations + 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, GsmNode[T]]] = defaultdict(dict) diff --git a/aalpy/learning_algs/general_passive/Instrumentation.py b/aalpy/learning_algs/general_passive/Instrumentation.py index 37b5287819..33d08bed0d 100644 --- a/aalpy/learning_algs/general_passive/Instrumentation.py +++ b/aalpy/learning_algs/general_passive/Instrumentation.py @@ -124,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 @@ -146,16 +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 - self.log.append(("promote", new_red_prefix)) - elif node is None: + 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)) @@ -168,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 a80a580fba..daa76c3001 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -88,7 +88,7 @@ def has_local_compatibility(self) -> bool: :return bool: True if local_compatibility was overridden. """ - return self.__class__.local_compatibility is ScoreCalculation.local_compatibility + return self.__class__.local_compatibility is not ScoreCalculation.local_compatibility def has_score_function(self) -> bool: """ @@ -96,7 +96,7 @@ def has_score_function(self) -> bool: :return bool: True if score_function was overridden. """ - return self.__class__.score_function is ScoreCalculation.score_function + return self.__class__.score_function is not ScoreCalculation.score_function class SimpleScoreCalculation(ScoreCalculation): @@ -167,12 +167,15 @@ def __init__(self, """ if local_compatibility: if self.has_local_compatibility(): - raise ValueError("Exernal local compatibility is provided, but the class already defines a local compatibility criterion.") + 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 initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + if not first_pass: + return + if self.compatibility_on_pta and not isinstance(red.data, ShadowPTAData): raise TypeError("compatibility_on_pta is set but no PTA data is available") @@ -382,7 +385,11 @@ def default_aggregate_score(score_iterable: Iterable) -> list: :param Iterable score_iterable: Iterable of score results. :return list: List of the individual scores. """ - return list(score_iterable) + score_iterable = list(score_iterable) + for special_val in [SpecialScores.ImmediateReject, SpecialScores.ImmediateAccept, SpecialScores.NoScore]: + if all(x is special_val for x in score_iterable): + return special_val + return score_iterable def local_to_global_compatibility(local_fun: LocalCompatibilityFunction) -> ScoreFunction: diff --git a/tests/learning_algs/general_passive/test_gsm_node.py b/tests/learning_algs/general_passive/test_gsm_node.py index e325e36c29..12d5a84463 100644 --- a/tests/learning_algs/general_passive/test_gsm_node.py +++ b/tests/learning_algs/general_passive/test_gsm_node.py @@ -36,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): @@ -61,7 +61,7 @@ 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(), []) @@ -108,7 +108,7 @@ def test_is_tree_false_when_node_shared(self): def test_make_input_complete_adds_self_loops_for_missing_inputs(self): dh = NoOpDataHandler() - root = GsmNode((None, 'root_out'), None) + 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'] @@ -129,7 +129,7 @@ def test_lt_orders_by_prefix_length_then_lexicographically(self): 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') @@ -137,7 +137,7 @@ def test_resolve_unknown_prefix_output_only_updates_if_unknown(self): def test_add_labeled_sequence_sets_prefix_output_on_final_node(self): dh = NoOpDataHandler() - root = GsmNode((None, unknown_output), None) + 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. @@ -147,7 +147,7 @@ def test_add_labeled_sequence_sets_prefix_output_on_final_node(self): def test_add_labeled_sequence_raises_on_conflicting_label_for_same_sequence(self): dh= NoOpDataHandler() - root = GsmNode((None, unknown_output), None) + root = GsmNode((None, unknown_output), None, None) dh.add_labeled_sequence(root,(('a',), 'out1')) with self.assertRaises(ValueError): dh.add_labeled_sequence(root, (('a',), 'out2')) @@ -178,27 +178,27 @@ def test_deterministic_compatible_true_when_unknown_output_present(self): def test_is_moore_true_when_child_output_matches_transition_output(self): dh = NoOpDataHandler() - root = GsmNode((None, 'root_out'), None) + 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 = 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): @@ -253,22 +253,22 @@ def test_tree_format_rejects_non_tree(self): class TestGsmNodeToAutomaton(unittest.TestCase): def test_to_automaton_deterministic_moore(self): dh = NoOpDataHandler() - root = GsmNode((None, 0), None) + 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', 'different_output'), root) + 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): dh = NoOpDataHandler() - root = GsmNode((None, unknown_output), None) + 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 1b145a9404..55587f1376 100644 --- a/tests/learning_algs/general_passive/test_instrumentation.py +++ b/tests/learning_algs/general_passive/test_instrumentation.py @@ -36,7 +36,7 @@ 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 = GsmNode((None, True), None, None) root.transitions['a'][True] = root root.transitions['b'][True] = root return root @@ -59,7 +59,7 @@ 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. dh = NoOpDataHandler() - mismatched_ground_truth = GsmNode((None, True), None) + 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 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 4e893ffb51..ed3b63d52c 100644 --- a/tests/learning_algs/general_passive/test_score_functions_gsm.py +++ b/tests/learning_algs/general_passive/test_score_functions_gsm.py @@ -24,7 +24,7 @@ def node_with_counts(counts, prefix_access_pair=(None, unknown_output)): class TestScoreCalculationDefaults(unittest.TestCase): def test_default_local_compatibility_always_true(self): sc = SimpleScoreCalculation() - self.assertTrue(sc.local_compatibility(GsmNode((None, None), None), GsmNode((None, None), None))) + self.assertTrue(sc.local_compatibility(GsmNode((None, None), None, None), GsmNode((None, None), None, None))) def test_default_score_function_always_true(self): sc = SimpleScoreCalculation() @@ -70,9 +70,9 @@ def test_beyond_depth_k_is_always_compatible(self): 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.initialize_merge(root, blue_shallow, True) # first call establishes the depth offset at blue_shallow's depth (1) @@ -83,8 +83,8 @@ def test_beyond_depth_k_is_always_compatible(self): def test_within_depth_k_delegates_to_wrapped_score(self): 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) + 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)) @@ -95,8 +95,8 @@ def test_rejects_merge_between_sink_and_non_sink(self): is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink) - sink_node = GsmNode((None, 'sink'), None) - normal_node = GsmNode((None, 'normal'), None) + 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 @@ -107,8 +107,8 @@ def test_allows_merge_between_two_sinks_by_default(self): is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink) - 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)) @@ -117,8 +117,8 @@ def test_rejects_merge_between_two_sinks_when_disallowed(self): is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink, allow_sink_merge=False) - 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.assertFalse(wrapped.initialize_merge(sink_a, sink_b, True)) self.assertTrue(wrapped.local_compatibility(sink_a, sink_b)) @@ -127,8 +127,8 @@ def test_sink_check_only_applies_on_first_call(self): is_sink = lambda n: n.get_prefix_output() == 'sink' wrapped = ScoreWithSinks(always_true, sink_cond=is_sink, allow_sink_merge=False) - 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)) From 466a183b635089745f3adb8f595ee3bd8fd0ac74 Mon Sep 17 00:00:00 2001 From: Edi Muskardin <28546846+emuskardin@users.noreply.github.com> Date: Thu, 10 Sep 2026 19:42:36 +0200 Subject: [PATCH 58/70] GSM , the return of the bughunt --- aalpy/learning_algs/__init__.py | 2 +- .../GeneralizedStateMerging.py | 4 +- .../general_passive/ScoreFunctionsGSM.py | 22 +++--- .../test_generalized_state_merging.py | 16 +++++ .../general_passive/test_gsm_algorithms.py | 8 +++ .../general_passive/test_gsm_data_handler.py | 69 +++++++++++++++++++ .../test_score_functions_gsm.py | 40 +++++++++++ 7 files changed, 150 insertions(+), 11 deletions(-) create mode 100644 tests/learning_algs/general_passive/test_gsm_data_handler.py diff --git a/aalpy/learning_algs/__init__.py b/aalpy/learning_algs/__init__.py index 0dbbd4848a..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_GSM, 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/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 787a414088..b4768d7dcc 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -401,7 +401,9 @@ def get_partition_trans(part: GsmNode, in_symbol): if first_pass: local_compat = self.score_calc.local_compatibility(partition, blue) moore_check = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(partition, blue) - if local_compat is False or moore_check: + # determinism is a property of the result, not of the scoring: enforce it even if score_calc doesn't + det_check = self.transition_behavior == "deterministic" and not GsmNode.deterministic_compatible(partition, blue) + if local_compat is False or moore_check or det_check: partitioning.score = SpecialScores.ImmediateReject return if local_compat is None: diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index daa76c3001..b9d90004d4 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -378,18 +378,22 @@ def default_aggregate_compatibility(compatibility_iterable: Iterable) -> Any: 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. - """ - score_iterable = list(score_iterable) - for special_val in [SpecialScores.ImmediateReject, SpecialScores.ImmediateAccept, SpecialScores.NoScore]: - if all(x is special_val for x in score_iterable): - return special_val - return score_iterable + :return Any: The deciding special value, or the list of the individual scores. + """ + 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: 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 4e70d43121..dbb46fe894 100644 --- a/tests/learning_algs/general_passive/test_generalized_state_merging.py +++ b/tests/learning_algs/general_passive/test_generalized_state_merging.py @@ -6,6 +6,7 @@ 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 from aalpy.utils.HelperFunctions import dfa_from_moore from aalpy.utils.ModelChecking import bisimilar @@ -114,6 +115,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') diff --git a/tests/learning_algs/general_passive/test_gsm_algorithms.py b/tests/learning_algs/general_passive/test_gsm_algorithms.py index 01aa7ca57e..ddf56743c9 100644 --- a/tests/learning_algs/general_passive/test_gsm_algorithms.py +++ b/tests/learning_algs/general_passive/test_gsm_algorithms.py @@ -51,6 +51,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() 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..84d1b75243 --- /dev/null +++ b/tests/learning_algs/general_passive/test_gsm_data_handler.py @@ -0,0 +1,69 @@ +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, 'labeled_sequences', [(('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, 'labeled_sequences', [(('a', 'b'), True), (('a',), False)]) + for node in pta.get_all_nodes(): + self.assertIsNotNone(node.data) + + @unittest.expectedFailure + def test_labeled_sequences_and_io_traces_count_the_same_transitions(self): + # known gap: add_labeled_sequence aggregates intermediate steps under unknown_output and only + # remaps `transitions` when the output is resolved later, so the counts stay split + labeled = [(('a', 'b'), True), (('a',), False)] + io_traces = [[False, ('a', False), ('b', True)], [False, ('a', False)]] + pta_labeled = counting_pta(CountDataHandler(), 'labeled_sequences', labeled) + pta_traces = counting_pta(CountDataHandler(), 'io_traces', io_traces) + self.assertEqual(dict(pta_labeled.data.transition_count), dict(pta_traces.data.transition_count)) + + +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_score_functions_gsm.py b/tests/learning_algs/general_passive/test_score_functions_gsm.py index ed3b63d52c..105634b789 100644 --- a/tests/learning_algs/general_passive/test_score_functions_gsm.py +++ b/tests/learning_algs/general_passive/test_score_functions_gsm.py @@ -3,6 +3,7 @@ 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 ( + ScoreCalculation, AIC_score, EDSM_frequency_score, EDSM_score, SimpleScoreCalculation, ScoreCombinator, ScoreWithKTail, ScoreWithSinks, differential_info, hoeffding_compatibility, local_to_global_compatibility, lower_threshold, greedy_score, score_transformation, SpecialScores @@ -36,6 +37,45 @@ def test_custom_functions_are_detected_as_overridden(self): 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.initialize_merge(node, node, True), 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}) From 8234c000e7c697c61aa8b797b6276c8791745dac Mon Sep 17 00:00:00 2001 From: Edi Muskardin <28546846+emuskardin@users.noreply.github.com> Date: Thu, 10 Sep 2026 20:26:04 +0200 Subject: [PATCH 59/70] GSM , the return of the bughunt --- .../general_passive/AssociatedData.py | 3 +- .../GeneralizedStateMerging.py | 14 +- .../test_gsm_rpni_exhaustive.py | 142 ++++++++++++++++++ .../general_passive/test_gsm_algorithms.py | 18 +++ 4 files changed, 167 insertions(+), 10 deletions(-) create mode 100644 tests/learning_algs/deterministic_passive/test_gsm_rpni_exhaustive.py diff --git a/aalpy/learning_algs/general_passive/AssociatedData.py b/aalpy/learning_algs/general_passive/AssociatedData.py index 1b50143ee4..efe164388c 100644 --- a/aalpy/learning_algs/general_passive/AssociatedData.py +++ b/aalpy/learning_algs/general_passive/AssociatedData.py @@ -48,7 +48,8 @@ def merge(x: CountDict, y: CountDict) -> CountDict: for in_sym, y_o_dict in y.items(): x_o_dict = x.get(in_sym, None) if x_o_dict is None: - x[in_sym] = y_o_dict + # copy, don't alias: x must not keep changing if y's dict is mutated afterwards + x[in_sym] = y_o_dict.copy() continue for out_sym, count in y_o_dict.items(): int_dict_increment(x_o_dict, out_sym, count) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index b4768d7dcc..3cadd0fd7a 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -248,15 +248,11 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation blue_states.extend(best_candidate.child_iterator()) instrumentation.log_promote(best_candidate) - # check cached partitions - for partitioning in partition_candidates.values(): - updated_promoted_node = partitioning.full_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: - partitioning.new_blue.append(successor) + # any other cached candidate that speculatively touched this node (e.g. while resolving an + # unknown output through it) is now unsound to reuse: applying it would silently overwrite + # the just-promoted (now real, independently-decided) state with a stale speculative copy. + for key in [key for key, p in partition_candidates.items() if best_candidate in p.full_mapping]: + del partition_candidates[key] elif isinstance(best_candidate, Partitioning): # apply best merge candidate for real_node, partition_node in best_candidate.red_mapping.items(): 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_gsm_algorithms.py b/tests/learning_algs/general_passive/test_gsm_algorithms.py index ddf56743c9..9ad75e8201 100644 --- a/tests/learning_algs/general_passive/test_gsm_algorithms.py +++ b/tests/learning_algs/general_passive/test_gsm_algorithms.py @@ -2,6 +2,8 @@ import unittest from itertools import product +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 @@ -81,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) From b744e37242b9987fbb60b2392977df40ef8664a4 Mon Sep 17 00:00:00 2001 From: Edi Muskardin <28546846+emuskardin@users.noreply.github.com> Date: Thu, 10 Sep 2026 20:26:36 +0200 Subject: [PATCH 60/70] Add tests for associated_data --- .../general_passive/test_associated_data.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) create mode 100644 tests/learning_algs/general_passive/test_associated_data.py diff --git a/tests/learning_algs/general_passive/test_associated_data.py b/tests/learning_algs/general_passive/test_associated_data.py new file mode 100644 index 0000000000..7fb296cae5 --- /dev/null +++ b/tests/learning_algs/general_passive/test_associated_data.py @@ -0,0 +1,19 @@ +import unittest + +from aalpy.learning_algs.general_passive.AssociatedData import CountData + + +class TestCountDataMerge(unittest.TestCase): + def test_merge_does_not_alias_the_other_operands_dict(self): + # merge() used to do x[in_sym] = y_o_dict for an input unseen in x, aliasing y's inner dict + # instead of copying it, so later mutating y_o_dict would silently change x's counts too. + x = {} + y_o_dict = {'out1': 1} + y = {'a': y_o_dict} + merged = CountData.merge(x, y) + y_o_dict['out1'] = 100 + self.assertEqual(merged['a'], {'out1': 1}) + + +if __name__ == '__main__': + unittest.main() From f432db770015506821d7484db5f2927517c1ae47 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 11 Sep 2026 12:12:36 +0200 Subject: [PATCH 61/70] drop dead code --- .../learning_algs/general_passive/AssociatedData.py | 12 ------------ .../general_passive/test_associated_data.py | 12 ------------ 2 files changed, 24 deletions(-) diff --git a/aalpy/learning_algs/general_passive/AssociatedData.py b/aalpy/learning_algs/general_passive/AssociatedData.py index efe164388c..0a1c4c1138 100644 --- a/aalpy/learning_algs/general_passive/AssociatedData.py +++ b/aalpy/learning_algs/general_passive/AssociatedData.py @@ -43,18 +43,6 @@ def local_log_likelihood_contribution(self): def count(self): return sum(sum(trans.values()) for trans in self.transition_count.values()) - @staticmethod - def merge(x: CountDict, y: CountDict) -> CountDict: - for in_sym, y_o_dict in y.items(): - x_o_dict = x.get(in_sym, None) - if x_o_dict is None: - # copy, don't alias: x must not keep changing if y's dict is mutated afterwards - x[in_sym] = y_o_dict.copy() - continue - for out_sym, count in y_o_dict.items(): - int_dict_increment(x_o_dict, out_sym, count) - return x - def get_probabilities(self) -> ProbabilityDict: ret = dict() for in_sym, trans in self.transition_count.items(): diff --git a/tests/learning_algs/general_passive/test_associated_data.py b/tests/learning_algs/general_passive/test_associated_data.py index 7fb296cae5..9ccd4b6ded 100644 --- a/tests/learning_algs/general_passive/test_associated_data.py +++ b/tests/learning_algs/general_passive/test_associated_data.py @@ -3,17 +3,5 @@ from aalpy.learning_algs.general_passive.AssociatedData import CountData -class TestCountDataMerge(unittest.TestCase): - def test_merge_does_not_alias_the_other_operands_dict(self): - # merge() used to do x[in_sym] = y_o_dict for an input unseen in x, aliasing y's inner dict - # instead of copying it, so later mutating y_o_dict would silently change x's counts too. - x = {} - y_o_dict = {'out1': 1} - y = {'a': y_o_dict} - merged = CountData.merge(x, y) - y_o_dict['out1'] = 100 - self.assertEqual(merged['a'], {'out1': 1}) - - if __name__ == '__main__': unittest.main() From 8ead44d7dfcf0ec65e2ae9424096a9ac6c6af679 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 11 Sep 2026 12:13:49 +0200 Subject: [PATCH 62/70] make data handler limitations explicit --- aalpy/learning_algs/general_passive/DataHandler.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index 0cc598523a..b4acd91646 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -268,6 +268,8 @@ 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) @@ -284,6 +286,8 @@ def copy(self, x: CountOnPTAData) -> CountOnPTAData: 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) From d8d49a03a8d3c7efc5c46193071a839111c9043c Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 11 Sep 2026 12:14:58 +0200 Subject: [PATCH 63/70] fixed silent overwrite on nondeterminism for non-prefix-closed data --- aalpy/learning_algs/general_passive/DataHandler.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index b4acd91646..1ef8389bfc 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -104,9 +104,12 @@ def add_labeled_sequence(self, root_node: GsmNode[T], example: IOExample): node = GsmNode((in_sym, out_sym), curr_node, self.init_data()) transitions[out_sym] = node elif len(transitions) == 1: - node = next(iter(transitions.values())) + existing_out_sym, node = next(iter(transitions.keys())) + 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: - raise ValueError("Nondeterminism encountered for GSM with labeled_sequences. not supported") + 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 From 145f1f6232465924f29031c7e73f291e830e64c8 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Fri, 11 Sep 2026 12:16:11 +0200 Subject: [PATCH 64/70] add additional flag for overriding default checks --- .../general_passive/GeneralizedStateMerging.py | 12 +++++++----- .../general_passive/ScoreFunctionsGSM.py | 6 ++++++ 2 files changed, 13 insertions(+), 5 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 3cadd0fd7a..0673192e8a 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -395,11 +395,13 @@ def get_partition_trans(part: GsmNode, in_symbol): partitioning.nr_merged_states += 1 if first_pass: - local_compat = self.score_calc.local_compatibility(partition, blue) - moore_check = self.output_behavior == "moore" and self.transition_behavior == "deterministic" and not GsmNode.moore_compatible(partition, blue) - # determinism is a property of the result, not of the scoring: enforce it even if score_calc doesn't - det_check = self.transition_behavior == "deterministic" and not GsmNode.deterministic_compatible(partition, blue) - if local_compat is False or moore_check or det_check: + 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: diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index b9d90004d4..7f64bd4f7f 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -98,6 +98,12 @@ def has_score_function(self) -> bool: """ 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: From d7e36980ccb72bc690ac7f0743728bb4b208b673 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Mon, 14 Sep 2026 15:50:27 +0200 Subject: [PATCH 65/70] added warnings for unknown outputs and input completeness in final automaton + fix justification for cache wipe and add todo --- aalpy/learning_algs/general_passive/DataHandler.py | 2 +- .../general_passive/GeneralizedStateMerging.py | 14 +++++++++++--- aalpy/learning_algs/general_passive/GsmNode.py | 10 ++++++++++ 3 files changed, 22 insertions(+), 4 deletions(-) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index 1ef8389bfc..b0bab9400b 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -104,7 +104,7 @@ def add_labeled_sequence(self, root_node: GsmNode[T], example: IOExample): 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.keys())) + 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: diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 0673192e8a..0dd722b58f 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -248,11 +248,19 @@ def run(self, data: Any, convert: bool = True, instrumentation: Instrumentation blue_states.extend(best_candidate.child_iterator()) instrumentation.log_promote(best_candidate) - # any other cached candidate that speculatively touched this node (e.g. while resolving an - # unknown output through it) is now unsound to reuse: applying it would silently overwrite - # the just-promoted (now real, independently-decided) state with a stale speculative copy. + # 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(): diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index 95e7ab69e4..f6dc18e77a 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -1,6 +1,7 @@ # Generic prefix-tree / observation-tree node structure used by the general passive # (state-merging) learning algorithms, plus conversion to concrete AALpy automaton types. import pathlib +import warnings from collections import defaultdict from collections.abc import Callable, Iterable, Iterator, Sequence from typing import Any, TypeVar, Generic @@ -274,6 +275,15 @@ def to_automaton(self, output_behavior: OutputBehavior, transition_behavior: Tra 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): From 51b688ad7072ad4f7460369c4810c0759a71bc19 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 15 Sep 2026 09:07:37 +0200 Subject: [PATCH 66/70] remove testcase for unsupported feature --- .../general_passive/test_gsm_data_handler.py | 11 ----------- 1 file changed, 11 deletions(-) diff --git a/tests/learning_algs/general_passive/test_gsm_data_handler.py b/tests/learning_algs/general_passive/test_gsm_data_handler.py index 84d1b75243..bcbcb829e3 100644 --- a/tests/learning_algs/general_passive/test_gsm_data_handler.py +++ b/tests/learning_algs/general_passive/test_gsm_data_handler.py @@ -43,17 +43,6 @@ def test_labeled_sequences_initialize_node_data(self): for node in pta.get_all_nodes(): self.assertIsNotNone(node.data) - @unittest.expectedFailure - def test_labeled_sequences_and_io_traces_count_the_same_transitions(self): - # known gap: add_labeled_sequence aggregates intermediate steps under unknown_output and only - # remaps `transitions` when the output is resolved later, so the counts stay split - labeled = [(('a', 'b'), True), (('a',), False)] - io_traces = [[False, ('a', False), ('b', True)], [False, ('a', False)]] - pta_labeled = counting_pta(CountDataHandler(), 'labeled_sequences', labeled) - pta_traces = counting_pta(CountDataHandler(), 'io_traces', io_traces) - self.assertEqual(dict(pta_labeled.data.transition_count), dict(pta_traces.data.transition_count)) - - def _node_with(data): """Minimal stand-in for the source node of aggregate_data, which only accesses .data.""" From 514d0af7d0f633e93c395e080056098cf89d023a Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Tue, 15 Sep 2026 09:13:19 +0200 Subject: [PATCH 67/70] fix broken test cases. removed empty test file --- .../learning_algs/general_passive/test_associated_data.py | 7 ------- .../learning_algs/general_passive/test_gsm_data_handler.py | 4 ++-- 2 files changed, 2 insertions(+), 9 deletions(-) delete mode 100644 tests/learning_algs/general_passive/test_associated_data.py diff --git a/tests/learning_algs/general_passive/test_associated_data.py b/tests/learning_algs/general_passive/test_associated_data.py deleted file mode 100644 index 9ccd4b6ded..0000000000 --- a/tests/learning_algs/general_passive/test_associated_data.py +++ /dev/null @@ -1,7 +0,0 @@ -import unittest - -from aalpy.learning_algs.general_passive.AssociatedData import CountData - - -if __name__ == '__main__': - unittest.main() diff --git a/tests/learning_algs/general_passive/test_gsm_data_handler.py b/tests/learning_algs/general_passive/test_gsm_data_handler.py index bcbcb829e3..0b44f6a117 100644 --- a/tests/learning_algs/general_passive/test_gsm_data_handler.py +++ b/tests/learning_algs/general_passive/test_gsm_data_handler.py @@ -29,7 +29,7 @@ def test_copy_is_independent_of_original(self): class TestCountOnPTADataHandler(unittest.TestCase): def test_copy_keeps_pta_data(self): dh = CountOnPTADataHandler() - pta = counting_pta(dh, 'labeled_sequences', [(('a',), True), (('b',), False)]) + 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) @@ -39,7 +39,7 @@ def test_copy_keeps_pta_data(self): class TestCreatePTA(unittest.TestCase): def test_labeled_sequences_initialize_node_data(self): dh = CountOnPTADataHandler() - pta = counting_pta(dh, 'labeled_sequences', [(('a', 'b'), True), (('a',), False)]) + pta = counting_pta(dh, 'io_traces', [[True, ('a',True), ('b', False)]]) for node in pta.get_all_nodes(): self.assertIsNotNone(node.data) From 17d182c8c5848f3bc8d92e1ffa2d5c76da3e711b Mon Sep 17 00:00:00 2001 From: Edi Muskardin <28546846+emuskardin@users.noreply.github.com> Date: Tue, 15 Sep 2026 15:30:28 +0200 Subject: [PATCH 68/70] Address documentation inconsistencies and bugs in GSM --- .../general_passive/DataHandler.py | 7 +++-- .../GeneralizedStateMerging.py | 4 +-- .../general_passive/GsmAlgorithms.py | 26 ++++--------------- .../general_passive/ScoreFunctionsGSM.py | 4 +-- aalpy/utils/HelperFunctions.py | 19 ++++++++++++++ .../test_generalized_state_merging.py | 16 +++++++++++- .../general_passive/test_gsm_data_handler.py | 4 +++ .../test_score_functions_gsm.py | 4 +-- 8 files changed, 54 insertions(+), 30 deletions(-) diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index b0bab9400b..04c5e2a966 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -92,8 +92,11 @@ def add_labeled_sequence(self, root_node: GsmNode[T], example: IOExample): in_sym = None if len(inputs) == 0: - self.aggregate_data(None, no_op_input, output, root_node) 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): @@ -184,7 +187,7 @@ def abstract(self, in_val: Any, out_val: Any) -> tuple[Any, Any]: :param Any in_val: The input value. :param Any out_val: The output value. - :return tuple[Any, Any]: The abstract output value and the input symbols. + :return tuple[Any, Any]: The abstracted input and output symbols. """ return in_val, out_val diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index 0dd722b58f..d33e3a55ac 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -313,7 +313,7 @@ def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNo 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 self.output_behavior == "moore" and not GsmNode.moore_compatible(red, blue): + 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 @@ -471,7 +471,7 @@ def run_GSM(data: list, *, """ 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. diff --git a/aalpy/learning_algs/general_passive/GsmAlgorithms.py b/aalpy/learning_algs/general_passive/GsmAlgorithms.py index b07c455f97..f48562f79e 100644 --- a/aalpy/learning_algs/general_passive/GsmAlgorithms.py +++ b/aalpy/learning_algs/general_passive/GsmAlgorithms.py @@ -10,7 +10,7 @@ from aalpy.learning_algs.general_passive.Instrumentation import ProgressReport 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 +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, @@ -31,7 +31,7 @@ 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) @@ -45,7 +45,7 @@ def EDSM_score(part: dict[GsmNode, GsmNode]) -> int: evidence += 1 return evidence - score = SimpleScoreCalculation(local_compatibility=GsmNode.deterministic_compatible, 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 @@ -56,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 @@ -98,15 +90,7 @@ def run_k_tails(data: list, automaton_type: str, k: int, input_completeness: str 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 diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 7f64bd4f7f..9777495b4a 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -499,10 +499,10 @@ def score(part: dict[GsmNode, GsmNode]) -> Any: 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[CountData], GsmNode[CountData]]) -> Any: total_evidence = 0 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/general_passive/test_generalized_state_merging.py b/tests/learning_algs/general_passive/test_generalized_state_merging.py index dbb46fe894..8565cb730d 100644 --- a/tests/learning_algs/general_passive/test_generalized_state_merging.py +++ b/tests/learning_algs/general_passive/test_generalized_state_merging.py @@ -6,7 +6,7 @@ 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 +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) diff --git a/tests/learning_algs/general_passive/test_gsm_data_handler.py b/tests/learning_algs/general_passive/test_gsm_data_handler.py index 0b44f6a117..803f940a8b 100644 --- a/tests/learning_algs/general_passive/test_gsm_data_handler.py +++ b/tests/learning_algs/general_passive/test_gsm_data_handler.py @@ -43,6 +43,10 @@ def test_labeled_sequences_initialize_node_data(self): 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.""" 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 105634b789..fdb2edc018 100644 --- a/tests/learning_algs/general_passive/test_score_functions_gsm.py +++ b/tests/learning_algs/general_passive/test_score_functions_gsm.py @@ -271,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) From a1ee08bffe8d6c37d920e2ef0bb6e29793273a5a Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 16 Sep 2026 09:32:48 +0200 Subject: [PATCH 69/70] consistent string value for input completeness with self loops --- aalpy/learning_algs/general_passive/GsmNode.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/aalpy/learning_algs/general_passive/GsmNode.py b/aalpy/learning_algs/general_passive/GsmNode.py index f6dc18e77a..5ded550722 100644 --- a/aalpy/learning_algs/general_passive/GsmNode.py +++ b/aalpy/learning_algs/general_passive/GsmNode.py @@ -421,17 +421,17 @@ 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, target: 'GsmNode[T] | str' = "self-loop") -> list[tuple['GsmNode', Any, Any]]: + def make_input_complete(self, target: 'GsmNode[T] | str' = "self_loop") -> list[tuple['GsmNode', Any, Any]]: """ 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. + :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.") + 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} @@ -440,7 +440,7 @@ def make_input_complete(self, target: 'GsmNode[T] | str' = "self-loop") -> list[ for in_sym in inputs: transitions = node.transitions[in_sym] if len(transitions) == 0: - if target == "self-loop": + if target == "self_loop": successor = node else: successor = target From 0c4783e4150dca065bb8b1e05b63302620f3de76 Mon Sep 17 00:00:00 2001 From: Benjamin von Berg Date: Wed, 16 Sep 2026 10:33:09 +0200 Subject: [PATCH 70/70] split early score calculation and score initialization + homogenized naming of methods --- Examples.py | 4 +- .../general_passive/DataHandler.py | 11 +---- .../GeneralizedStateMerging.py | 7 ++- .../general_passive/ScoreFunctionsGSM.py | 49 ++++++++++++------- .../test_score_functions_gsm.py | 2 +- 5 files changed, 40 insertions(+), 33 deletions(-) diff --git a/Examples.py b/Examples.py index 0a6b02eee6..e662e99298 100644 --- a/Examples.py +++ b/Examples.py @@ -1271,9 +1271,9 @@ def __init__(self, eps: float): SimpleFutureBasedCompatibility.__init__(self, compatibility_on_pta=True) self.score = None - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: self.score = 0 - verdict = super().initialize_merge(red, blue, first_pass) + verdict = super().early_score(red, blue) if verdict is SpecialScores.ImmediateReject: return verdict return self.score diff --git a/aalpy/learning_algs/general_passive/DataHandler.py b/aalpy/learning_algs/general_passive/DataHandler.py index 04c5e2a966..713f4c3939 100644 --- a/aalpy/learning_algs/general_passive/DataHandler.py +++ b/aalpy/learning_algs/general_passive/DataHandler.py @@ -169,8 +169,7 @@ def createPTA(self, data: Any, output_behavior: OutputBehavior, data_format: Dat self.add_trace(root_node, trace) return root_node - @abstractmethod - def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): + 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. @@ -179,7 +178,7 @@ def init_merge(self, red: 'GsmNode[T]', blue: 'GsmNode[T]', first_pass: bool): :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]: """ @@ -244,9 +243,6 @@ def init_data(self) -> None: def aggregate_data(self, src_node: 'GsmNode[None]', in_sym, out_value, dst_node: 'GsmNode[None]'): pass - def init_merge(self, red: 'GsmNode[None]', blue: 'GsmNode[None]', first_pass: bool): - return None - def merge(self, x: None, y: None) -> None: return None @@ -255,9 +251,6 @@ def copy(self, x: None) -> None: class CountDataHandler(DataHandler[CountData]): - def init_merge(self, red: 'GsmNode[CountData]', blue: 'GsmNode[CountData]', first_pass: bool): - pass - 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] diff --git a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py index d33e3a55ac..7a20873201 100644 --- a/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py +++ b/aalpy/learning_algs/general_passive/GeneralizedStateMerging.py @@ -319,7 +319,7 @@ def _partition_from_merge(self, partitioning: Partitioning, red_nodes: set[GsmNo # check whether there is an early verdict and adapt helper functions accordingly # TODO maybe split init from early verdict - partitioning.score = self.score_calc.initialize_merge(red, blue, first_pass) + partitioning.score = self.score_calc.early_score(red, blue) if partitioning.score is not SpecialScores.NoScore: return partitioning.remaining_merges = [] @@ -356,8 +356,6 @@ def get_partition_trans(part: GsmNode, in_symbol): cow_set.add(id(trans)) return trans elif partitioning.remaining_merges is None or len(partitioning.remaining_merges) != 0: - self.score_calc.initialize_merge(red, blue, first_pass) - # best scoring merge candidate -> can manipulate nodes directly red_partitions = red_nodes def update_partition(red_node: GsmNode, blue_node: GsmNode | None) -> GsmNode: @@ -369,7 +367,8 @@ def get_partition_trans(part: GsmNode, in_symbol): # first pass already did all the work return - self.data_handler.init_merge(red, blue, first_pass) + 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: diff --git a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py index 9777495b4a..c92d5381c7 100644 --- a/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py +++ b/aalpy/learning_algs/general_passive/ScoreFunctionsGSM.py @@ -34,9 +34,20 @@ def __bool__(self): class ScoreCalculation(ABC): """Bundles a local compatibility check and a global score function used during state merging.""" - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: """ - Callback at the beginning of the evaluation of a merge candidate. + 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 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`. + """ + return None + + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): + """ + 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. @@ -44,7 +55,7 @@ def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any or the second pass (in which the partitioning is completed) :return: Either an early score for the merge candidate or `None`. """ - return None + pass def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool | None: """ @@ -178,10 +189,7 @@ def __init__(self, self.compatibility_on_pta = compatibility_on_pta self.depth_first = depth_first - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: - if not first_pass: - return - + 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") @@ -218,9 +226,9 @@ def __init__(self, eps: float, compat_on_pta: bool, compat_on_pta_data: bool, ed self.edsm = edsm self.score = None - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + def early_score(self, red: GsmNode, blue: GsmNode) -> Any: self.score = 0 - verdict = super().initialize_merge(red, blue, first_pass) + verdict = super().early_score(red, blue) if self.edsm is False or verdict is SpecialScores.ImmediateReject: return verdict return self.score @@ -239,8 +247,11 @@ def __init__(self, wrapped: ScoreCalculation): # if not hasattr(self, "initialized_merge"): # self.initialized_merge = wrapped.initialize_merge - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: - return self.wrapped.initialize_merge(red, blue, first_pass) + 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) @@ -273,9 +284,9 @@ def __init__(self, wrapped: ScoreCalculation, k: int) -> None: self.depth_offset = None - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): self.depth_offset = blue.get_prefix_length() - return self.wrapped.initialize_merge(red, blue, first_pass) + self.wrapped.initialize_merge(red, blue, first_pass) def local_compatibility(self, a: GsmNode, b: GsmNode) -> bool | None: """ @@ -309,11 +320,11 @@ def __init__(self, wrapped: ScoreCalculation, sink_cond: Callable[[GsmNode], boo self.sink_cond = sink_cond self.allow_sink_merge = allow_sink_merge - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: + 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.initialize_merge(red, blue, first_pass) + return self.wrapped.early_score(red, blue) class ScoreCombinator(ScoreCalculation): @@ -335,8 +346,12 @@ def __init__(self, scores: list[ScoreCalculation], aggregate_compatibility: Aggr self.aggregate_compatibility = aggregate_compatibility or self.default_aggregate_compatibility self.aggregate_score = aggregate_score or self.default_aggregate_score - def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool) -> Any: - scores = [score.initialize_merge(red, blue, first_pass) for score in self.scores] + def initialize_merge(self, red: GsmNode, blue: GsmNode, first_pass: bool): + for score in self.scores: + 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: 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 fdb2edc018..fd43d546ab 100644 --- a/tests/learning_algs/general_passive/test_score_functions_gsm.py +++ b/tests/learning_algs/general_passive/test_score_functions_gsm.py @@ -62,7 +62,7 @@ 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.initialize_merge(node, node, True), SpecialScores.NoScore) + self.assertIs(comb.early_score(node, node), SpecialScores.NoScore) def test_single_rejecting_sub_score_rejects(self): aggregate = ScoreCombinator.default_aggregate_score