This commit is contained in:
milner committed 2026-09-18 16:27:30 +00:00
commit 3170b7b398
96 files changed
+39927

No files matched your search

+355
View File
@@ -0,0 +1,355 @@
type tree = Node of string * tree list | Leaf of string
let rec tree_yield = function
| Leaf s -> [ s ]
| Node (_, children) -> List.concat_map tree_yield children
type prep = {
prods : Grammar.production array;
mutable probs : float array;
by_lhs : (string, int list) Hashtbl.t;
lhs_list : string list;
start : string;
}
let prepare (g : Grammar.t) =
let prods : Grammar.production array = Array.of_list g.productions in
let n = Array.length prods in
let probs = Array.make n 0.0 in
let by_lhs = Hashtbl.create 64 in
Array.iteri
(fun i (p : Grammar.production) ->
let l =
match Hashtbl.find_opt by_lhs p.lhs with Some l -> l | None -> []
in
Hashtbl.replace by_lhs p.lhs (i :: l))
prods;
Hashtbl.iter
(fun _ idxs ->
let tot =
List.fold_left (fun s i -> s +. prods.(i).count) 0.0 idxs
in
let k = List.length idxs in
List.iter
(fun i ->
probs.(i) <-
(if tot > 0.0 then prods.(i).count /. tot
else 1.0 /. float_of_int k))
idxs)
by_lhs;
{
prods;
probs;
by_lhs;
lhs_list = Hashtbl.fold (fun k _ acc -> k :: acc) by_lhs [];
start = g.start;
}
let prepare_uniform (g : Grammar.t) =
let prep = prepare g in
Array.fill prep.probs 0 (Array.length prep.probs) 0.0;
Hashtbl.iter
(fun _ idxs ->
let k = List.length idxs in
if k > 0 then
List.iter (fun i -> prep.probs.(i) <- 1.0 /. float_of_int k) idxs)
prep.by_lhs;
prep
let prob_of_prod prep (p : Grammar.production) =
let idxs =
Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs p.lhs)
in
let rec find = function
| [] -> 0.0
| i :: rest ->
if Grammar.equal_rhs prep.prods.(i).rhs p.rhs then prep.probs.(i)
else find rest
in
find idxs
let 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
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
inside_nt prep.start 0 n
let log_inside prep tokens =
let p = inside prep tokens in
if p > 0.0 then Some (log p) else None
let viterbi prep tokens =
let n = List.length tokens in
let w = Array.of_list tokens in
let vmemo = Hashtbl.create 1024 in
let smemo = Hashtbl.create 1024 in
let rec vnt nt i j =
if j <= i then (0.0, None)
else
match Hashtbl.find_opt vmemo (nt, i, j) with
| Some v -> v
| None ->
let idxs =
Option.value ~default:[] (Hashtbl.find_opt prep.by_lhs nt)
in
let bestp = ref 0.0 and bestt = ref None in
List.iter
(fun pi ->
let p = prep.prods.(pi) in
let sp, sch = seq_vit p.rhs i j in
if sp > 0.0 then begin
let v = prep.probs.(pi) *. sp in
if v > !bestp then begin
bestp := v;
bestt :=
(match sch with
| Some ch -> Some (Node (nt, ch))
| None -> None)
end
end)
idxs;
let res = (!bestp, !bestt) in
Hashtbl.replace vmemo (nt, i, j) res;
res
and seq_vit 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, Some []) else (0.0, None)
| Grammar.Term a :: rest ->
if i < n && String.equal w.(i) a then (
match seq_vit rest (i + 1) j with
| p, Some ts -> (p, Some (Leaf a :: ts))
| z -> z)
else (0.0, None)
| Grammar.Nonterm x :: rest ->
let bestp = ref 0.0 and bestch = ref None in
for k = i + 1 to j - List.length rest do
let px, tx = vnt x i k in
if px > 0.0 then begin
let pr, tr = seq_vit rest k j in
if pr > 0.0 then begin
let v = px *. pr in
if v > !bestp then begin
bestp := v;
bestch :=
(match (tx, tr) with
| Some tx, Some tr -> Some (tx :: tr)
| _ -> None)
end
end
end
done;
(!bestp, !bestch)
in
Hashtbl.replace smemo (rhs, i, j) v;
v
in
vnt prep.start 0 n
let tree_to_string ?(indent = 0) tree =
let rec go indent tree =
let pad = String.make (2 * indent) ' ' in
match tree with
| Leaf s -> pad ^ s
| Node (nt, children) ->
pad ^ nt ^ "\n"
^ String.concat "\n" (List.map (go (indent + 1)) children)
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 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)
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))
in
let add_o nt i j v =
Hashtbl.replace omemo (nt, i, j) (get_o nt i j +. v)
in
add_o prep.start 0 n 1.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
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
end)
(Option.value ~default:[]
(Hashtbl.find_opt prep.by_lhs b)))
prep.lhs_list
done
done;
let counts = Array.make np 0.0 in
Array.iteri
(fun pi (prod : Grammar.production) ->
let prob = prep.probs.(pi) in
let s = ref 0.0 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)
done
done;
counts.(pi) <- prob *. !s /. ptotal)
prep.prods;
(counts, ptotal)
end
let fit_em prep (corpus : Corpus.t) ~iters ~tol =
let np = Array.length prep.prods in
let last_ll = ref neg_infinity in
let ll = ref 0.0 in
let counts = ref (Array.make np 0.0) in
(try
for _ = 1 to iters do
let acc = Array.make np 0.0 in
ll := 0.0;
List.iter
(fun s ->
let cnt, p = expected_counts prep s.Corpus.tokens in
if p > 0.0 then begin
Array.iteri
(fun i x -> acc.(i) <- acc.(i) +. (s.Corpus.count *. x))
cnt;
ll := !ll +. (s.Corpus.count *. log p)
end)
corpus;
Hashtbl.iter
(fun _ idxs ->
let tot = List.fold_left (fun s i -> s +. acc.(i)) 0.0 idxs in
if tot > 0.0 then
List.iter (fun i -> prep.probs.(i) <- acc.(i) /. tot) idxs)
prep.by_lhs;
counts := acc;
if
!last_ll <> neg_infinity
&& abs_float (!ll -. !last_ll) < tol
then raise Exit;
last_ll := !ll
done
with Exit -> ());
(!counts, !ll)