From f8edd32f6b1f536c70911bf89999a065a206e294 Mon Sep 17 00:00:00 2001 From: sneeker Date: Wed, 23 Sep 2026 07:15:00 +0000 Subject: [PATCH] fix(parse): compute unary outside closure deterministically --- lib/parse.ml | 44 +++++++++++++++++++++++++++++++++++++++++++- test/test_scfg.ml | 15 +++++++++++++++ 2 files changed, 58 insertions(+), 1 deletion(-) diff --git a/lib/parse.ml b/lib/parse.ml index 3cfdb82..5dff148 100644 --- a/lib/parse.ml +++ b/lib/parse.ml @@ -12,6 +12,48 @@ type prep = { start : string; } +let unary_order prods by_lhs = + let children = Hashtbl.create 64 in + let indegree = Hashtbl.create 64 in + Hashtbl.iter (fun lhs _ -> Hashtbl.replace indegree lhs 0) by_lhs; + Array.iter + (fun (p : Grammar.production) -> + match p.rhs with + | [ Grammar.Nonterm child ] when Hashtbl.mem by_lhs child -> + let outgoing = + Option.value ~default:[] (Hashtbl.find_opt children p.lhs) + in + if not (List.mem child outgoing) then begin + Hashtbl.replace children p.lhs (child :: outgoing); + Hashtbl.replace indegree child (Hashtbl.find indegree child + 1) + end + | _ -> ()) + prods; + let ready = + Hashtbl.fold + (fun lhs degree acc -> if degree = 0 then lhs :: acc else acc) + indegree [] + |> List.sort String.compare + in + let rec visit order ready = + match ready with + | [] -> List.rev order + | lhs :: rest -> + let ready = ref rest in + List.iter + (fun child -> + let degree = Hashtbl.find indegree child - 1 in + Hashtbl.replace indegree child degree; + if degree = 0 then + ready := List.sort_uniq String.compare (child :: !ready)) + (Option.value ~default:[] (Hashtbl.find_opt children lhs)); + visit (lhs :: order) !ready + in + let order = visit [] ready in + if List.length order <> Hashtbl.length by_lhs then + invalid_arg "unit-production cycles must be normalised before parsing"; + order + let prepare (g : Grammar.t) = let prods : Grammar.production array = Array.of_list g.productions in let n = Array.length prods in @@ -41,7 +83,7 @@ let prepare (g : Grammar.t) = prods; probs; by_lhs; - lhs_list = Hashtbl.fold (fun k _ acc -> k :: acc) by_lhs []; + lhs_list = unary_order prods by_lhs; start = g.start; } diff --git a/test/test_scfg.ml b/test/test_scfg.ml index a3cc896..8d09bff 100644 --- a/test/test_scfg.ml +++ b/test/test_scfg.ml @@ -95,6 +95,21 @@ let () = check "unit productions are preserved" (has_prod gn "S" [ nt "A" ]); check "unit productions chain preserved" (has_prod gn "A" [ nt "B" ]) +let () = + List.iter + (fun (a, b, c) -> + let gu = + g "S" + [ prod "S" [ nt a ]; prod a [ nt b ]; prod b [ nt c ]; + prod c [ tm "a" ] ] + in + let counts, probability = Parse.expected_counts (Parse.prepare gu) [ "a" ] in + check "unit chain probability" (approx probability 1.0); + Array.iter + (fun count -> check "unit chain expected count" (approx count 1.0)) + counts) + [ ("A", "B", "C"); ("X", "Y", "Z"); ("Q", "R", "T") ] + let () = let cfg = { Scoring.default_config with alpha = 1.0; em_iters = 5 } in let g1 = g "S" [ prod "S" [ tm "a" ] ] in