From 697801a0ee0b4209c9897a4345540691cf31abc7 Mon Sep 17 00:00:00 2001 From: milner Date: Wed, 23 Sep 2026 11:20:00 +0000 Subject: [PATCH] fix(parse): move inside outside inference into log space --- lib/parse.ml | 170 ++++++++++++++++++++-------------------------- lib/scoring.ml | 27 +++++--- test/test_scfg.ml | 13 ++++ 3 files changed, 104 insertions(+), 106 deletions(-) diff --git a/lib/parse.ml b/lib/parse.ml index 5dff148..8a3d936 100644 --- a/lib/parse.ml +++ b/lib/parse.ml @@ -110,13 +110,19 @@ let prob_of_prod prep (p : Grammar.production) = in 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 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 + if j <= i then neg_infinity else match Hashtbl.find_opt imemo (nt, i, j) with | Some v -> v @@ -127,9 +133,12 @@ let inside prep tokens = let v = List.fold_left (fun acc pi -> - acc - +. (prep.probs.(pi) *. seq_inside prep.prods.(pi).rhs i j)) - 0.0 idxs + let probability = prep.probs.(pi) in + if probability <= 0.0 then acc + else + log_add acc + (log probability +. seq_inside prep.prods.(pi).rhs i j)) + neg_infinity idxs in Hashtbl.replace imemo (nt, i, j) v; v @@ -139,28 +148,32 @@ let inside prep tokens = | None -> let v = 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 -> if i < n && String.equal w.(i) a then seq_inside rest (i + 1) j - else 0.0 + else neg_infinity | Grammar.Nonterm x :: rest -> - - let acc = ref 0.0 in + let acc = ref neg_infinity 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) + if a <> neg_infinity then + acc := log_add !acc (a +. seq_inside rest k j) done; !acc in Hashtbl.replace smemo (rhs, i, j) v; v in - inside_nt prep.start 0 n + (n, inside_nt, seq_inside) let log_inside prep tokens = - let p = inside prep tokens in - if p > 0.0 then Some (log p) else None + let n, inside_nt, _ = make_log_inside prep tokens in + 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 n = List.length tokens in @@ -244,100 +257,59 @@ let tree_to_string ?(indent = 0) tree = in go indent tree -let expected_counts prep tokens = - let n = List.length 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 expected_counts_log prep tokens = + let n, inside_nt, seq_inside = make_log_inside prep tokens in let ptotal = inside_nt prep.start 0 n 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 let omemo = Hashtbl.create 1024 in 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 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 - add_o prep.start 0 n 1.0; + add_o prep.start 0 n 0.0; for len = n downto 1 do for p = 0 to n - len do let q = p + len in List.iter (fun b -> let ob = get_o b p q in - if ob > 0.0 then + if ob <> neg_infinity then List.iter (fun pi -> let prod = prep.prods.(pi) in if String.equal prod.lhs b then begin let prob = prep.probs.(pi) in - let arr = Array.of_list prod.rhs in - let m = Array.length arr in - Array.iteri - (fun k s -> - match s with - | Grammar.Nonterm x -> - let left = - Array.to_list (Array.sub arr 0 k) - in - let right = - Array.to_list - (Array.sub arr (k + 1) (m - k - 1)) - in - for i = p to q do - let lw = seq_inside left p i in - if lw > 0.0 then - for j = i to q do - let rw = seq_inside right j q in - if rw > 0.0 then - add_o x i j - (ob *. prob *. lw *. rw) - done - done - | Grammar.Term _ -> ()) - arr + if prob > 0.0 then begin + let arr = Array.of_list prod.rhs in + let m = Array.length arr in + Array.iteri + (fun k s -> + match s with + | Grammar.Nonterm x -> + let left = + Array.to_list (Array.sub arr 0 k) + in + let right = + Array.to_list + (Array.sub arr (k + 1) (m - k - 1)) + in + for i = p to q do + let lw = seq_inside left p i in + if lw <> neg_infinity then + for j = i to q do + let rw = seq_inside right j q in + if rw <> neg_infinity then + add_o x i j + (ob +. log prob +. lw +. rw) + done + done + | Grammar.Term _ -> ()) + arr + end end) (Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs b))) @@ -348,19 +320,25 @@ let expected_counts prep tokens = Array.iteri (fun pi (prod : Grammar.production) -> 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 j = i to n do let o = get_o prod.lhs i j in - if o > 0.0 then - s := !s +. (o *. seq_inside prod.rhs i j) + if o <> neg_infinity then + s := log_add !s (o +. seq_inside prod.rhs i j) 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; (counts, ptotal) 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 np = Array.length prep.prods in let last_ll = ref neg_infinity in @@ -372,12 +350,14 @@ let fit_em prep (corpus : Corpus.t) ~iters ~tol = ll := 0.0; List.iter (fun s -> - let cnt, p = expected_counts prep s.Corpus.tokens in - if p > 0.0 then begin + let cnt, log_probability = + expected_counts_log prep s.Corpus.tokens + in + if log_probability <> neg_infinity then begin Array.iteri (fun i x -> acc.(i) <- acc.(i) +. (s.Corpus.count *. x)) cnt; - ll := !ll +. (s.Corpus.count *. log p) + ll := !ll +. (s.Corpus.count *. log_probability) end) corpus; Hashtbl.iter diff --git a/lib/scoring.ml b/lib/scoring.ml index 9712b57..0189945 100644 --- a/lib/scoring.ml +++ b/lib/scoring.ml @@ -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) diff --git a/test/test_scfg.ml b/test/test_scfg.ml index 1ad3dbe..386c39f 100644 --- a/test/test_scfg.ml +++ b/test/test_scfg.ml @@ -43,6 +43,19 @@ let () = check "viterbi P(aa) = 1/8" (approx vp 0.125); 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 gr = g "S" [ prod "S" [ nt "A"; nt "B" ]; prod "A" [ tm "a" ]; prod "B" [ tm "b" ] ] in let prep = Parse.prepare gr in