flashinfer.cake_sampling

cake_sampling is a thread-block-cluster implementation of top-k-then-top-p sampling from probabilities for Hopper and newer GPUs (compute capability 9.x, 10.x, 11.x and 12.x). It is measured on H100 (9.0), B200 (10.0), B300 / GB300 (10.3) and Rubin R200 (10.7); 11.x and 12.x are compile targets that have not been run on hardware. It fuses the three stages of flashinfer.sampling.top_k_top_p_sampling_from_probs() with filter_apply_order="top_k_first" into at most two kernels per call:

  1. a thread-block-cluster radix select that writes the exact per-row top-k slab: one 11-bit cluster pass reduces the 2048-bucket histograms through distributed shared memory, the entries of the selected bucket (at most 2048 cluster-wide) are gathered into every CTA as 64-bit (key, ~index) composites and the remaining key bits and boundary ties are resolved with local passes (rows whose bucket overflows that capacity take the exact three-pass cluster path); and

  2. a programmatic-dependent-launch sparse top-p kernel that sorts the slab prefix with one composite-key bitonic network (registers, warp shuffles and a cross-warp shared-memory exchange), keeps the shortest prefix whose exclusive mass is below top_p times the top-k mass, renormalizes, and draws one token per row by inverse CDF from curand_init(seed, row, offset). For top_k_max <= 64 (fused_tail_kcap in the manifest) this stage runs inside the stage-1 kernel on one warp of the cluster’s first CTA, so the call is a single launch; the outputs are bitwise identical to the two-launch form. Every variant the dispatcher can pick for such a top-k carries the tail (manifest fused_tail); the (8, 48) resident, a large-k pick, is built without it. In the two-launch form the dispatcher also decides where the dependent kernel lands: stage 1 signals griddepcontrol.launch_dependents before its first pass only when the batch fits on the SMs its last wave leaves free (one stage-1 CTA per SM), so the stage-2/3 CTAs are never packed onto the few SMs free mid-flight; larger batches let the dependent launch as stage 1 exits. A streaming variant signals at that early point only on Blackwell and Rubin (compute capability 10.x); on Hopper it signals after its filter pass, once the whole row has been read, where the round-3 kernels did. All three decisions travel in the stage-1 launch_flags argument (bit 0 fused tail, bit 1 early trigger, bit 2 stream pre-pass point) and none changes any output.

Semantics (support, tie-breaking toward the lower vocabulary index, Philox stream advancement through the generator) follow the top_k_first route with two extra guarantees:

  • Strict determinism. The support is exactly the first k entries of lexsort(-prob, index); every top-p and sampling decision is made on exact 64-bit fixed-point prefix sums of the sorted slab, and no atomic decides an output. Identical inputs give bitwise identical samples, renorm_out and slab for every kernel variant, stream launch or CUDA-graph replay. renorm_out and the workspace slab are emitted in sorted order (descending probability, ascending index).

  • NaN / Inf. NaN, negative and -0.0 probabilities are treated as +0 and never sampled. A row containing +inf keeps exactly its +inf entries, samples uniformly among them and renormalizes them to 1/m. A row whose top-k mass is zero returns its smallest slab index with an all-zero renorm_out. Every output is finite. Per-row top_p is clamped to (0, 1] and per-row top_k to [1, min(vocab, 1024)].

Requests the frozen kernels cannot serve are dispatched to the top_k_first route (deterministic=True): top-k disabled or k >= vocab, k > 1024, non-float32 or non-contiguous rows, or a device outside the build targets. The build targets are FlashInfer’s CompilationContext targets (FLASHINFER_CUDA_ARCH_LIST when set, as in AOT builds on hosts without a GPU; otherwise the capabilities of the visible devices) restricted to the supported majors 9-12; a device whose architecture is not among them takes the top_k_first route instead of failing at launch. Large batch * vocab launches run on the streaming stage-1 variants (a cluster of 1-8 CTAs walks the row in register chunks, samples one 11-bit pass to bound the candidate range, filters the row into a per-CTA candidate list and finishes on a gathered copy of that list with local passes; the exact fallback for a list that overflowed or came up short streams the row again through a compact runtime radix loop, kept small so the kernel’s code stays resident in the SM instruction cache next to the stage-2/3 kernel), so there is no size-based fallback. cake_sampling_route() reports the decision without launching.

The checked-in source product lives in csrc/cake_sampling/generated/ as one translation unit plus a manifest that records every frozen variant’s launch resources; FlashInfer verifies the source hash and compiles it once into a single fatbin with one -gencode per build target (the kernels use only clusters, distributed shared memory, programmatic dependent launch and redux.sync, so no per-architecture source exists). Stage-1 variants whose dynamic shared memory exceeds the device’s opt-in limit are not dispatch candidates: on 12.x (99 KB) the streaming variants drop out, so vocabularies above the register-resident capacity (196608) take the top_k_first route there.

The stage-1 variant is chosen per call by a cost model whose single-wave CTA capacity table and cost constants are keyed by the device’s SM count (148 for B200 / B300, 132 for H100, 212 for Rubin R200; other devices use the nearest measured table); the constants were fitted per table on per-variant sweeps of all four architectures so that every measured cell picks its fastest frozen variant.

Measured performance

python benchmarks/bench_cake_sampling.py --cupti --skip-joint --batches 1,2,4,8,16,32,64,128 --vocabs 32768,128256,151936,262144 with --top-k 10 / 50 / 1000 and with --cuda-graph (flashinfer.testing.bench_gpu_time, CUPTI kernel time, p = 0.9, median µs), round-4 bundle (streaming stage 1 with a gathered candidate list, fused stage 2/3 for top_k_max <= 64, per-table k-aware dispatch) on the four supported architectures. The summary compares every cell with FlashInfer’s default top_k_first route and with the previous frozen bundle (#5607) measured in the same run.

Round-4 bundle, 192 cells per architecture (8 batches × 4 vocabularies × 3 k × eager/graph)

GPU

cells slower than #5607 by > 2 %

speedup vs top_k_first (min / median / max)

gain vs #5607, B ≤ 16 (min / median / max)

gain vs #5607, B ≥ 32 (min / median)

H100 SXM (sm_90a, 132 SMs)

2 of 192

1.65x / 3.04x / 7.50x

0.98x / 1.11x / 1.56x

0.98x / 1.15x

B200 (sm_100a, 148 SMs)

0 of 192

1.56x / 3.33x / 7.01x

0.98x / 1.10x / 1.58x

1.00x / 1.11x

GB300 (sm_103a, 152 SMs)

2 of 192

1.39x / 3.32x / 11.89x

0.97x / 1.12x / 2.01x

1.00x / 1.13x

VR200 R200 (sm_107a, 212 SMs)

0 of 192

1.58x / 3.19x / 5.55x

0.99x / 1.09x / 1.50x

1.00x / 1.09x

Cells that gain less than 15 % over #5585 are the ones that run the unchanged streaming stage-1 template (V = 262144 at every batch, and V = 128256 / 151936 at B = 16 on H100 and GB300) and the k = 1000 rows at V = 151936, B ≤ 8, where the 1024-entry slab sort already dominated before this round; every other B ≤ 16 cell gains 15-75 %.

H100 SXM (sm_90a, 132 SMs), CUPTI median µs, top_k_first → cake_sampling

V

B

k = 50

k = 1000

k = 50, CUDA graph

k = 1000, CUDA graph

32768

1

45.7 → 11.0 (4.15x)

45.7 → 16.4 (2.79x)

47.8 → 21.8 (2.19x)

50.7 → 28.0 (1.81x)

32768

8

49.3 → 11.5 (4.29x)

50.1 → 16.7 (3.00x)

54.2 → 22.7 (2.39x)

54.0 → 28.8 (1.88x)

32768

16

52.9 → 12.3 (4.30x)

53.1 → 17.4 (3.05x)

54.1 → 23.2 (2.33x)

54.0 → 29.2 (1.85x)

32768

32

54.7 → 14.7 (3.72x)

54.9 → 20.8 (2.64x)

55.7 → 25.7 (2.17x)

57.3 → 33.2 (1.73x)

32768

64

55.9 → 14.8 (3.78x)

56.4 → 22.9 (2.46x)

67.3 → 25.6 (2.63x)

61.4 → 36.0 (1.71x)

32768

128

58.0 → 17.6 (3.30x)

58.7 → 25.6 (2.29x)

65.3 → 28.6 (2.28x)

66.9 → 38.3 (1.75x)

128256

1

111.1 → 15.3 (7.26x)

65.8 → 25.3 (2.60x)

82.6 → 26.3 (3.14x)

87.1 → 36.4 (2.39x)

128256

8

112.5 → 16.9 (6.66x)

81.5 → 25.0 (3.26x)

84.1 → 28.2 (2.98x)

110.7 → 37.1 (2.98x)

128256

16

111.7 → 19.0 (5.88x)

94.6 → 27.3 (3.47x)

88.1 → 30.3 (2.91x)

106.5 → 40.0 (2.66x)

128256

32

114.4 → 25.3 (4.52x)

106.1 → 33.6 (3.16x)

92.8 → 36.8 (2.52x)

123.3 → 46.3 (2.66x)

128256

64

114.2 → 31.4 (3.64x)

168.0 → 39.9 (4.21x)

102.0 → 43.0 (2.37x)

181.4 → 53.2 (3.41x)

128256

128

114.5 → 48.8 (2.35x)

236.7 → 57.5 (4.12x)

121.2 → 60.2 (2.01x)

247.6 → 70.6 (3.51x)

151936

1

111.9 → 17.4 (6.43x)

75.0 → 34.0 (2.21x)

87.7 → 29.0 (3.02x)

97.7 → 46.8 (2.09x)

151936

8

115.2 → 18.6 (6.19x)

91.7 → 32.9 (2.79x)

89.9 → 30.5 (2.95x)

109.8 → 45.6 (2.41x)

151936

16

113.5 → 20.5 (5.54x)

105.5 → 30.9 (3.41x)

93.0 → 32.5 (2.86x)

137.2 → 43.8 (3.13x)

151936

32

114.3 → 28.8 (3.97x)

123.2 → 37.7 (3.27x)

97.7 → 40.9 (2.39x)

115.2 → 50.8 (2.27x)

151936

64

116.9 → 36.2 (3.23x)

195.6 → 45.0 (4.35x)

114.8 → 48.2 (2.38x)

201.3 → 58.8 (3.42x)

151936

128

156.9 → 55.9 (2.81x)

278.8 → 63.1 (4.42x)

172.9 → 67.7 (2.55x)

300.8 → 77.3 (3.89x)

262144

1

112.7 → 19.1 (5.90x)

92.7 → 26.7 (3.47x)

89.2 → 31.0 (2.88x)

124.9 → 40.2 (3.11x)

262144

8

114.2 → 20.9 (5.46x)

120.2 → 27.4 (4.39x)

91.3 → 33.0 (2.77x)

150.5 → 42.3 (3.56x)

262144

16

114.0 → 26.1 (4.37x)

148.5 → 32.1 (4.63x)

97.3 → 38.5 (2.53x)

163.5 → 46.8 (3.49x)

262144

32

128.1 → 40.6 (3.16x)

236.2 → 46.1 (5.12x)

145.3 → 53.2 (2.73x)

224.2 → 60.2 (3.72x)

262144

64

135.7 → 52.5 (2.58x)

316.7 → 58.2 (5.44x)

151.0 → 65.9 (2.29x)

335.1 → 73.6 (4.55x)

262144

128

158.4 → 83.2 (1.90x)

470.3 → 90.8 (5.18x)

174.9 → 95.6 (1.83x)

478.0 → 104.7 (4.57x)

B200 (sm_100a, 148 SMs), CUPTI median µs, top_k_first → cake_sampling

V

B

k = 50

k = 1000

k = 50, CUDA graph

k = 1000, CUDA graph

32768

1

42.2 → 10.8 (3.91x)

42.8 → 16.0 (2.67x)

45.5 → 20.4 (2.23x)

41.6 → 26.7 (1.56x)

32768

8

46.4 → 11.3 (4.11x)

47.3 → 16.4 (2.88x)

51.4 → 21.2 (2.42x)

47.4 → 27.9 (1.70x)

32768

16

49.0 → 11.5 (4.26x)

49.5 → 16.7 (2.96x)

46.3 → 21.2 (2.18x)

51.5 → 28.1 (1.83x)

32768

32

50.5 → 12.2 (4.14x)

51.3 → 17.4 (2.95x)

52.2 → 21.9 (2.38x)

52.2 → 29.1 (1.79x)

32768

64

52.2 → 13.7 (3.81x)

53.3 → 19.9 (2.68x)

53.7 → 23.7 (2.27x)

56.6 → 31.3 (1.81x)

32768

128

53.2 → 14.8 (3.59x)

54.5 → 22.6 (2.41x)

56.4 → 24.2 (2.33x)

54.8 → 34.7 (1.58x)

128256

1

101.2 → 14.7 (6.88x)

65.2 → 20.0 (3.26x)

82.8 → 24.6 (3.37x)

68.6 → 31.2 (2.20x)

128256

8

101.1 → 15.4 (6.56x)

81.0 → 21.1 (3.84x)

83.9 → 25.6 (3.28x)

87.4 → 32.8 (2.66x)

128256

16

101.2 → 17.6 (5.75x)

92.9 → 25.5 (3.64x)

82.2 → 27.6 (2.98x)

92.7 → 36.5 (2.54x)

128256

32

104.5 → 18.8 (5.56x)

100.3 → 26.6 (3.77x)

87.2 → 29.1 (3.00x)

107.1 → 38.5 (2.78x)

128256

64

104.4 → 25.1 (4.16x)

144.4 → 32.8 (4.40x)

88.6 → 35.5 (2.50x)

150.9 → 44.8 (3.37x)

128256

128

105.3 → 37.0 (2.85x)

192.8 → 45.0 (4.28x)

94.9 → 46.9 (2.02x)

213.3 → 57.4 (3.72x)

151936

1

101.7 → 16.2 (6.28x)

73.4 → 22.4 (3.28x)

83.9 → 26.4 (3.18x)

74.7 → 33.6 (2.22x)

151936

8

102.6 → 17.0 (6.04x)

92.3 → 23.8 (3.88x)

91.7 → 27.3 (3.36x)

98.2 → 36.0 (2.73x)

151936

16

102.2 → 19.3 (5.30x)

101.8 → 25.3 (4.02x)

88.1 → 29.5 (2.99x)

106.0 → 37.4 (2.83x)

151936

32

102.7 → 20.5 (5.01x)

113.8 → 26.2 (4.34x)

89.7 → 30.9 (2.90x)

134.0 → 38.6 (3.47x)

151936

64

105.7 → 28.9 (3.66x)

164.7 → 35.9 (4.59x)

92.8 → 39.8 (2.33x)

170.6 → 48.0 (3.55x)

151936

128

105.6 → 43.2 (2.44x)

221.6 → 50.0 (4.43x)

104.4 → 53.4 (1.96x)

242.2 → 62.2 (3.89x)

262144

1

101.9 → 17.9 (5.69x)

92.0 → 26.4 (3.48x)

89.1 → 28.3 (3.15x)

86.4 → 37.6 (2.30x)

262144

8

103.3 → 18.9 (5.47x)

118.4 → 25.7 (4.61x)

93.2 → 29.7 (3.14x)

204.6 → 38.6 (5.30x)

262144

16

104.0 → 24.3 (4.28x)

140.5 → 30.5 (4.61x)

91.5 → 35.0 (2.61x)

193.2 → 42.3 (4.57x)

262144

32

114.7 → 25.8 (4.45x)

200.5 → 31.5 (6.37x)

133.0 → 36.5 (3.64x)

218.5 → 45.0 (4.86x)

262144

64

107.3 → 39.7 (2.70x)

267.2 → 45.2 (5.91x)

123.8 → 50.7 (2.44x)

276.3 → 58.9 (4.69x)

262144

128

123.5 → 65.3 (1.89x)

376.9 → 71.9 (5.24x)

143.8 → 75.7 (1.90x)

397.8 → 85.1 (4.67x)

GB300 (sm_103a, 152 SMs), CUPTI median µs, top_k_first → cake_sampling

V

B

k = 50

k = 1000

k = 50, CUDA graph

k = 1000, CUDA graph

32768

1

66.2 → 11.1 (5.96x)

74.6 → 23.9 (3.12x)

42.3 → 22.7 (1.86x)

42.4 → 30.5 (1.39x)

32768

8

72.3 → 11.5 (6.29x)

70.8 → 22.2 (3.19x)

46.0 → 22.7 (2.03x)

48.2 → 30.6 (1.58x)

32768

16

71.6 → 11.7 (6.12x)

73.2 → 22.2 (3.30x)

46.7 → 23.2 (2.01x)

54.6 → 30.7 (1.78x)

32768

32

74.5 → 12.4 (6.01x)

75.1 → 22.2 (3.38x)

54.5 → 24.1 (2.26x)

52.2 → 32.4 (1.61x)

32768

64

76.8 → 14.1 (5.45x)

76.9 → 22.1 (3.48x)

57.9 → 25.6 (2.26x)

55.7 → 35.6 (1.56x)

32768

128

76.9 → 14.8 (5.20x)

79.1 → 23.7 (3.34x)

57.1 → 26.3 (2.17x)

54.1 → 36.7 (1.47x)

128256

1

179.2 → 15.2 (11.79x)

82.2 → 22.2 (3.70x)

77.6 → 27.1 (2.86x)

70.0 → 34.6 (2.02x)

128256

8

173.4 → 15.6 (11.12x)

101.5 → 22.6 (4.49x)

83.7 → 27.4 (3.05x)

100.1 → 35.2 (2.84x)

128256

16

172.9 → 17.8 (9.71x)

107.4 → 26.8 (4.01x)

84.1 → 29.8 (2.82x)

89.9 → 40.3 (2.23x)

128256

32

179.6 → 18.6 (9.66x)

114.0 → 27.6 (4.13x)

85.5 → 30.6 (2.79x)

110.4 → 41.6 (2.65x)

128256

64

178.2 → 24.7 (7.21x)

139.1 → 33.6 (4.14x)

88.4 → 36.4 (2.43x)

151.6 → 47.4 (3.20x)

128256

128

181.0 → 36.0 (5.03x)

184.5 → 45.1 (4.09x)

92.0 → 47.4 (1.94x)

192.6 → 58.3 (3.30x)

151936

1

176.5 → 16.6 (10.63x)

86.3 → 24.1 (3.58x)

85.9 → 28.7 (2.99x)

103.2 → 38.6 (2.67x)

151936

8

175.0 → 17.2 (10.17x)

106.6 → 25.1 (4.25x)

89.3 → 29.4 (3.04x)

96.6 → 38.4 (2.52x)

151936

16

175.1 → 19.3 (9.07x)

114.6 → 26.6 (4.31x)

89.7 → 31.9 (2.81x)

118.9 → 39.9 (2.98x)

151936

32

176.6 → 20.2 (8.74x)

122.5 → 27.1 (4.52x)

90.2 → 32.5 (2.78x)

102.7 → 41.8 (2.46x)

151936

64

179.7 → 28.3 (6.35x)

157.1 → 35.4 (4.44x)

91.9 → 40.5 (2.27x)

162.1 → 49.3 (3.29x)

151936

128

179.3 → 41.8 (4.29x)

212.9 → 49.4 (4.31x)

101.0 → 53.6 (1.88x)

230.5 → 63.2 (3.65x)

262144

1

175.3 → 18.2 (9.63x)

106.7 → 28.2 (3.78x)

82.7 → 30.4 (2.72x)

90.5 → 42.8 (2.11x)

262144

8

175.2 → 19.1 (9.17x)

137.1 → 27.2 (5.04x)

90.2 → 31.6 (2.85x)

123.1 → 41.6 (2.96x)

262144

16

174.4 → 24.2 (7.21x)

150.1 → 31.4 (4.78x)

90.2 → 36.9 (2.44x)

169.0 → 45.5 (3.71x)

262144

32

173.7 → 25.4 (6.84x)

190.7 → 32.3 (5.90x)

131.9 → 38.1 (3.46x)

207.7 → 46.6 (4.46x)

262144

64

177.3 → 38.6 (4.59x)

254.3 → 45.4 (5.60x)

120.4 → 51.2 (2.35x)

271.4 → 59.4 (4.57x)

262144

128

178.6 → 63.4 (2.82x)

361.2 → 71.3 (5.07x)

143.2 → 75.9 (1.89x)

375.9 → 84.9 (4.43x)

VR200 R200 (sm_107a, 212 SMs), CUPTI median µs, top_k_first → cake_sampling

V

B

k = 50

k = 1000

k = 50, CUDA graph

k = 1000, CUDA graph

32768

1

30.0 → 9.6 (3.12x)

30.9 → 12.9 (2.40x)

32.6 → 16.9 (1.93x)

34.4 → 21.7 (1.59x)

32768

8

33.4 → 10.3 (3.24x)

33.5 → 13.8 (2.43x)

41.5 → 18.0 (2.31x)

42.0 → 23.0 (1.83x)

32768

16

35.3 → 10.8 (3.27x)

35.5 → 14.0 (2.54x)

42.2 → 18.5 (2.28x)

43.5 → 23.2 (1.88x)

32768

32

37.1 → 11.3 (3.28x)

37.1 → 14.2 (2.61x)

44.4 → 18.5 (2.40x)

41.5 → 22.9 (1.81x)

32768

64

38.2 → 12.6 (3.03x)

38.4 → 15.9 (2.42x)

44.0 → 20.1 (2.19x)

43.7 → 24.5 (1.78x)

32768

128

39.3 → 13.4 (2.93x)

39.5 → 18.5 (2.14x)

45.0 → 20.7 (2.17x)

51.4 → 28.5 (1.80x)

128256

1

70.2 → 13.1 (5.36x)

56.6 → 16.0 (3.54x)

70.5 → 20.6 (3.42x)

62.4 → 25.5 (2.45x)

128256

8

70.1 → 13.9 (5.04x)

67.9 → 17.4 (3.90x)

74.0 → 21.6 (3.43x)

81.1 → 26.0 (3.12x)

128256

16

70.0 → 14.8 (4.73x)

76.4 → 18.0 (4.24x)

73.0 → 22.4 (3.26x)

85.3 → 27.5 (3.10x)

128256

32

72.4 → 16.5 (4.39x)

84.4 → 21.5 (3.93x)

75.2 → 23.9 (3.15x)

97.7 → 31.0 (3.15x)

128256

64

72.5 → 22.5 (3.22x)

92.9 → 27.4 (3.39x)

76.9 → 30.0 (2.56x)

115.1 → 36.7 (3.14x)

128256

128

72.3 → 32.7 (2.21x)

137.1 → 37.1 (3.70x)

78.5 → 40.0 (1.96x)

145.3 → 46.4 (3.13x)

151936

1

70.0 → 14.5 (4.83x)

64.2 → 18.1 (3.55x)

70.2 → 22.0 (3.19x)

68.8 → 27.3 (2.52x)

151936

8

70.6 → 15.2 (4.64x)

77.1 → 19.3 (3.99x)

77.4 → 22.8 (3.39x)

104.3 → 28.0 (3.73x)

151936

16

70.4 → 15.9 (4.43x)

87.3 → 20.1 (4.34x)

78.6 → 23.3 (3.37x)

87.8 → 29.5 (2.98x)

151936

32

70.3 → 18.0 (3.91x)

96.4 → 21.3 (4.53x)

80.4 → 25.7 (3.13x)

111.5 → 30.1 (3.70x)

151936

64

72.5 → 25.5 (2.84x)

105.9 → 30.5 (3.47x)

79.4 → 33.1 (2.40x)

110.4 → 39.3 (2.81x)

151936

128

73.5 → 37.6 (1.95x)

161.0 → 41.6 (3.87x)

89.2 → 44.9 (1.99x)

175.6 → 50.3 (3.49x)

262144

1

70.0 → 16.1 (4.35x)

79.1 → 20.1 (3.94x)

75.1 → 23.7 (3.17x)

83.5 → 29.8 (2.80x)

262144

8

70.3 → 17.0 (4.14x)

102.8 → 20.9 (4.92x)

77.2 → 24.4 (3.16x)

121.2 → 30.3 (4.00x)

262144

16

70.2 → 18.0 (3.90x)

119.4 → 21.5 (5.55x)

78.3 → 25.2 (3.11x)

123.9 → 30.8 (4.02x)

262144

32

70.5 → 22.4 (3.15x)

134.3 → 25.6 (5.25x)

80.2 → 30.1 (2.66x)

136.1 → 35.4 (3.84x)

262144

64

87.1 → 34.5 (2.52x)

195.6 → 37.6 (5.20x)

108.0 → 41.8 (2.58x)

209.8 → 46.9 (4.47x)

262144

128

106.7 → 55.5 (1.92x)

293.9 → 58.7 (5.01x)

123.4 → 63.0 (1.96x)

333.6 → 67.9 (4.91x)

top_k_top_p_sampling_from_probs(probs, ...)

Fused top-k-then-top-p sampling from probabilities (thread-block-cluster radix pipeline).

top_k_probs_to_slab(probs, top_k, *[, ...])

Stage 1 alone: exact per-row top-k into a [batch, 1024] slab.

cake_sampling_route(probs, top_k[, top_k_max])

"pipeline" when the frozen kernels serve this request, else "fallback:<reason>".