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

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