fix(eval): count unparseable heldout mass in cross entropy

This commit is contained in:
milner committed 2026-09-23 15:00:00 +00:00
1 parent 46982a23d9
commit 086495847a
4 files changed
+76 -41

No files matched your search

+1 -14
View File
@@ -1,14 +1 @@
An implementation of stochastic context free grammar induction, following Stolcke and Omohundro's [Reproducing Bayesian Model Merging for Stochastic Context Free Grammar Induction](https://gibsonsec.net/~karlsson/scfg-induction.pdf)
[Inducing Probabilistic Grammars by Bayesian Model Merging](https://arxiv.org/abs/cmp-lg/9409010) (ICGI 1994).
<p align="center">
<img src="meta/paper-and-induced.svg" alt="the paper's Figure 2 and Table 1, and the eleven induced grammars" width="900">
</p>
We start from the most specific grammar the data permits: every sample contributes its own production, and every terminal that occurs gets a corresponding nonterminal. At this stage there's effectively no sharing between samples, so the grammar is just memorising the corpus rather than generalising beyond it.
From there, we generalise with two operators, merging and chunking. Merging takes a pair of nonterminals and folds them into a single nonterminal containing the union of their productions; chunking replaces a contiguous sequence of symbols with a fresh nonterminal. Chunking doesn't itself change the language the grammar generates, but it changes the internal structure in a way that can expose useful merges which weren't previously available.
We rank candidate grammars with a structural description length and an approximate evidence objective. The default scorer fits production probabilities with expectation maximisation, computes fractional expected production counts, and applies a symmetric Dirichlet correction. Maximum likelihood and a variational evidence lower bound are also available.
We explore the resulting grammar space using either beam or best first search, then fit the parameters by expectation maximisation once the grammar structure is fixed. Parsing is a generalised CYK or inside computation over spans, which lets the grammar be scored directly without an intermediate conversion into Chomsky normal form.
+23 -16
View File
@@ -128,7 +128,8 @@ type result = {
nts : int; nts : int;
prods : int; prods : int;
heldout_avg : float; heldout_avg : float;
unparseable_test : int; heldout_coverage : float;
unparseable_test : float;
covers_target : bool; covers_target : bool;
precision : float; precision : float;
exact_language : bool; exact_language : bool;
@@ -195,8 +196,11 @@ let run_one ?(export_dir = None) ~config ~beam_width ~max_steps ~patience ~seed
search.Search.steps); search.Search.steps);
let final_grammar = search.Search.best in let final_grammar = search.Search.best in
let heldout_avg, _, _, unparseable = let heldout = Scoring.heldout ~config final_grammar ~train test in
Scoring.heldout ~config final_grammar ~train test let heldout_coverage =
if heldout.Scoring.total_mass > 0.0 then
heldout.Scoring.parsed_mass /. heldout.Scoring.total_mass
else 0.0
in in
let covers, precision, exact = language_sets target final_grammar max_len in let covers, precision, exact = language_sets target final_grammar max_len in
{ {
@@ -209,8 +213,9 @@ let run_one ?(export_dir = None) ~config ~beam_width ~max_steps ~patience ~seed
final_score = search.Search.best_score; final_score = search.Search.best_score;
nts = List.length (Grammar.nonterminals final_grammar); nts = List.length (Grammar.nonterminals final_grammar);
prods = List.length final_grammar.Grammar.productions; prods = List.length final_grammar.Grammar.productions;
heldout_avg; heldout_avg = heldout.Scoring.average_loglik_per_token;
unparseable_test = unparseable; heldout_coverage;
unparseable_test = heldout.Scoring.unparseable_mass;
covers_target = covers; covers_target = covers;
precision; precision;
exact_language = exact; exact_language = exact;
@@ -220,14 +225,15 @@ let run_one ?(export_dir = None) ~config ~beam_width ~max_steps ~patience ~seed
} }
let header () = let header () =
"group\tname\tnts\tprods\tinitial_score\tfinal_score\theldout_avg\tunparseable\t\ "group\tname\tnts\tprods\tinitial_score\tfinal_score\theldout_avg\t\
covers\tprecision\texact\trecursive\tsteps\titerations" heldout_coverage\tunparseable_mass\tcovers\tprecision\texact\trecursive\tsteps\t\
iterations"
let row r = let row r =
Printf.sprintf Printf.sprintf
"%s\t%s\t%d\t%d\t%.6g\t%.6g\t%.6g\t%d\t%b\t%.4f\t%b\t%b\t%d\t%d" "%s\t%s\t%d\t%d\t%.6g\t%.6g\t%.6g\t%.6g\t%.6g\t%b\t%.4f\t%b\t%b\t%d\t%d"
r.target.group r.target.name r.nts r.prods r.initial_score r.final_score r.target.group r.target.name r.nts r.prods r.initial_score r.final_score
r.heldout_avg r.unparseable_test r.covers_target r.precision r.heldout_avg r.heldout_coverage r.unparseable_test r.covers_target r.precision
r.exact_language r.recursive r.steps r.iterations r.exact_language r.recursive r.steps r.iterations
let to_tsv results = String.concat "\n" (header () :: List.map row results) let to_tsv results = String.concat "\n" (header () :: List.map row results)
@@ -238,21 +244,22 @@ let report_markdown results =
Buffer.add_string b Buffer.add_string b
"Scores are natural-log posterior values. `heldout_avg` is the average \ "Scores are natural-log posterior values. `heldout_avg` is the average \
log-likelihood per token on a held-out sample drawn from the same target \ log-likelihood per token on a held-out sample drawn from the same target \
grammar (higher is better). `precision` is the fraction of strings the \ grammar. It is negative infinity when any weighted test mass is \
induced grammar generates up to the comparison length that belong to the \ unparseable. `coverage` is the fraction of weighted held-out mass that the \
target language.\n\n"; grammar parses. `precision` is the fraction of strings the induced grammar \
generates up to the comparison length that belong to the target language.\n\n";
Buffer.add_string b Buffer.add_string b
"| group | target | nonterms | prods | initial | final | heldout/token | \ "| group | target | nonterms | prods | initial | final | heldout/token | \
covers | precision | exact | recursive | steps |\n"; coverage | covers | precision | exact | recursive | steps |\n";
Buffer.add_string b Buffer.add_string b
"|---|---|---|---|---|---|---|---|---|---|---|---|\n"; "|---|---|---|---|---|---|---|---|---|---|---|---|---|\n";
List.iter List.iter
(fun r -> (fun r ->
Buffer.add_string b Buffer.add_string b
(Printf.sprintf (Printf.sprintf
"| %s | %s | %d | %d | %.2f | %.2f | %.3f | %b | %.3f | %b | %b | %d |\n" "| %s | %s | %d | %d | %.2f | %.2f | %.3f | %.3f | %b | %.3f | %b | %b | %d |\n"
r.target.group r.target.name r.nts r.prods r.initial_score r.target.group r.target.name r.nts r.prods r.initial_score
r.final_score r.heldout_avg r.covers_target r.precision r.final_score r.heldout_avg r.heldout_coverage r.covers_target r.precision
r.exact_language r.recursive r.steps)) r.exact_language r.recursive r.steps))
results; results;
Buffer.add_string b "\n## Final grammars\n\n"; Buffer.add_string b "\n## Final grammars\n\n";
+32 -8
View File
@@ -267,20 +267,44 @@ let details ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) =
num_productions = List.length g.Grammar.productions; num_productions = List.length g.Grammar.productions;
} }
let heldout ?(config = default_config) (g : Grammar.t) ~train (test : Corpus.t) = type heldout_result = {
average_loglik_per_token : float;
total_loglik : float;
total_tokens : float;
parsed_mass : float;
total_mass : float;
unparseable_mass : float;
}
let heldout ?(config = default_config) (g : Grammar.t) ~train
(test : Corpus.t) =
let prep = Parse.prepare g in let prep = Parse.prepare g in
ignore (Parse.fit_em prep train ~iters:config.em_iters ~tol:config.em_tol); ignore (Parse.fit_em prep train ~iters:config.em_iters ~tol:config.em_tol);
let total = ref 0.0 in let total = ref 0.0 in
let ntok = ref 0.0 in let ntok = ref 0.0 in
let unparseable = ref 0 in let parsed_mass = ref 0.0 in
let total_mass = ref 0.0 in
let unparseable_mass = ref 0.0 in
List.iter List.iter
(fun s -> (fun s ->
total_mass := !total_mass +. s.Corpus.count;
ntok :=
!ntok +. (s.Corpus.count *. float_of_int (List.length s.Corpus.tokens));
match Parse.log_inside prep s.Corpus.tokens with match Parse.log_inside prep s.Corpus.tokens with
| Some probability -> | Some probability ->
total := !total +. (s.Corpus.count *. probability); parsed_mass := !parsed_mass +. s.Corpus.count;
ntok := !ntok +. (s.Corpus.count *. float_of_int (List.length s.Corpus.tokens)) total := !total +. (s.Corpus.count *. probability)
| None -> incr unparseable) | None -> unparseable_mass := !unparseable_mass +. s.Corpus.count)
test; test;
if !ntok > 0.0 then let total_loglik =
(!total /. !ntok, !total, !ntok, !unparseable) if !unparseable_mass > 0.0 then neg_infinity else !total
else (neg_infinity, !total, 0.0, !unparseable) in
{
average_loglik_per_token =
(if !ntok > 0.0 then total_loglik /. !ntok else neg_infinity);
total_loglik;
total_tokens = !ntok;
parsed_mass = !parsed_mass;
total_mass = !total_mass;
unparseable_mass = !unparseable_mass;
}
+20 -3
View File
@@ -274,9 +274,26 @@ let () =
g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ] g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ]
in in
let corpus = [ sample [ "a"; "b" ] 1.0 ] in let corpus = [ sample [ "a"; "b" ] 1.0 ] in
let avg, _, _, unp = Scoring.heldout gr ~train:corpus corpus in let heldout = Scoring.heldout gr ~train:corpus corpus in
check "heldout loglik per token = 0" (approx avg 0.0); check "heldout loglik per token = 0"
check "heldout parseable" (unp = 0) (approx heldout.Scoring.average_loglik_per_token 0.0);
check "heldout parseable" (heldout.Scoring.unparseable_mass = 0.0)
let () =
let gr = g "S" [ prod "S" [ tm "a" ] ] in
let train = [ sample [ "a" ] 1.0 ] in
let test = [ sample [ "a" ] 1.0; sample [ "b" ] 7.0 ] in
let heldout = Scoring.heldout gr ~train test in
check "heldout counts weighted unparseable mass"
(approx heldout.Scoring.unparseable_mass 7.0);
check "heldout counts all token mass"
(approx heldout.Scoring.total_tokens 8.0);
check "heldout reports partial coverage"
(approx
(heldout.Scoring.parsed_mass /. heldout.Scoring.total_mass)
0.125);
check "unparseable heldout mass has infinite cross entropy"
(heldout.Scoring.average_loglik_per_token = neg_infinity)
let contains hay needle = let contains hay needle =
let n = String.length hay and m = String.length needle in let n = String.length hay and m = String.length needle in