Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion Makefile
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@

.PHONY: build lib doc clean install uninstall test gen gen_ragel gen_metaocaml archive
.PHONY: build lib doc clean install uninstall test benchs gen gen_ragel gen_metaocaml archive

OCAMLBUILD=ocamlbuild -use-ocamlfind -no-links -j 0

Expand Down Expand Up @@ -27,6 +27,9 @@ top:
test:
dune runtest $(DUNEFLAGS)

benchs:
RUN_BENCHS=1 dune exec --profile bench $(DUNEFLAGS) bench/bench_sharded_hash_trie.exe

doc:
dune build $(DUNEFLAGS) @doc

Expand Down
159 changes: 159 additions & 0 deletions bench/bench_sharded_hash_trie.ml
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
(* Compare ExtThread.ShardedHashTrie (64 and 256 shards) with Saturn.Htbl, a flat bucketed assoc list,
Hashtbl + Mutex, and (on one domain) a plain Hashtbl.

dune exec --release bench/bench_sharded_hash_trie.exe *)

open Devkit

type ('k, 'v) ops = { find : 'k -> 'v; get_or_create : 'k -> f:('k -> 'v) -> 'v }
type 'k impl = { name : string; make : unit -> ('k, int Atomic.t) ops }

module Impls (K : Hashtbl.HashedType) = struct
let devkit ?(shard_bits = 6) () = { name = Printf.sprintf "devkit/s%d" shard_bits; make = fun () ->
let t = ExtThread.ShardedHashTrie.create ~shard_bits ~hashable:{ equal = K.equal; hash = K.hash } () in
{ find = ExtThread.ShardedHashTrie.find t; get_or_create = (fun k ~f -> ExtThread.ShardedHashTrie.get_or_create t k ~f) } }

let saturn = { name = "saturn"; make = fun () ->
let t = Saturn.Htbl.create ~hashed_type:(module K) () in
let rec get_or_create k ~f =
match Saturn.Htbl.find_exn t k with
| v -> v
| exception Not_found -> let v = f k in if Saturn.Htbl.try_add t k v then v else get_or_create k ~f
in
{ find = Saturn.Htbl.find_exn t; get_or_create } }

(* fixed buckets of assoc lists, the simplest lock-free option *)
let flat = { name = "flat128"; make = fun () ->
let b = Array.init 128 (fun _ -> Atomic.make []) in
let bucket k = b.(K.hash k land 127) in
let rec assoc k = function [] -> raise_notrace Not_found | (k', v) :: tl -> if K.equal k k' then v else assoc k tl in
let find k = assoc k (Atomic.get (bucket k)) in
let rec get_or_create k ~f =
let b = bucket k in
let l = Atomic.get b in
match assoc k l with
| v -> v
| exception Not_found -> let v = f k in if Atomic.compare_and_set b l ((k, v) :: l) then v else get_or_create k ~f
in
{ find; get_or_create } }

let mutex = { name = "hashtbl+mutex"; make = fun () ->
let module T = Hashtbl.Make (K) in
let t = T.create 16 and m = Mutex.create () in
{ find = (fun k -> Mutex.protect m (fun () -> T.find t k));
get_or_create = (fun k ~f -> Mutex.protect m (fun () ->
match T.find t k with v -> v | exception Not_found -> let v = f k in T.add t k v; v)) } }

(* not domain-safe: baseline for 1-domain runs only *)
let plain = { name = "hashtbl (unsafe)"; make = fun () ->
let module T = Hashtbl.Make (K) in
let t = T.create 16 in
{ find = T.find t;
get_or_create = (fun k ~f -> match T.find t k with v -> v | exception Not_found -> let v = f k in T.add t k v; v) } }

let tries = [ devkit (); devkit ~shard_bits:8 () ]
let all = tries @ [ saturn; flat; mutex ]
let growing = tries @ [ saturn; mutex ] (* flat128 degrades to long lists *)
end

module S = Impls (struct include String let hash = Hashtbl.hash end)
module I = Impls (struct type t = int let equal = Int.equal let hash = Hashtbl.hash end)

let counter _ = Atomic.make 0

let on_domains n f = List.init n (fun d -> Domain.spawn (fun () -> f d)) |> List.iter Domain.join

let report ~title ~ops_per_call samples =
Printf.printf "\n== %s\n" title;
List.iter (fun (name, ts) ->
let rates = List.map (fun (t : Benchmark.t) -> Int64.to_float t.iters *. float ops_per_call /. t.wall /. 1e6) ts in
let rates = List.sort compare rates in
Printf.printf " %-16s %8.1f Mops/s (median of %d)\n%!" name (List.nth rates (List.length rates / 2)) (List.length rates))
samples

(* the unsynchronized Hashtbl only makes sense on one domain *)
let with_plain domains impls plain = if domains = 1 then impls @ [ plain ] else impls

let bench ~title ~ops_per_call impls run =
let samples = Benchmark.throughputN ~style:Benchmark.Nil ~repeat:5 1
(List.map (fun impl -> impl.name, run, impl) impls) in
report ~title ~ops_per_call samples

(* Cache.Count: a few constant keys, hit on every call *)
let count_like domains =
let keys = Array.init 20 (Printf.sprintf "metric_name_%d") in
let per_domain = 1_000_000 in
bench ~title:(Printf.sprintf "count-like: 20 string keys, hits + incr, %d domain(s)" domains)
~ops_per_call:(per_domain * domains) (with_plain domains S.all S.plain)
(fun impl ->
let t = impl.make () in
Array.iter (fun k -> ignore (t.get_or_create k ~f:counter)) keys;
on_domains domains (fun d ->
for i = 0 to per_domain - 1 do
Atomic.incr (t.get_or_create keys.((i + d) mod 20) ~f:counter)
done))

(* growing table: every call is a miss and an insert *)
let inserts domains =
let n = 200_000 in
bench ~title:(Printf.sprintf "inserts: %d fresh int keys, %d domain(s)" n domains)
~ops_per_call:n (with_plain domains I.all I.plain)
(fun impl ->
let t = impl.make () in
on_domains domains (fun d ->
let i = ref d in
while !i < n do ignore (t.get_or_create !i ~f:counter); i := !i + domains done))

(* read-mostly: 10k warm keys, one insert of a new key every [insert_every] ops, finds otherwise.
The table grows by [per_domain / insert_every] keys per domain and call. *)
let mixed ~label ~insert_every domains =
let warm = 10_000 and per_domain = 2_000_000 in
bench ~title:(Printf.sprintf "%s: %d warm int keys, 1 insert per %d ops, %d domain(s)" label warm insert_every domains)
~ops_per_call:(per_domain * domains) (with_plain domains I.growing I.plain)
(fun impl ->
let t = impl.make () in
for k = 0 to warm - 1 do ignore (t.get_or_create k ~f:counter) done;
on_domains domains (fun d ->
let fresh = ref (warm + d) in
for i = 0 to per_domain - 1 do
if i mod insert_every = 0 then (ignore (t.get_or_create !fresh ~f:counter); fresh := !fresh + domains)
else ignore (t.find ((i * 7919) mod warm))
done))

let mixes = [ "mixed90", 10; "mixed99", 100; "mixed99.9", 1000 ]
let all_mixed () = List.iter (fun (label, insert_every) -> List.iter (mixed ~label ~insert_every) [ 1; 4; 8 ]) mixes

(* sanity check every implementation before timing it *)
let self_check () =
List.iter (fun impl ->
let t = impl.make () in
let n = 100_000 in
on_domains 4 (fun d -> let i = ref d in while !i < n do ignore (t.get_or_create !i ~f:(fun k -> Atomic.make k)); i := !i + 4 done);
for k = 0 to n - 1 do if Atomic.get (t.find k) <> k then failwith (impl.name ^ ": self-check failed") done)
I.all

let alloc_per_hit () =
Printf.printf "\n== minor words per hit (20 string keys, 1 domain)\n";
let keys = Array.init 20 (Printf.sprintf "metric_name_%d") in
List.iter (fun impl ->
let t = impl.make () in
Array.iter (fun k -> ignore (t.get_or_create k ~f:counter)) keys;
let n = 1_000_000 in
let before = Gc.minor_words () in
for i = 0 to n - 1 do ignore (t.get_or_create keys.(i mod 20) ~f:counter) done;
Printf.printf " %-16s %6.2f\n%!" impl.name ((Gc.minor_words () -. before) /. float n))
S.all

let () =
self_check ();
match Sys.argv with
| [| _; "mixed" |] -> all_mixed ()
| [| _; "single" |] ->
count_like 1;
inserts 1;
List.iter (fun (label, insert_every) -> mixed ~label ~insert_every 1) mixes
| _ ->
alloc_per_hit ();
List.iter count_like [ 1; 4; 8 ];
List.iter inserts [ 1; 4 ];
all_mixed ()
6 changes: 6 additions & 0 deletions bench/dune
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
; optional: needs saturn and benchmark. Build with make benchs.
(executable
(name bench_sharded_hash_trie)
(optional)
(enabled_if (= %{profile} bench))
(libraries devkit saturn benchmark unix))
4 changes: 3 additions & 1 deletion devkit.opam
Original file line number Diff line number Diff line change
Expand Up @@ -12,9 +12,11 @@ build: [
]
depends: [
"ocaml" {>= "5.0"}
"dune" {>= "2.0"}
"dune" {>= "2.3"}
("extlib" {>= "1.7.1"} | "extlib-compat" {>= "1.7.1"})
"ounit2"
"qcheck-core" {with-test}
"benchmark" {with-dev-setup}
"camlzip"
"libevent" {>= "0.8.0"}
"curl" {>= "0.10.0"}
Expand Down
2 changes: 1 addition & 1 deletion dune-project
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
(lang dune 2.0)
(lang dune 2.3)
(name devkit)
(implicit_transitive_deps false)
117 changes: 117 additions & 0 deletions extThread.ml
Original file line number Diff line number Diff line change
Expand Up @@ -175,3 +175,120 @@ module Pool = struct
end

end

(* Writes copy the path from the shard root to the leaf, then CAS the root. *)
module ShardedHashTrie = struct
type 'k hashable = { equal : 'k -> 'k -> bool; hash : 'k -> int }

(* [compare], not [(=)], as in Hashtbl: [nan] must be equal to itself *)
let poly_hashable = { equal = (fun a b -> compare a b = 0); hash = Hashtbl.hash }

let level_bits = 4
let level_width = 1 lsl level_bits
let level_mask = level_width - 1

(* assoc list (all keys have same hash) *)
type ('k, 'v) bindings =
| Nil
| Cons of { key : 'k; value : 'v; rest : ('k, 'v) bindings }

type ('k, 'v) tree =
| Empty
| Leaf of { h : int; key : 'k; value : 'v; next : ('k, 'v) bindings }
(** [h]: hash bits remaining at this depth; [next]: other keys with the same hash *)
| Node of ('k, 'v) tree array (** [level_width] children *)

type ('k, 'v) t = {
hashable : 'k hashable;
shard_bits : int;
shard_mask : int;
shards : ('k, 'v) tree Atomic.t array;
}

let create ?(hashable = poly_hashable) ?(shard_bits = 6) () =
if shard_bits < 0 || shard_bits > 16 then invalid_arg "ShardedHashTrie.create: shard_bits must be in [0,16]";
{ hashable; shard_bits; shard_mask = 1 lsl shard_bits - 1;
shards = Array.init (1 lsl shard_bits) (fun _ -> Atomic.make Empty) }

let rec assoc equal k = function
| Nil -> raise_notrace Not_found
| Cons { key; value; rest } -> if equal k key then value else assoc equal k rest

let rec find_tree equal h k = function
| Empty -> raise_notrace Not_found
| Leaf { h = h2; key; value; next } ->
if h <> h2 then raise_notrace Not_found
else if equal k key then value
else assoc equal k next
| Node a -> find_tree equal (h lsr level_bits) k a.(h land level_mask)

(* shard and remaining hash computed inline: a helper returning both would allocate a tuple *)
let find_notrace t k =
let h = t.hashable.hash k in
find_tree t.hashable.equal (h lsr t.shard_bits) k (Atomic.get t.shards.(h land t.shard_mask))

(* re-raise so that callers get a backtrace *)
let find t k = try find_notrace t k with Not_found -> raise Not_found

let find_opt t k = match find_notrace t k with v -> Some v | exception Not_found -> None
let mem t k = match find_notrace t k with _ -> true | exception Not_found -> false

(* Pure: returns a new tree. [k] must not be bound in [tree]. *)
let rec insert h k v tree =
match tree with
| Empty -> Leaf { h; key = k; value = v; next = Nil }
| Leaf { h = h2; key; value; next } when h = h2 ->
Leaf { h; key = k; value = v; next = Cons { key; value; rest = next } }
| Leaf { h = h2; key; value; next } ->
(* different hashes: replace the leaf with a node holding it one level
down, and insert there. Terminates since [h] and [h2] differ in some bit. *)
let a = Array.make level_width Empty in
a.(h2 land level_mask) <- Leaf { h = h2 lsr level_bits; key; value; next };
let i = h land level_mask in
a.(i) <- insert (h lsr level_bits) k v a.(i);
Node a
| Node a ->
let i = h land level_mask in
let a = Array.copy a in
a.(i) <- insert (h lsr level_bits) k v a.(i);
Node a

let get_or_create t k ~f =
let h = t.hashable.hash k in
let shard = t.shards.(h land t.shard_mask) in
let h = h lsr t.shard_bits in
let equal = t.hashable.equal in
let tree = Atomic.get shard in
match find_tree equal h k tree with
| v -> v
| exception Not_found ->
let v = f k in
(* invariant: [k] is not bound in [tree] *)
let rec add tree =
if Atomic.compare_and_set shard tree (insert h k v tree) then v
else
let tree = Atomic.get shard in
match find_tree equal h k tree with
| v' -> v' (* another domain bound [k] first: drop [v] *)
| exception Not_found -> add tree
in
add tree

let rec fold_bindings f acc = function
| Nil -> acc
| Cons { key; value; rest } -> fold_bindings f (f key value acc) rest

let rec fold_tree f acc = function
| Empty -> acc
| Leaf { key; value; next; _ } -> fold_bindings f (f key value acc) next
| Node a -> Array.fold_left (fold_tree f) acc a

let fold t f acc =
let snapshot = Array.map Atomic.get t.shards in
Array.fold_left (fold_tree f) acc snapshot

let iter t f = fold t (fun k v () -> f k v) ()
let to_list t = fold t (fun k v acc -> (k, v) :: acc) []
let length t = fold t (fun _ _ n -> n + 1) 0
let clear t = Array.iter (fun shard -> Atomic.set shard Empty) t.shards
end
56 changes: 56 additions & 0 deletions extThread.mli
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,62 @@ val check_main_domain : string -> unit
(** [check_main_domain name] raises [Failure] if not called from the main domain.
Use it to guard process-wide setup and configuration. [name] identifies the caller in the error. *)

(** Domain-safe hash table, fast for reads.

All operations may be called concurrently from any domain.
Lookups are wait-free and do not allocate. Writes are lock-free: they copy
a short path of the underlying trie and retry on contention, so inserts are
about twice as slow as with [Hashtbl]. Meant for read-mostly tables, eg. a
set of keys that is looked up far more often than it grows.

There is no resizing: the table is split into a fixed number of shards
(see [shard_bits] in {!create}), each holding a hash trie of branching
factor 16 that deepens as it grows. *)
module ShardedHashTrie : sig
type 'k hashable = { equal : 'k -> 'k -> bool; hash : 'k -> int }
(** [equal a b] implies [hash a = hash b]. Keys with equal hashes are kept in a
list, so a hash with few distinct values makes the table slow. *)

val poly_hashable : 'k hashable
(** Same as [Hashtbl]: [compare a b = 0] and [Hashtbl.hash]. For string keys, prefer
[{ equal = String.equal; hash = Hashtbl.hash }]. *)

type ('k, 'v) t

val create : ?hashable:'k hashable -> ?shard_bits:int -> unit -> ('k, 'v) t
(** @param hashable defaults to {!poly_hashable}
@param shard_bits log2 of the number of shards, in [\[0,16\]]. Default 6
(64 shards, about 1.5KB), fine for up to a few thousand keys; use 8 or more for larger tables.
@raise Invalid_argument if [shard_bits] is out of range *)

val find : ('k, 'v) t -> 'k -> 'v
(** @raise Not_found if the key is not bound *)

val find_opt : ('k, 'v) t -> 'k -> 'v option
val mem : ('k, 'v) t -> 'k -> bool

val get_or_create : ('k, 'v) t -> 'k -> f:('k -> 'v) -> 'v
(** Return the value bound to the key, binding it to [f key] first if absent.
Concurrent callers for the same absent key may each call [f], but only one
result is stored and all of them get that one. *)

(** {2 Iteration}

Iteration works on a snapshot of the shards taken at the start, so writes
made during iteration are not observed. The snapshot is not atomic across
shards. Order is unspecified. *)

val iter : ('k, 'v) t -> ('k -> 'v -> unit) -> unit
val fold : ('k, 'v) t -> ('k -> 'v -> 'acc -> 'acc) -> 'acc -> 'acc
val to_list : ('k, 'v) t -> ('k * 'v) list

val length : ('k, 'v) t -> int
(** Linear time: counts the bindings of a snapshot, like {!fold} *)

val clear : ('k, 'v) t -> unit
(** Not atomic across shards *)
end

val locked : Mutex.t -> (unit -> 'a) -> 'a

type 'a t
Expand Down
Loading
Loading