Skip to content

Commit 8b7d318

Browse files
committed
some refactoring
1 parent 2984795 commit 8b7d318

1 file changed

Lines changed: 21 additions & 11 deletions

File tree

grape/pruning/obs_equiv_pruner.py

Lines changed: 21 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,12 @@
11
from collections import defaultdict
22
import math
3-
from typing import Callable
4-
from grape.automaton.spec_manager import specialize
3+
from typing import Any, Callable
4+
from grape.automaton.spec_manager import (
5+
despecialize,
6+
is_specialized,
7+
specialize,
8+
type_request_from_specialized,
9+
)
510
from grape.dsl import DSL
611
from grape.enumerator import Enumerator
712
from grape.evaluator import Evaluator
@@ -61,7 +66,7 @@ def __get_base_grammar__(
6166
max_size: int,
6267
base_dfta: DFTA | None,
6368
type_req: str,
64-
):
69+
) -> tuple[DFTA[Any, Program], dict[int, int]]:
6570
base_grammar = grammar_by_saturation(dsl, type_req)
6671
if base_dfta is None:
6772
commutatives = commutativity_pruner.prune(dsl, evaluator, manager)
@@ -71,17 +76,21 @@ def __get_base_grammar__(
7176
[commutativity_constraint(dsl, commutatives, type_req)],
7277
)
7378
else:
74-
base_grammar = dsl.map_to_variants(base_dfta)
79+
base_grammar = base_dfta
80+
if is_specialized(base_grammar):
81+
tr = type_request_from_specialized(base_dfta, dsl)
82+
base_grammar = despecialize(base_dfta, tr)
83+
base_grammar = dsl.map_to_variants(base_grammar)
7584
base_grammar = specialize(base_grammar, type_req, dsl)
76-
base_grammar = base_grammar.map_alphabet(
77-
lambda x: Variable(int(x[len("var") :]))
85+
# alphabet is potentially str so convert it
86+
grammar = base_grammar.map_alphabet(
87+
lambda x: Variable(int(str(x)[len("var") :]))
7888
if str(x).startswith("var")
79-
else Primitive(x)
89+
else Primitive(str(x))
8090
)
81-
grammar = base_grammar
91+
8292
base_trees_by_size = base_grammar.trees_by_size(max_size)
83-
enum_ntrees = grammar.trees_until_size(max_size)
84-
return grammar, base_trees_by_size, enum_ntrees
93+
return grammar, base_trees_by_size
8594

8695

8796
def prune(
@@ -96,14 +105,15 @@ def prune(
96105
type_req = __infer_mega_type_req__(
97106
dsl.primitives, rtype, max_size, set(evaluator.base_inputs.keys())
98107
)
99-
grammar, base_expected_trees, enum_ntrees = __get_base_grammar__(
108+
grammar, base_expected_trees = __get_base_grammar__(
100109
dsl,
101110
evaluator,
102111
manager,
103112
max_size,
104113
base_grammar,
105114
type_req,
106115
)
116+
enum_ntrees = grammar.trees_until_size(max_size)
107117
base_ntrees = sum(base_expected_trees.values())
108118

109119
enumerator = Enumerator(grammar)

0 commit comments

Comments
 (0)