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