fix(parse): compute unary outside closure deterministically
This commit is contained in:
1 parent
3170b7b398
commit
495aba78d4
2 files changed
+58
-1
No files matched your search
+43
-1
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in new issue
Block a user