fix(parse): compute unary outside closure deterministically
This commit is contained in:
2 files changed
+58
-1
No files matched your search
+43
-1
@@ -12,6 +12,48 @@ type prep = {
|
|||||||
start : string;
|
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 prepare (g : Grammar.t) =
|
||||||
let prods : Grammar.production array = Array.of_list g.productions in
|
let prods : Grammar.production array = Array.of_list g.productions in
|
||||||
let n = Array.length prods in
|
let n = Array.length prods in
|
||||||
@@ -41,7 +83,7 @@ let prepare (g : Grammar.t) =
|
|||||||
prods;
|
prods;
|
||||||
probs;
|
probs;
|
||||||
by_lhs;
|
by_lhs;
|
||||||
lhs_list = Hashtbl.fold (fun k _ acc -> k :: acc) by_lhs [];
|
lhs_list = unary_order prods by_lhs;
|
||||||
start = g.start;
|
start = g.start;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -95,6 +95,21 @@ let () =
|
|||||||
check "unit productions are preserved" (has_prod gn "S" [ nt "A" ]);
|
check "unit productions are preserved" (has_prod gn "S" [ nt "A" ]);
|
||||||
check "unit productions chain preserved" (has_prod gn "A" [ nt "B" ])
|
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 () =
|
||||||
let cfg = { Scoring.default_config with alpha = 1.0; em_iters = 5 } in
|
let cfg = { Scoring.default_config with alpha = 1.0; em_iters = 5 } in
|
||||||
let g1 = g "S" [ prod "S" [ tm "a" ] ] in
|
let g1 = g "S" [ prod "S" [ tm "a" ] ] in
|
||||||
|
|||||||
Reference in new issue
Block a user