354 lines
12 KiB
OCaml
354 lines
12 KiB
OCaml
open Scfg
|
|
|
|
let failures = ref 0
|
|
|
|
let check name cond =
|
|
if cond then Printf.printf "ok %s\n" name
|
|
else begin
|
|
incr failures;
|
|
Printf.printf "FAIL %s\n" name
|
|
end;
|
|
flush stdout
|
|
|
|
let approx ?(tol = 1e-9) a b = abs_float (a -. b) < tol
|
|
|
|
let nt s = Grammar.Nonterm s
|
|
let tm s = Grammar.Term s
|
|
let prod lhs rhs = { Grammar.lhs; rhs; count = 1.0 }
|
|
let g start ps = Grammar.make ~start ps
|
|
|
|
let has_prod g lhs rhs =
|
|
List.exists
|
|
(fun (p : Grammar.production) ->
|
|
String.equal p.lhs lhs && Grammar.equal_rhs p.rhs rhs)
|
|
g.Grammar.productions
|
|
|
|
let sample tokens count = { Corpus.tokens; count }
|
|
|
|
let () =
|
|
check "log_gamma(1)=0" (approx (Scoring.log_gamma 1.0) 0.0);
|
|
check "log_gamma(2)=0" (approx (Scoring.log_gamma 2.0) 0.0);
|
|
check "log_gamma(1/2)=log sqrt pi"
|
|
(approx (Scoring.log_gamma 0.5) (0.5 *. log Float.pi));
|
|
check "log_gamma(3)=log 2"
|
|
(approx (Scoring.log_gamma 3.0) (log 2.0))
|
|
|
|
let () =
|
|
let amb = g "S" [ prod "S" [ nt "S"; nt "S" ]; prod "S" [ tm "a" ] ] in
|
|
let prep = Parse.prepare amb in
|
|
check "P(a) = 1/2" (approx (Parse.inside prep [ "a" ]) 0.5);
|
|
check "P(aa) = 1/8" (approx (Parse.inside prep [ "a"; "a" ]) 0.125);
|
|
check "P(aaa) = 1/16" (approx (Parse.inside prep [ "a"; "a"; "a" ]) 0.0625);
|
|
let vp, vt = Parse.viterbi prep [ "a"; "a" ] in
|
|
check "viterbi P(aa) = 1/8" (approx vp 0.125);
|
|
check "viterbi tree exists" (vt <> None)
|
|
|
|
let () =
|
|
let gr = g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ] in
|
|
let prep = Parse.prepare gr in
|
|
let cnt, p = Parse.expected_counts prep [ "a"; "b" ] in
|
|
check "inside(ab)=1" (approx p 1.0);
|
|
check "E[count S->AB]=1" (approx cnt.(0) 1.0);
|
|
check "E[count A->a]=1" (approx cnt.(1) 1.0);
|
|
check "E[count B->b]=1" (approx cnt.(2) 1.0)
|
|
|
|
let () =
|
|
let amb = g "S" [ prod "S" [ nt "S"; nt "S" ]; prod "S" [ tm "a" ] ] in
|
|
let prep = Parse.prepare amb in
|
|
let cnt, _ = Parse.expected_counts prep [ "a"; "a" ] in
|
|
check "E[count S->SS]=1" (approx cnt.(0) 1.0);
|
|
check "E[count S->a]=2" (approx cnt.(1) 2.0)
|
|
|
|
let () =
|
|
let gd =
|
|
g "S"
|
|
[ prod "S" [ tm "a" ];
|
|
{ Grammar.lhs = "S"; rhs = [ tm "a" ]; count = 2.0 } ]
|
|
in
|
|
let gn = Grammar.normalize gd in
|
|
check "dedupe keeps one production"
|
|
(List.length gn.Grammar.productions = 1);
|
|
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
|
|
check "self-loop removed"
|
|
(not (has_prod gn "S" [ nt "S" ]));
|
|
check "self-loop grammar keeps S->a" (has_prod gn "S" [ tm "a" ])
|
|
|
|
let () =
|
|
let gc =
|
|
Grammar.make ~start:"A"
|
|
[ prod "A" [ nt "B" ]; prod "B" [ nt "A" ]; prod "A" [ tm "a" ] ]
|
|
in
|
|
let gn = Grammar.normalize gc in
|
|
check "unit cycle removed" (not (has_prod gn "A" [ nt "B" ]));
|
|
check "unit cycle grammar keeps A->a" (has_prod gn "A" [ tm "a" ])
|
|
|
|
let () =
|
|
let gu =
|
|
g "S" [ prod "S" [ nt "A" ]; prod "A" [ nt "B" ]; prod "B" [ tm "a" ] ]
|
|
in
|
|
let gn = Grammar.normalize gu in
|
|
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
|
|
let m1, _, _, _ =
|
|
Scoring.marginal_loglik ~config:cfg g1 [ sample [ "a" ] 1.0 ]
|
|
in
|
|
check "marginal(S->a, {a}, alpha=1) = 0" (approx m1 0.0);
|
|
let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in
|
|
let m2, _, _, _ =
|
|
Scoring.marginal_loglik ~config:cfg g2 [ sample [ "a" ] 1.0 ]
|
|
in
|
|
check "marginal(S->a|b, {a}, alpha=1) = log 1/2"
|
|
(approx m2 (log 0.5))
|
|
|
|
let paper_initial () =
|
|
g "S"
|
|
[
|
|
prod "S" [ nt "A"; nt "B" ];
|
|
prod "S" [ nt "A"; nt "A"; nt "B"; nt "B" ];
|
|
prod "S" [ nt "A"; nt "A"; nt "A"; nt "B"; nt "B"; nt "B" ];
|
|
prod "A" [ tm "a" ];
|
|
prod "B" [ tm "b" ];
|
|
]
|
|
|
|
let () =
|
|
let g0 = paper_initial () in
|
|
let gc1 = Grammar.normalize (Transform.apply g0 (Transform.Chunk [ nt "A"; nt "B" ])) in
|
|
check "chunk(AB): S->X" (has_prod gc1 "S" [ nt "X" ]);
|
|
check "chunk(AB): S->AXB" (has_prod gc1 "S" [ nt "A"; nt "X"; nt "B" ]);
|
|
check "chunk(AB): S->AAXBB"
|
|
(has_prod gc1 "S" [ nt "A"; nt "A"; nt "X"; nt "B"; nt "B" ]);
|
|
check "chunk(AB): X->AB" (has_prod gc1 "X" [ nt "A"; nt "B" ]);
|
|
check "chunk preserves language"
|
|
(Grammar.language_up_to g0 6 = Grammar.language_up_to gc1 6)
|
|
|
|
let () =
|
|
let g0 = paper_initial () in
|
|
let gc1 = Grammar.normalize (Transform.apply g0 (Transform.Chunk [ nt "A"; nt "B" ])) in
|
|
let gc2 =
|
|
Grammar.normalize (Transform.apply gc1 (Transform.Chunk [ nt "A"; nt "X"; nt "B" ]))
|
|
in
|
|
check "chunk(AXB): S->Y" (has_prod gc2 "S" [ nt "Y" ]);
|
|
check "chunk(AXB): S->AYB" (has_prod gc2 "S" [ nt "A"; nt "Y"; nt "B" ]);
|
|
check "chunk(AXB): Y->AXB" (has_prod gc2 "Y" [ nt "A"; nt "X"; nt "B" ]);
|
|
let gm = Grammar.normalize (Transform.apply gc2 (Transform.Merge ("S", "Y"))) in
|
|
let gf = Grammar.normalize (Transform.apply gm (Transform.Merge ("S", "X"))) in
|
|
check "final: S->AB" (has_prod gf "S" [ nt "A"; nt "B" ]);
|
|
check "final: S->ASB" (has_prod gf "S" [ nt "A"; nt "S"; nt "B" ]);
|
|
let expected =
|
|
[
|
|
[ "a"; "b" ];
|
|
[ "a"; "a"; "b"; "b" ];
|
|
[ "a"; "a"; "a"; "b"; "b"; "b" ];
|
|
[ "a"; "a"; "a"; "a"; "b"; "b"; "b"; "b" ];
|
|
]
|
|
|> List.sort compare
|
|
in
|
|
check "final grammar generates a^n b^n"
|
|
(Grammar.language_up_to gf 8 = expected)
|
|
|
|
let () =
|
|
let g0 =
|
|
g "S"
|
|
[ prod "S" [ nt "A" ]; prod "S" [ nt "B" ]; prod "A" [ tm "a" ];
|
|
prod "B" [ tm "b" ] ]
|
|
in
|
|
let gm = Grammar.normalize (Transform.apply g0 (Transform.Merge ("A", "B"))) in
|
|
check "merge unifies terminals into one nonterminal"
|
|
(List.length (Grammar.terminals gm) = 2
|
|
&& List.length
|
|
(List.filter (fun p -> Grammar.symbol_is_nonterm (List.hd p.Grammar.rhs))
|
|
gm.Grammar.productions)
|
|
>= 1);
|
|
check "merge keeps both terminal productions" (has_prod gm "A" [ tm "a" ] && has_prod gm "A" [ tm "b" ])
|
|
|
|
let () =
|
|
let cfg =
|
|
{
|
|
Scoring.default_config with
|
|
em_iters = 5;
|
|
max_chunk = 3;
|
|
prior_weight = 1.0;
|
|
}
|
|
in
|
|
let corpus = [ sample [ "a"; "b" ] 1.0; sample [ "a"; "a"; "b"; "b" ] 1.0; sample [ "a"; "a"; "a"; "b"; "b"; "b" ] 1.0 ] in
|
|
let result =
|
|
Search.run ~config:cfg ~beam_width:3 ~max_steps:12 ~patience:3
|
|
(paper_initial ()) corpus
|
|
in
|
|
check "search improves the posterior"
|
|
(result.Search.best_score > result.Search.initial_score);
|
|
check "search returns accepted steps"
|
|
(List.length result.Search.steps > 0);
|
|
check "search output generates training strings"
|
|
(List.for_all
|
|
(fun s ->
|
|
List.mem s (Grammar.language_up_to result.Search.best 8))
|
|
[ [ "a"; "b" ]; [ "a"; "a"; "b"; "b" ]; [ "a"; "a"; "a"; "b"; "b"; "b" ] ])
|
|
|
|
let () =
|
|
let g0 =
|
|
g "S"
|
|
[ prod "S" [ nt "A"; nt "S"; nt "A" ]; prod "S" [ tm "c" ];
|
|
prod "S" [ nt "A"; nt "A"; nt "S"; nt "A"; nt "A" ];
|
|
prod "A" [ tm "a" ] ]
|
|
in
|
|
let gn = Grammar.normalize g0 in
|
|
check "redundant production pruned"
|
|
(not (has_prod gn "S" [ nt "A"; nt "A"; nt "S"; nt "A"; nt "A" ]));
|
|
check "recursive production kept" (has_prod gn "S" [ nt "A"; nt "S"; nt "A" ]);
|
|
check "pruning preserves the language"
|
|
(Grammar.language_up_to g0 5 = Grammar.language_up_to gn 5)
|
|
|
|
let () =
|
|
let gr =
|
|
g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ]
|
|
in
|
|
let prep = Parse.prepare gr in
|
|
let _, tree = Parse.viterbi prep [ "a"; "b" ] in
|
|
match tree with
|
|
| Some t ->
|
|
check "parse tree yield" (Parse.tree_yield t = [ "a"; "b" ]);
|
|
check "parse tree root" (match t with Parse.Node ("S", _) -> true | _ -> false)
|
|
| None -> check "parse tree yield" false
|
|
|
|
let () =
|
|
let gr =
|
|
g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ]
|
|
in
|
|
let corpus = [ sample [ "a"; "b" ] 1.0 ] in
|
|
let avg, _, _, unp = Scoring.heldout gr ~train:corpus corpus in
|
|
check "heldout loglik per token = 0" (approx avg 0.0);
|
|
check "heldout parseable" (unp = 0)
|
|
|
|
let contains hay needle =
|
|
let n = String.length hay and m = String.length needle in
|
|
let rec go i = if i + m > n then false else if String.sub hay i m = needle then true else go (i + 1) in
|
|
m = 0 || go 0
|
|
|
|
let () =
|
|
let g0 = g "S" [ prod "S" [ nt "S"; nt "S" ]; prod "S" [ tm "a" ] ] in
|
|
let d = Dot.dot_grammar g0 in
|
|
check "dot numbers repeated rhs symbols in order"
|
|
(contains d "pr_0 -> nt_S [label=\"1\"];"
|
|
&& contains d "pr_0 -> nt_S [label=\"2\"];");
|
|
check "dot labels production probability" (contains d "p=0.500");
|
|
check "dot marks the start symbol" (contains d "shape=doublecircle");
|
|
let tree_dot = Dot.dot_tree (Parse.Node ("S", [ Parse.Leaf "a" ])) in
|
|
check "parse tree dot has an edge" (contains tree_dot "->")
|
|
|
|
let () =
|
|
check "digamma(1)" (approx (Scoring.digamma 1.0) (-0.5772156649015329) ~tol:1e-9);
|
|
check "digamma(2)" (approx (Scoring.digamma 2.0) 0.42278433509846713 ~tol:1e-9)
|
|
|
|
let () =
|
|
let mk mode =
|
|
{ Scoring.default_config with alpha = 1.0; em_iters = 10; mode }
|
|
in
|
|
let corpus = [ sample [ "a" ] 1.0 ] in
|
|
let g1 = g "S" [ prod "S" [ tm "a" ] ] in
|
|
let d1 = Scoring.details ~config:(mk Scoring.Variational) g1 corpus in
|
|
check "variational ELBO (single production) = 0"
|
|
(approx d1.Scoring.posterior 0.0 ~tol:1e-6);
|
|
let g2 = g "S" [ prod "S" [ tm "a" ]; prod "S" [ tm "b" ] ] in
|
|
let dm = Scoring.details ~config:(mk Scoring.Marginal) g2 corpus in
|
|
let dv = Scoring.details ~config:(mk Scoring.Variational) g2 corpus in
|
|
check "variational matches marginal (unambiguous)"
|
|
(approx dv.Scoring.posterior dm.Scoring.posterior ~tol:1e-6)
|
|
|
|
let () =
|
|
let g0 =
|
|
g "S"
|
|
[ prod "S" [ nt "A"; nt "B"; nt "A"; nt "B" ]; prod "A" [ tm "a" ];
|
|
prod "B" [ tm "b" ] ]
|
|
in
|
|
let all =
|
|
Grammar.normalize (Transform.apply g0 (Transform.Chunk [ nt "A"; nt "B" ]))
|
|
in
|
|
let first =
|
|
Grammar.normalize
|
|
(Transform.apply ~occurrence:Scoring.First_occurrence g0
|
|
(Transform.Chunk [ nt "A"; nt "B" ]))
|
|
in
|
|
check "chunk all replaces every occurrence"
|
|
(has_prod all "S" [ nt "X"; nt "X" ]);
|
|
check "chunk first replaces only the leftmost"
|
|
(has_prod first "S" [ nt "X"; nt "A"; nt "B" ])
|
|
|
|
let () =
|
|
let g0 =
|
|
g "S"
|
|
[ prod "S" [ nt "A"; nt "S"; nt "A" ]; prod "S" [ tm "c" ];
|
|
prod "S" [ nt "A"; nt "A"; nt "A"; nt "S"; nt "A"; nt "A"; nt "A" ];
|
|
prod "A" [ tm "a" ] ]
|
|
in
|
|
let gn = Grammar.normalize g0 in
|
|
check "deeply redundant production pruned"
|
|
(not
|
|
(has_prod gn "S"
|
|
[ nt "A"; nt "A"; nt "A"; nt "S"; nt "A"; nt "A"; nt "A" ]));
|
|
check "deep redundancy preserves language"
|
|
(Grammar.language_up_to g0 7 = Grammar.language_up_to gn 7)
|
|
|
|
let () =
|
|
let g0 =
|
|
g "S"
|
|
[ prod "S" [ nt "A"; nt "B" ]; prod "S" [ nt "A"; nt "S"; nt "B" ];
|
|
prod "S" [ nt "A"; nt "A"; nt "B" ]; prod "A" [ tm "a" ];
|
|
prod "B" [ tm "b" ] ]
|
|
in
|
|
let gn = Grammar.normalize g0 in
|
|
check "non-redundant production kept"
|
|
(has_prod gn "S" [ nt "A"; nt "A"; nt "B" ])
|
|
|
|
let () =
|
|
if !failures = 0 then Printf.printf "\nAll tests passed.\n"
|
|
else begin
|
|
Printf.printf "\n%d test(s) failed.\n" !failures;
|
|
exit 1
|
|
end
|