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

+20 -3
View File
@@ -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