fix(eval): count unparseable heldout mass in cross entropy
This commit is contained in:
1 parent
46982a23d9
commit
086495847a
4 files changed
+76
-41
No files matched your search
@@ -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
@@ -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
@@ -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
@@ -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
|
||||||
|
|||||||
Reference in new issue
Block a user