Specialize scalar helpers without variable capture

This commit is contained in:
milner committed 2017-02-25 12:48:00 +00:00
1 parent 35d9cebcc3
commit c0dc550578
5 files changed
+344 -3

No files matched your search

+5 -3
View File
@@ -91,11 +91,13 @@ let resolve_source path = Resolve.program (parse_source path)
let infer_source path = Infer.program (resolve_source path)
let specialize_source path = Specialize.program (infer_source path)
let frontend_unavailable () =
Diagnostic.error Location.none "the delta pipeline beyond parsing is not implemented in this revision"
let check path =
ignore (infer_source path)
ignore (specialize_source path)
let dump stage path =
let program = infer_source path in
@@ -104,12 +106,12 @@ let dump stage path =
| _ -> frontend_unavailable ()
let emit path output =
ignore (infer_source path);
ignore (specialize_source path);
ignore output;
frontend_unavailable ()
let build path output =
ignore (infer_source path);
ignore (specialize_source path);
ignore output;
frontend_unavailable ()
+222
View File
@@ -0,0 +1,222 @@
type context = {
helpers : (int * Typed.expr) list;
node_limit : int;
mutable nodes : int;
mutable error : (Location.span * string) option;
}
let make_context helpers node_limit = { helpers = helpers; node_limit = node_limit; nodes = 0; error = None }
let count ctx span =
ctx.nodes <- ctx.nodes + 1;
if ctx.nodes > ctx.node_limit && ctx.error = None then
ctx.error <-
Some
( span,
"specializing the helper functions expanded this query beyond the supported size" )
let ident table stamp =
let rec search = function
| [] -> None
| (candidate, replacement) :: rest -> if candidate = stamp then Some replacement else search rest
in
search table
let rec subst table expr =
match expr.Typed.te with
| Typed.TVar ident' -> (
match ident table (Ident.stamp ident') with
| Some replacement -> replacement
| None -> expr)
| Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TSource _ -> expr
| Typed.TLet (ident', bound, body) ->
let table = List.filter (fun (stamp, _) -> stamp <> Ident.stamp ident') table in
{ expr with Typed.te = Typed.TLet (ident', subst table bound, subst table body) }
| Typed.TLambda (ident', body) ->
let table = List.filter (fun (stamp, _) -> stamp <> Ident.stamp ident') table in
{ expr with Typed.te = Typed.TLambda (ident', subst table body) }
| Typed.TApp (fn, argument) ->
{ expr with Typed.te = Typed.TApp (subst table fn, subst table argument) }
| Typed.TIf (condition, then_branch, else_branch) ->
{
expr with
Typed.te = Typed.TIf (subst table condition, subst table then_branch, subst table else_branch);
}
| Typed.TBinop (operator, left, right) ->
{ expr with Typed.te = Typed.TBinop (operator, subst table left, subst table right) }
| Typed.TTuple items -> { expr with Typed.te = Typed.TTuple (List.map (subst table) items) }
| Typed.TRecord (name, fields) ->
{
expr with
Typed.te = Typed.TRecord (name, List.map (fun (label, value) -> (label, subst table value)) fields);
}
| Typed.TField (record, label) -> { expr with Typed.te = Typed.TField (subst table record, label) }
| Typed.TFilter (collection, predicate) ->
{ expr with Typed.te = Typed.TFilter (subst table collection, subst table predicate) }
| Typed.TMap (collection, projection) ->
{ expr with Typed.te = Typed.TMap (subst table collection, subst table projection) }
| Typed.TSum collection -> { expr with Typed.te = Typed.TSum (subst table collection) }
| Typed.TCount collection -> { expr with Typed.te = Typed.TCount (subst table collection) }
let rec freshen table expr =
match expr.Typed.te with
| Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TSource _ -> expr
| Typed.TVar ident' -> (
match ident table (Ident.stamp ident') with
| Some replacement -> { expr with Typed.te = Typed.TVar replacement }
| None -> expr)
| Typed.TLet (ident', bound, body) ->
let replacement = Ident.fresh (Ident.name ident') (Ident.span ident') in
{
expr with
Typed.te =
Typed.TLet
( replacement,
freshen table bound,
freshen ((Ident.stamp ident', replacement) :: table) body );
}
| Typed.TLambda (ident', body) ->
let replacement = Ident.fresh (Ident.name ident') (Ident.span ident') in
{
expr with
Typed.te = Typed.TLambda (replacement, freshen ((Ident.stamp ident', replacement) :: table) body);
}
| Typed.TApp (fn, argument) -> { expr with Typed.te = Typed.TApp (freshen table fn, freshen table argument) }
| Typed.TIf (condition, then_branch, else_branch) ->
{
expr with
Typed.te =
Typed.TIf (freshen table condition, freshen table then_branch, freshen table else_branch);
}
| Typed.TBinop (operator, left, right) ->
{ expr with Typed.te = Typed.TBinop (operator, freshen table left, freshen table right) }
| Typed.TTuple items -> { expr with Typed.te = Typed.TTuple (List.map (freshen table) items) }
| Typed.TRecord (name, fields) ->
{
expr with
Typed.te =
Typed.TRecord (name, List.map (fun (label, value) -> (label, freshen table value)) fields);
}
| Typed.TField (record, label) -> { expr with Typed.te = Typed.TField (freshen table record, label) }
| Typed.TFilter (collection, predicate) ->
{ expr with Typed.te = Typed.TFilter (freshen table collection, freshen table predicate) }
| Typed.TMap (collection, projection) ->
{ expr with Typed.te = Typed.TMap (freshen table collection, freshen table projection) }
| Typed.TSum collection -> { expr with Typed.te = Typed.TSum (freshen table collection) }
| Typed.TCount collection -> { expr with Typed.te = Typed.TCount (freshen table collection) }
let helper_of ctx ident =
let rec search = function
| [] -> None
| (stamp, body) :: rest -> if stamp = Ident.stamp ident then Some body else search rest
in
search ctx.helpers
let is_collection expr = Types.is_collection expr.Typed.ty
let rec specialize ctx expr =
count ctx expr.Typed.tspan;
match expr.Typed.te with
| Typed.TVar ident' -> (
match helper_of ctx ident' with
| Some body -> specialize ctx (freshen [] body)
| None -> expr)
| Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TSource _ -> expr
| Typed.TLambda (ident', body) -> { expr with Typed.te = Typed.TLambda (ident', specialize ctx body) }
| Typed.TApp (fn, argument) ->
let fn = specialize ctx fn in
let argument = specialize ctx argument in
(match fn.Typed.te with
| Typed.TLambda (ident', body) ->
specialize ctx (subst [ (Ident.stamp ident', argument) ] body)
| _ -> { expr with Typed.te = Typed.TApp (fn, argument) })
| Typed.TLet (ident', bound, body) ->
let bound = specialize ctx bound in
if is_collection bound then { expr with Typed.te = Typed.TLet (ident', bound, specialize ctx body) }
else specialize ctx (subst [ (Ident.stamp ident', bound) ] body)
| Typed.TIf (condition, then_branch, else_branch) ->
{
expr with
Typed.te =
Typed.TIf (specialize ctx condition, specialize ctx then_branch, specialize ctx else_branch);
}
| Typed.TBinop (operator, left, right) ->
{ expr with Typed.te = Typed.TBinop (operator, specialize ctx left, specialize ctx right) }
| Typed.TTuple items -> { expr with Typed.te = Typed.TTuple (List.map (specialize ctx) items) }
| Typed.TRecord (name, fields) ->
{
expr with
Typed.te =
Typed.TRecord (name, List.map (fun (label, value) -> (label, specialize ctx value)) fields);
}
| Typed.TField (record, label) -> { expr with Typed.te = Typed.TField (specialize ctx record, label) }
| Typed.TFilter (collection, predicate) ->
let collection = specialize ctx collection in
let predicate = specialize ctx predicate in
check_function ctx predicate "predicate";
{ expr with Typed.te = Typed.TFilter (collection, predicate) }
| Typed.TMap (collection, projection) ->
let collection = specialize ctx collection in
let projection = specialize ctx projection in
check_function ctx projection "mapping function";
{ expr with Typed.te = Typed.TMap (collection, projection) }
| Typed.TSum collection -> { expr with Typed.te = Typed.TSum (specialize ctx collection) }
| Typed.TCount collection -> { expr with Typed.te = Typed.TCount (specialize ctx collection) }
and check_function ctx expr what =
match expr.Typed.te with
| Typed.TLambda (parameter, body) ->
let free = free_variables [ Ident.stamp parameter ] body [] in
(match free with
| [] -> ()
| (ident, span) :: _ ->
Diagnostic.error span
"the %s cannot be incrementalized because it uses `%s` from the surrounding scope" what
(Ident.display ident))
| _ ->
Diagnostic.error expr.Typed.tspan
"the %s must be a function literal or a helper that specializes to one, but this expression is a function value"
what
and free_variables bound expr acc =
match expr.Typed.te with
| Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit -> acc
| Typed.TSource ident -> (ident, expr.Typed.tspan) :: acc
| Typed.TVar ident ->
if List.exists (fun stamp -> stamp = Ident.stamp ident) bound then acc
else (ident, expr.Typed.tspan) :: acc
| Typed.TLet (ident, bound_expr, body) ->
let acc = free_variables bound bound_expr acc in
free_variables (Ident.stamp ident :: bound) body acc
| Typed.TLambda (ident, body) -> free_variables (Ident.stamp ident :: bound) body acc
| Typed.TApp (fn, argument) -> free_variables bound argument (free_variables bound fn acc)
| Typed.TIf (condition, then_branch, else_branch) ->
free_variables bound else_branch
(free_variables bound then_branch (free_variables bound condition acc))
| Typed.TBinop (_, left, right) -> free_variables bound right (free_variables bound left acc)
| Typed.TTuple items -> List.fold_left (fun acc item -> free_variables bound item acc) acc items
| Typed.TRecord (_, fields) ->
List.fold_left (fun acc (_, value) -> free_variables bound value acc) acc fields
| Typed.TField (record, _) -> free_variables bound record acc
| Typed.TFilter (collection, predicate) ->
free_variables bound predicate (free_variables bound collection acc)
| Typed.TMap (collection, projection) ->
free_variables bound projection (free_variables bound collection acc)
| Typed.TSum collection | Typed.TCount collection -> free_variables bound collection acc
let program typed =
let ctx =
make_context
(List.map (fun helper -> (Ident.stamp helper.Typed.th_ident, helper.Typed.th_body)) typed.Typed.tp_helpers)
20000
in
let helpers =
List.map
(fun helper -> { helper with Typed.th_body = specialize ctx helper.Typed.th_body })
typed.Typed.tp_helpers
in
let query_body = specialize ctx typed.Typed.tp_query_body in
(match ctx.error with
| Some (span, message) -> Diagnostic.error span "%s" message
| None -> ());
{ typed with Typed.tp_helpers = helpers; tp_query_body = query_body }
+1
View File
@@ -0,0 +1 @@
val program : Typed.program -> Typed.program
+1
View File
@@ -374,6 +374,7 @@ let () =
Test_harness.run_suite "program" program_cases;
Test_harness.run_suite "resolve" resolve_cases;
Test_harness.run_suite "types" Test_type.cases;
Test_harness.run_suite "specialize" Test_type.specialize_cases;
Test_harness.run_suite "interpret" Test_incremental.cases;
Printf.printf "%d cases, %d failures\n" (Test_harness.case_count ()) (Test_harness.failure_count ());
exit (if Test_harness.failure_count () = 0 then 0 else 1)
+115
View File
@@ -209,3 +209,118 @@ let cases =
check_equal_string "query type" "collection int"
(Types.pp program.Typed.tp_query_body.Typed.ty) );
]
let rec binder_stamps expr acc =
let walk = binder_stamps in
match expr.Typed.te with
| Typed.TLet (ident, bound, body) -> walk body (walk bound (Ident.stamp ident :: acc))
| Typed.TLambda (ident, body) -> walk body (Ident.stamp ident :: acc)
| Typed.TApp (fn, argument) -> walk argument (walk fn acc)
| Typed.TIf (a, b, c) -> walk c (walk b (walk a acc))
| Typed.TBinop (_, a, b) -> walk b (walk a acc)
| Typed.TTuple items -> List.fold_left (fun acc item -> walk item acc) acc items
| Typed.TRecord (_, fields) -> List.fold_left (fun acc (_, value) -> walk value acc) acc fields
| Typed.TField (record, _) -> walk record acc
| Typed.TFilter (collection, predicate) -> walk predicate (walk collection acc)
| Typed.TMap (collection, projection) -> walk projection (walk collection acc)
| Typed.TSum collection | Typed.TCount collection -> walk collection acc
| Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TVar _ | Typed.TSource _ -> acc
let stamp_unique stamps =
let seen = Hashtbl.create 32 in
let rec check = function
| [] -> true
| stamp :: rest -> if Hashtbl.mem seen stamp then false else (Hashtbl.add seen stamp (); check rest)
in
check stamps
let rec collect_map_functions expr acc =
match expr.Typed.te with
| Typed.TMap (collection, projection) ->
collect_map_functions collection (("map", projection) :: acc)
| Typed.TFilter (collection, predicate) ->
collect_map_functions collection (("filter", predicate) :: acc)
| Typed.TSum collection | Typed.TCount collection -> collect_map_functions collection acc
| Typed.TLet (_, bound, body) -> collect_map_functions body (collect_map_functions bound acc)
| _ -> acc
let function_is_lambda expr =
match expr.Typed.te with Typed.TLambda _ -> true | _ -> false
let rec has_let expr =
match expr.Typed.te with
| Typed.TLet (_, _, _) -> true
| Typed.TLambda (_, body) -> has_let body
| Typed.TApp (fn, argument) -> has_let fn || has_let argument
| Typed.TIf (a, b, c) -> has_let a || has_let b || has_let c
| Typed.TBinop (_, a, b) -> has_let a || has_let b
| Typed.TTuple items -> List.exists has_let items
| Typed.TRecord (_, fields) -> List.exists (fun (_, value) -> has_let value) fields
| Typed.TField (record, _) -> has_let record
| Typed.TFilter (collection, predicate) -> has_let collection || has_let predicate
| Typed.TMap (collection, projection) -> has_let collection || has_let projection
| Typed.TSum collection | Typed.TCount collection -> has_let collection
| Typed.TInt _ | Typed.TBool _ | Typed.TString _ | Typed.TUnit | Typed.TVar _ | Typed.TSource _ -> false
let specialize text = Specialize.program (infer text)
let specialize_cases =
[
( "inlining a helper twice does not duplicate binder stamps",
fun () ->
Ident.reset ();
let program =
specialize
"input rows : collection int\nlet scale n = n * 2\nlet quad n = scale (scale n)\nquery q = rows |> map (fun r -> quad r + quad r)\n"
in
check "all binders are distinct"
(stamp_unique (binder_stamps program.Typed.tp_query_body [])) );
( "a helper used as a mapping function specializes to a lambda",
fun () ->
let program =
specialize
"type order = { total : int }\ninput rows : collection order\nlet value o = o.total\nquery q = rows |> map value\n"
in
let functions = collect_map_functions program.Typed.tp_query_body [] in
check "map function is a lambda"
(List.for_all (fun (_, fn) -> function_is_lambda fn) functions) );
( "a helper used as a predicate specializes to a lambda",
fun () ->
let program =
specialize
"type order = { total : int }\ninput rows : collection order\nlet big o = o.total > 1000\nquery q = rows |> filter big |> count\n"
in
let functions = collect_map_functions program.Typed.tp_query_body [] in
check_equal_int "one collection function" 1 (List.length functions);
check "filter predicate is a lambda"
(List.for_all (fun (_, fn) -> function_is_lambda fn) functions) );
( "higher order helpers specialize at their call sites",
fun () ->
let program =
specialize
"input rows : collection int\nlet twice f x = f (f x)\nlet inc n = n + 1\nquery q = rows |> map (fun r -> twice inc r) |> sum\n"
in
check_equal_int "no helper references remain" 0
(List.length (collect_map_functions program.Typed.tp_query_body []) - 1) );
( "an escaping function value is rejected",
fun () ->
check_message "escaping function"
"the mapping function must be a function literal or a helper that specializes to one, but this expression is a function value"
(fun () ->
specialize
"input rows : collection int\nlet choose b = if b then (fun n -> n) else (fun n -> n + 1)\nquery q = rows |> map (choose true)\n") );
( "a predicate that uses the input collection is rejected",
fun () ->
check_message "captured collection"
"the predicate cannot be incrementalized because it uses `rows` from the surrounding scope"
(fun () ->
specialize
"input rows : collection int\nquery q = rows |> filter (fun r -> r > count rows) |> count\n") );
( "plain let bindings are inlined into the query",
fun () ->
let program =
specialize
"input rows : collection int\nquery q = let base = 10 in rows |> map (fun r -> r + base) |> sum\n"
in
check "no let bindings survive" (not (has_let program.Typed.tp_query_body)) );
]