fix(parse): move inside outside inference into log space
This commit is contained in:
3 files changed
+82
-84
No files matched your search
+53
-73
@@ -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,75 +257,33 @@ 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
|
||||||
|
if prob > 0.0 then begin
|
||||||
let arr = Array.of_list prod.rhs in
|
let arr = Array.of_list prod.rhs in
|
||||||
let m = Array.length arr in
|
let m = Array.length arr in
|
||||||
Array.iteri
|
Array.iteri
|
||||||
@@ -328,16 +299,17 @@ let expected_counts prep tokens =
|
|||||||
in
|
in
|
||||||
for i = p to q do
|
for i = p to q do
|
||||||
let lw = seq_inside left p i in
|
let lw = seq_inside left p i in
|
||||||
if lw > 0.0 then
|
if lw <> neg_infinity then
|
||||||
for j = i to q do
|
for j = i to q do
|
||||||
let rw = seq_inside right j q in
|
let rw = seq_inside right j q in
|
||||||
if rw > 0.0 then
|
if rw <> neg_infinity then
|
||||||
add_o x i j
|
add_o x i j
|
||||||
(ob *. prob *. lw *. rw)
|
(ob +. log prob +. lw +. rw)
|
||||||
done
|
done
|
||||||
done
|
done
|
||||||
| Grammar.Term _ -> ())
|
| Grammar.Term _ -> ())
|
||||||
arr
|
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
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in new issue
Block a user