fix(grammar): make structural canonicalisation order independent

This commit is contained in:
milner committed 2026-09-23 09:05:00 +00:00
1 parent 2bdb8aede6
commit 38ab0ea6eb
2 files changed
+116 -35

No files matched your search

+95 -35
View File
@@ -84,43 +84,103 @@ let probabilities g =
tbl
let canonical_string g =
let seen = Hashtbl.create 64 in
let order = ref [] in
let q = Queue.create () in
Queue.add g.start q;
Hashtbl.replace seen g.start ();
while not (Queue.is_empty q) do
let a = Queue.pop q in
order := a :: !order;
let nts = nonterminals g in
let signature colors nt =
let symbol = function
| Term value -> Printf.sprintf "t%d:%s" (String.length value) value
| Nonterm value -> Printf.sprintf "n%d" (Hashtbl.find colors value)
in
let productions =
productions_of g nt
|> List.map (fun p ->
"[" ^ String.concat ";" (List.map symbol p.rhs) ^ "]")
|> List.sort String.compare
in
Printf.sprintf "%d|%b|%s" (Hashtbl.find colors nt)
(String.equal nt g.start) (String.concat "|" productions)
in
let refine colors =
let rec loop colors =
let signatures = List.map (fun nt -> (nt, signature colors nt)) nts in
let unique =
signatures |> List.map snd |> List.sort_uniq String.compare
in
let ids = Hashtbl.create (List.length unique) in
List.iteri (fun index key -> Hashtbl.replace ids key index) unique;
let next = Hashtbl.create (List.length nts) in
List.iter
(fun (nt, key) -> Hashtbl.replace next nt (Hashtbl.find ids key))
signatures;
if
List.for_all
(fun nt -> Hashtbl.find colors nt = Hashtbl.find next nt)
nts
then next
else loop next
in
loop colors
in
let encode colors =
let order =
List.sort
(fun left right ->
compare (Hashtbl.find colors left) (Hashtbl.find colors right))
nts
in
let names = Hashtbl.create (List.length order) in
List.iteri
(fun index nt -> Hashtbl.replace names nt (Printf.sprintf "N%d" index))
order;
let symbol = function
| Term value -> Printf.sprintf "t%d:%s" (String.length value) value
| Nonterm value -> Hashtbl.find names value
in
g.productions
|> List.map (fun p ->
Printf.sprintf "%s -> %s" (Hashtbl.find names p.lhs)
(String.concat " " (List.map symbol p.rhs)))
|> List.sort String.compare |> String.concat "\n"
in
let rec canonical colors =
let colors = refine colors in
let classes = Hashtbl.create (List.length nts) in
List.iter
(fun p ->
if String.equal p.lhs a then
List.iter
(function
| Nonterm x ->
if not (Hashtbl.mem seen x) then begin
Hashtbl.replace seen x ();
Queue.add x q
end
| Term _ -> ())
p.rhs)
g.productions
done;
let order = List.rev !order in
let idx = Hashtbl.create 64 in
List.iteri (fun i a -> Hashtbl.replace idx a i) order;
let name a =
match Hashtbl.find_opt idx a with
| Some i -> Printf.sprintf "N%d" i
| None -> a
(fun nt ->
let color = Hashtbl.find colors nt in
let members =
Option.value ~default:[] (Hashtbl.find_opt classes color)
in
Hashtbl.replace classes color (nt :: members))
nts;
let ambiguous =
Hashtbl.fold
(fun color members acc ->
if List.length members > 1 then (color, members) :: acc else acc)
classes []
|> List.sort (fun (left, _) (right, _) -> compare left right)
in
match ambiguous with
| [] -> encode colors
| (_, members) :: _ ->
let fresh =
List.fold_left
(fun highest nt -> max highest (Hashtbl.find colors nt))
(-1) nts
+ 1
in
members
|> List.map (fun chosen ->
let branch = Hashtbl.copy colors in
Hashtbl.replace branch chosen fresh;
canonical branch)
|> List.sort String.compare |> List.hd
in
let sym = function Term s -> "t:" ^ s | Nonterm s -> name s in
let pstr p =
Printf.sprintf "%s -> %s" (name p.lhs)
(String.concat " " (List.map sym p.rhs))
in
g.productions |> List.map pstr |> List.sort String.compare
|> String.concat "\n"
let colors = Hashtbl.create (List.length nts) in
List.iter
(fun nt ->
Hashtbl.replace colors nt (if String.equal nt g.start then 0 else 1))
nts;
canonical colors
let equal_structure a b = String.equal (canonical_string a) (canonical_string b)
+21
View File
@@ -71,6 +71,27 @@ let () =
check "dedupe sums counts"
(approx (List.hd gn.Grammar.productions).Grammar.count 3.0)
let () =
let left =
g "S"
[ prod "S" [ nt "A" ]; prod "S" [ nt "B" ]; prod "A" [ tm "a" ];
prod "B" [ tm "b" ] ]
in
let reordered =
g "S"
[ prod "S" [ nt "B" ]; prod "S" [ nt "A" ]; prod "B" [ tm "b" ];
prod "A" [ tm "a" ] ]
in
let renamed =
g "Root"
[ prod "Root" [ nt "Right" ]; prod "Left" [ tm "a" ];
prod "Root" [ nt "Left" ]; prod "Right" [ tm "b" ] ]
in
check "canonical form ignores production order"
(Grammar.equal_structure left reordered);
check "canonical form ignores nonterminal names"
(Grammar.equal_structure left renamed)
let () =
let gs = g "S" [ prod "S" [ nt "S" ]; prod "S" [ tm "a" ] ] in
let gn = Grammar.normalize gs in