type symbol = Term of string | Nonterm of string type production = { lhs : string; rhs : symbol list; count : float; } type t = { start : string; productions : production list; } module SSet = Set.Make (String) let symbol_to_string = function Term s -> s | Nonterm s -> s let symbol_key = function Term s -> "t:" ^ s | Nonterm s -> "n:" ^ s let symbol_is_nonterm = function Nonterm _ -> true | Term _ -> false let equal_symbol a b = match (a, b) with | Term x, Term y | Nonterm x, Nonterm y -> String.equal x y | _ -> false let equal_rhs a b = List.length a = List.length b && List.for_all2 equal_symbol a b let rhs_to_string rhs = String.concat " " (List.map symbol_to_string rhs) let production_to_string p = Printf.sprintf "%s -> %s" p.lhs (rhs_to_string p.rhs) let make ~start productions = { start; productions } let nonterminals g = let acc = ref (SSet.singleton g.start) in List.iter (fun p -> acc := SSet.add p.lhs !acc; List.iter (function Nonterm x -> acc := SSet.add x !acc | Term _ -> ()) p.rhs) g.productions; SSet.elements !acc let terminals g = let acc = ref SSet.empty in List.iter (fun p -> List.iter (function Term x -> acc := SSet.add x !acc | Nonterm _ -> ()) p.rhs) g.productions; SSet.elements !acc let productions_of g lhs = List.filter (fun p -> String.equal p.lhs lhs) g.productions let total_count prods = List.fold_left (fun s p -> s +. p.count) 0.0 prods let probabilities g = let by_lhs = Hashtbl.create 64 in List.iter (fun p -> let l = match Hashtbl.find_opt by_lhs p.lhs with Some l -> l | None -> [] in Hashtbl.replace by_lhs p.lhs (p :: l)) g.productions; let tbl = Hashtbl.create 64 in Hashtbl.iter (fun lhs ps -> let ps = List.rev ps in let tot = total_count ps in let k = List.length ps in let probs = if tot > 0.0 then List.map (fun p -> (p, p.count /. tot)) ps else List.map (fun p -> (p, 1.0 /. float_of_int k)) ps in Hashtbl.replace tbl lhs probs) by_lhs; tbl let canonical_string g = let nts = nonterminals g in let signature colors nt = let symbol = function | Term value -> Printf.sprintf "t%d:%s" (String.length value) value | Nonterm value -> Printf.sprintf "n%d" (Hashtbl.find colors value) in let productions = productions_of g nt |> List.map (fun p -> "[" ^ String.concat ";" (List.map symbol p.rhs) ^ "]") |> List.sort String.compare in Printf.sprintf "%d|%b|%s" (Hashtbl.find colors nt) (String.equal nt g.start) (String.concat "|" productions) in let refine colors = let rec loop colors = let signatures = List.map (fun nt -> (nt, signature colors nt)) nts in let unique = signatures |> List.map snd |> List.sort_uniq String.compare in let ids = Hashtbl.create (List.length unique) in List.iteri (fun index key -> Hashtbl.replace ids key index) unique; let next = Hashtbl.create (List.length nts) in List.iter (fun (nt, key) -> Hashtbl.replace next nt (Hashtbl.find ids key)) signatures; if List.for_all (fun nt -> Hashtbl.find colors nt = Hashtbl.find next nt) nts then next else loop next in loop colors in let encode colors = let order = List.sort (fun left right -> compare (Hashtbl.find colors left) (Hashtbl.find colors right)) nts in let names = Hashtbl.create (List.length order) in List.iteri (fun index nt -> Hashtbl.replace names nt (Printf.sprintf "N%d" index)) order; let symbol = function | Term value -> Printf.sprintf "t%d:%s" (String.length value) value | Nonterm value -> Hashtbl.find names value in g.productions |> List.map (fun p -> Printf.sprintf "%s -> %s" (Hashtbl.find names p.lhs) (String.concat " " (List.map symbol p.rhs))) |> List.sort String.compare |> String.concat "\n" in let rec canonical colors = let colors = refine colors in let classes = Hashtbl.create (List.length nts) in List.iter (fun nt -> let color = Hashtbl.find colors nt in let members = Option.value ~default:[] (Hashtbl.find_opt classes color) in Hashtbl.replace classes color (nt :: members)) nts; let ambiguous = Hashtbl.fold (fun color members acc -> if List.length members > 1 then (color, members) :: acc else acc) classes [] |> List.sort (fun (left, _) (right, _) -> compare left right) in match ambiguous with | [] -> encode colors | (_, members) :: _ -> let fresh = List.fold_left (fun highest nt -> max highest (Hashtbl.find colors nt)) (-1) nts + 1 in members |> List.map (fun chosen -> let branch = Hashtbl.copy colors in Hashtbl.replace branch chosen fresh; canonical branch) |> List.sort String.compare |> List.hd in let colors = Hashtbl.create (List.length nts) in List.iter (fun nt -> Hashtbl.replace colors nt (if String.equal nt g.start then 0 else 1)) nts; canonical colors let equal_structure a b = String.equal (canonical_string a) (canonical_string b) let copy_with g productions = { g with productions } let prods_of_list prods lhs = List.filter (fun p -> String.equal p.lhs lhs) prods let dedupe prods = let tbl = Hashtbl.create 64 in List.iter (fun p -> match Hashtbl.find_opt tbl (p.lhs, p.rhs) with | Some c -> Hashtbl.replace tbl (p.lhs, p.rhs) (c +. p.count) | None -> Hashtbl.replace tbl (p.lhs, p.rhs) p.count) prods; Hashtbl.fold (fun (lhs, rhs) count acc -> { lhs; rhs; count } :: acc) tbl [] let remove_unit_cycles prods = let unit_edges = Hashtbl.create 64 in List.iter (fun (p : production) -> match p.rhs with | [ Nonterm b ] -> let l = match Hashtbl.find_opt unit_edges p.lhs with | Some l -> l | None -> [] in Hashtbl.replace unit_edges p.lhs (b :: l) | _ -> ()) prods; let reaches a target = let seen = Hashtbl.create 16 in let rec go x = if String.equal x target then true else if Hashtbl.mem seen x then false else begin Hashtbl.replace seen x (); List.exists go (Option.value ~default:[] (Hashtbl.find_opt unit_edges x)) end in go a in List.filter (fun (p : production) -> match p.rhs with | [ Nonterm b ] -> not (String.equal p.lhs b) && not (reaches b p.lhs) | _ -> true) prods let productive_set prods = let tbl = Hashtbl.create 64 in let changed = ref true in while !changed do changed := false; List.iter (fun p -> if not (Hashtbl.mem tbl p.lhs) then begin let ok = List.for_all (function Term _ -> true | Nonterm x -> Hashtbl.mem tbl x) p.rhs in if ok then begin Hashtbl.replace tbl p.lhs (); changed := true end end) prods done; tbl let remove_nonproductive prods = let prod = productive_set prods in List.filter (fun p -> Hashtbl.mem prod p.lhs && List.for_all (function Term _ -> true | Nonterm x -> Hashtbl.mem prod x) p.rhs) prods let remove_unreachable start prods = let tbl = Hashtbl.create 64 in Hashtbl.replace tbl start (); let changed = ref true in while !changed do changed := false; List.iter (fun p -> if Hashtbl.mem tbl p.lhs then List.iter (function | Nonterm x -> if not (Hashtbl.mem tbl x) then begin Hashtbl.replace tbl x (); changed := true end | Term _ -> ()) p.rhs) prods done; List.filter (fun p -> Hashtbl.mem tbl p.lhs) prods let derives_form g start target = let n = List.length target in let arr = Array.of_list target in let key = symbol_key in let d = Array.make_matrix (n + 1) (n + 1) SSet.empty in for i = 0 to n - 1 do d.(i).(i + 1) <- SSet.singleton (key arr.(i)) done; let rhs_matches rhs i j = let rec go rhs p = match rhs with | [] -> p = j | [ last ] -> p < j && SSet.mem (key last) d.(p).(j) | s :: rest -> let rem = List.length rest in let rec try_q q = if q > j - rem then false else if SSet.mem (key s) d.(p).(q) && go rest q then true else try_q (q + 1) in try_q (p + 1) in List.length rhs > 0 && go rhs i in for len = 1 to n do for i = 0 to n - len do let j = i + len in let changed = ref true in while !changed do changed := false; List.iter (fun (p : production) -> let lhs_key = key (Nonterm p.lhs) in if not (SSet.mem lhs_key d.(i).(j)) then begin let ok = match p.rhs with | [ y ] -> SSet.mem (key y) d.(i).(j) | _ -> rhs_matches p.rhs i j in if ok then begin d.(i).(j) <- SSet.add lhs_key d.(i).(j); changed := true end end) g.productions done done done; SSet.mem (key (Nonterm start)) d.(0).(n) let prune_redundant g = let prods = ref (dedupe g.productions) in let changed = ref true in while !changed do changed := false; let victim = ref None in (try List.iter (fun (p : production) -> let g' = { g with productions = List.filter (fun q -> not (q == p)) !prods } in if derives_form g' p.lhs p.rhs then begin victim := Some p; raise Exit end) !prods with Exit -> ()); match !victim with | Some p -> prods := List.filter (fun q -> not (q == p)) !prods; changed := true | None -> () done; { g with productions = !prods } let normalize g = let step prods = prods |> dedupe |> remove_unit_cycles |> remove_nonproductive |> remove_unreachable g.start in let p1 = step g.productions in let p2 = step p1 in let g1 = { g with productions = p2 } in let g2 = prune_redundant g1 in let g3 = step g2.productions in { g with productions = g3 } let language_up_to g max_len = let tbl = Hashtbl.create 256 in let get x len = Option.value ~default:[] (Hashtbl.find_opt tbl (x, len)) in for len = 1 to max_len do let changed = ref true in while !changed do changed := false; List.iter (fun x -> let acc = ref (get x len) in List.iter (fun (p : production) -> let rhs = p.rhs in let m = List.length rhs in let rec gen rhs l = match rhs with | [] -> if l = 0 then [ [] ] else [] | s :: rest -> let out = ref [] in for k = 1 to l - (List.length rest) do let ss = match s with | Term a -> if k = 1 then [ [ a ] ] else [] | Nonterm y -> get y k in let rs = gen rest (l - k) in List.iter (fun a -> List.iter (fun b -> out := (a @ b) :: !out) rs) ss done; !out in if len >= m then acc := gen rhs len @ !acc) (productions_of g x); let acc = List.sort_uniq compare !acc in if List.length acc > List.length (get x len) then begin Hashtbl.replace tbl (x, len) acc; changed := true end) (nonterminals g) done done; let all = List.concat (List.init max_len (fun i -> get g.start (i + 1))) in List.sort_uniq compare all let initial_grammar ?(start = "S") samples = let term_nts = Hashtbl.create 64 in let prods = ref [] in List.iter (fun (tokens, count) -> let rhs = List.map (fun tok -> match Hashtbl.find_opt term_nts tok with | Some nt -> Nonterm nt | None -> let nt = "T_" ^ tok in Hashtbl.replace term_nts tok nt; prods := { lhs = nt; rhs = [ Term tok ]; count = 1.0 } :: !prods; Nonterm nt) tokens in prods := { lhs = start; rhs; count } :: !prods) samples; make ~start (List.rev !prods) let is_recursive g = let deps = Hashtbl.create 64 in List.iter (fun x -> let l = List.concat_map (fun p -> List.filter_map (function Nonterm y -> Some y | Term _ -> None) p.rhs) (productions_of g x) in Hashtbl.replace deps x l) (nonterminals g); List.exists (fun a -> let seen = Hashtbl.create 64 in let rec go x = List.exists (fun y -> if String.equal y a then true else if Hashtbl.mem seen y then false else begin Hashtbl.replace seen y (); go y end) (Option.value ~default:[] (Hashtbl.find_opt deps x)) in go a) (nonterminals g) let to_string ?(decimals = 3) g = let probs = probabilities g in let fmt p = let pr = match Hashtbl.find_opt probs p.lhs with | None -> 0.0 | Some l -> ( match List.find_opt (fun (q, _) -> q == p) l with | Some (_, x) -> x | None -> 0.0) in Printf.sprintf "%-8s -> %-24s [%.*f]" p.lhs (rhs_to_string p.rhs) decimals pr in let header = Printf.sprintf "# start = %s, %d nonterminals, %d terminals, %d productions" g.start (List.length (nonterminals g)) (List.length (terminals g)) (List.length g.productions) in String.concat "\n" (header :: List.map fmt g.productions) let is_nonterminal_name s = String.length s > 0 && let c = s.[0] in (c >= 'A' && c <= 'Z') || c = '_' || c = '<' let find_sub s sub = let n = String.length s and m = String.length sub in let rec go i = if i + m > n then None else if String.equal (String.sub s i m) sub then Some i else go (i + 1) in go 0 let parse_start line = match find_sub line "start =" with | None -> None | Some i -> let rest = String.sub line (i + 7) (String.length line - i - 7) |> String.trim in let stop = match (String.index_opt rest ' ', String.index_opt rest ',') with | Some a, Some b -> min a b | Some a, None | None, Some a -> a | None, None -> String.length rest in let name = String.sub rest 0 stop in if String.length name = 0 then None else Some name let of_string text = let lines = String.split_on_char '\n' text in let prods = ref [] in let start = ref None in List.iter (fun raw -> let trimmed = String.trim raw in if String.length trimmed > 0 && trimmed.[0] = '#' then match parse_start trimmed with | Some s -> start := Some s | None -> () else begin let line = match String.index_opt raw '#' with | Some i -> String.sub raw 0 i | None -> raw in let toks = String.split_on_char ' ' (String.trim line) |> List.filter (fun s -> String.length s > 0) in match toks with | [] -> () | lhs :: arrow :: rhs_toks when String.equal arrow "->" -> let rhs_toks, prob = match List.rev rhs_toks with | last :: rest_rev when String.length last > 2 && last.[0] = '[' -> ( match float_of_string_opt (String.sub last 1 (String.length last - 2)) with | Some p -> (List.rev rest_rev, Some p) | None -> (rhs_toks, None)) | _ -> (rhs_toks, None) in let rhs = List.map (fun s -> if is_nonterminal_name s then Nonterm s else Term s) rhs_toks in if !start = None then start := Some lhs; prods := { lhs; rhs; count = Option.value ~default:1.0 prob } :: !prods | _ -> () end) lines; let start = Option.value ~default:"S" !start in make ~start (List.rev !prods)