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

+75 -95
View File
@@ -110,13 +110,19 @@ let prob_of_prod prep (p : Grammar.production) =
in in
find idxs find idxs
let inside prep tokens = let log_add left right =
if left = neg_infinity then right
else if right = neg_infinity then left
else if left >= right then left +. log1p (exp (right -. left))
else right +. log1p (exp (left -. right))
let make_log_inside prep tokens =
let n = List.length tokens in let n = List.length tokens in
let w = Array.of_list tokens in let w = Array.of_list tokens in
let imemo = Hashtbl.create 1024 in let imemo = Hashtbl.create 1024 in
let smemo = Hashtbl.create 1024 in let smemo = Hashtbl.create 1024 in
let rec inside_nt nt i j = let rec inside_nt nt i j =
if j <= i then 0.0 if j <= i then neg_infinity
else else
match Hashtbl.find_opt imemo (nt, i, j) with match Hashtbl.find_opt imemo (nt, i, j) with
| Some v -> v | Some v -> v
@@ -127,9 +133,12 @@ let inside prep tokens =
let v = let v =
List.fold_left List.fold_left
(fun acc pi -> (fun acc pi ->
acc let probability = prep.probs.(pi) in
+. (prep.probs.(pi) *. seq_inside prep.prods.(pi).rhs i j)) if probability <= 0.0 then acc
0.0 idxs else
log_add acc
(log probability +. seq_inside prep.prods.(pi).rhs i j))
neg_infinity idxs
in in
Hashtbl.replace imemo (nt, i, j) v; Hashtbl.replace imemo (nt, i, j) v;
v v
@@ -139,28 +148,32 @@ let inside prep tokens =
| None -> | None ->
let v = let v =
match rhs with match rhs with
| [] -> if i = j then 1.0 else 0.0 | [] -> if i = j then 0.0 else neg_infinity
| Grammar.Term a :: rest -> | Grammar.Term a :: rest ->
if i < n && String.equal w.(i) a then if i < n && String.equal w.(i) a then
seq_inside rest (i + 1) j seq_inside rest (i + 1) j
else 0.0 else neg_infinity
| Grammar.Nonterm x :: rest -> | Grammar.Nonterm x :: rest ->
let acc = ref neg_infinity in
let acc = ref 0.0 in
for k = i + 1 to j - List.length rest do for k = i + 1 to j - List.length rest do
let a = inside_nt x i k in let a = inside_nt x i k in
if a > 0.0 then acc := !acc +. (a *. seq_inside rest k j) if a <> neg_infinity then
acc := log_add !acc (a +. seq_inside rest k j)
done; done;
!acc !acc
in in
Hashtbl.replace smemo (rhs, i, j) v; Hashtbl.replace smemo (rhs, i, j) v;
v v
in in
inside_nt prep.start 0 n (n, inside_nt, seq_inside)
let log_inside prep tokens = let log_inside prep tokens =
let p = inside prep tokens in let n, inside_nt, _ = make_log_inside prep tokens in
if p > 0.0 then Some (log p) else None let probability = inside_nt prep.start 0 n in
if probability = neg_infinity then None else Some probability
let inside prep tokens =
match log_inside prep tokens with None -> 0.0 | Some probability -> exp probability
let viterbi prep tokens = let viterbi prep tokens =
let n = List.length tokens in let n = List.length tokens in
@@ -244,100 +257,59 @@ let tree_to_string ?(indent = 0) tree =
in in
go indent tree go indent tree
let expected_counts prep tokens = let expected_counts_log prep tokens =
let n = List.length tokens in let n, inside_nt, seq_inside = make_log_inside prep tokens in
let w = Array.of_list tokens in
let imemo = Hashtbl.create 1024 in
let smemo = Hashtbl.create 1024 in
let rec inside_nt nt i j =
if j <= i then 0.0
else
match Hashtbl.find_opt imemo (nt, i, j) with
| Some v -> v
| None ->
let idxs =
Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs nt)
in
let v =
List.fold_left
(fun acc pi ->
acc
+. (prep.probs.(pi) *. seq_inside prep.prods.(pi).rhs i j))
0.0 idxs
in
Hashtbl.replace imemo (nt, i, j) v;
v
and seq_inside rhs i j =
match Hashtbl.find_opt smemo (rhs, i, j) with
| Some v -> v
| None ->
let v =
match rhs with
| [] -> if i = j then 1.0 else 0.0
| Grammar.Term a :: rest ->
if i < n && String.equal w.(i) a then seq_inside rest (i + 1) j
else 0.0
| Grammar.Nonterm x :: rest ->
let acc = ref 0.0 in
for k = i + 1 to j - List.length rest do
let a = inside_nt x i k in
if a > 0.0 then acc := !acc +. (a *. seq_inside rest k j)
done;
!acc
in
Hashtbl.replace smemo (rhs, i, j) v;
v
in
let ptotal = inside_nt prep.start 0 n in let ptotal = inside_nt prep.start 0 n in
let np = Array.length prep.prods in let np = Array.length prep.prods in
if ptotal <= 0.0 then (Array.make np 0.0, 0.0) if ptotal = neg_infinity then (Array.make np 0.0, neg_infinity)
else begin else begin
let omemo = Hashtbl.create 1024 in let omemo = Hashtbl.create 1024 in
let get_o nt i j = let get_o nt i j =
Option.value ~default:0.0 (Hashtbl.find_opt omemo (nt, i, j)) Option.value ~default:neg_infinity (Hashtbl.find_opt omemo (nt, i, j))
in in
let add_o nt i j v = let add_o nt i j v =
Hashtbl.replace omemo (nt, i, j) (get_o nt i j +. v) Hashtbl.replace omemo (nt, i, j) (log_add (get_o nt i j) v)
in in
add_o prep.start 0 n 1.0; add_o prep.start 0 n 0.0;
for len = n downto 1 do for len = n downto 1 do
for p = 0 to n - len do for p = 0 to n - len do
let q = p + len in let q = p + len in
List.iter List.iter
(fun b -> (fun b ->
let ob = get_o b p q in let ob = get_o b p q in
if ob > 0.0 then if ob <> neg_infinity then
List.iter List.iter
(fun pi -> (fun pi ->
let prod = prep.prods.(pi) in let prod = prep.prods.(pi) in
if String.equal prod.lhs b then begin if String.equal prod.lhs b then begin
let prob = prep.probs.(pi) in let prob = prep.probs.(pi) in
let arr = Array.of_list prod.rhs in if prob > 0.0 then begin
let m = Array.length arr in let arr = Array.of_list prod.rhs in
Array.iteri let m = Array.length arr in
(fun k s -> Array.iteri
match s with (fun k s ->
| Grammar.Nonterm x -> match s with
let left = | Grammar.Nonterm x ->
Array.to_list (Array.sub arr 0 k) let left =
in Array.to_list (Array.sub arr 0 k)
let right = in
Array.to_list let right =
(Array.sub arr (k + 1) (m - k - 1)) Array.to_list
in (Array.sub arr (k + 1) (m - k - 1))
for i = p to q do in
let lw = seq_inside left p i in for i = p to q do
if lw > 0.0 then let lw = seq_inside left p i in
for j = i to q do if lw <> neg_infinity then
let rw = seq_inside right j q in for j = i to q do
if rw > 0.0 then let rw = seq_inside right j q in
add_o x i j if rw <> neg_infinity then
(ob *. prob *. lw *. rw) add_o x i j
done (ob +. log prob +. lw +. rw)
done done
| Grammar.Term _ -> ()) done
arr | Grammar.Term _ -> ())
arr
end
end) end)
(Option.value ~default:[] (Option.value ~default:[]
(Hashtbl.find_opt prep.by_lhs b))) (Hashtbl.find_opt prep.by_lhs b)))
@@ -348,19 +320,25 @@ let expected_counts prep tokens =
Array.iteri Array.iteri
(fun pi (prod : Grammar.production) -> (fun pi (prod : Grammar.production) ->
let prob = prep.probs.(pi) in let prob = prep.probs.(pi) in
let s = ref 0.0 in let s = ref neg_infinity in
for i = 0 to n do for i = 0 to n do
for j = i to n do for j = i to n do
let o = get_o prod.lhs i j in let o = get_o prod.lhs i j in
if o > 0.0 then if o <> neg_infinity then
s := !s +. (o *. seq_inside prod.rhs i j) s := log_add !s (o +. seq_inside prod.rhs i j)
done done
done; done;
counts.(pi) <- prob *. !s /. ptotal) counts.(pi) <-
if prob <= 0.0 || !s = neg_infinity then 0.0
else exp (log prob +. !s -. ptotal))
prep.prods; prep.prods;
(counts, ptotal) (counts, ptotal)
end end
let expected_counts prep tokens =
let counts, probability = expected_counts_log prep tokens in
(counts, exp probability)
let fit_em prep (corpus : Corpus.t) ~iters ~tol = let fit_em prep (corpus : Corpus.t) ~iters ~tol =
let np = Array.length prep.prods in let np = Array.length prep.prods in
let last_ll = ref neg_infinity in let last_ll = ref neg_infinity in
@@ -372,12 +350,14 @@ let fit_em prep (corpus : Corpus.t) ~iters ~tol =
ll := 0.0; ll := 0.0;
List.iter List.iter
(fun s -> (fun s ->
let cnt, p = expected_counts prep s.Corpus.tokens in let cnt, log_probability =
if p > 0.0 then begin expected_counts_log prep s.Corpus.tokens
in
if log_probability <> neg_infinity then begin
Array.iteri Array.iteri
(fun i x -> acc.(i) <- acc.(i) +. (s.Corpus.count *. x)) (fun i x -> acc.(i) <- acc.(i) +. (s.Corpus.count *. x))
cnt; cnt;
ll := !ll +. (s.Corpus.count *. log p) ll := !ll +. (s.Corpus.count *. log_probability)
end) end)
corpus; corpus;
Hashtbl.iter Hashtbl.iter
+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 marginal_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corpus.t) =
let uprep = Parse.prepare_uniform g in 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) then (neg_infinity, neg_infinity, [||], Parse.prepare g)
else begin else begin
let prep = Parse.prepare g in let prep = Parse.prepare g in
@@ -107,8 +111,9 @@ let marginal_loglik ?(config = default_config) (g : Grammar.t) (corpus : Corpus.
let ll = let ll =
List.fold_left List.fold_left
(fun acc s -> (fun acc s ->
let p = Parse.inside prep s.Corpus.tokens in match Parse.log_inside prep s.Corpus.tokens with
if p > 0.0 then acc +. (s.Corpus.count *. log p) else neg_infinity) | Some probability -> acc +. (s.Corpus.count *. probability)
| None -> neg_infinity)
0.0 corpus 0.0 corpus
in in
if ll = neg_infinity then (neg_infinity, neg_infinity, [||], prep) 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 let log_z = ref 0.0 in
List.iter List.iter
(fun s -> (fun s ->
let z = Parse.inside prep s.Corpus.tokens in match Parse.log_inside prep s.Corpus.tokens with
if z > 0.0 then log_z := !log_z +. (s.Corpus.count *. log z) | Some probability ->
else log_z := neg_infinity) log_z := !log_z +. (s.Corpus.count *. probability)
| None -> log_z := neg_infinity)
corpus; corpus;
!log_z -. !kl !log_z -. !kl
@@ -266,12 +272,11 @@ let heldout ?(config = default_config) (g : Grammar.t) ~train (test : Corpus.t)
let unparseable = ref 0 in let unparseable = ref 0 in
List.iter List.iter
(fun s -> (fun s ->
let p = Parse.inside prep s.Corpus.tokens in match Parse.log_inside prep s.Corpus.tokens with
if p > 0.0 then begin | Some probability ->
total := !total +. (s.Corpus.count *. log p); total := !total +. (s.Corpus.count *. probability);
ntok := !ntok +. (s.Corpus.count *. float_of_int (List.length s.Corpus.tokens)) ntok := !ntok +. (s.Corpus.count *. float_of_int (List.length s.Corpus.tokens))
end | None -> incr unparseable)
else incr unparseable)
test; test;
if !ntok > 0.0 then if !ntok > 0.0 then
(!total /. !ntok, !total, !ntok, !unparseable) (!total /. !ntok, !total, !ntok, !unparseable)
+13
View File
@@ -43,6 +43,19 @@ let () =
check "viterbi P(aa) = 1/8" (approx vp 0.125); check "viterbi P(aa) = 1/8" (approx vp 0.125);
check "viterbi tree exists" (vt <> None) check "viterbi tree exists" (vt <> None)
let () =
let recursive =
g "S"
[ { Grammar.lhs = "S"; rhs = [ tm "a"; nt "S" ]; count = 1.0 };
{ Grammar.lhs = "S"; rhs = [ tm "a" ]; count = 999.0 } ]
in
let tokens = List.init 120 (fun _ -> "a") in
match Parse.log_inside (Parse.prepare recursive) tokens with
| Some probability ->
check "log inside survives probability underflow"
(Float.is_finite probability && probability < -700.0)
| None -> check "log inside survives probability underflow" false
let () = let () =
let gr = g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ] in let gr = g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ] in
let prep = Parse.prepare gr in let prep = Parse.prepare gr in