Memoizing F# expressions with computation expressions
Here is another approach to memoization in F#. The goal is to mark sections of a function for memoization without explicitly listing the values used as cache keys:
let doWork x y =
// ...
let result = memo {
// memoize a computation that
// depends on the values of x and y
return x + y
}
// ...
result + 1
Let’s rewrite the example as follows:
let doWork' x y =
// ...
let f = (fun() ->
// memoize a computation that
// depends on the values of x and y
x + y)
let result = f ()
// ...
result + 1
We wrap the expression in a unit -> 'T lambda and invoke it immediately. Inspecting this code in Reflector reveals that the F# compiler generates a class derived from FSharpFunc<TArg, TResult>:
[Serializable]
internal class f@44 : FSharpFunc<Unit, int>
{
public int x;
public int y;
internal f@44(int x, int y)
{
this.x = x;
this.y = y;
}
public override int Invoke(Unit unitVar0)
{
return (this.x + this.y);
}
}
All the memoization inputs, the captured values, become fields of this class. Its instances already contain the complete set of inputs, so why not use them as cache keys?
The missing pieces are equality and hashing based on those fields: the generated class uses reference identity by default. However, System.Collections.Generic.Dictionary<TKey, TValue> accepts a custom key comparer implementing IEqualityComparer<'T>. We therefore need to create a comparer at runtime for the class generated by F# to represent the function value. We can do this by building expression trees with System.Linq.Expressions and compiling them into delegates.
To have the compiler wrap the expression in (fun() -> …) automatically, we can define a Delay(f: unit -> 'T) method on a computation expression builder. Consider this code:
let result = memo {
return x + y
}
With a Return method that simply returns its input, the compiler translation is effectively:
let result =
memo.Delay(fun() -> x + y)
Here is the signature of ComparerCompiler, a module that compiles a comparer for a type and a specified set of its fields:
module ComparerCompiler
open System.Collections.Generic
open System.Reflection
[<RequiresExplicitTypeArguments>]
val compile: FieldInfo[] -> IEqualityComparer<'T>
And its implementation:
/// Functions for compiling comparers for instances
/// of a given type using a specified set of fields
module ComparerCompiler
open System
open System.Collections.Generic
open System.Linq.Expressions
open System.Reflection
/// Cached reflection metadata
let eqComparerType = typedefof<_ EqualityComparer>
let eqComparerIface = typedefof<_ IEqualityComparer>
let getHashMethod = typeof<obj>.GetMethod "GetHashCode"
let func2xType = typedefof<Func<_,_>>
let func3xType = typedefof<Func<_,_,_>>
/// Compiles Equals and GetHashCode delegates
/// for type t using the supplied fields
let emit (t: Type) (fields: FieldInfo[]) =
// create the delegate parameter expressions
let x = Expression.Parameter(t, "x")
let y = Expression.Parameter(t, "y")
// for each field, build a pair of expressions
// for equality comparison and hash calculation
fields |> Array.map (fun field ->
let typ = field.FieldType // get the field type
let comparer = // get the default comparer
eqComparerType // for the field type
.MakeGenericType([| typ |])
.GetProperty("Default")
.GetValue(null, null)
let equalsMethod = // get the Equals method
eqComparerIface // from the comparer interface
.MakeGenericType([| typ |])
.GetMethod("Equals")
// build the field access expression
let fieldAccess = Expression.Field(x, field)
// compare the two field values
// using the default comparer
Expression.Call(
Expression.Constant(comparer),
equalsMethod, fieldAccess,
Expression.Field(y, field)) :> Expression,
// build a call to fieldValue.GetHashCode()
let hashCall: Expression =
upcast Expression.Call(fieldAccess, getHashMethod)
// add a null check for reference types
if typ.IsValueType then hashCall
else upcast Expression.Condition(
Expression.Equal( // if (fieldValue = null)
fieldAccess, Expression.Constant(null, typ)),
Expression.Constant(0), // then 0
hashCall)) // else fieldValue.GetHashCode()
|> function // check how many expression pairs were produced
| [| |] -> raise (ArgumentOutOfRangeException "fields")
| [| x |] -> x
| list -> // combine the expressions if there is more than one pair
list |> Array.reduce (fun (eq1, hash1) (eq2, hash2) ->
// combine equality checks with short-circuiting &&
upcast Expression.AndAlso(eq1, eq2),
// combine hash codes as (h1 << 5) ^ h2
upcast Expression.ExclusiveOr(
Expression.LeftShift(hash1, Expression.Constant(5)), hash2))
|> fun (eqBody, hashBody) -> // compile the expressions
// construct the delegate types
let eqType = func3xType.MakeGenericType(t, t, typeof<bool>)
let hashType = func2xType.MakeGenericType(t, typeof<int>)
// compile the expression trees into delegates
Expression.Lambda(eqType, eqBody, x, y).Compile(),
Expression.Lambda(hashType, hashBody, x).Compile()
/// Returns a comparer for instances of type 'T using the fields
/// in the fields array. The comparer also implements
/// System.Collections.IEqualityComparer
[<RequiresExplicitTypeArguments>]
let compile<'T> (fields: FieldInfo[]) =
if fields = null then
raise (ArgumentNullException "fields")
// compile the Equals and GetHashCode implementations
let eq, hash = emit typeof<'T> fields
// cast the results to their concrete delegate types
let equality : Func<_,_,_> = downcast eq
let hashCode : Func<_,_> = downcast hash
// return a comparer implemented with an object expression
{ new IEqualityComparer<'T> with
member __.Equals(x, y) = equality.Invoke(x, y)
member __.GetHashCode(x) = hashCode.Invoke(x)
// also implement the non-generic interface
interface Collections.IEqualityComparer with
member __.Equals(x, y) =
match x, y with
_ when obj.ReferenceEquals(x, y) -> true
| null, _ | _, null -> false
| (:? 'T as x),(:? 'T as y) -> equality.Invoke(x,y)
| _ -> raise (ArgumentException "invalid type")
member __.GetHashCode(x) =
match x with
null -> 0
| :? 'T as x -> hashCode.Invoke(x)
| _ -> raise (ArgumentException "invalid type") }
The memoization module has this signature:
module MemoBuilder
type MemoBuilder<'T> =
new: unit -> MemoBuilder<'T>
member inline Return: 'T -> 'T
member Delay: (unit -> 'T) -> 'T
val inline memo<'a> : MemoBuilder<'a>
And this implementation:
module MemoBuilder
open System
open System.Collections.Generic
let PrivateStatic =
Reflection.BindingFlags.NonPublic ||| Reflection.BindingFlags.Static
type MemoBuilder<'T>() =
// cache memoization functions by the runtime type of f
[<ThreadStatic>][<DefaultValue>]
static val mutable private cache: Dictionary<Type, (unit -> 'T) -> 'T>
// initialize the thread-local cache on first access
member __.FuncCache =
if MemoBuilder<'T>.cache = null then
MemoBuilder<'T>.cache <- Dictionary()
MemoBuilder<'T>.cache
// cache reflection metadata
static let selfType = typeof<MemoBuilder<'T>>
static let cacher = selfType.GetMethod("Cache", PrivateStatic)
// returning a value from memo { }
// requires no additional work
member inline __.Return(x: 'T) = x
// handle the delayed computation
member __.Delay(f: unit -> 'T) =
let typ = f.GetType() // get the runtime type of the function
match __.FuncCache.TryGetValue typ with
| true, memo -> memo f
| _ -> // call the memoizer factory with the function type
// as its generic type argument
let memo = downcast cacher.MakeGenericMethod(typ)
.Invoke(null, null)
__.FuncCache.Add(typ, memo) // cache the memoizer for this type
memo f // pass the function to the memoizer
// create a memoizer for function type 'F
static member Cache<'F when 'F :> FSharpFunc<unit, 'T>>() =
let t = typeof<'F> // get the function type
// select the fields containing captured values,
// excluding the builder itself
let fields = t.GetFields()
|> Array.filter (fun fi -> fi.FieldType <> selfType)
// compile a comparer for the closures
let comparer = ComparerCompiler.compile<'F> fields
// create a cache using that comparer
let cache = Dictionary<'F, 'T>(comparer)
// return the memoization function
fun (f: FSharpFunc<_,_>) ->
match cache.TryGetValue (downcast f) with
| true, result -> result
| _ -> let result = f.Invoke()
cache.Add(downcast f, result)
result
/// A builder for memoized expressions
let inline memo<'a> = MemoBuilder<'a>()
Finally, here is an example of using the memoizer:
open MemoBuilder
let func x y z =
printfn "func %d %d %d ->" x y z
let a = memo {
printfn " eval a = %d + %d" x y
// an expensive computation
// that captures x and y
return x + y
}
let b = memo {
printfn " eval b = %d + %d" y z
// an expensive computation
// that captures y and z
return y + z
}
let c = memo {
printfn " eval c = %d + %d" a b
// an expensive computation
// that captures a and b
return a + b
}
printfn " return %d\n" c
func 1 2 3
func 4 2 3
func 1 2 4
func 2 1 5
func 2 1 5
func 2 1 5
Output:
func 1 2 3 ->
eval a = 1 + 2
eval b = 2 + 3
eval c = 3 + 5
return 8
func 4 2 3 ->
eval a = 4 + 2
eval c = 6 + 5
return 11
func 1 2 4 ->
eval b = 2 + 4
eval c = 3 + 6
return 9
func 2 1 5 ->
eval a = 2 + 1
eval b = 1 + 5
return 9
func 2 1 5 ->
return 9
func 2 1 5 ->
return 9
This is a proof of concept. I would not recommend using this runtime code generation approach in production projects.