~/bend-docscommunity

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

type Two source · line 89 · raw

Type

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.