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

+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;
}