Specialize scalar helpers without variable capture
This commit is contained in:
5 files changed
+344
-3
No files matched your search
@@ -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)) );
|
||||
]
|
||||
Reference in new issue
Block a user