histogram.bend relies on unsafe/foreign
raw source on the hub · import 0xeef3a8486d410ea6f25709aac90dbf40/histogram.bend as Histogram
histogram.bend: a parallel histogram of U32 keys into K buckets, with no atomics.
import ./histogram.bend as Hist Hist.count(8n, 256, keys, Hist.Bufs.new()) # (keys, bufs) Hist.Bufs{counts, table} = bufs # counts[c]: the keys in bucket c
k is the bucket count K (0 counts as 1), any U32. A key at or past K
counts as K - 1. counts holds 2^ceil(log2 K) entries: the K counts, then
zeros. b is the fork depth, as in scan.bend: 2^b blocks of keys (b is
lowered to log2 of the key count when larger), and every b gives the
same counts. A call under ! runs its forks on the GPU.
Bufs carries the two buffers from call to call: the count table (2^b
rows of 2^ceil(log2 K) entries, so 2^(b + ceil(log2 K)) U32s) and the
counts. A call keeps a buffer that is big enough and allocates a new one
otherwise. Array.new fills on one lane, which under ! costs about 76 ms
per million entries on the GPU we measured (see README), so size the
buffers once with Bufs.sized outside the ! and hand them back each
call.
Two passes: 1. per block of keys, in parallel: zero the block's row of the table, then count its keys into it (the block owns its row); 2. the column sums. Below 2^12 rows, per group of buckets, in parallel: a bucket's count is the sum of its column over the rows (the group owns its counts). From 2^12 rows up, in two steps: the 2^b rows are cut into 2^(b/2) groups of rows; a) per (group of rows, group of buckets), in parallel: each bucket's sum over the group's rows into the group's last row; b) per group of buckets, in parallel: a bucket's count is the sum of those. Every lane walks 2^(b/2) rows at most, not 2^b. After a call the table holds scratch (from 2^12 rows up, partial sums in the groups' last rows), not per-block counts. ref.bend's histogram is the sequential reference.
Sharing: as in scan.bend, lanes share arrays through Base's @unsafe Array.fork and Array.join; no def here is @unsafe itself, and none may be evaluated in the checker (call them from an IO main).
2 imports
import Base import ./scan.bend as Scan
Types
type Bufs source · line 46 · raw
Type
Bufs@counts:Array<U32> -> @table:Array<U32> -> Bufs
type Two source · line 89 · raw
Type
Two@a:Array<U32> -> @b:Array<U32> -> Two
Definitions
def Bufs.new source · line 50 · raw
Bufs
Empty buffers: the first call allocates what it needs.
def kbits.of source · line 53 · raw
@big:Bool -> @+k:U32 -> Nat
def kbits source · line 61 · raw
@+k:U32 -> Nat
ceil(log2 K), for K >= 1.
def Bufs.sized source · line 65 · raw
@+b:Nat -> @+k:U32 -> @+lg:Nat -> Bufs
Buffers for 2^lg keys at fork depth b and k buckets, allocated now.
def fit.pick source · line 69 · raw
@ok:Bool -> @+d:Nat -> @t:Array<U32> -> Array<U32>
def fit.eq source · line 77 · raw
@+d:Nat -> @r:Pair(Array<U32>, U32) -> Array<U32>
t when it holds exactly 2^d entries, else a new one.
def fit.ge source · line 82 · raw
@+d:Nat -> @r:Pair(Array<U32>, U32) -> Array<U32>
t when it holds at least 2^d entries, else a new one.
def two.join source · line 92 · raw
@l:Two -> @r:Two -> Two
def two.rejoin source · line 97 · raw
@a2:Array<U32> -> @b2:Array<U32> -> @r:Two -> Two
def zero.go source · line 105 · raw
@+n:Nat -> @+at:U32 -> @t:Array<U32> -> Array<U32>
n + 1 zeros from at.
def bump.at source · line 112 · raw
@+i:U32 -> @r:Pair(Array<U32>, U32) -> Array<U32>
def bump source · line 117 · raw
@+i:U32 -> @t:Array<U32> -> Array<U32>
One more at t[i].
def count.go source · line 122 · raw
@+n:Nat -> @+j:U32 -> @+row:U32 -> @+k1:U32 -> @t:Array<U32> -> @r:Pair(Array<U32>, U32) -> Two
r holds the keys and the key x read at j: one more in x's bucket of the row at row, then the same for the n keys after j. k1 is K - 1.
def rows.leaf source · line 134 · raw
@+i:U32 -> @+e:Nat -> @+kb:Nat -> @+k:U32 -> @ks:Array<U32> -> @t:Array<U32> -> Two
Block i (2^e keys from i * 2^e): its row, at i * 2^kb, zeroed, then its keys counted into it. k is K (at least 1).
def rows source · line 144 · raw
@+d:Nat -> @+i:U32 -> @+e:Nat -> @+kb:Nat -> @+k:U32 -> @fk:Pair(Array<U32>, Array<U32>) -> @ft:Pair(Array<U32>, Array<U32>) -> Two
Pass 1's forks: d levels down to the blocks under block index i. fk and ft are two handles each to the keys and the table. Block i reads only its keys and reads and writes only its row (entries i * 2^kb to i * 2^kb + K - 1): no two lanes touch one slot.
def small source · line 167 · raw
@+d:Nat -> Bool
d is below 12: the table is small, and its rows are not split.
def col.go source · line 172 · raw
@+n:Nat -> @+at:U32 -> @+step:U32 -> @+s:U32 -> @r:Pair(Array<U32>, U32) -> Pair(Array<U32>, U32)
r holds the table and the entry x read at at: adds x and the n entries after it, each step further on, to s.
def ins.go source · line 183 · raw
@+n:Nat -> @+at:U32 -> @+step:U32 -> @+s:U32 -> @r:Pair(Array<U32>, U32) -> Array<U32>
As col.go, but each entry becomes the sum through it (in place).
def grp.put source · line 194 · raw
@+at:U32 -> @r:Pair(Array<U32>, U32) -> Array<U32>
def grp.one source · line 201 · raw
@lt:Bool -> @+sc:Bool -> @+c:U32 -> @+top:U32 -> @+last:U32 -> @+step:U32 -> @+nm:Nat -> @t:Array<U32> -> Array<U32>
Column c of a group (rows from top, nm + 1 of them, step apart), when c < K: sc, each entry becomes the group's sum through it; else the sum goes into the last row (at last).
def grp.go source · line 213 · raw
@+n:Nat -> @+sc:Bool -> @+c:U32 -> @+k:U32 -> @+top:U32 -> @+last:U32 -> @+step:U32 -> @+nm:Nat -> @t:Array<U32> -> Array<U32>
Columns c to c + n of a group.
def grp.leaf source · line 222 · raw
@+sc:Bool -> @+i:U32 -> @+gc:Nat -> @+w:Nat -> @+m:Nat -> @+kb:Nat -> @+k:U32 -> @t:Array<U32> -> Array<U32>
Lane i: rows group i / 2^gc (2^m rows of 2^kb), columns group i % 2^gc (2^w columns).
def grp source · line 233 · raw
@+d:Nat -> @+sc:Bool -> @+i:U32 -> @+gc:Nat -> @+w:Nat -> @+m:Nat -> @+kb:Nat -> @+k:U32 -> @ft:Pair(Array<U32>, Array<U32>) -> Array<U32>
Pass 2a's forks (sort.bend's too, with sc): d levels down to the lanes under lane i, with two handles to the table. Lane i reads and writes only its group's rows in its columns: no two lanes touch one slot.
def col.put source · line 245 · raw
@+c:U32 -> @cs:Array<U32> -> @r:Pair(Array<U32>, U32) -> Two
def col.one source · line 251 · raw
@lt:Bool -> @+c:U32 -> @+o:U32 -> @+step:U32 -> @+nb:Nat -> @p:Two -> Two
Column c's count into cs[c]: the sum of its nb + 1 entries from o + c, step apart (the groups' last rows), when c < K; else 0.
def cols.go source · line 261 · raw
@+n:Nat -> @+c:U32 -> @+k:U32 -> @+o:U32 -> @+step:U32 -> @+nb:Nat -> @p:Two -> Two
Columns c to c + n.
def cols source · line 272 · raw
@+d:Nat -> @+j:U32 -> @+w:Nat -> @+k:U32 -> @+o:U32 -> @+step:U32 -> @+nb:Nat -> @ft:Pair(Array<U32>, Array<U32>) -> @fc:Pair(Array<U32>, Array<U32>) -> Two
Pass 2b's forks: d levels down to groups of 2^w columns under group index j. ft and fc are two handles each to the table and the counts. Group j reads only its columns of the table and writes only its counts (entries j * 2^w to j * 2^w + 2^w - 1): no two lanes write one slot.
def count.fin source · line 289 · raw
@keys:Array<U32> -> @r:Two -> Pair(Array<U32>, Bufs)
def count.cols source · line 295 · raw
@+r:Nat -> @+m:Nat -> @+g:Nat -> @+k:U32 -> @+kb:Nat -> @keys:Array<U32> -> @counts:Array<U32> -> @table:Array<U32> -> Pair(Array<U32>, Bufs)
Pass 2b over the 2^kb columns, 2^min(d, kb) groups: each sums its columns over the 2^r groups' last rows.
def count.grp source · line 301 · raw
@lo:Bool -> @+d:Nat -> @+k:U32 -> @+kb:Nat -> @counts:Array<U32> -> @p:Two -> Pair(Array<U32>, Bufs)
Pass 2a over 2^r groups of 2^m rows by 2^min(d, kb) groups of columns, then 2b; below 2^12 rows, 2b alone, over 2^d groups of one row.
def count.at source · line 313 · raw
@+b:Nat -> @+k:U32 -> @bufs:Bufs -> @r:Pair(Array<U32>, U32) -> Pair(Array<U32>, Bufs)
def count source · line 328 · raw
@+b:Nat -> @+k:U32 -> @keys:Array<U32> -> @bufs:Bufs -> Pair(Array<U32>, Bufs)
The histogram of keys into k buckets at fork depth b. Returns the keys, unchanged, and the buffers, whose counts field holds the counts.