fix(parse): move inside outside inference into log space

This commit is contained in:
milner committed 2026-09-23 11:20:00 +00:00
1 parent 38ab0ea6eb
commit 697801a0ee
3 files changed
+104 -106

No files matched your search

+16 -11
View File
@@ -98,7 +98,11 @@ let structural_logprior ?(config = default_config) (g : Grammar.t) =
let marginal_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) =
let uprep = Parse.prepare_uniform g in
if not (List.for_all (fun s -> Parse.inside uprep s.Corpus.tokens > 0.0) corpus)
if
not
(List.for_all
(fun s -> Option.is_some (Parse.log_inside uprep s.Corpus.tokens))
corpus)
then (neg_infinity, neg_infinity, [||], Parse.prepare g)
else begin
let prep = Parse.prepare g in
@@ -107,8 +111,9 @@ let marginal_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corpus.
let ll =
List.fold_left
(fun acc s ->
let p = Parse.inside prep s.Corpus.tokens in
if p > 0.0 then acc +. (s.Corpus.count *. log p) else neg_infinity)
match Parse.log_inside prep s.Corpus.tokens with
| Some probability -> acc +. (s.Corpus.count *. probability)
| None -> neg_infinity)
0.0 corpus
in
if ll = neg_infinity then (neg_infinity, neg_infinity, [||], prep)
@@ -211,9 +216,10 @@ let variational_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corp
let log_z = ref 0.0 in
List.iter
(fun s ->
let z = Parse.inside prep s.Corpus.tokens in
if z > 0.0 then log_z := !log_z +. (s.Corpus.count *. log z)
else log_z := neg_infinity)
match Parse.log_inside prep s.Corpus.tokens with
| Some probability ->
log_z := !log_z +. (s.Corpus.count *. probability)
| None -> log_z := neg_infinity)
corpus;
!log_z -. !kl
@@ -266,12 +272,11 @@ let heldout ?(config = default_config) (g : Grammar.t) ~train (test : Corpus.t)
let unparseable = ref 0 in
List.iter
(fun s ->
let p = Parse.inside prep s.Corpus.tokens in
if p > 0.0 then begin
total := !total +. (s.Corpus.count *. log p);
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))
end
else incr unparseable)
| None -> incr unparseable)
test;
if !ntok > 0.0 then
(!total /. !ntok, !total, !ntok, !unparseable)