11from collections import defaultdict
22import 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+ )
510from grape .dsl import DSL
611from grape .enumerator import Enumerator
712from 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
8796def 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