From 321bf25cb232a4e4f889ea01922ce70f8c658143 Mon Sep 17 00:00:00 2001 From: sneeker Date: Wed, 23 Sep 2026 15:00:00 +0000 Subject: [PATCH] fix(eval): count unparseable heldout mass in cross entropy --- README.md | 15 +-------------- lib/experiment.ml | 39 +++++++++++++++++++++++---------------- lib/scoring.ml | 40 ++++++++++++++++++++++++++++++++-------- test/test_scfg.ml | 23 ++++++++++++++++++++--- 4 files changed, 76 insertions(+), 41 deletions(-) diff --git a/README.md b/README.md index 210929e..117d922 100644 --- a/README.md +++ b/README.md @@ -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). - -

- the paper's Figure 2 and Table 1, and the eleven induced grammars -

- -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) diff --git a/lib/experiment.ml b/lib/experiment.ml index 81b1990..47de599 100644 --- a/lib/experiment.ml +++ b/lib/experiment.ml @@ -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"; diff --git a/lib/scoring.ml b/lib/scoring.ml index 05a6152..79ce2b9 100644 --- a/lib/scoring.ml +++ b/lib/scoring.ml @@ -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; + } diff --git a/test/test_scfg.ml b/test/test_scfg.ml index e6f7cad..cb21547 100644 --- a/test/test_scfg.ml +++ b/test/test_scfg.ml @@ -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