fix(grammar): make structural canonicalisation order independent

This commit is contained in:
milner committed 2026-09-23 09:05:00 +00:00
1 parent 495aba78d4
commit 6023c46560
2 files changed
+114 -33

No files matched your search

+93 -33
View File
@@ -84,43 +84,103 @@ let probabilities g =
tbl tbl
let canonical_string g = let canonical_string g =
let seen = Hashtbl.create 64 in let nts = nonterminals g in
let order = ref [] in let signature colors nt =
let q = Queue.create () in let symbol = function
Queue.add g.start q; | Term value -> Printf.sprintf "t%d:%s" (String.length value) value
Hashtbl.replace seen g.start (); | Nonterm value -> Printf.sprintf "n%d" (Hashtbl.find colors value)
while not (Queue.is_empty q) do in
let a = Queue.pop q in let productions =
order := a :: !order; 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 List.iter
(fun p -> (fun (nt, key) -> Hashtbl.replace next nt (Hashtbl.find ids key))
if String.equal p.lhs a then signatures;
List.iter if
(function List.for_all
| Nonterm x -> (fun nt -> Hashtbl.find colors nt = Hashtbl.find next nt)
if not (Hashtbl.mem seen x) then begin nts
Hashtbl.replace seen x (); then next
Queue.add x q else loop next
end in
| Term _ -> ()) loop colors
p.rhs) 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 g.productions
done; |> List.map (fun p ->
let order = List.rev !order in Printf.sprintf "%s -> %s" (Hashtbl.find names p.lhs)
let idx = Hashtbl.create 64 in (String.concat " " (List.map symbol p.rhs)))
List.iteri (fun i a -> Hashtbl.replace idx a i) order; |> List.sort String.compare |> String.concat "\n"
let name a =
match Hashtbl.find_opt idx a with
| Some i -> Printf.sprintf "N%d" i
| None -> a
in in
let sym = function Term s -> "t:" ^ s | Nonterm s -> name s in let rec canonical colors =
let pstr p = let colors = refine colors in
Printf.sprintf "%s -> %s" (name p.lhs) let classes = Hashtbl.create (List.length nts) in
(String.concat " " (List.map sym p.rhs)) List.iter
(fun nt ->
let color = Hashtbl.find colors nt in
let members =
Option.value ~default:[] (Hashtbl.find_opt classes color)
in in
g.productions |> List.map pstr |> List.sort String.compare Hashtbl.replace classes color (nt :: members))
|> String.concat "\n" 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 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) 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" check "dedupe sums counts"
(approx (List.hd gn.Grammar.productions).Grammar.count 3.0) (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 () =
let gs = g "S" [ prod "S" [ nt "S" ]; prod "S" [ tm "a" ] ] in let gs = g "S" [ prod "S" [ nt "S" ]; prod "S" [ tm "a" ] ] in
let gn = Grammar.normalize gs in let gn = Grammar.normalize gs in