1:- use_module(library(pairs)).    2:- use_module(library(edcg)).    3:- use_module(library(debug)).    4
    5edcg:acc_info(egraph, Node, In, Out, add_node(Node, In, Out), t(black('', A, B, ''), black('', A, B, '')), _).
    6edcg:acc_info(right, Node, In, Out, quick_add_node(Node, In, Out), t(black('', A, B, ''), black('', A, B, '')), _).
    7edcg:acc_info(nodes, T, In, Out, In = [T | Out]).
    8edcg:acc_info(unifs, T, In, Out, In = [T | Out]).
    9edcg:acc_info(worklist, T, In, Out, Out = [T | In], [], _).
   10edcg:pass_info(left).
   11edcg:pass_info(index).
   12edcg:pass_info(rules).
   13
   14edcg:pred_info(add_term, 2, [egraph]).
   15edcg:pred_info(add_terms, 2, [egraph]).
   16edcg:pred_info(add_ac_term, 2, [dcg, egraph]).
   17edcg:pred_info(saturate, 2, [egraph]).
   18edcg:pred_info(match_goals, 2, [left, index, right, unifs]).
   19edcg:pred_info(rebuild, 1, [egraph]).
   20edcg:pred_info(rebuild_egraph, 2, [unifs]).
   21edcg:pred_info(rebuild_egraph, 3, [unifs]).
   22edcg:pred_info(rebuild_nodes, 1, [nodes, unifs]).
   23edcg:pred_info(congruence_closure, 0, [nodes, unifs]).
   24edcg:pred_info(merge_groups, 2, [unifs]).
   25edcg:pred_info(merge_ids, 2, [unifs]).
   26edcg:pred_info(critical_pairs, 1, [worklist]).
   27edcg:pred_info(critical_pairs, 2, [worklist]).
   28edcg:pred_info(ac_canonize_rule, 3, [rules]).
   29edcg:pred_info(ac_closure, 1, [unifs, worklist]).
   30edcg:pred_info(ac_closure, 3, [unifs, worklist]).
   31edcg:pred_info(interreduce, 3, [worklist]).
   32
   33lookup(Node-Class, [N-C | L]) :-
   34    (   Node == N
   35    ->  Class = C
   36    ;   lookup(Node-Class, L)
   37    ).
   38
   39add_node(ac(F, Args)-Id, In, Out) =>
   40    msort(Args, Sort),
   41    (   rb_lookup(F/ac, NodesIn, In)
   42    ->  true
   43    ;   NodesIn = []
   44    ),
   45    (   lookup(ac(F, Sort)-Id, NodesIn)
   46    ->  Out = In
   47    ;   ord_add_element(NodesIn, ac(F, Sort)-Id, NodesOut),
   48        rb_insert(In, F/ac, NodesOut, Out)
   49    ).
   50add_node(Node-Id, In, Out) =>
   51    (   compound(Node)
   52    ->  compound_name_arity(Node, Name, Arity),
   53        Key = Name/Arity
   54    ;   Key = Node
   55    ),
   56    (   rb_lookup(Key, NodesIn, In)
   57    ->  true
   58    ;   NodesIn = []
   59    ),
   60    (   lookup(Node-Id, NodesIn)
   61    ->  Out = In
   62    ;   ord_add_element(NodesIn, Node-Id, NodesOut),
   63        rb_insert(In, Key, NodesOut, Out)
   64    ).
   65
   66quick_add_node(ac(F, Args)-Id, In, Out) =>
   67    (   rb_update(In, F/ac, Nodes, [ac(F, Args)-Id | Nodes], Out)
   68    ->  true
   69    ;   rb_insert(In, F/ac, [ac(F, Args)-Id], Out)
   70    ).
   71quick_add_node(Node-Id, In, Out) =>
   72    (   compound(Node)
   73    ->  compound_name_arity(Node, Name, Arity),
   74        Key = Name/Arity
   75    ;   Key = Node
   76    ),
   77    (   rb_update(In, Key, Nodes, [Node-Id | Nodes], Out)
   78    ->  true
   79    ;   rb_insert(In, Key, [Node-Id], Out)
   80    ).
   81
   82add_term(A*B, Id) ==>>
   83    add_ac_term(*, A*B):[dcg(Ids, []), egraph],
   84    [ac(*, Ids)-Id]:egraph.
   85add_term(A+B, Id) ==>>
   86    add_ac_term(+, A+B):[dcg(Ids, []), egraph],
   87    [ac(+, Ids)-Id]:egraph.
   88add_term(Term, Id), ? compound(Term) ==>>
   89    compound_name_arguments(Term, F, Args),
   90    foldl(add_term, Args, Ids):egraph,
   91    compound_name_arguments(Node, F, Ids),
   92    [Node-Id]:egraph.
   93add_term(Term, Id) ==>>
   94    [Term-Id]:egraph.
   95
   96add_ac_term(*, A*B) ==>>
   97    add_ac_term(*, A):[dcg, egraph],
   98    add_ac_term(*, B):[dcg, egraph].
   99add_ac_term(+, A+B) ==>>
  100    add_ac_term(+, A):[dcg, egraph],
  101    add_ac_term(+, B):[dcg, egraph].
  102add_ac_term(_, Term) ==>>
  103    add_term(Term, Id):egraph,
  104    [Id]:dcg.
  105
  106:- meta_predicate saturate(:, +, ?, ?).  107
  108saturate(Module:Goals, N) -->>
  109    insert(In, Right):egraph, variant_hash(In, L1),
  110    get_time(T1),
  111    make_index(In, Index),
  112    match_goals(Module, Goals):[left(In), index(Index), right(In, Right), unifs(Unifs, [])],
  113    rebuild(Unifs):egraph,
  114    Out/egraph, variant_hash(Out, L2),
  115    get_time(T2),
  116    T is T2 - T1,
  117    debug(saturate, "saturate ~p", [L1-L2-T-Unifs]),
  118    (   ( L1 == L2 ; N =< 0)
  119    ->  []
  120    ;   ( N =:= inf -> N1 = N ; N1 is N - 1),
  121        saturate(Module:Goals, N1):egraph
  122    ).
  123
  124make_index(EGraph, Index) :-
  125    rb_fold(append_nodes, EGraph, [], Nodes),
  126    sort(Nodes, Sort),
  127    group_pairs_by_key(Sort, Groups),
  128    ord_list_to_rbtree(Groups, Index).
  129
  130append_nodes(_Key-Nodes, In, Out) :-
  131    foldl(append_node, Nodes, In, Out).
  132
  133append_node(Node-Id, In, [Id-Node | In]).
  134
  135match_goals(_Module, []) -->> [].
  136match_goals(Module, [Goal | Goals]) -->>
  137    call(Module:Goal):[left, index, right, unifs],
  138    match_goals(Module, Goals):[left, index, right, unifs].
  139
  140rebuild([A=B | Unifs]) -->>
  141    A = B,
  142    rebuild(Unifs):egraph.
  143rebuild([]) -->>
  144    get_time(T1),
  145    rebuild_egraph:[egraph, unifs(Unifs, [])],
  146    get_time(T2),
  147    T is T2 - T1,
  148    debug(rebuild, "rebuild ~p", [T]),
  149    (   Unifs == []
  150    ->  []
  151    ;   rebuild(Unifs):egraph
  152    ).
  153
  154rebuild_egraph(t(Nil, Tree), NewTree2) ==>>
  155    NewTree2 = t(Nil, NewTree), 
  156    rebuild_egraph(Tree, NewTree, Nil):unifs.
  157rebuild_egraph(black('', _, _, ''), Nil0, Nil) ==>> Nil0 = Nil.
  158rebuild_egraph(red(L, K, V, R), NewTree, Nil) ==>>
  159    NewTree = red(NL, K, NV, NR), 
  160    rebuild_nodes(K):[nodes(V, NV), unifs],
  161    rebuild_egraph(L, NL, Nil):unifs, 
  162    rebuild_egraph(R, NR, Nil):unifs.
  163rebuild_egraph(black(L, K, V, R), NewTree, Nil) ==>>
  164    NewTree = black(NL, K, NV, NR), 
  165    rebuild_nodes(K):[nodes(V, NV), unifs],
  166    rebuild_egraph(L, NL, Nil):unifs, 
  167    rebuild_egraph(R, NR, Nil):unifs.
  168
  169rebuild_nodes(_F/ac) ==>>
  170    maplist(ac_canonize_node):nodes,
  171    sort:nodes,
  172    congruence_closure:[nodes, unifs],
  173    Nodes/nodes,
  174    convlist(ac_node_rule, Nodes, Rules),
  175    critical_pairs(Rules):worklist([], Worklist),
  176    ac_closure(Rules):[unifs, worklist(Worklist, [])].
  177rebuild_nodes(_) ==>>
  178    sort:nodes,
  179    congruence_closure:[nodes, unifs].
  180
  181congruence_closure -->>
  182    group_pairs_by_key:nodes,
  183    merge_groups:[nodes, unifs].
  184
  185merge_groups([], []) -->> [].
  186merge_groups([Node-[Id | Ids] | Groups], [Node-Id | Out]) -->>
  187    merge_ids(Ids, Id):unifs,
  188    merge_groups(Groups, Out):unifs.
  189
  190merge_ids([], _) -->> [].
  191merge_ids([B | Ids], A) -->>
  192    [A=B]:unifs,
  193    merge_ids(Ids, A):unifs.
  194
  195ac_canonize_node(ac(F, Ids)-C, ac(F, Sort)-C) :-
  196    msort(Ids, Sort).
  197
  198ac_node_rule(ac(_F, [Id | Ids])-C, [Id | Ids]-[C]).
  199
  200critical_pairs([]) -->> [].
  201critical_pairs([Rule | Rules]) -->>
  202    critical_pairs(Rules, Rule):worklist,
  203    critical_pairs(Rules):worklist.
  204
  205critical_pairs([], _) -->> [].
  206critical_pairs([L2-R2 | Rules], L1-R1) -->>
  207    (   ord_intersect(L1, L2)
  208    ->  ord_union(L1, L2, Peak),
  209        apply_rule(L1-R1, Peak, C1),
  210        apply_rule(L2-R2, Peak, C2),
  211        [C1-C2]:worklist
  212    ;   []
  213    ),
  214    critical_pairs(Rules, L1-R1):worklist.
  215    
  216apply_rule(Left-Right, In, Out) :-
  217    ord_subtract(In, Left, Sub),
  218    % multiset union
  219    append(Sub, Right, Add),
  220    msort(Add, Out).
  221
  222ac_closure(_) ==>> []/worklist.
  223ac_closure(Rules) ==>>
  224    insert([C1-C2 | L], L):worklist,
  225    ac_canonize_rule(Rules, C1, T1):rules(Rules),
  226    ac_canonize_rule(Rules, C2, T2):rules(Rules),
  227    ac_closure(Rules, T1, T2):[unifs, worklist].
  228ac_closure(Rules, T, T) ==>>
  229    ac_closure(Rules):[unifs, worklist].
  230ac_closure(Rules, T1, T2) ==>>
  231    reverse_grevlex(T1, T2, Rule),
  232    interreduce(Rule, Rules, NewRules):worklist,
  233    critical_pairs(NewRules, Rule):worklist,
  234    (   T1 = [A], T2 = [B]
  235    ->  [A=B]:unifs
  236    ;   []
  237    ),
  238    ac_closure([Rule | NewRules]):[unifs, worklist].
  239
  240ac_canonize_rule([], T1, T2) ==>> T2 = T1.
  241ac_canonize_rule([Left-Right | _], T1, T), ? Left = [_|_], ? ord_subset(Left, T1) ==>>
  242    apply_rule(Left-Right, T1, T2),
  243    AllRules/rules,
  244    ac_canonize_rule(AllRules, T2, T):rules.
  245ac_canonize_rule([_ | Rules], T1, T) ==>>
  246    ac_canonize_rule(Rules, T1, T):rules.
  247
  248reverse_grevlex(T1, T2, Rule) :-
  249    length(T1, L1),
  250    length(T2, L2),
  251    compare(O, L1, L2),
  252    (   O == (>)
  253    ->  Rule = T1-T2
  254    ;   O == (<)
  255    ->  Rule = T2-T1
  256    ;   reverse(T1, R1),
  257        reverse(T2, R2),
  258        lex_order(R1, R2, Order),
  259        (   Order == (>)
  260        ->  Rule = T1-T2
  261        ;   Rule = T2-T1
  262        )
  263    ).
  264lex_order([A | T1], [B | T2], Order) :-
  265    compare(O, A, B),
  266    (   O == (=)
  267    ->  lex_order(T1, T2, Order)
  268    ;   O = Order
  269    ).
  270
  271interreduce(_, [], RemainingRules) ==>> RemainingRules = [].
  272interreduce(L1-R1, [L2-R2 | Rules], RemainingRules), ? ord_subset(L1, L2) ==>>
  273    [L2-R2]:worklist,
  274    interreduce(L1-R1, Rules, RemainingRules):worklist.
  275interreduce(L1-R1, [L2-R2 | Rules], RemainingRules) ==>>
  276    RemainingRules = [L2-R2 | R],
  277    interreduce(L1-R1, Rules, R):worklist.
  278
  279edcg:pred_info(cd_equals_be, 2, [egraph]).
  280edcg:pred_info(ab_d, 0, [left, index, right, unifs]).
  281edcg:pred_info(ab_d, 3, [left, index, right, unifs]).
  282edcg:pred_info(ab_d_, 5, [left, index, right, unifs]).
  283edcg:pred_info(ab_d__, 5, [left, index, right, unifs]).
  284
  285edcg:pred_info(ac_e, 0, [left, index, right, unifs]).
  286edcg:pred_info(ac_e, 3, [left, index, right, unifs]).
  287edcg:pred_info(ac_e_, 5, [left, index, right, unifs]).
  288edcg:pred_info(ac_e__, 5, [left, index, right, unifs]).
  289
  290ab_d -->>
  291    (   rb_lookup('+'/ac, Nodes):left, rb_lookup(a, [a-A]):left, rb_lookup(b, [b-B]):left
  292    ->  (   A @=< B
  293        ->  ab_d(Nodes, A, B):[left, index, right, unifs]
  294        ;   ab_d(Nodes, B, A):[left, index, right, unifs]
  295        )
  296    ;   []
  297    ).
  298ab_d([], _, _) -->> [].
  299ab_d([ac(+, Node)-Id | Nodes], A, B) -->>
  300    ab_d_(Node, Node, A, B, Id):[left, index, right, unifs],
  301    ab_d(Nodes, A, B):[left, index, right, unifs].
  302ab_d_([], _, _, _, _) ==>> [].
  303ab_d_([A | R], Ids, A, B, Class) ==>>
  304    ab_d__([A | R], Ids, A, B, Class):[left, index, right, unifs],
  305    ab_d_(R, Ids, A, B, Class):[left, index, right, unifs].
  306ab_d_([_ | R], Ids, A, B, Class) ==>>
  307    ab_d_(R, Ids, A, B, Class):[left, index, right, unifs].
  308ab_d__([], _, _, _, _) ==>> [].
  309ab_d__([B | R], Ids, A, B, Class) ==>>
  310    ord_subtract(Ids, [A, B], Rest),
  311    (   Rest == []
  312    ->  [d-D]:right,
  313        [Class=D]:unifs
  314    ;   (   Rest = [RestId]
  315        ->  [ac(+, [A, B])-AB, ac(+, [AB, RestId])-ABRest, d-D]:right
  316        ;   [ac(+, [A, B])-AB, ac(+, Rest)-RestId, ac(+, [AB, RestId])-ABRest, d-D]:right
  317        ),
  318        [ABRest=Class, AB=D]:unifs
  319    ),
  320    ab_d__(R, Ids, A, B, Class):[left, index, right, unifs].
  321ab_d__([_ | R], Ids, A, B, Class) ==>>
  322    ab_d__(R, Ids, A, B, Class):[left, index, right, unifs].
  323
  324ac_e -->>
  325    (   rb_lookup('+'/ac, Nodes):left, rb_lookup(a, [a-A]):left, rb_lookup(c, [c-C]):left
  326    ->  (   A @=< C
  327        ->  ac_e(Nodes, A, C):[left, index, right, unifs]
  328        ;   ac_e(Nodes, C, A):[left, index, right, unifs]
  329        )
  330    ;   []
  331    ).
  332ac_e([], _, _) -->> [].
  333ac_e([ac(+, Node)-Id | Nodes], A, C) -->>
  334    ac_e_(Node, Node, A, C, Id):[left, index, right, unifs],
  335    ac_e(Nodes, A, C):[left, index, right, unifs].
  336ac_e_([], _, _, _, _) ==>> [].
  337ac_e_([A | R], Ids, A, C, Class) ==>>
  338    ac_e__([A | R], Ids, A, C, Class):[left, index, right, unifs],
  339    ac_e_(R, Ids, A, C, Class):[left, index, right, unifs].
  340ac_e_([_ | R], Ids, A, C, Class) ==>>
  341    ac_e_(R, Ids, A, C, Class):[left, index, right, unifs].
  342ac_e__([], _, _, _, _) ==>> [].
  343ac_e__([C | R], Ids, A, C, Class) ==>>
  344    ord_subtract(Ids, [A, C], Rest),
  345    (   Rest == []
  346    ->  [e-E]:right,
  347        [Class=E]:unifs
  348    ;   (   Rest = [RestId]
  349        ->  [ac(+, [A, C])-AC, ac(+, [AC, RestId])-ACRest, e-E]:right
  350        ;   [ac(+, [A, C])-AC, ac(+, Rest)-RestId, ac(+, [AC, RestId])-ACRest, e-E]:right
  351        ),
  352        [ACRest=Class, AC=E]:unifs
  353    ),
  354    ac_e__(R, Ids, A, C, Class):[left, index, right, unifs].
  355ac_e__([_ | R], Ids, A, C, Class) ==>>
  356    ac_e__(R, Ids, A, C, Class):[left, index, right, unifs].
  357
  358cd_equals_be(CD, BE) -->>
  359    add_term(a+b, _AB):egraph,
  360    add_term(a+c, _AC):egraph,
  361    add_term(c+d, CD):egraph,
  362    add_term(b+e, BE):egraph,
  363    saturate([ab_d, ac_e], inf):egraph.
  364
  365edcg:pred_info(dist, 0, [left, index, right, unifs]).   % dist/6
  366edcg:pred_info(dist_, 1, [index, right, unifs]).        % dist_/6
  367edcg:pred_info(dist__, 3, [index, right, unifs]).       % dist__/8
  368edcg:pred_info(dist___, 3, [right, unifs]).             % dist___/7
  369edcg:pred_info(dist____, 3, [right]).                   % dist____/5
  370
  371dist -->>
  372    (   rb_lookup('*'/ac, Nodes):left
  373    ->  dist_(Nodes):[index, right, unifs]
  374    ;   []
  375    ).
  376dist_([]) ==>> [].
  377dist_([ac(*, Node)-Mul | Nodes]) ==>>
  378    dist__(Node, [], Mul):[index, right, unifs],
  379    dist_(Nodes):[index, right, unifs].
  380dist__([], _, _) ==>> [].
  381dist__([Sum | Ms], Left, Mul) ==>>
  382    (   rb_lookup(Sum, SumNodes):index
  383    ->  append(Left, Ms, Context),
  384        dist___(SumNodes, Context, Mul):[right, unifs]
  385    ;   []
  386    ),
  387    dist__(Ms, [Sum | Left], Mul):[index, right, unifs].
  388dist___([], _, _) ==>> [].
  389dist___([ac(+, SumArgs) | SumNodes], Context, Mul) ==>>
  390    (   Context == []
  391    ->  domain_error(2, Context)
  392    ;   []
  393    ),
  394    dist____(Context, SumArgs, MulArgs):right,
  395    [ac(+, MulArgs)-Add]:right,
  396    [Mul=Add]:unifs,
  397    dist___(SumNodes, Context, Mul):[right, unifs].
  398dist___([_ | SumNodes], Context, Mul) ==>>
  399    dist___(SumNodes, Context, Mul):[right, unifs].
  400dist____(_, [], []) -->> [].
  401dist____(C, [Sum | SumArgs], [Mul | MulArgs])