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)