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 a4db48b636
commit e8dcc8bcb6
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
[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.
[Reproducing Bayesian Model Merging for Stochastic Context Free Grammar Induction](https://gibsonsec.net/~karlsson/scfg-induction.pdf)
+23 -16
View File
@@ -128,7 +128,8 @@ type result = {
nts : int;
prods : int;
heldout_avg : float;
unparseable_test : int;
heldout_coverage : float;
unparseable_test : float;
covers_target : bool;
precision : float;
exact_language : bool;
@@ -195,8 +196,11 @@ let run_one ?(export_dir = None) ~config ~beam_width ~max_steps ~patience ~seed
search.Search.steps);
let final_grammar = search.Search.best in
let heldout_avg, _, _, unparseable =
Scoring.heldout ~config final_grammar ~train test
let heldout = Scoring.heldout ~config final_grammar ~train test in
let heldout_coverage =
if heldout.Scoring.total_mass > 0.0 then
heldout.Scoring.parsed_mass /. heldout.Scoring.total_mass
else 0.0
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;
nts = List.length (Grammar.nonterminals final_grammar);
prods = List.length final_grammar.Grammar.productions;
heldout_avg;
unparseable_test = unparseable;
heldout_avg = heldout.Scoring.average_loglik_per_token;
heldout_coverage;
unparseable_test = heldout.Scoring.unparseable_mass;
covers_target = covers;
precision;
exact_language = exact;
@@ -220,14 +225,15 @@ let run_one ?(export_dir = None) ~config ~beam_width ~max_steps ~patience ~seed
}
let header () =
"group\tname\tnts\tprods\tinitial_score\tfinal_score\theldout_avg\tunparseable\t\
covers\tprecision\texact\trecursive\tsteps\titerations"
"group\tname\tnts\tprods\tinitial_score\tfinal_score\theldout_avg\t\
heldout_coverage\tunparseable_mass\tcovers\tprecision\texact\trecursive\tsteps\t\
iterations"
let row r =
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.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
let to_tsv results = String.concat "\n" (header () :: List.map row results)
@@ -238,21 +244,22 @@ let report_markdown results =
Buffer.add_string b
"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 \
grammar (higher is better). `precision` is the fraction of strings the \
induced grammar generates up to the comparison length that belong to the \
target language.\n\n";
grammar. It is negative infinity when any weighted test mass is \
unparseable. `coverage` is the fraction of weighted held-out mass that the \
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
"| 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
"|---|---|---|---|---|---|---|---|---|---|---|---|\n";
"|---|---|---|---|---|---|---|---|---|---|---|---|---|\n";
List.iter
(fun r ->
Buffer.add_string b
(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.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))
results;
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;
}
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
ignore (Parse.fit_em prep train ~iters:config.em_iters ~tol:config.em_tol);
let total = 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
(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
| Some probability ->
total := !total +. (s.Corpus.count *. probability);
ntok := !ntok +. (s.Corpus.count *. float_of_int (List.length s.Corpus.tokens))
| None -> incr unparseable)
parsed_mass := !parsed_mass +. s.Corpus.count;
total := !total +. (s.Corpus.count *. probability)
| None -> unparseable_mass := !unparseable_mass +. s.Corpus.count)
test;
if !ntok > 0.0 then
(!total /. !ntok, !total, !ntok, !unparseable)
else (neg_infinity, !total, 0.0, !unparseable)
let total_loglik =
if !unparseable_mass > 0.0 then neg_infinity else !total
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" ] ]
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 heldout = Scoring.heldout gr ~train:corpus corpus in
check "heldout loglik per token = 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 n = String.length hay and m = String.length needle in