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
+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)