initial
This commit is contained in:
96 files changed
+39927
No files matched your search
@@ -0,0 +1,4 @@
|
||||
(test
|
||||
(name test_scfg)
|
||||
(flags (:standard -warn-error -a))
|
||||
(libraries scfg))
|
||||
@@ -0,0 +1,317 @@
|
||||
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 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 () =
|
||||
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
|
||||
Reference in new issue
Block a user