From 38ab0ea6eb6db4da9f4b12bf6d2e9d186b6d933d Mon Sep 17 00:00:00 2001 From: milner Date: Wed, 23 Sep 2026 09:05:00 +0000 Subject: [PATCH] fix(grammar): make structural canonicalisation order independent --- lib/grammar.ml | 130 +++++++++++++++++++++++++++++++++------------- test/test_scfg.ml | 21 ++++++++ 2 files changed, 116 insertions(+), 35 deletions(-) diff --git a/lib/grammar.ml b/lib/grammar.ml index 51b9433..cfe2cd2 100644 --- a/lib/grammar.ml +++ b/lib/grammar.ml @@ -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) diff --git a/test/test_scfg.ml b/test/test_scfg.ml index 8d09bff..1ad3dbe 100644 --- a/test/test_scfg.ml +++ b/test/test_scfg.ml @@ -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