This commit is contained in:
sneeker committed 2026-09-18 16:27:30 +00:00
commit 100e5a8239
96 files changed
+39927

No files matched your search

+4
View File
@@ -0,0 +1,4 @@
(test
(name test_scfg)
(flags (:standard -warn-error -a))
(libraries scfg))
+317
View File
@@ -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