initial
This commit is contained in:
commit
3170b7b398
96 files changed
+39927
No files matched your search
+522
@@ -0,0 +1,522 @@
|
||||
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 seen = Hashtbl.create 64 in
|
||||
let order = ref [] in
|
||||
let q = Queue.create () in
|
||||
Queue.add g.start q;
|
||||
Hashtbl.replace seen g.start ();
|
||||
while not (Queue.is_empty q) do
|
||||
let a = Queue.pop q in
|
||||
order := a :: !order;
|
||||
List.iter
|
||||
(fun p ->
|
||||
if String.equal p.lhs a then
|
||||
List.iter
|
||||
(function
|
||||
| Nonterm x ->
|
||||
if not (Hashtbl.mem seen x) then begin
|
||||
Hashtbl.replace seen x ();
|
||||
Queue.add x q
|
||||
end
|
||||
| Term _ -> ())
|
||||
p.rhs)
|
||||
g.productions
|
||||
done;
|
||||
let order = List.rev !order in
|
||||
let idx = Hashtbl.create 64 in
|
||||
List.iteri (fun i a -> Hashtbl.replace idx a i) order;
|
||||
let name a =
|
||||
match Hashtbl.find_opt idx a with
|
||||
| Some i -> Printf.sprintf "N%d" i
|
||||
| None -> a
|
||||
in
|
||||
let sym = function Term s -> "t:" ^ s | Nonterm s -> name s in
|
||||
let pstr p =
|
||||
Printf.sprintf "%s -> %s" (name p.lhs)
|
||||
(String.concat " " (List.map sym p.rhs))
|
||||
in
|
||||
g.productions |> List.map pstr |> List.sort String.compare
|
||||
|> String.concat "\n"
|
||||
|
||||
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)
|
||||
Reference in new issue
Block a user