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
+23
-16
@@ -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";
|
||||
|
||||
Reference in new issue
Block a user