122 Commits
Author SHA1 Message Date
ViperEkura 432dfec3c2 refactor: collapse fp8 recipe hierarchy and state property layers
- Merge DelayedScaling/DynamicScaling and the abstract FP8Recipe base into one FP8Recipe dataclass with a dynamic flag; dispatch now reads cfg.recipe.dynamic instead of isinstance checks
- Drop the _ActiveOrDefault descriptor and the FP8State property views; the persistent defaults are plain default_* attributes and get_weight_meta takes the active recipe explicitly
- Convert FP8TensorMeta to a NamedTuple of the three per-operand rings
- Update tests to the new API; the autocast context test now asserts _active_config push/restore directly
2026-08-31 14:24:51 +08:00
ViperEkura e3c3e28a11 docs: fix stale developer documentation claims
- Move task_alloc/task_free/task_extend/task_cached/task_record_hashes and bind from the PagePool card to a new TaskCacheManager card matching pool.py
- Drop the nonexistent Executor tokenizer attribute and association, add task_cache instead
- Add AllocationStrategy/ContiguousStrategy/PagedStrategy cards and point Allocator/RadixCache composition at PagedStrategy
- Add TaskCacheManager and the allocation strategies to the module overview, add _task_cache to InferenceScheduler
- Fix the design-pattern count in the table of contents (15 -> 16)
- Rewrite the FlashAttnBackend class docstring: packed decode gathers flat K/V via req_to_token and calls flash_attn_varlen_func; dense prefill uses flash_attn_func (no flash_attn_with_kvcache exists)
- Apply the same correction to the backend bullets in internals.md and cuda_kernels.md
- Rename the stale fp8_mma_test.cu reference to fp8_test.cu in cuda_kernels.md
2026-08-31 14:24:51 +08:00
ViperEkura a7d4cb25c5 docs: scope trainer environment variables per job
- Add a Per-Job Environment section explaining that runtime.environment reaches only the GPUs declared in the same job YAML, with one-YAML-per-GPU-group examples for local, cross-PCIe workaround, and NVSwitch NVLink tuning setups
- Replace the NCCL workaround pair in the runtime schema example with ASTR_LOG_LEVEL and ASTR_BACKEND and document value semantics (str() rendering, null exports empty, no host-shell passthrough)
- Comment out the blanket NCCL exports in the get-started multi-GPU example so they are opt-in per docs/guides/distributed.md
- Add a hard rule against copying NCCL workarounds into every training config
2026-08-31 14:24:51 +08:00
ViperEkura 0546331637 fix: skip gradient checkpointing log when no modules configured
- GradientCheckpointingCallback.on_train_begin returns early on empty module list
- previously logged "Gradient checkpointing enabled" even when checkpointing was inactive, misleading profiling
2026-08-31 14:24:51 +08:00
ViperEkura 962c10c52b perf: fold the delayed-scaling ring update into the quantize kernel
- the kernel's last block folds amax into the history window and publishes the next scale in-kernel (atomicAdd ticket + fences), replacing the host update chain
- quantize bindings split into quantize(transposed) / quantize_dual with fixed arities and a QuantLayout enum; the python adapter becomes a thin attention-style wrapper over pybind (Optional ring_state at the boundary, no torch.library custom_ops)
- tests: in-kernel fold vs host reference (exact), dual/transposed orientation byte-equality

Benchmark: L20 (sm_89), 1.2B model, full train step. Per-linear fixed overhead 28.8us -> 8.8us; fp8 vs bf16: M=512 77.5ms, M=2048 144.5ms (1.15x), M=8192 527.4ms (1.28x); losses bit-identical.
2026-08-31 14:24:51 +08:00
ViperEkura 1cf7d6c76b perf: fill steady-state decode input ids via d2d copy
- add InferenceWorkspace.fill_input_ids_from_device copying device tokens straight into the fixed-address input_ids buffer
- cache each decode step's sampled tokens on-device in DecodeSteadyState.last_tokens; when the task signature is unchanged the next step reuses them, replacing the tolist -> python list -> elementwise host fill -> pageable h2d round-trip
- _sample_logits returns (host payload, device tokens); prefill discards the device tensor
- signature change (task join/leave/first decode) still takes the host path; both dispatch paths covered by tests

Benchmark: NVIDIA L20, BF16, 1B model + 0.11B test model (4 layers, hidden 512), contiguous KV cache, CUDA Graph, greedy, prompt 512, generation 256, engine decode via scripts/tools/benchmark.py (alternating A/B, 2-4 paired runs)
- 0.11B batch 32: 21429 -> 24415 tok/s mean (1.14x, +13.9%), 4/4 paired runs faster
- 1B batch 32: 4242 -> 4388 tok/s (1.034x, +3.4%), 7.54 -> 7.29 ms/step
- batch 1: no measurable change (<0.5%)
2026-08-31 14:24:51 +08:00
ViperEkura 36e39496d4 perf: vectorize tiled fp8 transpose quantize and arm amax via memset
- tiled transpose quantize becomes one 64x32-tile kernel: native pair loads (128B warp reads) with in-kernel scalar fallback at unaligned or ragged rows, so odd widths and misaligned bases no longer route to a separate kernel
- the old 32x32 scalar tiled kernel and its launcher correctness branch are gone; grid sizing simplifies to 1 + total / (vec * threads) since both elementwise loops are grid-stride
- quantize arms the amax buffer with cudaMemsetAsync instead of the zeros() fill kernel, dropping one tensor-op dispatch and kernel launch per call
- byte-exact parity holds over 1404 golden records (13 shapes x 3 dtypes x 3 scales x 2 formats x 3 layouts x aligned/misaligned) and tests/extension passes 65/65
- elementwise quantize kernel left unchanged: 16B-store pairing, __ldcs streaming hints and amax tree reduction all measured neutral at its ~52% DRAM ceiling and were reverted

Benchmark: L20 (sm_89), profiler kernel time with L2 flushed between calls.
- transposed quantize (layout 1): 230 -> 294 GB/s on 2048x1536 (+28%), 245 -> 299 on 2048x1536 weights (+22%); dual-layout (layout 2) 248 -> 329 (+33%) on the same shapes
- DRAM-saturated sizes (~10.6M elements) regress ~5% (404 -> 384 GB/s on 8192x1536), ~0.02% of a training step; accepted for the single-kernel shape after scalar-path and geometry variants both measured the same
- amax init fill kernel 3.0us -> memset 0.9us; quantize call CPU wall 18.5 -> 13.4us on 128x1536
2026-08-31 14:24:51 +08:00
ViperEkura a1a1a6bf0f perf: pack gqa q-heads per prefill block to reuse kv tiles
- pack HB = min(G, WARPS) q heads per block; K/V tiles stream once per block instead of once per q head
- G=1 keeps the old grid; paged path splits 64-row host Q tiles into HB blocks along grid.x (host maps unchanged)

Benchmark: NVIDIA RTX 6000D, short-q/long-kv prefill 1.4-3.4x (G=8 B=16 q=16 kv=16k 4.22 -> 1.26 ms); full prefill/MHA/paged unchanged (compute-bound); verified vs SDPA G in {1,2,3,4,8,32}, 99 tests pass
2026-08-31 13:39:47 +08:00
ViperEkura 7dd184a4e5 refactor: split fp8 gemm device code into layered headers
- split gemm.cuh into gemm/{policy,load,scheduler,mainloop,epilogue}.cuh (humming/CUTLASS-style layering, files 28-336 lines); the umbrella keeps the kernel orchestrator, host planning and the gemm<> entry so ops.cu and the C tests build unchanged
- move the measured design essays (swizzle derivation, ring-depth barrier invariant, launch crossovers, NN swap) into an FP8 design-notes section in docs/developer/cuda_kernels.md, leaving one-line constraints at each symbol
- refresh the doc's FP8 file table and layout tree (fix stale mm.cu / fp8_mma_test.cu names)

- structure-only change: extension rebuilds identical, C tests all pass, tests/extension 65 passed, quantize layouts byte-exact, NT routing torch.equal, e2e M=8192 530.6ms / 1.26x unchanged
2026-08-28 17:29:13 +08:00
ViperEkura bf239d194c refactor: dedupe fp8 kernel helpers and trim comments
- merge the quantize launchers into one Tiled template; extract shared cvt_fp8/publish_amax helpers and replace the dtype x format ladder with two-level template dispatch
- fold the gemm interior/generic operand loads into one kInterior template and the fast/generic async loads into load_async<kFast>; Policy carries the smem budget
- compress kernel comments to the load-bearing invariants, dropping measured-number essays; Policy signature and kernel code unchanged

Benchmark: NVIDIA L20, 1.2B model train step fwd+bwd+CE
- M=8192: fp8 532.2 -> 530.4 ms (1.26x, noise); tests/extension 65 passed, quantize layouts byte-exact, NT routing diff 0.0
2026-08-28 16:45:21 +08:00
ViperEkura 8a353117ea perf: transpose-quantize backward operands to route all gemms nt
- quantize gains out_layout (0 row-major / 1 transposed / 2 single-read dual-write); modes 1/2 run a new 32x32 smem-tile transpose kernel
- backward feeds g8/w8T and g8T/x8T to trans_b=True gemms, dropping the NN-swap and TT crosswise kernels from training; fp8 weights keep the swap fallback
- a 64x64 tile variant tied on the real step mix and was reverted; noted in the kernel header

Benchmark: NVIDIA L20, 1.2B model, full train step fwd+bwd+CE
- M=8192: fp8 551.8 -> 532.2 ms, 1.21x -> 1.26x vs bf16; M=2048 0.90x -> 0.95x
- kernel-level grad_x +3.7..12.4%, grad_w +13.8..20.8%; layouts byte-exact, fp8 tests 36/36
2026-08-28 16:12:15 +08:00
ViperEkura 04a8e2517a perf: split fp8 gemm cta plan by operand layout
- pass the crosswise operand count from gemm into plan_gemm so congruous and crosswise problems stop sharing one threshold ladder
- congruous grids past one big-cta wave pick big vs narrow by the wave cost ceil(tiles/sm) * T_tile with T_narrow ~= 0.53 * T_big, reproducing every measured crossover
- crosswise problems run the small 64x64 s3 cta up to ~1.5 waves of 128x128 tiles; the narrow cta never wins there (loses to small below the band, to big above it)
- keep the sub-wave congruous ladder and the padding rules unchanged

Benchmark: NVIDIA L20 (92 SMs, sm_89), CUDA 12.8, fp8 e4m3 -> bf16, interleaved A/B against the previous ladder
- NT Mx4096x4096: M=384 114.5 -> 134.5 TF (+17.4%), M=512 152.5 -> 171.7 (+12.6%), M=768 153.5 -> 165.4 (+7.7%); all other NT shapes unchanged
- NN/TN M=64..512 +3.2..+17.4%, 1024^3 +13.7%/+12.8%; M>=640 and 2048^3+ unchanged
- TT 1024^3 +25.5%, M=256 +25.4%; TT M=512 -4.6% at the 1.5-wave boundary that favors TN/NN
2026-08-28 12:44:34 +08:00
ViperEkura c4f7f82725 refactor: drop dead fp8 gemm knobs and dedupe ring depth logic
- remove the LeanRing knob: every production Policy already ran full kStages+1 rings (the lean variant measured slower, 1280³ +5..9%), so the barrier-4 branch, the kInterleave condition and the ring-depth ternaries collapse to a single kRingDepth in Fp8GemmSmem, now the single source the mainloop reads
- remove the always-true grouped field from Fp8GemmPlan: every layout canonicalize_gemm produces is grouped-raster, so plan_gemm drops the parameter; the plain-raster experiment knob stays available via launch_plan's GroupRaster template parameter
- extract load_b_frags for the duplicated B-fragment fill (initial + double-buffer next-seg sites)
- device_sm_count: fold the out-of-range branch into one cached query path
- Fp8GemmPolicy goes 12 -> 11 template parameters; fp8_test's CasePolicy follows

Benchmark: NVIDIA L20 (sm_89, 92 SMs), kernel bench and the 1204M bf16 model e2e training step both unchanged (fp8 step 503.7 -> 503.9 ms, 1.23x vs bf16; per-shape TFLOPS within +-2%); fp8_test All PASS, tests/extension/test_fp8_mma.py 36 passed.
2026-08-28 01:55:30 +08:00
ViperEkura fac9d07542 refactor: fp8 gemm policy layering with swap-NN and narrow-N ctas
Kernel restructured CUTLASS-style: Fp8GemmPolicy as the kernel's single template parameter (traits + operand layouts + scheduling knobs), the body split into Fp8GemmTileScheduler / Fp8CollectiveMainloop / Fp8CollectiveEpilogue collectives, and the entry split into canonicalize_gemm -> plan_gemm -> launch_plan behind fp8::gemm.

- NN (dual-N-contiguous) problems run as their transpose: the swap in canonicalize_gemm plus an out-transposed epilogue removes one kernel instantiation per (format, tile config)
- new 128x64 narrow CTA (8 warps of 32x32) serves the sub-wave band once its grid passes ~3/8 of a wave: +7..77% there (128x4096x4096 116->131T, 1024^3 131->174T, 4096x384x4096 147->242T, 8192x128x4096 131->233T); decode, the padding band and multi-wave shapes unchanged
- launch_with_smem no longer swallows cudaFuncSetAttribute failures
- fp8_test: GPU-side fp32 reference (O(m*n) compare instead of O(m*n*k) host loop), production-dispatch cases for the NN swap and the plan selection; dead transpose_layout trait removed

Device: NVIDIA RTX 6000D (sm_120, 156 SMs), CUDA 13.1, torch 2.11.0+cu130. Kernel-only bench vs CUTLASS 4.8.0 sm120 dense fp8: ahead up to 1.68x below one wave (512^3 44 vs 26T, 64x4096x4096 95 vs 62T), within ~7% in the DRAM-streaming regime (8192^3 248 vs 266T).
2026-08-28 01:21:55 +08:00
ViperEkura bbb2d95256 fix: report true rank in logs via dist-aware helpers
- get_rank/get_world_size fall back to RANK/WORLD_SIZE env instead of hardcoded 0/1, matching torchrun's env-before-init contract
- log filter reuses the helpers so initialized groups show real ranks; local-spawn children previously logged rank=0/8
2026-08-27 13:39:38 +08:00
ViperEkura 1c04a0b9fa refactor: move runtime parsers into scripts/docker
- serve_runtime.py and train_runtime.py are host-side Docker helpers, so they join train-entrypoint.sh and lib/ under scripts/docker/
- scripts/tools/ now contains only in-container CLIs
- update wrapper call sites, test import, and docker guide references
2026-08-27 12:33:32 +08:00
ViperEkura ba8beb81be fix: serve and train reuse the built image and expose GPUs correctly
- serve.sh/train.sh no longer pass --build on up/run; the build subcommand is the only path that rebuilds
- compose services pin image: astrai:latest so run reuses the existing image instead of triggering a rebuild
- runtime parsers leave CUDA_VISIBLE_DEVICES unset for gpu.devices: all; an empty string hid every GPU inside the container
- server service reserves count: all GPUs so CUDA_VISIBLE_DEVICES performs the only filtering, matching the trainer
- wrapper compose() strips an empty host CUDA_VISIBLE_DEVICES before invoking docker compose
2026-08-27 12:22:02 +08:00
ViperEkura 7cfcc6c86a Merge pull request #26 from Cytosine-code/fix/cuda-kernel-install
fix: build CUDA kernels during editable installation (sm89-gated fp8_ops) plus doc corrections
2026-08-27 05:07:44 +08:00
ViperEkura f4c44ebf1c docs: correct fp8 kernel descriptions in cuda_kernels guide
- keep the separate fp8/quantize.cuh row (quantize kernel lives there, not in gemm.cuh)
- gemm.cuh now dispatches 64x64/128x128 CTAs at runtime via prefer_small_cta
- quantize takes the quantization multiplier (strategy passes scale.reciprocal())
- python primitives are fp8_quantize/fp8_gemm plus quantize/mm_fp8 wrappers
- cmake builds the five base targets plus fp8_ops on sm89+
2026-08-27 05:07:41 +08:00
Cytosine f86f605f5f fix: make CUDA kernel installation reliable 2026-08-27 02:34:16 +08:00
ViperEkura 4c82d5d84b perf: dispatch non-128-divisible shapes to the 64x64 cta
- a 128x128 CTA that is not exactly tiled (m or n not a multiple of 128) runs its edge tiles on the predicated generic path, and with a single in-flight wave the runtime is the slowest CTA — the edge tiles drag the whole shape down, so prefer_small_cta now also takes n and returns true when 64 divides both dims but 128 does not (measured sweep: 1088^3 big-CTA 76T vs 64x64 93T)
- the 64x64 grid tiles exactly on such shapes and overlaps waves (CTA count jumps 64->100->144 in the 1024-1536 band; the non-divisible points like 1088/1216 sit in deep sawtooth valleys that the divisibility rule lifts to the flat ~100-111T plateau)
- double-non-divisible shapes (e.g. 1000^3) stay on the big CTA: measured 64x64 36.7T vs 128x128 40.0T — both grids carry edge tiles there and the big CTA's efficiency wins
- m <= 64 or n <= 64 also takes the small CTA (a 128-wide CTA wastes more than half its columns on narrow N)

Benchmark: L20, e2e CUDA graph (GPU 7, same GPU as all comparisons): 960^3 60.2->100.1T, 1088^3 77.6->101.6T, 1216^3 69.7->103.5T, 1344^3 84.9->111.1T; 128-divisible shapes unchanged within noise (1024^3 118.4, 1152^3 150.6, 1280^3 106.8, 1536^3 153.7, 2048^3 192.2). 596 pytests pass.
2026-08-26 19:02:50 +08:00
ViperEkura 76aa4edc9f perf: skip grad-bias reduce for bias-free linears
- _LinearFp8.backward computed g2.sum(0) unconditionally and dropped it when needs_input_grad[2] was false; now the column-sum only runs when the bias actually requires grad
- saves ~327 reduce kernels per train step on bias-free LLMs (215M GQA: end-to-end 1.08x -> 1.13x vs bf16)
2026-08-26 18:51:18 +08:00
ViperEkura 6354dbe8bc perf: interleave prefetch into the mma phase
- our 128x128 fast loop padded 26 NOPs between the 32 QMMAs while all four LDGSTS sat bunched at the loop tail: ptxas had no independent instructions to fill the tensor-pipe issue gaps, the exact structure the decompiled cuBLAS loop (166 i, NOP=0) and CUTLASS MmaMultistage avoid by issuing cp.async in small groups inside the MMA phase (copy_tiles_and_advance per warp-tile batch)
- the steady-state prefetch is now a loop-carried register pair per congruous operand (PrefetchCarry: swizzled stage offset + global source, constructed once from the same (r, c0) mapping as the interior loader), whose chunks ride after the first and last k_seg MMA batches — SASS: 282 -> 110 instructions, 0 branches, 0 UIMAD.WIDE magic-divisions, LDGSTS interleaved inside the QMMA range, 26 -> 19 NOPs, still 128 regs (2 CTAs/SM)
- the wait-count dispatch ladder (16 instructions of ISETP/SEL picking DEPBAR immediates) and the per-k-tile (tile % ring) * stage_bytes recomputation (UIMAD.WIDE by 0x55555555) are gone: the prologue commits unconditionally so the wait_group<kStages-1> immediate is valid for every iteration, and both read and write stage addresses advance as carried pointers with an equality wrap
- cp_async.cuh splits the emitter from its policies: one raw PTX site (cp_async_16_raw) plus wrappers for unconditional/predicated and pointer/offset destinations, and the now-unused wait_group_dispatch ladder is deleted; the dispatch flip: with the stall gone the big CTA wins the whole former dip band (1280^3 kernel-level 98.2->104.4T), so prefer_small_cta keeps only the sub-5/8-wave band and the single-wave s3 variant is retired

Benchmark: L20 (sm_89), kernel-level sweep 128s2ff: 1024^3 104.2->114.0T (cuBLAS 151.2), 1152^3 130.4->145.9T (157.2), 1280^3 98.2->104.7T (163.8), 1536^3 141.4->150.4T (182.3), 2048^3 178.1->189.6T (202.8), 4096^3 197.6->208.8T (223.7), 8192^3 207.6->219.1T (227.5). CUDA-graph e2e: 1024^3 108.2->118.4T, 1152^3 137.7->150.7T, 1536^3 143.4->153.7T, 2048^3 179.4->191.9T, 8192^3 199.0->208.9T, 1280^3 108.3->106.6T (old dip-band rule re-measured 106.1T — within noise). Four-layout C++ suite and 596 pytests pass.
2026-08-26 17:47:24 +08:00
ViperEkura f45230fb2c perf: pair b fragments into ldmatrix x4 loads
- fold the two adjacent nt B fragments of each pair into one ldmatrix.x4: lanes 0-7/8-15 address rows n0..n7 chunks c/c+1, lanes 16-23/24-31 the same chunks of rows n8..n15, so {r0,r1} feed the even nt mma and {r2,r3} the odd nt — 4 x4 B loads per k-tile instead of 8 x2 (12 LDSM total, matching the decompiled cuBLAS and CUTLASS loop shapes)
- the +8-row half never reaches the XOR-swizzle source bits for kK <= 64 (row[2:1]), so the pairing rides the existing per-lane address closure with one extra term (rh16 * 8 * kK); kK=128 swizzles on row[2:0] and keeps the x2 path
- decompilation trail: nsys shows cuBLAS never split-Ks on the gap shapes (grid.z=1, no atomics; it fills waves with 64x128/64x64 tiles instead), and a CUTLASS 3.8 reference at our exact 128x128 s3 geometry reaches 196.5T at 2048^3 vs our 177.9T with NOP=0 and 12 LDSM — proving the loop shape is reachable from CUDA C++ (see perf/fp8_next_ideas.md F/C)

Benchmark: L20 (sm_89), kernel-level sweep: 2048^3 176.7->178.1T, 4096^3 196.3->197.2T, 8192^3 ->208.7T, 4096x512x4096 137.3->139.1T, 896x1152x4096 121->123.4T. CUDA-graph e2e: 512^3 54.2->55.3T, 1536^3 142.3->143.4T, 2048^3 177.9->179.4T, 8192^3 198.2->199.0T. Four-layout C++ suite, 596 pytests pass.
2026-08-26 15:43:09 +08:00
ViperEkura d4534be8ca perf: single-wave big cta takes a stage deeper pipeline
- dispatch 128x128 CTAs at kStages=3 when the grid fits one wave (tiles <= SM count): with no second wave to overlap the drain, latency hiding comes only from the pipeline depth
- multi-wave grids keep kStages=2 — the shorter prologue wins once retiring CTAs overlap (measured 4096^3: s2 196T vs s3 175T)
- geometry sweep across the mid band (128x64, 64x128, kK=128, 256-row CTAs) measured and rejected: all lose to the 128x128 fast loop; the remaining mid-band gap concentrates in the 1.0-1.4 wave dip (1280^3-class shapes, ~105T vs cuBLAS 186T), which is a scheduling problem (split-K), not a geometry one

Benchmark: L20 (sm_89), CUDA-graph e2e: 1024^3 106.3->107.9T, 1152^3 133.3->137.1T, others unchanged (512^3 54.2T, 2048^3 177.9T, 4096^3 192.9T, 8192^3 198.2T). C++ four-layout suite and 114 targeted pytests pass.
2026-08-26 15:02:00 +08:00
ViperEkura 1d57588d27 fix: free default ubuntu uid/gid before creating astrai user
- ubuntu:24.04 base image ships a default 'ubuntu' user/group at uid/gid 1000, so the default USER_UID/USER_GID build args collided and docker build failed at groupadd with "GID '1000' already exists"
- remove the default ubuntu user/group first (tolerating images without it) so the astrai non-root user is created with the host uid/gid as intended
2026-08-26 15:02:00 +08:00
ViperEkura a92bf79295 perf: fast interior loop on the big cta and fused epilogue bias
- re-enable kFastLoop on the 128x128 CTA for congruous layouts: the base-pair fragment addressing freed the registers the old offset tables spilled, and the predication-free interior loop now wins across the band (fast body 142 SASS instr with zero predicated fallback vs 719/136 generic; 128 regs, no spill)
- move the big/small CTA dispatch boundary from 3/4 to 5/8 wave: with the fast big-CTA loop the crossover sits between 49 and 63 tiles (63-tile rect +8%, 1024^3 now takes the big CTA)
- fuse the linear bias into the GEMM epilogue: FP8Params.bias_ptr adds in fp32 before the single bf16 rounding, replacing the separate out + bias elementwise pass; guarded loads keep N tails exact and batch broadcast falls out of the row-major layout
- resolve Python None bias in the pybind layer (py::object + cast) so ops/fp8.py and fp8.py pass the argument through untouched; drop the _empty_bias sentinel machinery
- add fused-bias tests covering odd N tails, no-bias parity and batched broadcast

Benchmark: L20 (sm_89), CUDA-graph e2e. Big-CTA fast loop + dispatch: 1024^3 102.6->106.3T, 1152^3 128.5->133.3T, 2048^3 173.8->178.2T, 3072^3 180.2->185.3T, 8192^3 196.2->197.7T. Bias fusion (with-bias GEMM vs unfused out + bias): 1024^3 90.5->106.1T (+17%), 2048^3 162.2->178.3T (+10%), 4096^3 178.2->191.1T (+7%). Fused bias differs from the split path by <=1 bf16 ulp and is closer to the fp64 reference. 596 tests pass.
2026-08-26 14:52:06 +08:00
ViperEkura f7d96455a5 perf: base-pair fragment addressing and full-ring small cta
- replace the a_off/b_off per-lane offset tables with two loop-invariant lane bases; ldmatrix fragments now address [base + immediate] with the k_seg step as a single XOR (0x20), mirroring cuBLAS SASS mechanism 1
- drop the 16-register offset table that pushed the kernel past the 128-reg budget and forced per-k-tile address rematerialization; hot-loop integer instructions 444 -> 359 (128x128), immediate-addressed LDSM 8 -> 13/16
- small CTA switches from the lean ring (two __syncthreads per k-tile) to the full ring (one barrier, cuBLAS's structure): s2/24KB below one 3-CTA wave, s3/32KB above
- remove the kAheadFrag cross-k-tile fragment pipeline after measurement (neutral to -8%); mechanism recorded in perf/fp8_gemm_optimization.md

Benchmark: NVIDIA L20 (92 SM, sm_89), torch 2.11.0+cu128, kernel-level event timing on one idle GPU, extension rebuilt from source before each run.
- 2048^3 171.0 -> 172.7 TF (+1.0%), 4096^3 191.0 -> 195.9 (+2.6%), 8192^3 ~190 -> 202.5 (+4.7%)
- 512^3 48.3 -> 50.2 (+3.9%), 1024^3 99.1 -> 101.4 (+2.3%), 1280^3 102.2 -> 106.9 (+4.6%)
- e2e mm_fp8 CUDA-graph: 1280^3 107.6 T, 2048^3 173.8 T, 8192^3 196.2 T
- numerics unchanged: accumulation order identical, per-shape precision equal to the committed baseline (594 pytest, 4-layout C++ suite, short-K and ragged repros all pass)
2026-08-26 14:10:14 +08:00
ViperEkura 8cfe7536ea perf: retarget fp8 gemm tile dispatch and prune dead configs
- replace the sm_count*14/3 small-shape threshold (calibrated on a 24-SM part, so 429 tiles on the 92-SM L20) with a wave-quantization-aware rule: 128x128 CTA for tiles in [3/4, 1] wave or >= 1.4 waves, 64x64 below and inside the just-past-one-wave dip where the finer grid fills the tail
- add a predication-free fast interior loop (kFastLoop) for the 64x64 small CTA: XOR-folded chunk addresses cut ~9 to ~3 instructions per loaded chunk
- delete the dead staged-B pipeline family and launcher dead branches (gemm.cuh 982 -> 801 lines), unused since 5745c2f

Benchmark: L20 (92 SM, sm_89), e2e CUDA-graph TF/s vs prior dispatch: 1152^3 103.1 -> 131.2 (+27%), 1536^3 112.1 -> 138.5 (+24%), 2048^3 123.4 -> 173.0 (+40%); 512/768/1024/1280/3072/4096 cubes unchanged within 1%; 594 pytest + C four-layout tests pass.
2026-08-26 08:24:51 +08:00
ViperEkura a8b63fa362 chore: untrack fp8 sweep tool
- csrc/tests/fp8_sweep.cu is a local measurement tool, not shipped code; git rm --cached keeps the working-tree copy
- the file shows up untracked in git status (allowlist gitignore is untouched by design); never stage it
2026-08-26 07:22:50 +08:00
ViperEkura 4d6a244093 perf: fp8 batched gemm and measured dispatch table
- mm_fp8 accepts 3D operands through the same signature: grid.z slices by batch strides, size-1 batches broadcast (stride 0), inner .t() views fold into the layout tag at zero copy
- fix _LinearFp8 backward crash on 3D [B,L,d] training inputs (flatten before mm_fp8, reduce grad_b over leading dims)
- expose kRasterGroup/kStreamOut as template knobs; drop the 64x128 mid CTA and staged crosswise-B path from dispatch (direct wins everywhere re-measured, including DRAM-streamed B)
- dispatch thresholds grounded in fresh sweeps: m<=64 -> 64x64 CTA (+27% at 64x8192x2048), small-CTA crossover at SM*14/3 total tiles (+13% at 96 tiles), threshold counts batch x per-matrix tiles (+31% at 64x512^3 bmm, +25% at 8x1024x2048)
- remove scripts/tools/bench_fp8_gemm.py (superseded by csrc/tests/fp8_sweep.cu for kernel-level tuning)

Benchmark: NVIDIA L20, E4M3, NT pre-quantized, median of 100-200 iters
- 64x8192x2048: 29.1 -> 22.8 us (94 TF/s)
- 1024x1536x2048: 67.4 -> 59.6 us (108 TF/s)
- bmm 64x512^3: 139.8 -> 106.7 us; bmm 8x1024x2048: 186 TF/s
- regression-free: 4096^3 192 TF/s, 8192^3 200 TF/s, 512^3 unchanged
2026-08-26 06:49:25 +08:00
ViperEkura 01eacbde51 perf: speed up fp8 gemm across small and large shapes
- parameterize warp tile (WarpM/WarpN) in Fp8GemmTraits; MMA loops, fragment arrays and epilogue scale with kMt/kNt instead of the fixed 64x32/4x4, enabling cuBLAS-style 64x64 CTAs of 32x32 warps
- dispatch by output tiling (grid-searched via csrc/tests/fp8_sweep.cu): fewer than 48 output tiles take 64x64/32x32 with a lean ring (4 CTAs/SM fill the wave-quantization gap: 512^3 goes 16 -> 64 CTAs); larger shapes keep 128x128 with the kStages+1 ring
- kStages+1 canonic ring rotation drops the post-compute barrier on the congruous path (one __syncthreads per k-tile); LeanRing keeps the kStages ring for the small CTA; direct-crosswise operands always rotate kStages+1 (their prefetch issues right after barrier 1 and would race a lean ring - caught by the pure C layout suite)
- stage the bf16 epilogue through the reclaimed operand smem: swizzled scatter + barrier + coalesced 16B copy-out replaces 8 disjoint 16B per-warp segments (~50% write efficiency before)
- hoist per-lane ldmatrix swizzle offsets out of the mainloop (stage-relative table + ring-base add) so the innermost loop stops recomputing IMAD/LOP3 address chains
- bypass the torch.library dispatch for real CUDA tensors in quantize/mm_fp8 wrappers (~5us/call, ~40% of a 512-wide call's wall time); fake/subclass tensors keep the custom_op route

vs the previous kernel + python path, wall clock on NT squares: 512^3 52 -> 13us (4.0x, 5.2 -> 20.5 TF, now 1.36x cuBLAS _scaled_mm), 1024^3 1.05x, 2048^3 1.02x (46.9 -> 48.2 TF kernel-only); correctness: 4 layouts x 6 shapes pure C suite PASS, 588 pytest PASS
2026-08-25 22:24:51 +08:00
ViperEkura 057c0d33df refactor: split quantize into multi-type primitive
- split quantize into quantize.cuh, templated on input type (bf16/fp16/fp32)
- rename pybind entry quantize_bf16 to quantize; validate the fmt enum
- fix fp8x2 packing: one 32-bit word packs two pairs (halves were dropped)
- drop the dead OutFp8 template param; GEMM output is always bf16
- fp8_state.reset() restores recipe/format defaults too (test state leak)
- rewrite tests for the two-primitive API with fp32-domain amax references
2026-08-25 20:07:40 +08:00
ViperEkura 3e57cc8069 feat: fp8 backward reuses pre-quantized operands
- g/x/w may each be bf16 or pre-quantized fp8 matching fmt; a pre-quantized operand skips its quantize kernel
- snapshot sx/sw/sg before the ring finalize overwrites the aliased scale slot so the gemm dequantizes with the quantize scale
- forward carries its scale to backward so gradients reuse the forward's scale
- grad_input/grad_weight forced bf16; a pre-quantized g dequantizes before the bias-sum
- regression test: two delayed steps with a changing amax must not leak the scale ratio
2026-08-25 17:24:20 +08:00
ViperEkura 2eeac02d70 perf: overlap direct crosswise loads with fp8 gemm mma
- Rotate kStages+1 canonical buffers for direct-crosswise operands so their load issues right after barrier 1 and its LDG+PRMT latency overlaps the MMA phase, instead of stalling the post-compute inter-barrier window (ncu on the production 24L/dim1536/ffn6912 shapes: barrier 3.7-4.1 + long-scoreboard 1.6-1.8 stalls per issue before).
- Split the stage loaders into load_async (cp.async operands, committed after the post-compute barrier) and load_direct (ring-indexed by tile, not stage, since the rings differ in depth).

On the production shapes vs cuBLASLt _scaled_mm: dX 1.11-1.24x (from 1.27-1.42x), dW 1.32-1.35x (from 1.26-1.49x); generic 4096/11008 shapes also improve (dX 37.6->40.8 TF, dW 36.0->37.3 TF). fp8 train step vs bf16: ~1.2x at both 512 and 2048 tokens.
2026-08-25 15:39:44 +08:00
ViperEkura 5e76fbd1bf perf: fp8 rings, lean autocast, gemm staging
- Finalize scale rings inside the quantize kernels: a last-block epilogue (threadfence + counter elect) folds amax into hist, reduces the window and publishes the next scale on device, zero extra launches; _ScaleRing packs [hist | scale | counter] into one CUDA buffer.
- Split FP8QuantizeParams out of FP8Params so each operator owns its fields; linear_forward/backward_fp8 take optional ring arguments.
- Drop the inference weight-quantization cache; the optimizer bumps the weight version every step, so a cache would miss anyway.
- Zero amax scratch via empty + cudaMemsetAsync instead of torch::zeros, cutting a ~50us fill_ dispatch per quantize.
- Stage crosswise-B operands K-major with cp.async (contract >= 8192) and PRMT-transpose per k_seg region in smem, interleaved with the MMAs; the sync LDG + byte-scatter path it replaces was long-scoreboard bound (ncu 4.6 vs 0.4 stalls/issue).
- Load crosswise-A direct with an in-register PRMT transpose; its operands are typically L2-resident and the staging round trip measured as a net loss.
- Enable grouped rasterization for the congruous NT forward (shared B stripe keeps the weight operand hot in L2) and make the smem budget layout-aware (Fp8GemmSmem) while holding two CTAs per SM.
- Annotate ops/fp8.py return types; drop weight-cache and decorator tests, hoist their imports to module level.

e2e 12L/dim1024/B4xT512 fused AdamW: fp8 137.8ms/step vs bf16 210.3ms, 1.53x. Kernel vs cuBLASLt _scaled_mm: fwd 1.03-1.09x, dX 1.33-1.47x, dW 1.30-1.39x (from 1.10/1.42-1.49/1.52-1.56x), before the pre-transposed copies cuBLASLt needs for dX/dW. fp8 train step vs bf16: 1.34x at 2048 tokens (was 1.25x), 1.08x at 512.
2026-08-25 14:24:11 +08:00
ViperEkura 4dc5e923e0 perf: finalize fp8 scale rings inside quantize kernels
- last-block epilogue (threadfence + counter elect) folds amax into hist[idx], reduces the window and publishes the next scale on device — zero extra launches per linear layer
- _ScaleRing packs [hist | scale | counter] into one CUDA buffer; the eager hist-write / max / scale-copy chain and update() are gone
- split FP8QuantizeParams out of FP8Params so each operator owns its fields; linear_forward/backward_fp8 take optional ring arguments
- e2e 12L/dim1024/B4xT512 (fused AdamW): fp8 137.8ms/step vs bf16 210.3ms, 1.53x; fwd 1.82x, bwd 1.50x
2026-08-25 11:12:14 +08:00
ViperEkura 998b443aa3 refactor: dedupe fp8 meta state into per-operand rings
- collapse FP8TensorMeta's 12 slots + 6 copy-paste methods into three _ScaleRing objects (hist/idx/scale/initialized + update/seed)
- skip meta allocation entirely on the DynamicScaling path (zero rings, scales measured inline)
- drop write-only FP8State._last_device and unused E4M3_MAX alias
2026-08-24 21:19:45 +08:00
ViperEkura cebdd45d3a perf: batch crosswise stage loads in fp8 gemm
- load_operand_tile ColMajor path issued one LDG then immediately scattered 16 byte-granular shared stores, so every store waited on the preceding global load; the runs of one row group now batch into registers first (v[kPasses]) and scatter after, overlapping the LDG latencies
- hoist pass-invariant predicates: the alignment check folds to one uniform (base | ld) & 15 test since r0 is always a multiple of 16, and rows_full leaves the per-pass condition; the contract tail zero-fills without global traffic
- RowMajor path hoists the row bound and the (invariant) chunk-alignment check out of the per-chunk loop
- measured (cuda events, old/new interleaved): crosswise bwd gemms +5-9%, RowMajor and fwd NT within noise; model step unchanged (in the 563-617 ms band)
- file passed through clang-format with the new .clang-format config
2026-08-24 19:49:10 +08:00
ViperEkura 7da1439c9e feat: static fp8 weights and bias with fused epilogue
- linear_forward_fp8 accepts pre-quantized w8 (matching fmt) and skips the weight quantize; amax_w returns 0 on that path since no bf16 values are seen
- bias is now fused into the GEMM epilogue for both dtypes, replacing the separate torch-level add (one elementwise kernel per linear removed)
- FP8Params.bias becomes void* with a new bias_scale slot: null scale = raw bf16 bias, non-null = fp8 storage dequantized in the epilogue after the operand scaling and before any output quantization
- ops/fp8.py relaxes the w dtype check to bf16-or-fp8 and passes bias_scale through
- regression test covers w8/b8, w8/bf16-bias and the amax_w = 0 contract vs an explicit quantization reference
2026-08-24 19:25:23 +08:00
ViperEkura 29e5f571af fix: own fp8 linear backward via autograd Function
- backward used to read the global fp8 flag at loss.backward() time, so calling it outside fp8_autocast silently fell back to bf16 mm (953 ms cublas per step, 49.9% of the model step)
- _LinearFp8(torch.autograd.Function) now owns the fwd/bwd pair: forward captures fmt/recipe/meta on ctx inside the autocast region, backward reads only ctx (scales from the meta rings, masks from ctx.needs_input_grad), so backward is fp8 wherever it runs
- register the aten::linear impl on AutogradCUDA (replaces torch's generated linear formula that calls aten::linear_backward into the bf16 fallback) and keep the CUDA key for inference_mode
- drop the aten::linear_backward override and fp8_linear_backward (dead paths)
- regression test asserts the fp8 backward fires outside the autocast region and grads match the bf16 reference by direction/norm (E5M2 noise)
- model step (0.67B, CE loss, batch 4x1024): backward GEMMs 953 -> 618 ms (1.54x), full step ~1.2x
2026-08-24 19:03:01 +08:00
ViperEkura 74e694921c perf: speed up fp8 gemm tiles and scheduling
- K tile 32->64 (new default): fewer barriers, more MMA per stage; generalize tile_at swizzle and load_operand_tile accordingly
- 64x128 small-M CTA for m<=64 (2x at 64x4096x4096)
- L2 rasterization for crosswise-A layouts (+6..21%)
- micro-bench: NT 4096^3 +35%; linear fwd 1.24-1.76x, bwd 1.71-2.27x vs bf16
- add csrc/tests/fp8_test.cu (single MMA demo + GEMM layouts x K-tiles vs CPU reference)
2026-08-24 18:29:55 +08:00
ViperEkura d5067af064 refactor: harden param PODs and CUTLASS-style fp8 layout tags
- NSDMI null/-1 defaults for AttentionParams/FP8Params pointer+flag members: partially packed structs can no longer hold garbage non-null pointers that gate optional paths (root cause class of the paged test bug); still aggregates, still trivially copyable
- move per-lane ldmatrix wrappers (ldsm_x2/x4) from fp8/gemm.cuh to common/mma.cuh as ldmatrix_x2_lane/x4_lane, next to the single-address variants
- DEVICE_FORCEINLINE macro in common/mma.cuh (matches layout_policies.cuh, internal linkage)
- frag_addr now delegates to tile_at: the swizzle math has one source
- operand layouts as CUTLASS-style RowMajor/ColMajor tags threaded from launch_fp8_gemm through the kernel to load_operand_tile; B's operand view via transpose_layout_t; call sites read <Fmt, false, RowMajor, ColMajor> instead of <Fmt, false, false, true>
2026-08-24 15:27:52 +08:00
ViperEkura f6db546578 fix: zero-init AttentionParams in pure C tests
- paged decode test left new_k_ptr/new_v_ptr as stack garbage; PagedKV::decode_addr then took the new-KV path on wild pointers (illegal access or wrong last-token K/V)
- value-init the POD (= {}) at every construction site
2026-08-24 14:56:19 +08:00
ViperEkura 31ca357c61 refactor: namespace csrc kernels and extract common helpers
- attention family -> astrai::attention; fp8 family -> astrai::fp8
- new common/reduce.cuh (warp/group reductions, atomic_max_float)
- new common/cp_async.cuh (predicated cp_async_16, commit/wait group)
- move MAX_SPLITS into attention/common.h; delete warp_utils.cuh
- .cu bindings and pure C tests open family namespaces via using
2026-08-24 14:49:55 +08:00
ViperEkura 34471252ab perf: pipeline fp8 gemm fragment loads and pack bf16 epilogue
- software-pipeline A-fragment ldmatrix: row mt+1 loads hide behind row mt MMAs
- bf16 epilogue packs two columns into one bfloat162 store (half the stores)
- fp8 vs bf16 linear: fwd 1.15x@512, 1.5x@2048, 2.8x@4096; bwd up to 2.7x, peak 35 TFLOPS
2026-08-23 21:32:20 +08:00
ViperEkura aa08479285 perf: widen fp8 gemm tile and load fragments with ldmatrix
- 128x128 CTA of 8 warps x 64x32 warp tiles: 16 mma.sync per warp per K-segment (was 8)
- ldmatrix.x4/x2 with per-lane swizzled addresses replaces 36 scalar LDS per warp-tile step
- __launch_bounds__(256, 2) caps registers at 124 so two CTAs fit per SM
- fp8 linear vs bf16 cuBLAS: fwd 1.07x->2.65x, bwd 1.41x->2.65x by size, peak 31-33 TFLOPS
2026-08-23 21:13:09 +08:00
ViperEkura 4b10d3ca37 perf: vectorize fp8 quantize and swizzle gemm smem 2026-08-23 20:31:44 +08:00
ViperEkura 2bc4d2b8a8 refactor: unify kernel module loading and packaging
- loader.py: lazy/cached import; is_available defers the actual load; get_module raises on unavailable
- ops/{attention,rotary,fp8}: use get_module instead of touching private _modules or their own _mod() cache
- package-data: ship astrai.extension.lib *.so in built wheels (non-editable installs previously lost every kernel)
2026-08-23 15:57:04 +08:00
ViperEkura 4244df2785 perf: pure FP8 fwd/bwd and lean non-transposed GEMM
- drop the fused kernel; forward/backward are quantize + a pre-quantized GEMM
- rename module fp8_mm -> fp8_ops (mm.cu -> ops.cu)
- kernels/launchers fp8_gemm_kernel / launch_fp8_gemm; drop PqTraits/gather_trans/pack_fp8x4_vector
- remove the in-kernel transposed-operand branches (TransA/TransB)
- backward: quantize g once (amax_g here), explicit fp8 transposes, fast non-transposed GEMMs (dX = g@w^T, dW = g^T@x^T)
- each pass uses a single FP8 format (E4M3 fwd / E5M2 bwd)
2026-08-23 15:38:30 +08:00
ViperEkura a29bdfae46 refactor: rework attention backend resolution
- explicit attn_backend() context wins over ASTR_BACKEND env
- polymorphic available()/supports_call() replace isinstance dispatch
- cache singleton backend instances to avoid hot-path allocation
- training (fwd=None) resolves cuda > flash > torch by capability
- flash dense supports mask-free calls only; masked training falls back to torch
2026-08-23 14:47:02 +08:00
ViperEkura 10fec8dca1 docs: fix stale docs and align with code
- update cuda_kernels layout, arch flags, and add FP8 section
- fix install docs: kernels auto-build when nvcc + CUDA detected
- mark ignored OpenAI request params and complete KVCache fields
- add docker docs to indexes and astrai.optim to module overview
- refresh document update timestamps
2026-08-23 14:23:28 +08:00
ViperEkura 75304d084d refactor: unify CUDA skip guards in tests
- add skip_no_fp8 (CUDA + fp8_mm kernel + cc 8.9+) to tests/conftest.py
- use skip_no_cuda / skip_no_kernel / skip_no_fp8 directly in test modules
- drop _GPU alias and tests.extension.conftest re-exports
- remove unused imports (Union in hf_adapter, make_grpo_config in data conftest)
2026-08-22 21:08:06 +08:00
ViperEkura 16a55bb474 refactor: reorganize CUDA kernels into per-family directories
- move attention kernels to csrc/kernels/attention/ and rotary to rotary/
- add shared common/mma.cuh (mma_sync, ldmatrix) and device.cuh (sm checks)
- split fp8_mm into three-layer fp8/common.h, gemm.cuh, mm.cu
- fix fused FP8 GEMM ldmatrix lane indexing to fix OOB shared reads
- update extension ops, loader, and kernel tests
2026-08-22 20:40:31 +08:00
ViperEkura cb21af38ba feat: unify Docker serving configuration in YAML
- Add server.py --config serve.yaml; explicit CLI flags override YAML
- Add scripts/serve.sh and serve_runtime.py for the Compose lifecycle
- Template server/cpu ports and param mounts in docker-compose.yml
- Document schema in docs/developer/docker-serving.md and params guide
- Add tests for runtime parsing and server CLI merge logic
2026-08-21 23:16:45 +08:00
ViperEkura dcc96de12a test: refactor tests and fix inference edge cases
- Convert protocol and MoE test classes to plain functions
- Add real server/engine integration and generate_async tests
- Isolate test model per test and use pytest tmp_path
- Reset FastAPI engine state after inference tests
- Fix generate_async StopIteration handling on Python 3.12
- Fix HF adapter MoE dense/shared and Gemma qk_norm mapping
- Correct dev dependency httpx2 to httpx
2026-08-21 22:59:51 +08:00
ViperEkura 7d27f3e078 feat: load HuggingFace checkpoints via key/config conversion
- Add astrai.serialization.hf_adapter mapping LLaMA-style HF keys to AstrAI names (input_layernorm, gate_proj, MoE experts/shared_experts) with config aliases for dense and MoE (Mixtral/DeepSeek-V3) layouts; reject biased projections, mismatched head_dim and MLA
- Give AutoModel.from_pretrained weights_format=auto|astrai|hf with auto-detection; read sharded safetensors via model.safetensors.index.json
- Adapt preloaded weights/config in train_context and benchmark CLI
2026-08-20 11:34:59 +08:00
ViperEkura 84753d3e08 refactor: deduplicate and restructure test suite
- extract preprocessing config factories into tests/data/factories.py
- keep conftest.py fixtures-only; stop importing builders from it
- promote temp_dir fixture to root conftest for cross-directory reuse
- unify duplicate BPE tokenizer builders into build_test_tokenizer
- merge grpo/dpo online e2e tests into one parametrized integration test
- extract engine mock factory and shared model batch builders
- drop local tempfile usage in favor of shared fixtures

No behavior change: 519 tests pass.
2026-08-20 01:53:28 +08:00
ViperEkura 53a7149577 feat: add distributed rank to logs 2026-08-20 01:37:00 +08:00
ViperEkura c79d34eee1 refactor: simplify training and inference interfaces
- avoid constructing model_fn more than once when reading config
- keep inference package exports focused on public entry points
- rename extra strategy arguments to strategy_kwargs
2026-08-19 20:55:13 +08:00
ViperEkura 398e8a3ea3 refactor: deduplicate low-risk code paths 2026-08-19 16:17:40 +08:00
ViperEkura f252af495c refactor: remove dead code and deduplicate scheduler setup 2026-08-19 14:58:37 +08:00
ViperEkura 00c2c80c8f feat: unify Docker training configuration in YAML 2026-08-19 14:04:35 +08:00
ViperEkura a6c6a54ace feat: read host training vars from YAML infra section
- scripts/train.sh load_infra() parses the top-level infra: section of TRAIN_CONFIG_FILE
- exports TRAIN_JOB_NAME/DATA/MODEL/CHECKPOINT_DIR/TRAIN_GPU_COUNT/CUDA_VISIBLE_DEVICES
- infra overrides .env.train via compose interpolation precedence; keys absent fall back
- train.yaml is now the single per-job config: host mounts, GPU filter, and hyperparameters
- requires host python3 with PyYAML when TRAIN_CONFIG_FILE is set; errors fail fast
- docs: docker-training.md documents the infra overrides and precedence
2026-08-19 01:27:57 +08:00
ViperEkura 3d3ea47d37 refactor: standardize packed 3d inference
- keep training attention on dense 4d tensors
- use packed 3d tensors with KV cache for inference
- extend CUDA rotary embedding to packed 3d inputs
- adapt torch, CUDA and FlashAttention backend dispatch

Benchmark: NVIDIA L20, BF16, 1B model, paged KV cache, CUDA Graph, prompt 512, generation 128 (median of 3 alternating runs)
- batch 1: 234.5 -> 242.6 tok/s (1.034x, +3.4%)
- batch 8: 1243.1 -> 1286.6 tok/s (1.035x, +3.5%)
2026-08-19 00:36:53 +08:00
ViperEkura f7f14d0e5f test: add online GRPO end-to-end training test 2026-08-19 00:07:16 +08:00
ViperEkura 7580d80d45 perf: accelerate FP8 backward with fused fast kernel
- route dX/dW through the fused 128x64 fast kernel via contiguous transposes
- drop the legacy 64x64 kernel, cutting dX 1.55->0.38 ms and dW 1.28->0.26 ms
- sync all threads after cp.async.wait_group to fix sporadic NaN in large GEMMs
- add fp8_mm_prequant_fp8 custom op for FP8-in/FP8-out GEMM
2026-08-18 23:46:43 +08:00
ViperEkura cb51a3587b perf: optimize fused FP8 GEMM kernel 2026-08-18 19:56:15 +08:00
ViperEkura 1bcd8f53ab perf: precompute ragged Q tile scheduling 2026-08-16 23:32:46 +08:00
ViperEkura 0d0dc64884 docs: explain extension layer boundaries 2026-08-16 21:30:03 +08:00
ViperEkura 6ac3b51496 refactor: separate extension ops and backends 2026-08-16 21:15:52 +08:00
ViperEkura 3406157431 refactor: standardize packed 3d inference
- keep training attention on dense 4d tensors
- use packed 3d tensors with KV cache for inference
- extend CUDA rotary embedding to packed 3d inputs
- adapt torch, CUDA and FlashAttention backend dispatch
2026-08-16 13:24:02 +08:00
ViperEkura 0dd9a417b7 refactor: separate KV token address resolution 2026-08-15 22:59:35 +08:00
ViperEkura a01c1fd427 perf: bypass L1 for attention tile loads 2026-08-15 21:23:20 +08:00
ViperEkura f8d9ab344d refactor: remove unused streaming dataset 2026-08-15 20:55:08 +08:00
ViperEkura 3fb4b8ab13 perf: use int32 paged KV indices
- store page-table, request-row, and cache-location indices as int32
- preserve CUDA graph replay with bit-exact logits and KV cache coverage
- improve B=1 decode latency by 1-6% across 1K-32K contexts on L20
2026-08-15 13:17:06 +08:00
ViperEkura b5afe3d7a4 perf: optimize small-head causal prefill
- map D=32 and D=64 causal prefill to BC=64 tiles

- add small-head correctness and benchmark coverage
2026-08-14 23:25:49 +08:00
ViperEkura 69f35c46e0 fix: quantize amax from raw values, not scaled fp8 values
- amax for delayed scale was the quantized max (always ~448), so scale collapsed to 1
- this made fp8 gradients diverge (cosine 0.05) and training stall
- stop w/x transpose-quantize amax from polluting the grad scale
2026-08-14 14:26:16 +08:00
ViperEkura 0378e62e17 refactor: split fp8 into fp8_ops adapter and fp8 policy module
- fp8_ops is the only module touching the pybind (kernel interface)
- fp8.py keeps scaling state, delayed amax and aten::linear dispatch
- remove circular imports between old fp8_ops/fp8_state/fp8_dispatch
2026-08-14 12:25:29 +08:00
ViperEkura 5244f1a8fc feat: add te-style scaled fp8 training via fp8_autocast
- per-tensor scales applied inside cublasLt via A_SCALE/B_SCALE
- delayed scaling: weight amax history ring, refresh every 16 steps
- quantize kernels emit atomic amax, device-side scale updates
- fp8_autocast context toggles aten::linear dispatch like torch.autocast
- fallback to bf16 when M/N not 16-aligned (fp8 gemm constraint)
- x/g scales delayed one step, reuse free atomic amax (no abs/max reduce)
2026-08-14 12:14:04 +08:00
ViperEkura 5104638447 perf: use fp8 tensor-core gemm in linear backward
- dX/dW run as fp8 cublasLt gemms via fused transpose-cast
- shared (m,k,n) algo cache for fwd/bwd, mutex-protected
- bias add in-place on bf16 output, drop output copy
2026-08-14 10:43:03 +08:00
ViperEkura a711d9f478 perf: eliminate gemm output transpose via A/B swap
- pass w as param A (op=T) and x as param B (op=N) so the col-major [N,M] output storage is row-major C[M,N] directly, zero copy
- transpose_bias_cast kernel becomes a plain bias+write kernel
- fp8 e2e now beats bf16: 1.09x at M=4096, 1.06x at M=8192 (was 0.88x)
2026-08-14 01:42:09 +08:00
ViperEkura 15862d4b56 perf: fuse fp8 linear fwd and bwd into single kernel calls
- fp8_linear_forward: cast + cublasLt GEMM + transpose + bias in one call
- fp8_linear_backward: scale-free, dtype derived from input tensor
- drops per-op Python dispatch (was ~6-8 launches per linear) and amax syncs
- 1024x1024 linear: 6.8x slow -> 0.67x (36.7us vs 24.8us bf16)
- small-model e2e still 1.71x slow; 15bt estimate ~0.78x (linear-heavy)
2026-08-14 01:24:37 +08:00
ViperEkura f9efb705b8 perf: output fp8 gemm in bf16 instead of fp32
- cublasLt C layout and buffer switched to CUDA_R_16BF, halving output bandwidth
- downstream ops (RMSNorm etc.) keep matching bf16 dtype, fused kernels stay
- numeric error unchanged (0.19% vs fp32 ref on quantized inputs)
2026-08-14 01:08:26 +08:00
ViperEkura c6a82a5029 refactor: align linear backward dtype with weight
- cast gradients and inputs to weight.dtype instead of hardcoded bf16
- single code path covers bf16 and fp32 models, no branch needed
- gradient dtype now matches the leaf parameter dtype exactly
2026-08-14 01:01:58 +08:00
ViperEkura a5b238dd86 feat: add fp8 training via cublasLt dispatch
- fp8_mm kernel (csrc): cublasLt fp8 e4m3 gemm, TN layout mapped zero-copy
- custom::fp8_mm custom op: meta/cuda/cpu kernels + scale-corrected bf16 autograd
- aten::linear and linear_backward dispatch on CUDA key, zero model changes
- per-tensor scale or raw cast; single-GPU smoke loss matches bf16
2026-08-14 00:39:49 +08:00
ViperEkura da6d94492d fix: parse yaml floats with yaml 1.2 schema
- register yaml 1.2 float resolver so scientific notation (2e-5) becomes float, not str
- replaces the decimal-point workaround in train configs
- add containerized training doc under docs/developer
2026-08-13 23:28:06 +08:00
ViperEkura 71b6e3aaaf feat: rework docker workflow for gpu-first training
- rewrite docker.sh with gpu default and --no-gpu override
- inject host uid/gid via ASTRAI_UID/GID in train.sh compose()
- filter readonly UID/GID lines when sourcing .env.train
- build image user via USER_UID/USER_GID args matching host uid/gid
- pass all GPUs (count: all) and filter by CUDA_VISIBLE_DEVICES inside the container
- forward NCCL vars through compose environment
2026-08-13 22:51:13 +08:00
ViperEkura f95722a277 feat: add containerized training workflow
- add a GPU trainer Compose profile with mounted data, models, and checkpoints
- add host commands for preflight, lifecycle, logs, status, and checkpoint cleanup
- resume from the latest complete checkpoint with external config or CLI arguments
2026-08-12 20:21:33 +08:00
ViperEkura 9f48cb8928 refactor: streamline Q block mapping
- bypass shared mapping for contiguous attention
- centralize paged Q tile broadcast in KV policy helpers
2026-08-10 08:40:18 +08:00
ViperEkura 9b58fef222 refactor: extract QTileMapper for prefill tile dispatch
- wrap one-thread map + shared broadcast + early exit
- both scalar and MMA prefill kernels use the shared helper
2026-08-09 23:18:03 +08:00
ViperEkura c5fba9c238 perf: flatten paged prefill tile dispatch
- remove the host-provided max_q_len argument
- dispatch only the ragged prefill tile upper bound
- validate the rebuilt CUDA backend end to end
2026-08-09 23:12:53 +08:00
ViperEkura cd31f1f62f refactor: tidy attention params and launcher interfaces
- rename output pointer field o to o_ptr for consistency with q_ptr/k_ptr/v_ptr
- regroup AttentionParams fields by responsibility and fix misleading comments
- drop unused max_seq_len/total_q fields and paged decode max_seq_len arg
- drop redundant group_size param from decode launchers (computed from p)
2026-08-09 20:52:06 +08:00
ViperEkura a5a3cc1fc2 refactor: unify attention param field names
- rename q_stride_* to q_*_stride to match mask stride convention
- rename mask_q_stride to mask_l_stride for consistent l-dim naming
- merge k/v and k_cache/v_cache into k_ptr/v_ptr; rename q to q_ptr
- KVSource policy selects contiguous vs paged mode at compile time
2026-08-09 20:23:58 +08:00
ViperEkura d565d44c43 fix: harden attention kernel boundaries
- fix scalar prefill head_dim=32 out-of-bounds via G=4 dispatch
- fix MMA decode 4D mask head indexing and invalid-row mask access
- add q_head/kv_head divisibility and head-dim contiguity checks
- validate split-KV scratch and decode out_buf layout in bindings
- set max dynamic shared memory for scalar decode D=256
- cover scalar prefill D=32 in pure C test
2026-08-09 14:53:24 +08:00
ViperEkura 596c35fd71 fix: report gradient snr in db 2026-08-09 13:40:27 +08:00
ViperEkura 47b3ed4e44 feat: propagate attention backend across scheduler threads
- InferenceEngine/Scheduler accept an explicit backend
- capture request-level attn_backend context onto Task
- split prefill/decode batches by backend instance
- ASTR_BACKEND env overrides ContextVar as process-wide policy
- report resolved backend and CUDA-graph state in benchmark
2026-08-09 13:32:40 +08:00
ViperEkura c1d05ae11d perf: benchmark decode via real inference engine
- route decode benchmark through InferenceEngine generate path
- add enable_cuda_graph toggle to engine, scheduler, and executor
- make benchmark --cuda-graph/--no-cuda-graph control the toggle
- hoist local time imports to module top
2026-08-09 11:47:14 +08:00
ViperEkura cf4f5ab9f6 feat: add persistent DataLoader workers
- Keep training workers alive between epochs when enabled.
- Avoid invalid prefetch settings for single-process loading.
2026-08-09 11:38:50 +08:00
ViperEkura 3416f98c58 fix: wire benchmark cache selection 2026-08-09 10:56:05 +08:00
ViperEkura d28552f878 refactor: use C++17 struct dispatch in csrc tests, tighten paged tolerances to 0.01
- Replace C++20 explicit lambda template parameters with file-scope structs (DecodeDispatch/PrefillDispatch etc.)
- Remove unused gs variable in run_decode_test
- Tighten paged test atol/rtol from 0.02 to 0.01 to match contiguous tests
2026-08-09 10:23:12 +08:00
ViperEkura be90dfe2bd fix: isolate continuous batch decode state
- Match steady-state metadata to the active task IDs
- Rebuild request mappings for cached prefix pages
- Add regressions for batch refill and prefix reuse
2026-08-09 01:01:41 +08:00
ViperEkura a33ca04f60 fix: synchronize final decode async copy
- wait for the final split-KV tile before reading shared memory
- cover long decode with production context capacity
2026-08-09 00:31:47 +08:00
ViperEkura 7f0e8bb8c2 fix: let flash backend handle 4D causal prefill mask
- Treat 4D masks as causal (flash handles it natively), keep rejecting custom non-causal masks
- Enables flash backend in benchmark --compare and real prefill path
2026-08-08 23:51:27 +08:00
ViperEkura 0c1b7664c1 refactor: split infer core into subpackages by concern
- Eliminate core/ directory into cache/, runtime/, network/ subpackages plus flat modules
- Split cache.py (647 lines) into cache/{buffer,strategy,pool}.py by layer
- Add explicit ContiguousStrategy, make AllocationStrategy a real ABC
- Move TaskCacheState to cache/strategy.py, drop string forward references
- Rename api/ to network/, server.py to app.py
- Move sample.py into runtime/ alongside executor and graph
- Simplify TaskCacheManager.__init__ to single pool param
- Expose pool.strategy and pool.req_pool as public properties
- Fix KVCache import in attention_backend.py (TYPE_CHECKING guard)
- Fix steady-state decode reading uninitialized position_ids on first step
2026-08-08 23:43:05 +08:00
ViperEkura 3fa7e66676 refactor: decouple task cache from PagePool and unify steady-state detection
- TaskCacheRegistry -> TaskCacheManager (independent, held by scheduler)
- TaskCacheState co-locates 5 parallel dicts into one dataclass
- AllocationStrategy base class + PagedStrategy subclass (page_size is a parameter)
- _rollback() helper for unified cleanup (no duplicate free paths)
- Task._kv_len + prefill_done property (explicit, no output_tokens proxy)
- Steady-state detection single-sourced in TaskCacheManager.bind()
- PagePool is now pure physical layer (no task knowledge)
- Removed dead _page_to_hash dict in RadixCache
2026-08-08 22:43:06 +08:00
ViperEkura ca50fe4721 refactor: remove inference redundancy and fix cache leaks
- drop Executor unused tokenizer field, _head_dim, stale metrics docstring
- unify greedy sampling via SamplingPipeline.sample, drop top-level duplicate
- drop Task.flush_remaining no-op and unreachable prompt-length branch
- drop ProtocolHandler redundant chunks list (reuse body)
- fix page_size=1 token-slot leak on task_free
- clear _task_pages/_task_slots on alloc-failure paths
- reset _bind_state on task_free to avoid stale steady-state reuse
- remove unreachable contiguous branches in paged-only helpers
2026-08-08 21:45:01 +08:00
ViperEkura d9240ab149 refactor: split train context build steps
- separate checkpoint, model, data, and strategy setup\n- keep build orchestration concise and readable
2026-08-08 18:15:14 +08:00
ViperEkura d7cd69fef5 feat: add streaming IterableDataset for pretraining
- StreamingSeqDataset yields windows sequentially through each shard
- Shard-level shuffle, distributed and multi-worker shard partitioning
- __len__ returns total window count for scheduler total_steps
- Better OS page-cache locality than random-access map-style datasets
2026-08-08 16:15:10 +08:00
ViperEkura 9bff61fb91 perf: use cudaEvent for precise GPU timing in debug logs
- cudaEvent.elapsed_time gives microsecond precision vs perf_counter
- cudaEvent measures actual GPU execution, not just kernel launch
- falls back to time.perf_counter on CPU-only devices
2026-08-08 13:18:11 +08:00
ViperEkura 0b661bae85 fix: remove blocking cleanup from streaming generator
- stream finally froze main thread on cache.task_free
- scheduler handles cleanup in next loop iteration instead
2026-08-08 13:12:40 +08:00
ViperEkura ae9fd546ef perf: merge prefill warmup into _warmup_cuda_graphs
- 64-token prefill forward triggers cuBLAS auto-tuning at init
- reduces first-chat prefill from ~520ms to ~27ms
- warmup decode also drops from ~215ms to ~71ms
2026-08-08 13:10:58 +08:00
ViperEkura e3ea850dc9 fix: default backend race, raise on explicit fallback
- _default_backend lazy init protected with threading.Lock
- attention() raises when explicit backend cannot handle call
- FlashAttnBackend rejects prefill with non-None attn_mask
- training test uses TORCH_NATIVE backend directly
2026-08-08 13:00:58 +08:00
ViperEkura 6e5088cc7d refactor: remove prefill from CUDA graph warmup
- decode capture works without pre-filled KV values
- reduces init time and eliminates unused prefill forward
2026-08-08 12:49:07 +08:00
ViperEkura cbc584470d refactor: centralize logging in astrai.logging, replace ASTRAI_TIMED with log level
- move setup_logging to astrai/logging.py
- timed() now uses logger.isEnabledFor(DEBUG) instead of separate env var
- enable ASTR_LOG_LEVEL=DEBUG to see per-step timing logs
- call setup_logging() in stream_chat.py
2026-08-08 12:39:27 +08:00
ViperEkura cb60713a72 feat: add per-task throughput and latency metrics
- extract TaskTiming + MetricsCollector out of Task/TaskManager
- unify prefill/decode timing into single record() context manager
- expose avg_ttft_ms, avg_decode_tps, avg_e2e_latency_ms via /stats
2026-08-08 12:10:06 +08:00
ViperEkura c52a2487ae fix: rename CUDA wheels with tag suffix to avoid upload clash
- all three release builds (pure, cu128, cu130) produce the same .whl filename, causing uploads to overwrite each other
- append the CUDA tag as a local version label (e.g. +cu128, +cu130)
2026-08-08 00:49:17 +08:00
ViperEkura 49aaa9a714 version: bump to 1.3.13 2026-08-07 23:56:01 +08:00
ViperEkura 056c1382ff docs: sync all documentation with current codebase
- remove GenerationRequest and generate_with_request references (class deleted)
- document cuda>flash>torch default priority and FlashAttnBackend
- add ASTR_BACKEND env var to backend docs, TorchNativeBackend (default) → (fallback)
- fix JsonlStore transform routing → DatasetFactory ownership
- fix CudaBackend fallback chain description (FlashAttn → TorchNative)
- add FlashAttnBackend to architecture strategy table
- add router_stats to DecoderOutput/FFNOutput TypedDict diagrams
- add decode_o_part/ml_part/decode_out to KVCache diagram
- add --append_eos/--no-append_eos to IFD evaluation parameter table
- update get-started CUDA kernel note (no longer requires explicit attn_backend activation)
- fix python -m scripts.tools.server (no __init__.py) → direct script call
2026-08-07 23:52:27 +08:00
ViperEkura f163520fff refactor: break JsonlStore→preprocessing circular dependency
- move JSONL transform auto-creation from JsonlStore.load to DatasetFactory.load via _build_jsonl_transform helper
- remove TokenizeTransform and PipelineConfig imports from storage module
- JsonlStore.load now requires explicit transform= for eager mode
- DatasetFactory.load remains the public API with identical convenience behavior
2026-08-07 23:22:32 +08:00
ViperEkura 1b1f1a0707 fix: add dtype guard to FlashAttnBackend capability check
- _backend_supports now rejects fp32 for FlashAttnBackend (flash-attn only supports fp16/bf16), preventing runtime crash on fallback chain
- rename test_default_backend_is_torch_native to reflect multi-backend reality
- scheduler test fixture uses bf16 model (matches production, avoids unnecessary 3-step fallback chain)
2026-08-07 23:08:22 +08:00
ViperEkura 184fbbce5c refactor: extract shared steady-state increment detection
- add _BindState dataclass and _is_steady_increment() to cache.py
- replace _bind_sig/_bind_seq_lens dual fields with single _bind_state
- replace DecodeSteadyState bare tuple with named dataclass
- use _is_steady_increment() in both PagePool.bind_tasks and Executor.execute_decode
2026-08-07 23:00:25 +08:00
155 changed files with 12224 additions and 4142 deletions
+3
View File
@@ -54,6 +54,9 @@ jobs:
- name: Build wheel (with CUDA kernels) - name: Build wheel (with CUDA kernels)
run: | run: |
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/ CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
for f in dist/*.whl; do
mv "$f" "dist/$(basename "$f" .whl)+${{ matrix.cuda_tag }}.whl"
done
- uses: actions/upload-artifact@v4 - uses: actions/upload-artifact@v4
with: with:
+3 -2
View File
@@ -39,11 +39,12 @@ ruff format . # re-format after fix
python -u -m pytest tests/ -v python -u -m pytest tests/ -v
``` ```
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed. > Failed tests may leave orphan tempdirs under the system temp directory
> (`$TMPDIR` on Linux/macOS, `%TEMP%` on Windows). Clean them manually if needed.
### 4. (Optional) Full pre-commit check script ### 4. (Optional) Full pre-commit check script
If you have Git Bash available: If you have `bash` available (Git Bash on Windows works too):
```bash ```bash
bash scripts/pre_commit.sh bash scripts/pre_commit.sh
+11 -2
View File
@@ -57,8 +57,17 @@ COPY docs/ ./docs/
COPY pyproject.toml . COPY pyproject.toml .
COPY README.md . COPY README.md .
# Create non-root user # Create non-root user matching the host uid/gid (passed via build args).
RUN useradd -m astrai && chown -R astrai:astrai /app # ubuntu:24.04 ships a default 'ubuntu' user/group at uid/gid 1000, so remove
# it first to free those ids before creating astrai.
ARG USER_UID=1000
ARG USER_GID=1000
RUN userdel -r ubuntu 2>/dev/null || true \
&& groupdel ubuntu 2>/dev/null || true \
&& groupadd -g "${USER_GID}" astrai \
&& useradd -m -u "${USER_UID}" -g astrai astrai \
&& chown -R astrai:astrai /app
ENV HOME=/home/astrai
USER astrai USER astrai
ENV PYTHONUNBUFFERED=1 \ ENV PYTHONUNBUFFERED=1 \
+9 -3
View File
@@ -51,7 +51,7 @@ AstrAI is an end-to-end Transformer framework for building, training, evaluating
| **Data** | Declarative JSON preprocessing, configurable masking and packing, binary/JSONL storage, and streaming datasets | | **Data** | Declarative JSON preprocessing, configurable masking and packing, binary/JSONL storage, and streaming datasets |
| **Inference** | Continuous batching, paged KV cache, radix prefix caching, streaming generation, and Torch/CUDA/FlashAttention backends | | **Inference** | Continuous batching, paged KV cache, radix prefix caching, streaming generation, and Torch/CUDA/FlashAttention backends |
| **Serving** | FastAPI server with OpenAI and Anthropic chat completion protocols, including SSE streaming and tool calls | | **Serving** | FastAPI server with OpenAI and Anthropic chat completion protocols, including SSE streaming and tool calls |
| **Evaluation** | Perplexity, MMLU, HumanEval, IFEval, IFD, and ROUGE evaluation tools | | **Evaluation** | Perplexity, MMLU, HumanEval, IFEval, IFD, ROUGE, and weight-analysis evaluation tools |
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components | | **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
### Getting Started ### Getting Started
@@ -65,8 +65,9 @@ AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `sc
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
pip install -e . # pure PyTorch (no CUDA kernels) pip install -e . # kernels auto-build when nvcc + CUDA are detected
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels # CSRC_KERNELS=false pip install -e . # skip kernels (pure PyTorch)
# CSRC_KERNELS=true pip install -e . --no-build-isolation # force the fused CUDA kernel build
# pip install -e ".[dev]" # dev dependencies (pytest, ruff) # pip install -e ".[dev]" # dev dependencies (pytest, ruff)
``` ```
@@ -191,6 +192,9 @@ docker compose up -d
# Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable) # Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
docker compose --profile cpu up -d docker compose --profile cpu up -d
# YAML-driven serving (see serve.yaml; up/run/down/logs/status...)
bash scripts/serve.sh up
``` ```
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`. > **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
@@ -236,6 +240,8 @@ See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error
| [Data Flow](./docs/developer/dataflow.md) | Data pipeline, storage backends & dataset architecture | | [Data Flow](./docs/developer/dataflow.md) | Data pipeline, storage backends & dataset architecture |
| [Internals](./docs/developer/internals.md) | Training internals: loss formulas, callback lifecycle, KV cache | | [Internals](./docs/developer/internals.md) | Training internals: loss formulas, callback lifecycle, KV cache |
| [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks | | [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks |
| [Docker Serving](./docs/developer/docker-serving.md) | YAML-driven containerized serving (`serve.yaml`, `serve.sh`) |
| [Docker Training](./docs/developer/docker-training.md) | YAML-driven containerized training (`train.yaml`, `train.sh`) |
### Contributing ### Contributing
+7 -36
View File
@@ -1,9 +1,6 @@
__version__ = "1.3.12" __version__ = "1.3.13"
__author__ = "ViperEkura" __author__ = "ViperEkura"
import logging
import os
from astrai.config import ( from astrai.config import (
AutoRegressiveLMConfig, AutoRegressiveLMConfig,
BaseModelConfig, BaseModelConfig,
@@ -20,14 +17,10 @@ from astrai.dataset import (
StoreFactory, StoreFactory,
) )
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.inference import ( from astrai.inference import InferenceEngine, get_app, run_server, sample
InferenceEngine, from astrai.inference.network import ProtocolHandler
ProtocolHandler, from astrai.inference.runtime.sample import SamplingPipeline
SamplingPipeline, from astrai.logging import setup_logging
get_app,
run_server,
sample,
)
from astrai.model import ( from astrai.model import (
AutoModel, AutoModel,
AutoRegressiveLM, AutoRegressiveLM,
@@ -55,30 +48,6 @@ from astrai.trainer import (
Trainer, Trainer,
) )
def setup_logging(level: str = "INFO"):
"""Attach a handler to the ``astrai`` logger (only, not root).
Call once per process, e.g. at the top of CLI scripts.
Set ``ASTR_LOG_LEVEL`` to override the default ``INFO``.
"""
_logger = logging.getLogger("astrai")
if _logger.handlers:
return
_level = getattr(
logging, os.environ.get("ASTR_LOG_LEVEL", level).upper(), logging.INFO
)
_logger.setLevel(_level)
_handler = logging.StreamHandler()
_handler.setFormatter(
logging.Formatter(
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
)
_logger.addHandler(_handler)
__all__ = [ __all__ = [
"AutoRegressiveLM", "AutoRegressiveLM",
"AutoRegressiveLMConfig", "AutoRegressiveLMConfig",
@@ -122,3 +91,5 @@ __all__ = [
"setup_logging", "setup_logging",
"spawn_parallel_fn", "spawn_parallel_fn",
] ]
setup_logging()
+16 -14
View File
@@ -11,10 +11,10 @@ from torch.utils.data import Dataset
from astrai.config.base import BaseConfig from astrai.config.base import BaseConfig
from astrai.model.components.lora import LoRAConfig from astrai.model.components.lora import LoRAConfig
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"}) TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"}) PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
_BACKENDS = frozenset({"nccl", "gloo"}) BACKENDS = frozenset({"nccl", "gloo"})
_START_METHODS = frozenset({"spawn", "fork", "forkserver"}) START_METHODS = frozenset({"spawn", "fork", "forkserver"})
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"}) _COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
@@ -48,6 +48,7 @@ class TrainConfig(BaseConfig):
random_seed (int): Random seed. Defaults to 3407. random_seed (int): Random seed. Defaults to 3407.
num_workers (int): Number of workers for dataloader. Defaults to 0. num_workers (int): Number of workers for dataloader. Defaults to 0.
prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None. prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None.
persistent_workers (bool): Keep DataLoader workers alive between epochs. Defaults to False.
pin_memory (bool): Pin memory for dataloader. Defaults to False. pin_memory (bool): Pin memory for dataloader. Defaults to False.
collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None. collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None.
nprocs (int): Number of processes for distributed training. Defaults to 1. nprocs (int): Number of processes for distributed training. Defaults to 1.
@@ -69,7 +70,7 @@ class TrainConfig(BaseConfig):
rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024. rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024.
reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None. reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None.
executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}. executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}.
extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}. strategy_kwargs (Dict[str, Any]): Extra strategy arguments. Defaults to {}.
""" """
model_fn: Callable[[], nn.Module] model_fn: Callable[[], nn.Module]
@@ -98,6 +99,7 @@ class TrainConfig(BaseConfig):
random_seed: int = 3407 random_seed: int = 3407
num_workers: int = 0 num_workers: int = 0
prefetch_factor: Optional[int] = None prefetch_factor: Optional[int] = None
persistent_workers: bool = False
pin_memory: bool = False pin_memory: bool = False
collate_fn: Optional[Callable[[List[Any]], Any]] = None collate_fn: Optional[Callable[[List[Any]], Any]] = None
@@ -123,35 +125,35 @@ class TrainConfig(BaseConfig):
reward_model_fn: Optional[Callable] = None reward_model_fn: Optional[Callable] = None
executor_kwargs: Dict[str, Any] = field(default_factory=dict) executor_kwargs: Dict[str, Any] = field(default_factory=dict)
extra_kwargs: Dict[str, Any] = field(default_factory=dict) strategy_kwargs: Dict[str, Any] = field(default_factory=dict)
@field_validator("strategy") @field_validator("strategy")
def _validate_strategy(cls, v: str) -> str: def _validate_strategy(cls, v: str) -> str:
if v not in _TRAIN_TYPES: if v not in TRAIN_TYPES:
raise ValueError( raise ValueError(
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}" f"strategy must be one of {sorted(TRAIN_TYPES)}, got {v!r}"
) )
return v return v
@field_validator("parallel_mode") @field_validator("parallel_mode")
def _validate_parallel_mode(cls, v: str) -> str: def _validate_parallel_mode(cls, v: str) -> str:
if v not in _PARALLEL_MODES: if v not in PARALLEL_MODES:
raise ValueError( raise ValueError(
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}" f"parallel_mode must be one of {sorted(PARALLEL_MODES)}, got {v!r}"
) )
return v return v
@field_validator("backend") @field_validator("backend")
def _validate_backend(cls, v: str) -> str: def _validate_backend(cls, v: str) -> str:
if v not in _BACKENDS: if v not in BACKENDS:
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}") raise ValueError(f"backend must be one of {sorted(BACKENDS)}, got {v!r}")
return v return v
@field_validator("start_method") @field_validator("start_method")
def _validate_start_method(cls, v: str) -> str: def _validate_start_method(cls, v: str) -> str:
if v not in _START_METHODS: if v not in START_METHODS:
raise ValueError( raise ValueError(
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}" f"start_method must be one of {sorted(START_METHODS)}, got {v!r}"
) )
return v return v
+41 -12
View File
@@ -25,20 +25,50 @@ function (pure ``record -> Dict[str, Tensor]``) is forwarded to
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from functools import partial from functools import partial
from pathlib import Path
from typing import Callable, Dict, List, Optional from typing import Callable, Dict, List, Optional
import torch import torch
from torch import Tensor from torch import Tensor
from torch.utils.data import Dataset from torch.utils.data import Dataset
from astrai.config.preprocess_config import PipelineConfig
from astrai.dataset.storage import ( from astrai.dataset.storage import (
Store, Store,
StoreFactory, StoreFactory,
detect_format, detect_format,
) )
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.preprocessing.transform import TokenizeTransform
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
_DEFAULT_MESSAGES_CONFIG = {
"version": 1,
"input": {"sections": [{"field": "messages", "action": "$role", "template": True}]},
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
"mask_default": "mask",
"output": {"position_ids_mode": "doc_reset"},
}
def _build_jsonl_transform(
path: str, tokenizer_path: Optional[str] = None
) -> Optional["TokenizeTransform"]:
"""Auto-build a TokenizeTransform for JSONL eager loading.
Reads ``dataset_config.json`` from the data dir if present, or
falls back to the built-in chatml SFT config when *tokenizer_path*
is provided.
"""
root = Path(path)
config_path = root / "dataset_config.json" if root.is_dir() else None
if config_path is not None and config_path.exists():
return TokenizeTransform.from_config_file(str(config_path))
if tokenizer_path:
config = PipelineConfig.from_dict(_DEFAULT_MESSAGES_CONFIG)
return TokenizeTransform(config, tokenizer_path)
return None
def dpo_tokenize( def dpo_tokenize(
record: dict, record: dict,
@@ -349,16 +379,18 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
) )
if processor is not None: if processor is not None:
store.load(load_path, processor=processor, **kwargs) store.load(load_path, processor=processor, **kwargs)
elif storage_type == "jsonl":
transform = _build_jsonl_transform(load_path, tokenizer_path)
if transform is None:
raise FileNotFoundError(
"JSONL dataset config not found. Expected "
"dataset_config.json alongside *.jsonl files, pass "
"tokenizer_path= for the built-in messages config, or "
"use processor= for lazy on-the-fly tokenisation."
)
store.load(load_path, transform=transform, **kwargs)
else: else:
load_kwargs = dict(kwargs) store.load(load_path, **kwargs)
if (
tokenizer_path is not None
and storage_type == "jsonl"
and train_type in ("seq", "sft")
and "tokenizer_path" not in load_kwargs
):
load_kwargs["tokenizer_path"] = tokenizer_path
store.load(load_path, **load_kwargs)
return cls.create(train_type, store=store) return cls.create(train_type, store=store)
@@ -460,9 +492,6 @@ class DPODataset(BaseDataset):
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"] required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
def make_processor(self, tokenizer, max_len: int):
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
def __getitem__(self, index: int) -> Dict[str, Tensor]: def __getitem__(self, index: int) -> Dict[str, Tensor]:
return { return {
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long), "chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
+5 -30
View File
@@ -55,9 +55,7 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.config.preprocess_config import PipelineConfig
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.preprocessing.transform import TokenizeTransform
from astrai.serialization import ( from astrai.serialization import (
load_bin, load_bin,
load_bin_offsets, load_bin_offsets,
@@ -219,7 +217,7 @@ class Store(ABC):
""" """
if self._window_size <= 0: if self._window_size <= 0:
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)") raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
if self._window_size <= 0 or self._length <= self._window_size: if self._length <= self._window_size:
raise IndexError( raise IndexError(
f"Data too short for window: token_count={self._length}, " f"Data too short for window: token_count={self._length}, "
f"window_size={self._window_size}" f"window_size={self._window_size}"
@@ -536,19 +534,8 @@ class JsonlStore(Store, Streamable, Recordable):
``len(store)`` returns ``num_records``; stream primitives raise. ``len(store)`` returns ``num_records``; stream primitives raise.
""" """
CONFIG_NAME = "dataset_config.json"
segments_are_records = True segments_are_records = True
_DEFAULT_MESSAGES_CONFIG = {
"version": 1,
"input": {
"sections": [{"field": "messages", "action": "$role", "template": True}]
},
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
"mask_default": "mask",
"output": {"position_ids_mode": "doc_reset"},
}
def __init__( def __init__(
self, self,
window_size: int = 0, window_size: int = 0,
@@ -569,22 +556,10 @@ class JsonlStore(Store, Streamable, Recordable):
return return
if transform is None: if transform is None:
root = Path(path) raise ValueError(
config_path = root / self.CONFIG_NAME if root.is_dir() else None "JsonlStore eager mode requires transform=. "
if config_path is not None and config_path.exists(): "Use DatasetFactory.load() which auto-constructs it."
transform = TokenizeTransform.from_config_file(str(config_path)) )
else:
tokenizer_path = kwargs.get("tokenizer_path")
if not tokenizer_path:
raise FileNotFoundError(
f"JSONL dataset config not found. Expected "
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
f"explicit transform, pass processor= for lazy "
f"on-the-fly tokenisation, or pass tokenizer_path= to "
f"use the built-in messages config."
)
config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
transform = TokenizeTransform(config, tokenizer_path)
transformed = transform.apply(records) transformed = transform.apply(records)
self._normalize(transformed) self._normalize(transformed)
+4 -4
View File
@@ -15,25 +15,25 @@ Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
SDPA is handled by the attention backend, not the wrapper functions. SDPA is handled by the attention backend, not the wrapper functions.
""" """
from astrai.extension.attention_backend import ( from astrai.extension.backend import (
ATTN_BACKEND, ATTN_BACKEND,
AttentionBackend, AttentionBackend,
AttentionBackendFactory, AttentionBackendFactory,
CudaBackend, CudaBackend,
FlashAttnBackend, FlashAttnBackend,
TorchNativeBackend, TorchNativeBackend,
apply_rotary_emb,
attention, attention,
attn_backend, attn_backend,
get_backend, get_backend,
) )
from astrai.extension.attention_ops import ( from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.ops import (
TensorLayout, TensorLayout,
attn_decode, attn_decode,
attn_paged_decode, attn_paged_decode,
attn_prefill, attn_prefill,
) )
from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.rotary_backend import apply_rotary_emb
__all__ = [ __all__ = [
"ATTN_BACKEND", "ATTN_BACKEND",
-671
View File
@@ -1,671 +0,0 @@
"""Attention backend abstraction with context-manager switching.
The backend encapsulates KV cache I/O and attention computation. The
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
and output projection; the backend handles everything from "write K/V
to cache" through "SDPA output".
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
from astrai.extension import attn_backend, ATTN_BACKEND
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
engine.generate("hello")
# or with an instance:
with attn_backend(TorchNativeBackend()):
...
# or the shorthand (instance is itself a context manager):
with TorchNativeBackend():
...
Thread-safe via ``contextvars`` — each scheduler thread gets its own
active backend. ``get_backend()`` returns the active one, falling back
to a process-wide default (cuda > flash > torch, overridable via
``ASTR_BACKEND``).
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
"""
import contextvars
import enum
import functools
import importlib
import os
from abc import ABC, abstractmethod
from contextlib import contextmanager
from typing import TYPE_CHECKING, Optional, Union
import torch
import torch.nn.functional as F
from torch import Tensor
from astrai.extension.attention_ops import (
attn_paged_decode,
attn_paged_prefill,
)
from astrai.extension.loader import is_available
from astrai.factory import BaseFactory
if TYPE_CHECKING:
from astrai.inference.core.cache import KVCache
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
"attn_backend"
)
@functools.lru_cache(maxsize=1)
def flash_attn_available() -> bool:
if not torch.cuda.is_available():
return False
fa = _get_flash_attn()
if fa is None:
return False
try:
major = int(fa.__version__.split(".")[0])
cc = torch.cuda.get_device_capability()
cc_num = cc[0] * 10 + cc[1]
except Exception:
major, cc_num = 0, 0
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
return False
try:
if not hasattr(fa, "flash_attn_func"):
return False
x = torch.zeros(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16)
out = fa.flash_attn_func(x, x, x, causal=True)
return bool(torch.isfinite(out).all().item())
except Exception:
return False
@functools.lru_cache(maxsize=1)
def _get_flash_attn():
try:
return importlib.import_module("flash_attn")
except Exception:
return None
class ATTN_BACKEND(enum.Enum):
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
TORCH_NATIVE = "torch_native"
CUDA = "cuda"
FLASH = "flash"
_default_backend: Optional["AttentionBackend"] = None
def _priority_backends() -> list["AttentionBackend"]:
"""Available backends in priority order: cuda -> flash -> torch."""
backends: list[AttentionBackend] = []
if is_available("attn_paged_decode") and is_available("attn_paged_prefill"):
backends.append(CudaBackend())
if flash_attn_available():
backends.append(FlashAttnBackend())
backends.append(TorchNativeBackend())
return backends
def _backend_supports(
backend: "AttentionBackend",
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
) -> bool:
"""Whether ``backend`` can run this attention call.
The CUDA kernels are bf16-only, support head_dim in 32/64/128/256, and
need a KV cache (decode/prefill); everything else falls back to torch.
"""
if isinstance(backend, CudaBackend):
return (
kv_cache is not None
and q.dtype == torch.bfloat16
and q.size(-1) in (32, 64, 128, 256)
)
if isinstance(backend, FlashAttnBackend):
if not flash_attn_available():
return False
if q.size(1) == 1 and kv_cache is not None:
return True
return not (attn_mask is not None and not is_causal)
return True
def _resolve_default_backend() -> "AttentionBackend":
"""Pick the highest-priority available backend (cuda -> flash -> torch).
Set ``ASTR_BACKEND`` to override: ``ASTR_BACKEND=cuda``, ``torch_native``,
or ``flash``. The value is the registered name (same as the
``ATTN_BACKEND`` enum value).
Resolved lazily on first ``get_backend()`` and cached. Per-call
capability fallback happens in ``attention()``, so the default is
safe for training and fp32 models.
"""
forced = os.environ.get("ASTR_BACKEND", "").strip().lower()
if forced:
try:
return AttentionBackendFactory.create(forced)
except (ValueError, RuntimeError):
pass
return _priority_backends()[0]
def get_backend() -> "AttentionBackend":
"""Return the active backend for the current thread/context.
Falls back to the highest-priority available backend (cuda -> flash ->
torch_native) when no backend has been activated via ``with``. Set
``ASTR_BACKEND`` to override the default.
"""
try:
return _current_backend.get()
except LookupError:
global _default_backend
if _default_backend is None:
_default_backend = _resolve_default_backend()
return _default_backend
@contextmanager
def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
"""Context manager to select an attention backend.
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
registered name, ``ATTN_BACKEND`` enum value, backend class, or instance.
Examples::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
...
with attn_backend(TorchNativeBackend):
...
with attn_backend(TorchNativeBackend()):
...
"""
if isinstance(backend, ATTN_BACKEND):
instance = AttentionBackendFactory.create(backend.value)
elif isinstance(backend, str):
instance = AttentionBackendFactory.create(backend)
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
instance = backend()
elif isinstance(backend, AttentionBackend):
instance = backend
else:
raise TypeError(
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
f"or instance, "
f"got {type(backend).__name__}"
)
token = _current_backend.set(instance)
try:
yield instance
finally:
_current_backend.reset(token)
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
"""Expand KV heads to match Q heads for GQA."""
bs, slen, n_heads, head_dim = x.shape
if n_rep == 1:
return x
return (
x[:, :, :, None, :]
.expand(bs, slen, n_heads, n_rep, head_dim)
.reshape(bs, slen, n_heads * n_rep, head_dim)
)
def _write_and_gather_kv(
kv_cache: "KVCache",
k: Tensor,
v: Tensor,
layer_id: int,
q: Tensor,
attn_mask: Optional[Tensor],
) -> tuple[Tensor, Tensor]:
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
max_len = kv_cache.max_len
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
if q.size(1) == 1 and attn_mask is not None and attn_mask.dim() == 4:
pos_mask = attn_mask[:, 0, 0]
else:
pos_mask = (
torch.arange(max_len, device=q.device)[None, :] < kv_cache.seq_lens[:, None]
)
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
return kv_cache.k_buffer[layer_id, indices], kv_cache.v_buffer[layer_id, indices]
def attention(
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"] = None,
layer_id: int = 0,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
Delegates to the active backend (set via ``with attn_backend(...)``).
Handles KV cache I/O, GQA head expansion, and causal masking so the
caller only needs to provide projected q/k/v.
Args:
q: [batch, q_len, n_heads, head_dim] (blhd)
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
kv_cache: cache dataclass, or None for training (no cache).
layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask (SDPA-compatible).
is_causal: whether to apply causal masking.
Returns:
[batch, q_len, n_heads * head_dim]
"""
backend = get_backend()
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal):
# The active backend cannot run this call (e.g. CUDA on a training /
# fp32 / unsupported-head_dim input) — fall back to the highest-
# priority backend that can, ending at torch SDPA.
for candidate in _priority_backends():
if isinstance(candidate, type(backend)):
continue
if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal):
backend = candidate
break
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
class AttentionBackend(ABC):
"""Abstract base for attention computation strategies.
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
``fwd_prefill`` (q_len > 1, with or without cache). The public
``forward`` method dispatches based on q_len.
Three equivalent ways to activate a backend::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
...
with attn_backend(TorchNativeBackend): # class
...
with TorchNativeBackend(): # instance
...
"""
def __enter__(self) -> "AttentionBackend":
self._token = _current_backend.set(self)
return self
def __exit__(self, *exc) -> None:
_current_backend.reset(self._token)
def forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Dispatch to decode or extend based on q_len.
Args:
q: [batch, q_len, n_heads, head_dim]
k: [batch, q_len, n_kv_heads, head_dim]
v: [batch, q_len, n_kv_heads, head_dim]
kv_cache: cache dataclass, or None for training (no cache).
layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask compatible with SDPA.
is_causal: whether to apply causal masking.
Returns:
[batch, q_len, n_heads * head_dim]
"""
if kv_cache is not None and q.size(1) == 1:
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
@abstractmethod
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Single-token decode with KV cache."""
@abstractmethod
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Multi-token prefill or training forward."""
@staticmethod
def supports_graph() -> bool:
"""Return True if this backend supports CUDA-graph capture.
Override in subclasses that can run under ``torch.cuda.graph``.
Called on the *active* backend instance (or its class) — a cheap
boolean check with no side-effects.
"""
return False
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
"""Factory for registered attention backends."""
@AttentionBackendFactory.register(ATTN_BACKEND.TORCH_NATIVE.value)
class TorchNativeBackend(AttentionBackend):
"""Reference backend using torch SDPA with indirect KV cache indexing.
Writes new K/V into the cache buffers, gathers the full sequence K/V
via ``req_to_token`` indirect indexing, then calls
``F.scaled_dot_product_attention``.
For training (``kv_cache is None``), skips cache I/O entirely and
runs SDPA directly on the projected q/k/v.
"""
@staticmethod
def supports(**kwargs) -> bool:
return True
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def _forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is not None:
k, v = _write_and_gather_kv(kv_cache, k, v, layer_id, q, attn_mask)
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
out = F.scaled_dot_product_attention(
q.permute(0, 2, 1, 3),
k.permute(0, 2, 1, 3),
v.permute(0, 2, 1, 3),
attn_mask,
is_causal=is_causal,
)
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
return out
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
class CudaBackend(AttentionBackend):
"""CUDA kernel backend with direct KV cache access.
Decode path: writes K/V to the flat pool, then calls
``attn_paged_decode`` with req_to_token + kv_indptr.
Prefill path: writes K/V to the flat pool, then calls
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
kv_indptr.
``kv_cache is None`` (training) raises — the per-call fallback to
torch SDPA for training / fp32 / unsupported head_dim happens in the
``attention()`` entry point.
Raises ``RuntimeError`` if the required kernel is not available.
"""
@staticmethod
def supports(**kwargs) -> bool:
head_dim = kwargs.get("head_dim", -1)
return (
torch.cuda.is_available()
and head_dim in (32, 64, 128, 256)
and is_available("attn_paged_decode")
and is_available("attn_paged_prefill")
)
@staticmethod
def supports_graph() -> bool:
return True
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
loc = kv_cache.out_cache_loc[:, 0]
kv_cache.k_buffer[layer_id].index_copy_(0, loc, k[:, 0])
kv_cache.v_buffer[layer_id].index_copy_(0, loc, v[:, 0])
q_3d = q.squeeze(1)
kv_indptr = kv_cache.kv_indptr
out = attn_paged_decode(
q_3d,
kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id],
kv_cache.req_to_token,
kv_cache.req_pool_indices,
kv_indptr,
kv_cache.max_len,
is_causal=True,
o_part_buf=kv_cache.decode_o_part,
ml_part_buf=kv_cache.decode_ml_part,
out_buf=kv_cache.decode_out,
)
return out.unsqueeze(1).flatten(2)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
loc = kv_cache.out_cache_loc.reshape(-1)
kv_cache.k_buffer[layer_id].index_copy_(
0, loc, k.reshape(-1, k.size(2), k.size(3))
)
kv_cache.v_buffer[layer_id].index_copy_(
0, loc, v.reshape(-1, v.size(2), v.size(3))
)
b = q.size(0)
q_len = q.size(1)
kv_indptr = kv_cache.kv_indptr
qo_indptr = kv_cache.qo_indptr
q_flat = q.reshape(b * q_len, q.size(2), q.size(3))
out = attn_paged_prefill(
q_flat,
kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id],
kv_cache.req_to_token,
kv_cache.req_pool_indices,
kv_indptr,
qo_indptr,
attn_mask,
q_len,
is_causal=is_causal,
)
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
class FlashAttnBackend(AttentionBackend):
"""FlashAttention backend via the optional ``flash-attn`` package.
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``,
which reads K/V directly from the flat pool via cache_batch_idx +
cache_seqlens — no materialized KV gather.
Prefill / non-contiguous decode: falls back to KV gather +
``flash_attn_func``.
"""
@staticmethod
def supports(**kwargs) -> bool:
return flash_attn_available()
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def _forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is not None:
if q.size(1) == 1 and kv_cache.k_buffer.size(
1
) == kv_cache.req_to_token.size(0) * kv_cache.req_to_token.size(1):
return self._decode_with_kvcache(q, k, v, kv_cache, layer_id)
k, v = _write_and_gather_kv(kv_cache, k, v, layer_id, q, attn_mask)
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
if attn_mask is not None and not is_causal:
raise ValueError(
"FlashAttnBackend does not support a custom attention mask; "
"use a causal mask or select TorchNativeBackend."
)
fa = _get_flash_attn()
if fa is None:
raise RuntimeError(
"FlashAttnBackend requires the optional 'flash-attn' package. "
"Install with `pip install flash-attn`."
)
out = fa.flash_attn_func(
q.contiguous(), k.contiguous(), v.contiguous(), causal=is_causal
)
return out.contiguous().flatten(2)
def _decode_with_kvcache(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: "KVCache",
layer_id: int,
) -> Tensor:
max_batch = kv_cache.req_to_token.size(0)
max_seq = kv_cache.req_to_token.size(1)
n_kv = k.size(2)
k_cache = kv_cache.k_buffer[layer_id].view(max_batch, max_seq, n_kv, k.size(3))
v_cache = kv_cache.v_buffer[layer_id].view(max_batch, max_seq, n_kv, v.size(3))
fa = _get_flash_attn()
out = fa.flash_attn_with_kvcache(
q=q,
k_cache=k_cache,
v_cache=v_cache,
k=k,
v=v,
cache_seqlens=(kv_cache.seq_lens - 1).to(torch.int32),
cache_batch_idx=kv_cache.req_pool_indices.to(torch.int32),
causal=True,
)
return out.flatten(2)
+27
View File
@@ -0,0 +1,27 @@
"""Backend selection, fallbacks, and execution policies."""
from astrai.extension.backend.attention import (
ATTN_BACKEND,
AttentionBackend,
AttentionBackendFactory,
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
attention,
attn_backend,
get_backend,
)
from astrai.extension.backend.rotary import apply_rotary_emb
__all__ = [
"ATTN_BACKEND",
"AttentionBackend",
"AttentionBackendFactory",
"CudaBackend",
"FlashAttnBackend",
"TorchNativeBackend",
"apply_rotary_emb",
"attention",
"attn_backend",
"get_backend",
]
+819
View File
@@ -0,0 +1,819 @@
"""Attention backend abstraction with context-manager switching.
The backend encapsulates KV cache I/O and attention computation. The
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
and output projection; the backend handles everything from "write K/V
to cache" through "SDPA output".
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
from astrai.extension import attn_backend, ATTN_BACKEND
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
engine.generate("hello")
# or with an instance:
with attn_backend(TorchNativeBackend()):
...
# or the shorthand (instance is itself a context manager):
with TorchNativeBackend():
...
Thread-safe via ``contextvars`` — each scheduler thread gets its own
active backend. Backend resolution follows a strict precedence:
1. explicit ``attn_backend(...)`` context (wins over everything),
2. the process-wide ``ASTR_BACKEND`` environment override,
3. an implicit default picked from the available backends
(cuda > flash > torch).
Capability is polymorphic: every backend declares ``available()``
(machine-level) and ``supports_call(...)`` (per-call), so adding a new
backend requires no changes to the resolution logic. Training calls
(``fwd=None``, no KV cache) resolve through the same priority list: the
CUDA cache kernels cannot run without a cache, so they fall back to
flash (when it can handle the call — mask-free/causal only) and finally
to the reference ``TorchNativeBackend``.
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
"""
import contextvars
import enum
import functools
import logging
import os
import threading
from abc import ABC, abstractmethod
from contextlib import contextmanager
from typing import TYPE_CHECKING, Dict, Optional, Tuple, Union
import torch
import torch.nn.functional as F
from torch import Tensor
from astrai.extension.loader import is_available
from astrai.extension.ops.attention import (
attn_paged_decode,
attn_paged_prefill,
)
from astrai.factory import BaseFactory
try:
import flash_attn as _flash_attn
except Exception:
_flash_attn = None
if TYPE_CHECKING:
from astrai.inference.cache import KVCache
logger = logging.getLogger(__name__)
_default_backend_lock = threading.Lock()
_env_backend_name: Optional[str] = None
_env_backend: Optional["AttentionBackend"] = None
_current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = (
contextvars.ContextVar("attn_backend", default=None)
)
# Backends are stateless — one canonical instance per class, created lazily
# and reused everywhere (resolution, fallback, context managers).
_singletons: Dict[type, "AttentionBackend"] = {}
@functools.lru_cache(maxsize=1)
def flash_attn_available() -> bool:
if not torch.cuda.is_available():
return False
fa = _flash_attn
if fa is None:
return False
try:
major = int(fa.__version__.split(".")[0])
cc = torch.cuda.get_device_capability()
cc_num = cc[0] * 10 + cc[1]
except Exception:
major, cc_num = 0, 0
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
return False
try:
if not hasattr(fa, "flash_attn_func"):
return False
x = torch.zeros(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16)
out = fa.flash_attn_func(x, x, x, causal=True)
return bool(torch.isfinite(out).all().item())
except Exception:
return False
class ATTN_BACKEND(enum.Enum):
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
TORCH_NATIVE = "torch_native"
CUDA = "cuda"
FLASH = "flash"
def _instance(backend_cls: type) -> "AttentionBackend":
"""Return the canonical singleton instance for a backend class.
Backends hold no per-instance state, so a single cached instance is
safe and avoids per-call allocation on the attention hot path.
"""
backend = _singletons.get(backend_cls)
if backend is None:
backend = backend_cls()
_singletons[backend_cls] = backend
return backend
@functools.lru_cache(maxsize=1)
def _priority_backends() -> Tuple["AttentionBackend", ...]:
"""Available backends in priority order: cuda -> flash -> torch.
Computed once (machine availability cannot change at runtime) and
cached forever; the tuple always ends with ``TorchNativeBackend``,
which is unconditionally available.
"""
return tuple(
_instance(cls)
for cls in (CudaBackend, FlashAttnBackend, TorchNativeBackend)
if cls.available()
)
def _resolve_default_backend() -> "AttentionBackend":
"""Pick the highest-priority available backend (cuda -> flash -> torch).
Resolved lazily on first use and cached via ``_priority_backends``.
Per-call capability fallback happens in ``attention()``, so the
default is safe for training and fp32 models.
"""
return _priority_backends()[0]
def _environment_backend() -> Optional["AttentionBackend"]:
"""Resolve the process-wide ``ASTR_BACKEND`` override, if configured."""
global _env_backend, _env_backend_name
name = os.environ.get("ASTR_BACKEND", "").strip().lower()
if not name:
return None
if name != _env_backend_name:
with _default_backend_lock:
if name != _env_backend_name:
try:
_env_backend = _resolve_backend(name)
except (ValueError, RuntimeError):
_env_backend = None
logger.warning(
"ASTR_BACKEND=%r is not a registered attention backend; "
"falling back to default resolution",
name,
)
_env_backend_name = name
return _env_backend
def _resolve_backend(
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
) -> "AttentionBackend":
"""Resolve a backend configuration to its canonical instance.
Accepts a registered name, ``ATTN_BACKEND`` enum value, backend class,
or instance. Names/classes resolve to the shared singleton; a caller
may still pass its own instance to opt out of sharing.
"""
if backend is not None:
if isinstance(backend, ATTN_BACKEND):
return _instance(AttentionBackendFactory.get_component_class(backend.value))
if isinstance(backend, str):
return _instance(AttentionBackendFactory.get_component_class(backend))
if isinstance(backend, type) and issubclass(backend, AttentionBackend):
return _instance(backend)
if isinstance(backend, AttentionBackend):
return backend
raise TypeError(
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
f"or instance, got {type(backend).__name__}"
)
return _resolve_default_backend()
def get_backend(
use_default: bool = True,
) -> Optional["AttentionBackend"]:
"""Resolve the active backend: explicit context > env > default.
An ``attn_backend(...)`` context is the caller's explicit choice and
always wins. ``ASTR_BACKEND`` is a process-wide override consulted
only when no context is set. Pass ``use_default=False`` at request
submission to retain only an environment override or the caller's
:func:`attn_backend` value.
"""
context_backend = _current_backend.get()
if context_backend is not None:
return context_backend
env_backend = _environment_backend()
if env_backend is not None:
return env_backend
return _resolve_default_backend() if use_default else None
@contextmanager
def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
"""Context manager to select an attention backend.
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
registered name, ``ATTN_BACKEND`` enum value, backend class, or instance.
Examples::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
...
with attn_backend(TorchNativeBackend):
...
with attn_backend(TorchNativeBackend()):
...
"""
instance = _resolve_backend(backend)
token = _current_backend.set(instance)
try:
yield instance
finally:
_current_backend.reset(token)
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
"""Expand KV heads to match Q heads for GQA."""
if n_rep == 1:
return x
n_heads, head_dim = x.shape[-2:]
return (
x.unsqueeze(-2)
.expand(*x.shape[:-2], n_heads, n_rep, head_dim)
.reshape(*x.shape[:-2], n_heads * n_rep, head_dim)
)
def attention(
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"] = None,
layer_id: int = 0,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
fwd: Optional[str] = None,
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
) -> Tensor:
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
Delegates to the active backend. ``backend`` (optional) is an explicit
escape hatch; when omitted the backend is resolved as
explicit context > ``ASTR_BACKEND`` env > default (cuda > flash > torch).
Handles KV cache I/O, GQA head expansion, and causal masking so the
caller only needs to provide projected q/k/v.
Training calls (``fwd=None``, ``kv_cache=None``) resolve through the
same capability chain — the CUDA cache kernels cannot run without a
cache, so they fall back to flash (mask-free/causal calls only) and
finally to torch SDPA. An explicitly-selected backend that cannot
handle the call raises — an implicit one falls back down the priority
list to the first capable backend.
Args:
q: [batch, q_len, n_heads, head_dim] (blhd)
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
kv_cache: cache dataclass, or None for training (no cache).
layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask (SDPA-compatible).
is_causal: whether to apply causal masking.
fwd: "prefill" / "decode" for inference, None for training.
backend: optional explicit backend (name, enum, class, or instance).
Returns:
[batch, q_len, n_heads * head_dim]
"""
if backend is not None:
selected = _resolve_backend(backend)
explicit = True
else:
context_backend = _current_backend.get()
explicit = context_backend is not None
# Resolve through the same chain as inference: explicit context >
# ASTR_BACKEND env > default. Training calls (fwd=None, no cache)
# land on the CUDA backend and fall back by capability below —
# flash when it can handle the call, else torch SDPA.
selected = get_backend()
assert selected is not None
if not selected.supports_call(q, kv_cache, attn_mask, is_causal, fwd):
if explicit:
raise RuntimeError(
f"Explicitly-set backend {type(selected).__name__} cannot "
f"handle this attention call (shape={q.shape}, "
f"dtype={q.dtype}, kv_cache={'none' if kv_cache is None else 'present'}, "
f"attn_mask={'none' if attn_mask is None else 'present'}). "
f"Remove the attn_backend() context or switch to a compatible backend."
)
selected = next(
(
candidate
for candidate in _priority_backends()
if candidate.supports_call(q, kv_cache, attn_mask, is_causal, fwd)
),
_instance(TorchNativeBackend),
)
return selected.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
class AttentionBackend(ABC):
"""Abstract base for attention computation strategies.
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
``fwd_prefill`` (q_len > 1, with or without cache). The public
``forward`` method dispatches based on q_len.
Capability contract — every backend declares:
* ``available()`` — machine-level: can this backend exist here
(kernel ``.so`` loaded, flash-attn present, GPU available)?
Used once to build the default priority list.
* ``supports_call(q, kv_cache, attn_mask, is_causal, fwd)`` — can this
backend run this *specific* call (shape/dtype/cache/mask)? Used by
``attention()`` for the per-call fallback. Resolution logic never
checks concrete backend types, so adding a backend requires no
changes outside its own class.
Three equivalent ways to activate a backend::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
...
with attn_backend(TorchNativeBackend): # class
...
with TorchNativeBackend(): # instance
...
"""
def __enter__(self) -> "AttentionBackend":
self._token = _current_backend.set(self)
return self
def __exit__(self, *exc) -> None:
_current_backend.reset(self._token)
@classmethod
@abstractmethod
def available(cls) -> bool:
"""Return True if this backend can run on the current machine.
Checks static availability only (compiled kernels, optional
packages, GPU presence) — not call-specific constraints.
"""
@abstractmethod
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
"""Return True if this backend can run this specific attention call.
Called on the canonical singleton instance (or a caller-provided
one); must be side-effect free.
"""
def forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor:
"""Dispatch to decode or extend based on q_len.
Args:
q: [batch, q_len, n_heads, head_dim]
k: [batch, q_len, n_kv_heads, head_dim]
v: [batch, q_len, n_kv_heads, head_dim]
kv_cache: cache dataclass, or None for training (no cache).
layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask compatible with SDPA.
is_causal: whether to apply causal masking.
Returns:
[batch, q_len, n_heads * head_dim]
"""
if fwd == "decode":
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
if fwd == "prefill" or fwd is None:
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
raise ValueError(f"unsupported attention forward mode: {fwd}")
@abstractmethod
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Single-token decode with KV cache."""
@abstractmethod
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Multi-token prefill or training forward."""
@staticmethod
def supports_graph() -> bool:
"""Return True if this backend supports CUDA-graph capture.
Override in subclasses that can run under ``torch.cuda.graph``.
Called on the *active* backend instance (or its class) — a cheap
boolean check with no side-effects.
"""
return False
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
"""Factory for registered attention backends."""
@AttentionBackendFactory.register(ATTN_BACKEND.TORCH_NATIVE.value)
class TorchNativeBackend(AttentionBackend):
"""Reference backend using torch SDPA with indirect KV cache indexing.
Writes new K/V into the cache buffers, gathers the full sequence K/V
via ``req_to_token`` indirect indexing, then calls
``F.scaled_dot_product_attention``.
For training (``kv_cache is None``), skips cache I/O entirely and
runs SDPA directly on the projected q/k/v.
"""
@classmethod
def available(cls) -> bool:
return True
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
return True
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def _forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if q.ndim == 4:
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
return (
F.scaled_dot_product_attention(
q.permute(0, 2, 1, 3),
k.permute(0, 2, 1, 3),
v.permute(0, 2, 1, 3),
attn_mask,
is_causal=is_causal,
)
.permute(0, 2, 1, 3)
.contiguous()
)
if kv_cache is None or kv_cache.qo_indptr is None:
raise ValueError("packed attention requires KV cache metadata")
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
outputs = []
n_rep = q.size(1) // k.size(1)
for i in range(kv_cache.req_pool_indices.numel()):
q_start = int(kv_cache.qo_indptr[i])
q_end = int(kv_cache.qo_indptr[i + 1])
indices = kv_cache.req_to_token[
kv_cache.req_pool_indices[i], : kv_cache.seq_lens[i]
]
k_i = kv_cache.k_buffer[layer_id, indices]
v_i = kv_cache.v_buffer[layer_id, indices]
if n_rep > 1:
k_i = repeat_kv(k_i, n_rep)
v_i = repeat_kv(v_i, n_rep)
q_len = q_end - q_start
kv_len = k_i.size(0)
q_pos = torch.arange(kv_len - q_len, kv_len, device=q.device)
causal_mask = q_pos[:, None] >= torch.arange(kv_len, device=q.device)
out = F.scaled_dot_product_attention(
q[q_start:q_end].transpose(0, 1).unsqueeze(0),
k_i.transpose(0, 1).unsqueeze(0),
v_i.transpose(0, 1).unsqueeze(0),
attn_mask=causal_mask,
)
outputs.append(out.squeeze(0).transpose(0, 1))
return torch.cat(outputs)
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
class CudaBackend(AttentionBackend):
"""CUDA kernel backend with direct KV cache access.
Decode path: writes K/V to the flat pool, then calls
``attn_paged_decode`` with req_to_token + kv_indptr.
Prefill path: writes K/V to the flat pool, then calls
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
kv_indptr.
``kv_cache is None`` (training) raises — the per-call fallback to
torch SDPA for training / fp32 / unsupported head_dim happens in the
``attention()`` entry point.
Raises ``RuntimeError`` if the required kernel is not available.
"""
# Head dims supported by the CUDA kernels (single source of truth).
HEAD_DIMS = (32, 64, 128, 256)
@classmethod
def available(cls) -> bool:
return (
torch.cuda.is_available()
and is_available("attn_paged_decode")
and is_available("attn_paged_prefill")
)
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
# The CUDA kernels are bf16-only, support head_dim in
# HEAD_DIMS, and need a KV cache (decode/prefill); everything
# else falls back down the priority list to torch.
return (
fwd in ("prefill", "decode")
and kv_cache is not None
and q.ndim == 3
and q.dtype == torch.bfloat16
and q.size(-1) in self.HEAD_DIMS
and is_available(f"attn_paged_{fwd}")
)
@staticmethod
def supports_graph() -> bool:
return True
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
kv_indptr = kv_cache.kv_indptr
out = attn_paged_decode(
q,
kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id],
kv_cache.req_to_token,
kv_cache.req_pool_indices,
kv_indptr,
new_k=k,
new_v=v,
is_causal=True,
o_part_buf=kv_cache.decode_o_part,
ml_part_buf=kv_cache.decode_ml_part,
out_buf=kv_cache.decode_out,
)
return out
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
loc = kv_cache.out_cache_loc
kv_cache.k_buffer[layer_id, loc] = k
kv_cache.v_buffer[layer_id, loc] = v
out = attn_paged_prefill(
q,
kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id],
kv_cache.req_to_token,
kv_cache.req_pool_indices,
kv_cache.kv_indptr,
kv_cache.qo_indptr,
kv_cache.q_tile_to_batch,
kv_cache.q_tile_to_index,
attn_mask,
is_causal=is_causal,
)
return out
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
class FlashAttnBackend(AttentionBackend):
"""FlashAttention backend via the optional ``flash-attn`` package.
Decode (q_len=1, contiguous cache): writes K/V to the pool, gathers
flat K/V via the ``req_to_token`` page table, and calls
``flash_attn_varlen_func`` over the ragged batch
(``qo_indptr``/``kv_indptr``).
Prefill: packed 3-D calls share the ``flash_attn_varlen_func`` path;
dense 4-D calls go through ``flash_attn_func`` (mask-free only).
"""
@classmethod
def available(cls) -> bool:
return flash_attn_available()
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
if not self.available():
return False
if q.dtype not in (torch.float16, torch.bfloat16):
return False
if fwd is not None:
return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
# Dense (training) path: flash_attn_func cannot apply a custom
# mask, so only mask-free calls are supported — ``is_causal`` is
# a flag, not a mask. Masked training (SFT/DPO/GRPO) must fall
# back to TorchNativeBackend instead of silently ignoring the mask.
return attn_mask is None
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward_packed(q, k, v, kv_cache, layer_id)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if q.ndim == 3:
return self._forward_packed(q, k, v, kv_cache, layer_id)
return self._forward_dense(q, k, v, attn_mask, is_causal)
def _forward_dense(
self,
q: Tensor,
k: Tensor,
v: Tensor,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
if attn_mask is not None:
raise ValueError(
"FlashAttnBackend cannot handle a custom attention mask; "
"use a causal mask or select TorchNativeBackend."
)
fa = _flash_attn
if fa is None:
raise RuntimeError(
"FlashAttnBackend requires the optional 'flash-attn' package. "
"Install with `pip install flash-attn`."
)
out = fa.flash_attn_func(
q.contiguous(),
k.contiguous(),
v.contiguous(),
causal=is_causal,
)
return out.contiguous()
def _forward_packed(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: "KVCache",
layer_id: int,
) -> Tensor:
fa = _flash_attn
if fa is None or not hasattr(fa, "flash_attn_varlen_func"):
raise RuntimeError("packed inference requires flash_attn_varlen_func")
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
page_table = kv_cache.req_to_token[
kv_cache.req_pool_indices, : kv_cache.max_len
]
positions = torch.arange(kv_cache.max_len, device=q.device)
indices = page_table[positions.unsqueeze(0) < kv_cache.seq_lens.unsqueeze(1)]
k_flat = kv_cache.k_buffer[layer_id, indices].contiguous()
v_flat = kv_cache.v_buffer[layer_id, indices].contiguous()
out = fa.flash_attn_varlen_func(
q.contiguous(),
k_flat,
v_flat,
kv_cache.qo_indptr,
kv_cache.kv_indptr,
int((kv_cache.qo_indptr[1:] - kv_cache.qo_indptr[:-1]).max()),
int(kv_cache.seq_lens.max()),
dropout_p=0.0,
causal=True,
)
return out
@@ -11,6 +11,7 @@ import torch
from torch import Tensor from torch import Tensor
from astrai.extension.loader import is_available from astrai.extension.loader import is_available
from astrai.extension.ops.rotary import rotary_emb as _cuda_rotary
_cache = {"available": None} _cache = {"available": None}
@@ -26,7 +27,7 @@ def _torch_apply(x: Tensor, freqs_cis: Tensor) -> Tensor:
dtype = x.dtype dtype = x.dtype
x_ = x.float().reshape(*x.shape[:-1], -1, 2) x_ = x.float().reshape(*x.shape[:-1], -1, 2)
x_complex = torch.view_as_complex(x_) x_complex = torch.view_as_complex(x_)
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(2) freqs_cis_complex = torch.complex(cos, sin).unsqueeze(-2)
x_rotated = x_complex * freqs_cis_complex x_rotated = x_complex * freqs_cis_complex
x_out = torch.view_as_real(x_rotated).flatten(-2) x_out = torch.view_as_real(x_rotated).flatten(-2)
return x_out.to(dtype) return x_out.to(dtype)
@@ -48,7 +49,5 @@ def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
and x.is_cuda and x.is_cuda
and x.dtype == torch.bfloat16 and x.dtype == torch.bfloat16
): ):
from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary
return _cuda_rotary(x, freqs_cis) return _cuda_rotary(x, freqs_cis)
return _torch_apply(x, freqs_cis) return _torch_apply(x, freqs_cis)
+450
View File
@@ -0,0 +1,450 @@
"""FP8 training: scaling recipes, per-tensor state, and aten::linear dispatch.
Layered (see ``ops/fp8.py`` for the CUDA interface adapter):
1. ``ops.fp8`` — the only module touching the pybind.
2. This module (strategy layer): scaling *recipes* (TE-style delayed scaling
or dynamic current-amax scaling), per-tensor scales + amax history, and the
``fp8_autocast`` context manager (like ``torch.autocast``).
3. aten::linear integration: registers the CUDA + AutogradCUDA impls.
Usage::
from astrai.extension.fp8 import fp8_autocast
with fp8_autocast(enabled=True, fp8_format="hybrid"):
logits = model(input_ids)
loss.backward() # fp8 backward runs anywhere; fwd captured state on the node
Format defaults follow the ecosystem consensus: E4M3 forward / E5M2 backward
("hybrid"); every operand's scale is a quantization step derived from its amax
history by the active recipe.
The context mirrors ``torch.autocast`` (``autocast_mode.py``): the active
``(enabled, recipe, fp8_format)`` triple is thread-local (a ``contextvars``
``ContextVar``, absent outside any region), and the manager is class-based and
reentrant with nested ``enabled=False`` disabling dispatch inside it. The module
targets *training*: every step quantizes x/w/g fresh (no weight-cast cache — the
optimizer bumps the weight version each step, so a torch-style cached_cast would
miss anyway), and the per-operand scales come from the delayed/dynamic recipe.
"""
import functools
from contextvars import ContextVar, Token
from dataclasses import dataclass
from enum import Enum
from typing import Dict, List, NamedTuple, Optional
import torch
from torch.library import Library
from astrai.extension.ops.fp8 import mm_fp8, quantize, quantize_dual
# Max representable value per FP8 format (E4M3: 448, E5M2: 57344).
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
class FP8Format(str, Enum):
"""Per-direction FP8 format. HYBRID = E4M3 forward / E5M2 backward."""
E4M3 = "e4m3"
E5M2 = "e5m2"
HYBRID = "hybrid"
def fwd(self) -> str:
return "e4m3" if self is FP8Format.HYBRID else self.value
def bwd(self) -> str:
return "e5m2" if self is FP8Format.HYBRID else self.value
@dataclass
class FP8Recipe:
"""Scale-from-amax policy: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``.
``dynamic=False`` (default) is TE-style delayed scaling: max over the
amax history window (amax from *previous* steps; the window trades
responsiveness against stability). ``dynamic=True`` is current-amax
scaling (torchao DYNAMIC): measure, then quantize — no history, at an
extra pass. ``scale_from_history`` receives the operand's amax tensor
(a ring window / the current amax) and returns the quantization step.
"""
history_len: int = 16
margin: int = 0
dynamic: bool = False
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
peak = amax.max()
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
class _ScaleRing:
"""One operand's delayed-scaling state: a float32 buffer
``[hist[n] | scale | legacy | amax | done]`` (views). The quantize
kernel folds its fused amax into ``hist[idx]`` and publishes the next
scale from the window in its own last block (``fold_args`` passes the
buffer + recipe constants); ``idx`` advances host-side each use. The
``amax``/``done`` tail slots are kernel scratch (self-cleaning across
launches); the legacy slot keeps state-buffer compatibility.
"""
__slots__ = ("recipe", "state", "hist", "scale", "idx", "initialized")
def __init__(self, device: torch.device, recipe: FP8Recipe):
self.recipe = recipe
n = recipe.history_len
self.state = torch.zeros(n + 4, device=device, dtype=torch.float32)
self.hist = self.state[:n]
self.scale = self.state[n : n + 1]
self.idx = 0
self.initialized = False
def advance(self) -> None:
"""Rotate to the next history slot after metadata update."""
self.idx = (self.idx + 1) % self.hist.numel()
def seed(self, t: torch.Tensor, fmt: str) -> None:
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
self.hist.fill_(amax)
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
self.initialized = True
def fold_args(self, fmt: str) -> dict:
"""Keyword arguments for quantize()'s in-kernel history fold."""
return {
"ring_state": self.state,
"hist_idx": self.idx,
"fp8_max": FP8_MAX[fmt],
"pow2_margin": float(2**self.recipe.margin),
}
class FP8TensorMeta(NamedTuple):
"""Per-weight delayed-scaling rings for ``w``, ``x`` and ``g``.
Dynamic scaling never allocates a meta; it measures the current amax inline.
"""
w: _ScaleRing
x: _ScaleRing
g: _ScaleRing
@dataclass(frozen=True)
class _ActiveConfig:
"""The immutable (enabled, recipe, format) triple of one open region."""
enabled: bool
recipe: FP8Recipe
fp8_format: FP8Format
# Thread-local active configuration (torch's autocast TLS analog): set by
# fp8_autocast on __enter__, absent outside any region. Autograd engine
# threads run backwards with their own empty context — fine, since backward
# only reads state captured on ctx at forward time.
_active_config: ContextVar[Optional[_ActiveConfig]] = ContextVar(
"astrai_fp8_active_config", default=None
)
class FP8State:
"""Global fp8 training state: per-tensor metas + out-of-region defaults.
The active ``(enabled, recipe, fp8_format)`` triple is a ``ContextVar``
set by ``fp8_autocast`` (see ``_active``/``_current_config``); these plain
attributes are the persistent defaults applied outside any region —
``fp8_linear_enable`` writes ``default_enabled``. The metas registry is
shared across threads (GIL-protected); fp8 backward runs on autograd
engine threads and only touches metas captured on ``ctx`` at forward time.
"""
def __init__(self):
self.default_enabled = False
self.default_recipe: FP8Recipe = FP8Recipe()
self.default_format: FP8Format = FP8Format.HYBRID
self._metas: Dict[tuple, FP8TensorMeta] = {}
def get_weight_meta(self, w: torch.Tensor, recipe: FP8Recipe) -> FP8TensorMeta:
key = (w.data_ptr(), w.shape, w.dtype)
meta = self._metas.get(key)
if meta is None:
meta = FP8TensorMeta(
_ScaleRing(w.device, recipe),
_ScaleRing(w.device, recipe),
_ScaleRing(w.device, recipe),
)
self._metas[key] = meta
return meta
def reset(self) -> None:
"""Restore construction defaults (switch, recipe, format) and drop all
per-weight metas — a full state reset for tests / reconfiguration."""
self.default_enabled = False
self.default_recipe = FP8Recipe()
self.default_format = FP8Format.HYBRID
self._metas.clear()
# Process-wide singleton; per-thread/per-region state lives in _active_config.
_state = FP8State()
def fp8_state() -> FP8State:
return _state
def _active() -> Optional[_ActiveConfig]:
"""The active config when fp8 dispatch is on, else ``None`` (fast guard).
A region config wins (honoring nested ``enabled=False`` regions); with no
region open this falls back to the persistent global switch
(``fp8_linear_enable``), so that flag still routes aten::linear to fp8.
"""
cfg = _active_config.get()
if cfg is not None:
return cfg if cfg.enabled else None
if _state.default_enabled:
return _ActiveConfig(True, _state.default_recipe, _state.default_format)
return None
def _current_config() -> _ActiveConfig:
"""Like ``_active()`` but always returns a config (disabled regions and
out-of-region direct calls resolve to the global defaults)."""
cfg = _active_config.get()
if cfg is not None:
return cfg
return _ActiveConfig(
_state.default_enabled, _state.default_recipe, _state.default_format
)
class fp8_autocast:
"""Autocast-style context: fp8 linear dispatch on this thread.
Mirrors ``torch.autocast`` — a class-based, reentrant, nestable context
over thread-local state::
with fp8_autocast(enabled=True, fp8_format="hybrid"):
logits = model(input_ids) # aten::linear -> fp8 path
loss.backward() # fp8 backward; state was captured at forward time
Nesting follows torch: each ``__enter__`` pushes the new active config, each
``__exit__`` restores the previous one, and a nested ``enabled=False`` region
simply disables dispatch inside it. The instance doubles as a decorator.
"""
def __init__(
self,
enabled: bool = True,
update_interval: int = 16,
recipe: Optional[FP8Recipe] = None,
fp8_format: str = "hybrid",
margin: int = 0,
):
if recipe is None:
recipe = FP8Recipe(history_len=update_interval, margin=margin)
self._config = _ActiveConfig(bool(enabled), recipe, FP8Format(fp8_format))
self._tokens: List[Token] = []
def __enter__(self) -> "fp8_autocast":
self._tokens.append(_active_config.set(self._config))
return self
def __exit__(self, exc_type, exc_val, exc_tb) -> bool:
token = self._tokens.pop()
_active_config.reset(token)
return False
def __call__(self, func):
@functools.wraps(func)
def decorate(*args, **kwargs):
with self:
return func(*args, **kwargs)
return decorate
# ---------------------------------------------------------------------------
# Strategy-level forward / backward (called from the aten::linear impl)
# ---------------------------------------------------------------------------
def _dynamic_scale(t: torch.Tensor, recipe: FP8Recipe, fmt: str) -> torch.Tensor:
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
return recipe.scale_from_history(amax, fmt)
def _is_fp8(dtype: torch.dtype) -> bool:
"""A pre-quantized weight takes the GEMM directly (no re-quantize)."""
return dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
def fp8_linear_forward(
x: torch.Tensor, w: torch.Tensor, bias=None, cfg: Optional[_ActiveConfig] = None
):
"""Scaled fp8 linear forward (called from the aten::linear impl).
Composed from the two stateless primitives: quantize x/w with the active
scales, run the pre-quantized GEMM with the bias fused into its epilogue.
Delayed scaling lets the quantize kernel fold the fused amax into the
history ring and publish the next scale in its own last block; dynamic
scaling measures the current amax itself. Training quantizes the weight
every step (the optimizer bumps its version, so there is no cast cache,
matching ``cached_cast``-less behavior).
"""
state = fp8_state()
if cfg is None:
cfg = _current_config()
fmt = cfg.fp8_format.fwd()
if cfg.recipe.dynamic:
sx = _dynamic_scale(x.reshape(-1, w.size(1)), cfg.recipe, fmt)
sw = _dynamic_scale(w, cfg.recipe, fmt)
x8, _ = quantize(x, sx.reciprocal(), fmt)
w8 = w if _is_fp8(w.dtype) else quantize(w, sw.reciprocal(), fmt)[0]
# Bias fuses into the GEMM epilogue (fp32 add before the single bf16
# rounding — one rounding fewer than the separate out + bias pass);
# None passes through to the kernel's no-bias path.
out = mm_fp8(
x8.reshape(-1, x8.size(-1)), w8, sx * sw, trans_b=True, bias=bias
).reshape(*x.shape[:-1], w.size(0))
return out, sx, sw
meta = state.get_weight_meta(w, cfg.recipe)
if not meta.w.initialized:
meta.w.seed(w, fmt)
if not meta.x.initialized:
meta.x.seed(x, fmt)
sx, sw = meta.x.scale.clone(), meta.w.scale.clone()
# The clones feed this call's kernels (stream-ordered before the in-kernel
# fold overwrites the ring scale slots); the fp8 quantize kernel folds the
# amax into the history window and publishes the next scale itself.
x8, _ = quantize(x, sx.reciprocal(), fmt, **meta.x.fold_args(fmt))
if _is_fp8(w.dtype):
w8 = w
else:
w8, _ = quantize(w, sw.reciprocal(), fmt, **meta.w.fold_args(fmt))
out = mm_fp8(
x8.reshape(-1, x8.size(-1)), w8, sx * sw, trans_b=True, bias=bias
).reshape(*x.shape[:-1], w.size(0))
meta.x.advance()
if not _is_fp8(w.dtype):
meta.w.advance()
return out, sx, sw
class _LinearFp8(torch.autograd.Function):
"""The fp8 linear forward/backward pair (standard Function style).
The forward runs inside ``fp8_autocast`` and captures the active
fmt/recipe/meta on ``ctx``; the backward reads only that captured state, so
``loss.backward()`` may run after the context exits. The gradient is
quantized once (E5M2 in hybrid) and both dX/dW GEMMs share it; the output
masks come from ``needs_input_grad``.
"""
@staticmethod
def forward(ctx, x, w, bias):
cfg = _current_config()
out, sx, sw = fp8_linear_forward(x, w, bias, cfg)
ctx.save_for_backward(x, w, sx, sw)
ctx.fmt_bwd = cfg.fp8_format.bwd()
ctx.recipe = cfg.recipe
ctx.is_dynamic = cfg.recipe.dynamic
ctx.meta = None if ctx.is_dynamic else _state.get_weight_meta(w, cfg.recipe)
return out
@staticmethod
@torch.autograd.function.once_differentiable
def backward(ctx, g):
x, w, _sx_fwd, _sw_fwd = ctx.saved_tensors
fmt = ctx.fmt_bwd
# Flatten leading dims (the forward GEMMs ran on [-1, N] / [-1, K]
# views; the kernels only accept 2D operands).
g2 = g.reshape(-1, g.size(-1))
if ctx.is_dynamic:
sg = _dynamic_scale(g2, ctx.recipe, fmt)
sw = _dynamic_scale(w, ctx.recipe, fmt)
sx = _dynamic_scale(x, ctx.recipe, fmt)
else:
meta = ctx.meta
if not meta.g.initialized:
meta.g.seed(g2, fmt)
sg = meta.g.scale.clone()
sw, sx = _sw_fwd, _sx_fwd
# Backward GEMMs route through the NT fast path via transposed
# quantize outputs: g8 [m,n] with w8T [k,n] (trans_b=True) gives
# grad_x, g8T [n,m] with x8T [k,m] gives grad_w — no NN-swap or TT
# crosswise kernel in the training path. g is consumed in both
# orientations, so quantize_dual's single pass feeds both.
# The g quantize folds the gradient amax into its ring in-kernel;
# the x8T/w8T orientation copies discard amax (those rings were
# folded at forward time).
g8, g8T, _ = quantize_dual(g2, sg.reciprocal(), fmt, **meta.g.fold_args(fmt))
x8T, _ = quantize(
x.reshape(-1, x.size(-1)), sx.reciprocal(), fmt, transposed=True
)
if _is_fp8(w.dtype):
# Pre-quantized weight has no transposed copy: keep the swap
# path for grad_x (grad_w is unaffected).
grad_x = mm_fp8(g8, w, sg * sw).reshape(x.shape)
else:
w8T, _ = quantize(w, sw.reciprocal(), fmt, transposed=True)
grad_x = mm_fp8(g8, w8T, sg * sw, trans_b=True).reshape(x.shape)
grad_w = mm_fp8(g8T, x8T, sg * sx, trans_b=True) # g8.T @ x8
# bias-free linears must not pay the column-sum
# reduce: g2.sum(0) is another full read of the gradient.
grad_b = g2.sum(0).to(torch.bfloat16) if ctx.needs_input_grad[2] else None
if not ctx.is_dynamic:
meta.g.advance()
return grad_x, grad_w, grad_b
# ---------------------------------------------------------------------------
# aten::linear integration
# ---------------------------------------------------------------------------
def fp8_linear_enable(enabled: bool = True) -> None:
"""Toggle fp8 dispatch for aten::linear globally (the out-of-region default;
``fp8_autocast`` regions override it thread-locally)."""
fp8_state().default_enabled = enabled
def fp8_linear_enabled() -> bool:
"""Whether fp8 dispatch is active right now (region config or global)."""
return _active() is not None
def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
"""Shape guard for the fp8 path. Unlike a strict 16-alignment requirement,
the kernels handle unaligned M/N via boundary checks (slower but correct) —
so no whole-call bf16 fallback for small decode batches. Only the K-dimension
contraction must match and the weight must be 2D."""
return x.dim() >= 2 and w.dim() == 2 and x.size(-1) == w.size(1)
def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
if (
_active() is not None
and x.dtype is torch.bfloat16
and w.dtype is torch.bfloat16
and _fp8_supported(x, w)
):
return _LinearFp8.apply(x, w, bias)
return torch.ops.aten.linear.default.redispatch(
torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd),
x,
w,
bias,
)
_lib = Library("aten", "IMPL", "CUDA")
_lib.impl("linear", _linear_cuda_impl)
# Also replace torch's generated linear autograd formula (which would call
# aten::linear_backward after the fp8_autocast region exits). The fp8 backward
# is owned by _LinearFp8 with state captured at forward time, so loss.backward()
# works wherever it is called; the CUDA registration still covers inference_mode.
_lib_autograd = Library("aten", "IMPL", "AutogradCUDA")
_lib_autograd.impl("linear", _linear_cuda_impl)
+63 -22
View File
@@ -1,42 +1,83 @@
"""Dynamic discovery and loading of compiled CUDA kernel modules. """Dynamic discovery and loading of compiled CUDA kernel modules.
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed Each kernel is built by the CMake build in ``csrc/CMakeLists.txt`` into a
in this package directory. On import we try to load each one; kernels that ``.so`` placed in ``astrai/extension/lib/`` — the module name equals the
failed to build (or are running on a CPU-only machine) are marked unavailable ``.so`` name equals the pybind name (e.g. ``attn_decode``, defined via
so the wrapper functions can fall back to ``torch`` SDPA. ``TORCH_EXTENSION_NAME``). ``KERNEL_NAMES`` is discovered automatically from
the ``.so`` files present, so adding a kernel to the CMake ``KERNELS``
registry needs no change here.
Loading is **lazy and centralized**: module names are discovered eagerly
(cheap glob), but each ``.so`` is imported on first use via the single
``get_module`` accessor, then cached. The wrapper modules (``ops/*.py``) never
touch the internals or keep their own caches — they call ``get_module(name)``
(or ``is_available(name)`` when a torch fallback is acceptable). A kernel that
failed to build (or is running on a CPU-only machine) is ``None`` in the cache,
so ``is_available`` returns ``False`` and ``get_module`` raises a clear error.
""" """
import glob
import importlib import importlib
import logging import logging
import os
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
KERNEL_NAMES = [ _LIB_DIR = os.path.join(os.path.dirname(__file__), "lib")
"attn_decode",
"attn_prefill",
"attn_paged_decode", def _discover_kernel_names() -> list[str]:
"attn_paged_prefill", """Return the module names of the compiled kernel ``.so`` files in lib/."""
"rotary_emb", names: list[str] = []
] for path in glob.glob(os.path.join(_LIB_DIR, "*.so")):
# strip the "<soabi>.so" suffix, e.g. attn_decode.cpython-312-...so
names.append(os.path.basename(path).split(".", 1)[0])
return sorted(names)
KERNEL_NAMES = _discover_kernel_names()
_available: dict[str, bool] = {} _available: dict[str, bool] = {}
_modules: dict[str, object] = {} _modules: dict[str, object] = {}
for _name in KERNEL_NAMES:
try: def _try_load(name: str) -> object:
_mod = importlib.import_module(f".lib.{_name}", package=__package__) """Import and cache the ``name`` kernel module (lazy, one attempt).
_available[_name] = True
_modules[_name] = _mod Returns the module, or ``None`` if it is unavailable. Cached so each
except ImportError: ``.so`` is imported at most once per process.
_available[_name] = False """
_modules[_name] = None if name not in _modules:
try:
_modules[name] = importlib.import_module(
f".lib.{name}", package=__package__
)
_available[name] = True
except ImportError:
logger.warning("kernel '%s' failed to import; marking unavailable", name)
_modules[name] = None
_available[name] = False
return _modules[name]
def is_available(name: str) -> bool: def is_available(name: str) -> bool:
"""Return ``True`` if the compiled kernel ``name`` was loaded.""" """Return ``True`` if the compiled kernel ``name`` could be loaded."""
if name not in _available:
_try_load(name)
return _available.get(name, False) return _available.get(name, False)
def get_module(name: str) -> object: def get_module(name: str) -> object:
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable.""" """Return the loaded kernel module for ``name``, importing it on first use.
return _modules.get(name)
Raises ``RuntimeError`` if the kernel is unavailable (not built, or failed
to import) — callers that can tolerate a torch fallback should check
``is_available(name)`` first instead.
"""
mod = _try_load(name)
if mod is None:
raise RuntimeError(
f"CUDA kernel '{name}' is not available. "
f"Build with CSRC_KERNELS=true (or use the torch-native fallback)."
)
return mod
+19
View File
@@ -0,0 +1,19 @@
"""Stateless wrappers around compiled extension kernels."""
from astrai.extension.ops.attention import (
TensorLayout,
attn_decode,
attn_paged_decode,
attn_paged_prefill,
attn_prefill,
)
from astrai.extension.ops.rotary import rotary_emb
__all__ = [
"TensorLayout",
"attn_decode",
"attn_paged_decode",
"attn_paged_prefill",
"attn_prefill",
"rotary_emb",
]
@@ -1,4 +1,4 @@
"""Attention kernel wrapper functions one entry point per compiled kernel. """Attention kernel wrapper functions - one entry point per compiled kernel.
Each wrapper calls its CUDA kernel directly. If the kernel is not Each wrapper calls its CUDA kernel directly. If the kernel is not
available, raises ``RuntimeError``. Fallback to torch SDPA is the available, raises ``RuntimeError``. Fallback to torch SDPA is the
@@ -17,7 +17,7 @@ from typing import Optional
import torch import torch
from astrai.extension.loader import _available, _modules from astrai.extension.loader import get_module
class TensorLayout(enum.IntEnum): class TensorLayout(enum.IntEnum):
@@ -30,14 +30,6 @@ class TensorLayout(enum.IntEnum):
BLHD = 1 # [batch, seq_len, n_heads, head_dim] BLHD = 1 # [batch, seq_len, n_heads, head_dim]
def _check_available(name: str):
if not _available.get(name):
raise RuntimeError(
f"CUDA kernel '{name}' is not available. "
f"Build with CSRC_KERNELS=true or use a torch-native backend."
)
def attn_decode( def attn_decode(
q: torch.Tensor, q: torch.Tensor,
k: torch.Tensor, k: torch.Tensor,
@@ -57,9 +49,9 @@ def attn_decode(
Returns: Returns:
[batch, 1, n_heads, head_dim] (blhd, bf16) [batch, 1, n_heads, head_dim] (blhd, bf16)
""" """
_check_available("attn_decode") mod = get_module("attn_decode")
causal_offset = (k.size(1) - 1) if is_causal else -1 causal_offset = (k.size(1) - 1) if is_causal else -1
return _modules["attn_decode"].attn_decode( return mod.attn_decode(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
) )
@@ -83,9 +75,9 @@ def attn_prefill(
Returns: Returns:
[batch, q_len, n_heads, head_dim] (blhd, bf16) [batch, q_len, n_heads, head_dim] (blhd, bf16)
""" """
_check_available("attn_prefill") mod = get_module("attn_prefill")
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1 causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
return _modules["attn_prefill"].attn_prefill( return mod.attn_prefill(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
) )
@@ -97,7 +89,8 @@ def attn_paged_decode(
req_to_token: torch.Tensor, req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
kv_indptr: torch.Tensor, kv_indptr: torch.Tensor,
max_seq_len: int, new_k: Optional[torch.Tensor] = None,
new_v: Optional[torch.Tensor] = None,
mask: Optional[torch.Tensor] = None, mask: Optional[torch.Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
o_part_buf: Optional[torch.Tensor] = None, o_part_buf: Optional[torch.Tensor] = None,
@@ -114,11 +107,12 @@ def attn_paged_decode(
q: [batch, n_heads, head_dim] (bf16, 3D no seq dim) q: [batch, n_heads, head_dim] (bf16, 3D no seq dim)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat) k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int64) token -> slot req_to_token: [num_reqs, max_context_len] (int32) token -> slot
req_pool_indices: [batch] (int64) rows into req_to_token req_pool_indices: [batch] (int32) rows into req_to_token
kv_indptr: [batch+1] (int32) prefix sum of per-request seq_lens kv_indptr: [batch+1] (int32) prefix sum of per-request seq_lens
max_seq_len: max per-request seq_len (Python int, for split computation) new_k: current-token K to append, [batch, n_kv_heads, head_dim]
mask: 2D [batch, max_seq_len] (bool, True=keep) or None new_v: current-token V to append, same shape as new_k
mask: 2D [batch, max_context_len] (bool, True=keep) or None
is_causal: apply causal mask is_causal: apply causal mask
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass) o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass) ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
@@ -127,16 +121,17 @@ def attn_paged_decode(
Returns: Returns:
[batch, n_heads, head_dim] (bf16, 3D) [batch, n_heads, head_dim] (bf16, 3D)
""" """
_check_available("attn_paged_decode") mod = get_module("attn_paged_decode")
causal_offset = 0 if is_causal else -1 causal_offset = 0 if is_causal else -1
return _modules["attn_paged_decode"].attn_paged_decode( return mod.attn_paged_decode(
q, q,
k_cache, k_cache,
v_cache, v_cache,
req_to_token, req_to_token,
req_pool_indices, req_pool_indices,
kv_indptr, kv_indptr,
max_seq_len, new_k=new_k,
new_v=new_v,
mask=mask, mask=mask,
causal_offset=causal_offset, causal_offset=causal_offset,
o_part_buf=o_part_buf, o_part_buf=o_part_buf,
@@ -153,8 +148,9 @@ def attn_paged_prefill(
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
kv_indptr: torch.Tensor, kv_indptr: torch.Tensor,
qo_indptr: torch.Tensor, qo_indptr: torch.Tensor,
q_tile_to_batch: torch.Tensor,
q_tile_to_index: torch.Tensor,
mask: Optional[torch.Tensor] = None, mask: Optional[torch.Tensor] = None,
max_q_len: int = 0,
is_causal: bool = False, is_causal: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
"""SGLang-style paged prefill (ragged batch, flat KV pool). """SGLang-style paged prefill (ragged batch, flat KV pool).
@@ -167,20 +163,21 @@ def attn_paged_prefill(
q: [total_q, n_heads, head_dim] (bf16, 3D flattened across requests) q: [total_q, n_heads, head_dim] (bf16, 3D flattened across requests)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat) k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int64) req_to_token: [num_reqs, max_context_len] (int32)
req_pool_indices: [batch] (int64) req_pool_indices: [batch] (int32)
kv_indptr: [batch+1] (int32) prefix sum of per-request kv_lens kv_indptr: [batch+1] (int32) prefix sum of per-request kv_lens
qo_indptr: [batch+1] (int32) prefix sum of per-request q_lens qo_indptr: [batch+1] (int32) prefix sum of per-request q_lens
q_tile_to_batch: [num_q_tiles] (int32) request index per Q tile
q_tile_to_index: [num_q_tiles] (int32) local Q tile index per request
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
max_q_len: max per-request q_len (Python int, for grid computation)
is_causal: apply causal mask is_causal: apply causal mask
Returns: Returns:
[total_q, n_heads, head_dim] (bf16, 3D) [total_q, n_heads, head_dim] (bf16, 3D)
""" """
_check_available("attn_paged_prefill") mod = get_module("attn_paged_prefill")
causal_offset = 0 if is_causal else -1 causal_offset = 0 if is_causal else -1
return _modules["attn_paged_prefill"].attn_paged_prefill( return mod.attn_paged_prefill(
q, q,
k_cache, k_cache,
v_cache, v_cache,
@@ -188,7 +185,8 @@ def attn_paged_prefill(
req_pool_indices, req_pool_indices,
kv_indptr, kv_indptr,
qo_indptr, qo_indptr,
q_tile_to_batch,
q_tile_to_index,
mask, mask,
max_q_len,
causal_offset=causal_offset, causal_offset=causal_offset,
) )
+116
View File
@@ -0,0 +1,116 @@
"""FP8 CUDA kernel interface adapter (the only module touching the pybind).
Attention-style thin wrappers: one Python entry per binding, called directly
— no torch.library dispatch layer. Optional arguments (``ring_state``,
``bias``) keep native Optional semantics at the pybind boundary, and
in-place buffer updates (the delayed-scaling ring fold, like attention's
KV-cache appends) happen on-stream without mutation declarations. CUDA-only:
non-CUDA or unsupported inputs raise from the binding's TORCH_CHECKs.
- ``quantize(x, scale, fmt, transposed=False) -> (x8|x8T, amax)`` — BF16/FP16/FP32
→ FP8 with fused amax (``transposed`` picks the orientation; arity is fixed)
- ``quantize_dual(x, scale, fmt) -> (x8, x8T, amax)`` — both orientations, one read
- ``mm_fp8(a8, b8, sa, sb) -> out`` — pre-quantized FP8 GEMM (BF16 output)
``scale`` is the quantization multiplier (device scalar); ``fmt`` is
``"e4m3"`` or ``"e5m2"``. ``amax`` values are *returned*, never passed as
output arguments.
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
this module is stateless.
"""
from typing import Optional, Tuple
import torch
from astrai.extension.loader import get_module
# fmt string -> kernel int (0 = E4M3, 1 = E5M2)
_FMT_TO_INT = {"e4m3": 0, "e5m2": 1}
def _fmt_int(fmt: str) -> int:
try:
return _FMT_TO_INT[fmt]
except KeyError:
raise ValueError(f"unsupported fp8 format {fmt!r} (expected 'e4m3' or 'e5m2')")
def quantize(
x: torch.Tensor,
scale: torch.Tensor,
fmt: str = "e4m3",
transposed: bool = False,
ring_state: Optional[torch.Tensor] = None,
hist_idx: int = 0,
fp8_max: float = 448.0,
pow2_margin: float = 1.0,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Float (bf16/fp16/fp32) -> FP8 quantize with fused amax.
``scale`` is the quantization multiplier (device scalar); ``fmt`` selects
E4M3 or E5M2. ``amax`` is a fresh 1-element float32 tensor.
``transposed=True`` swaps ``x8`` for ``x8T``, the ``[cols][rows]``
row-major transpose of the quantized input — the K-contiguous operand
orientation NT GEMMs want — at the same 2-tuple arity.
``ring_state`` (a 1D float32 CUDA buffer laid out
``[hist n | scale | legacy | amax | done]``) switches on the in-kernel
delayed-scaling fold: the kernel's last block folds the amax into
``hist[hist_idx]`` and publishes the next scale as
``max(hist) / fp8_max / pow2_margin`` — the returned ``amax`` is then the
self-cleaned persistent slot (reads zero). None keeps the classic
fresh-amax return.
"""
return get_module("fp8_ops").quantize(
x,
scale,
_fmt_int(fmt),
transposed,
ring_state,
hist_idx,
fp8_max,
pow2_margin,
)
def quantize_dual(
x: torch.Tensor,
scale: torch.Tensor,
fmt: str = "e4m3",
ring_state: Optional[torch.Tensor] = None,
hist_idx: int = 0,
fp8_max: float = 448.0,
pow2_margin: float = 1.0,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Dual-orientation quantize: one read of ``x`` produces both the
row-major ``x8`` and its transposed ``x8T`` (plus ``amax``), for tensors
consumed by GEMMs in both orientations (backward ``g``).
``ring_state`` switches on the in-kernel delayed-scaling fold exactly as
in :func:`quantize`.
"""
return get_module("fp8_ops").quantize_dual(
x, scale, _fmt_int(fmt), ring_state, hist_idx, fp8_max, pow2_margin
)
def mm_fp8(
a: torch.Tensor,
b: torch.Tensor,
scale: torch.Tensor,
trans_a: bool = False,
trans_b: bool = False,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Pre-quantized FP8 GEMM: ``a @ b * scale (+ bias)``.
``a``/``b`` must be FP8 tensors of the same format, 2D or 3D (batched,
matmul-style broadcast on the batch dim). Inner-transposed views (e.g.
``x.t()``) fold into the layout at zero copy. ``scale`` is their combined
dequantization scale. ``bias`` (CUDA bf16 1D of length n) adds inside the
kernel epilogue in fp32 — no separate elementwise pass. The result is
BF16; FP8 output is a separate quantize operation.
"""
return get_module("fp8_ops").mm_fp8(a, b, scale, trans_a, trans_b, bias)
+31
View File
@@ -0,0 +1,31 @@
"""Rotary embedding CUDA kernel wrapper.
Calls the compiled CUDA kernel directly. If the kernel is not available,
raises ``RuntimeError``. Fallback to torch complex multiply is the
responsibility of ``astrai.extension.backend.rotary.apply_rotary_emb``.
Layout: x is packed [tokens, n_heads, head_dim] or dense
[batch, seq_len, n_heads, head_dim]. ``freqs_cis`` has matching token axes.
"""
import torch
from astrai.extension.loader import get_module
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
"""Fused rotary embedding kernel.
Args:
x: packed 3D or dense 4D bf16 tensor.
freqs_cis: matching token axes followed by [head_dim/2, 2].
Returns:
Tensor with the same shape as ``x``.
"""
mod = get_module("rotary_emb")
if not x.is_contiguous():
x = x.contiguous()
if not freqs_cis.is_contiguous():
freqs_cis = freqs_cis.contiguous()
return mod.rotary_emb(x, freqs_cis)
-39
View File
@@ -1,39 +0,0 @@
"""Rotary embedding CUDA kernel wrapper.
Calls the compiled CUDA kernel directly. If the kernel is not available,
raises ``RuntimeError``. Fallback to torch complex multiply is the
responsibility of ``astrai.extension.rotary_backend.apply_rotary_emb``.
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16, contiguous).
freqs_cis is [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs.
"""
import torch
from astrai.extension.loader import _available, _modules
def _check_available():
if not _available.get("rotary_emb"):
raise RuntimeError(
"CUDA kernel 'rotary_emb' is not available. "
"Build with CSRC_KERNELS=true or use the torch fallback."
)
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
"""Fused rotary embedding kernel.
Args:
x: [batch, seq_len, n_heads, head_dim] (bf16, contiguous)
freqs_cis: [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs
Returns:
[batch, seq_len, n_heads, head_dim] (bf16)
"""
_check_available()
if not x.is_contiguous():
x = x.contiguous()
if not freqs_cis.is_contiguous():
freqs_cis = freqs_cis.contiguous()
return _modules["rotary_emb"].rotary_emb(x, freqs_cis)
+15 -76
View File
@@ -1,57 +1,23 @@
"""Inference module for continuous batching. """Inference module for continuous batching.
Layers: Subpackages:
- core/: Core inference loop (cache, executor, scheduler, task) - cache/: KV cache (buffers, strategies, pool)
- api/: HTTP orchestration (ProtocolHandler, server) - runtime/: Execution + sampling (executor, CUDA graph, sampling strategies)
- protocols/: Response builders (OpenAI, Anthropic) - task/: Request lifecycle + performance metrics
- transport/: SSE transport utilities - network/: HTTP protocol handling (server, protocol, OpenAI/Anthropic builders)
- engine.py: Facade (InferenceEngine)
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy) Modules:
- scheduler.py: Continuous batching loop
- workspace.py: Pre-allocated GPU buffers
- engine.py: Facade (InferenceEngine)
""" """
from astrai.inference.api import (
AnthropicMessage,
BaseToolParser,
ChatCompletionRequest,
ChatMessage,
FunctionDef,
GenContext,
MessagesRequest,
ProtocolHandler,
SimpleJsonToolParser,
StopChecker,
ToolDef,
ToolParserFactory,
get_app,
run_server,
)
from astrai.inference.api.anthropic import AnthropicResponseBuilder
from astrai.inference.api.openai import OpenAIResponseBuilder
from astrai.inference.core import (
STOP,
Allocator,
Executor,
InferenceScheduler,
KVCache,
KVStorage,
PagePool,
RadixCache,
ReqToTokenPool,
Task,
TaskManager,
TaskStatus,
page_hash,
)
from astrai.inference.engine import InferenceEngine from astrai.inference.engine import InferenceEngine
from astrai.inference.sample import ( from astrai.inference.network import get_app, run_server
BaseSamplingStrategy, from astrai.inference.runtime.executor import Executor
FrequencyPenaltyStrategy, from astrai.inference.runtime.sample import sample
SamplingPipeline, from astrai.inference.scheduler import InferenceScheduler
TemperatureStrategy, from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
TopKStrategy,
TopPStrategy,
sample,
)
__all__ = [ __all__ = [
"InferenceEngine", "InferenceEngine",
@@ -61,34 +27,7 @@ __all__ = [
"Task", "Task",
"TaskManager", "TaskManager",
"TaskStatus", "TaskStatus",
"Allocator",
"KVCache",
"KVStorage",
"PagePool",
"RadixCache",
"ReqToTokenPool",
"page_hash",
"sample", "sample",
"BaseSamplingStrategy",
"TemperatureStrategy",
"TopKStrategy",
"TopPStrategy",
"FrequencyPenaltyStrategy",
"SamplingPipeline",
"ProtocolHandler",
"StopChecker",
"GenContext",
"BaseToolParser",
"SimpleJsonToolParser",
"ToolParserFactory",
"OpenAIResponseBuilder",
"AnthropicResponseBuilder",
"ChatMessage",
"ChatCompletionRequest",
"FunctionDef",
"ToolDef",
"AnthropicMessage",
"MessagesRequest",
"get_app", "get_app",
"run_server", "run_server",
] ]
+27
View File
@@ -0,0 +1,27 @@
"""KV cache subsystem: buffers, strategies, pool management."""
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
from astrai.inference.cache.pool import PagePool, TaskCacheManager, page_hash
from astrai.inference.cache.strategy import (
AllocationStrategy,
Allocator,
ContiguousStrategy,
PagedStrategy,
RadixCache,
TaskCacheState,
)
__all__ = [
"KVCache",
"KVStorage",
"ReqToTokenPool",
"Allocator",
"RadixCache",
"TaskCacheState",
"AllocationStrategy",
"ContiguousStrategy",
"PagedStrategy",
"PagePool",
"TaskCacheManager",
"page_hash",
]
+96
View File
@@ -0,0 +1,96 @@
"""Physical KV cache buffers.
Layer 1 — ``KVStorage``: flat token-level K/V GPU buffers [n_layers, size, n_kv_heads, head_dim]
Layer 2 — ``ReqToTokenPool``: index table [req_idx, pos] → physical token slot
Layer 3 — ``KVCache``: pure dataclass passed to the model for direct buffer access
These classes have no knowledge of tasks, allocation policies, or scheduling.
They are the "dumb" physical storage layer.
"""
import threading
from dataclasses import dataclass
from typing import List, Optional
import torch
from torch import Tensor
class ReqToTokenPool:
"""Maps [req_idx, pos] → physical token slot in KV storage.
Each row is one request; each column is a sequence position. The value
at [req_idx, pos] is the flat index into the KV storage buffers.
"""
def __init__(self, size: int, max_context_len: int, device: torch.device):
self.size = size
self.max_context_len = max_context_len
self.req_to_token = torch.zeros(
(size, max_context_len), dtype=torch.int32, device=device
)
self.free_slots = list(range(size))
self._lock = threading.Lock()
def alloc(self, num_reqs: int) -> Optional[List[int]]:
with self._lock:
if num_reqs > len(self.free_slots):
return None
slots = self.free_slots[:num_reqs]
self.free_slots = self.free_slots[num_reqs:]
return slots
def free(self, req_indices: List[int]):
with self._lock:
self.free_slots.extend(req_indices)
def write(self, indices, values):
self.req_to_token[indices] = values
class KVStorage:
"""Token-level KV cache storage.
Buffers: ``[n_layers, size, n_kv_heads, head_dim]``. Each token occupies
one slot indexed by ``ReqToTokenPool``.
"""
def __init__(
self,
size: int,
n_layers: int,
n_kv_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.size = size
self.k_buffer = torch.empty(
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
)
self.v_buffer = torch.empty(
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
)
@dataclass
class KVCache:
"""Pure data struct passed to model for KV cache I/O.
The attention layer does raw buffer indexing — no methods, no abstraction.
"""
k_buffer: Tensor
v_buffer: Tensor
req_to_token: Tensor
req_pool_indices: Tensor
seq_lens: Tensor
out_cache_loc: Tensor
max_len: int = 0
kv_indptr: Optional[Tensor] = None
qo_indptr: Optional[Tensor] = None
q_tile_to_batch: Optional[Tensor] = None
q_tile_to_index: Optional[Tensor] = None
decode_o_part: Optional[Tensor] = None
decode_ml_part: Optional[Tensor] = None
decode_out: Optional[Tensor] = None
+382
View File
@@ -0,0 +1,382 @@
"""KV cache orchestration: PagePool + TaskCacheManager.
PagePool owns the physical buffers (``KVStorage`` + ``ReqToTokenPool``)
and wires them to an allocation strategy. It assembles the ``KVCache``
dataclass passed to the model forward.
TaskCacheManager owns the ``task_id`` → ``TaskCacheState`` mapping and
delegates physical slot allocation to the strategy, and KV bind to the pool.
See ``cache_buffer.py`` for the raw buffer primitives and ``cache_strategy.py``
for the allocation policies.
"""
from dataclasses import dataclass
from typing import Dict, List, Optional
import torch
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
from astrai.inference.cache.strategy import (
AllocationStrategy,
Allocator,
ContiguousStrategy,
PagedStrategy,
RadixCache,
TaskCacheState,
)
from astrai.inference.workspace import Q_TILE_ROWS, InferenceWorkspace
# Re-export everything so existing ``from astrai.inference.cache import ...``
# continues to work unchanged after the file split.
__all__ = [
"KVCache",
"KVStorage",
"ReqToTokenPool",
"Allocator",
"RadixCache",
"AllocationStrategy",
"ContiguousStrategy",
"PagedStrategy",
"PagePool",
"TaskCacheManager",
"TaskCacheState",
"page_hash",
]
# ---- helpers ----
def page_hash(
token_ids: List[int], page_idx: int, page_size: int, parent_hash: int = 0
) -> int:
start = page_idx * page_size
end = min(start + page_size, len(token_ids))
h = parent_hash
for i in range(start, end):
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
return h
def _is_steady_increment(
prev_sig: Optional[tuple],
prev_vals: Optional[List[int]],
cur_sig: tuple,
cur_vals: List[int],
) -> bool:
return (
prev_sig is not None
and prev_vals is not None
and prev_sig == cur_sig
and len(prev_vals) == len(cur_vals)
and all(c == p + 1 for c, p in zip(cur_vals, prev_vals))
)
# ---- task-scoped bind state ----
@dataclass
class _BindState:
"""Cached bind metadata for steady-state decode increment detection."""
sig: tuple
seq_lens: List[int]
# ---- pool + manager ----
class PagePool:
"""Physical KV cache: buffers + req-to-token table + allocation strategy + bind.
Does not know about tasks — task lifecycle is managed by
:class:`TaskCacheManager`, which holds a reference to this pool.
"""
def __init__(
self,
n_layers: int,
n_kv_heads: int,
head_dim: int,
max_batch_size: int,
max_seq_len: int,
device: torch.device,
dtype: torch.dtype,
page_size: int = 1,
n_tokens: Optional[int] = None,
):
self.page_size = page_size
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.device = device
self.dtype = dtype
self.n_layers = n_layers
self.n_kv_heads = n_kv_heads
self.head_dim = head_dim
self.contiguous = n_tokens is None
self.n_tokens = max_batch_size * max_seq_len if self.contiguous else n_tokens
if self.n_tokens > torch.iinfo(torch.int32).max:
raise ValueError("KV cache token count exceeds the int32 slot index limit")
self._storage = KVStorage(
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
)
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
if self.contiguous:
for i in range(max_batch_size):
self._req_pool.req_to_token[i] = torch.arange(
i * max_seq_len,
(i + 1) * max_seq_len,
dtype=torch.int32,
device=device,
)
self._strategy: AllocationStrategy = ContiguousStrategy()
else:
n_pages = self.n_tokens // page_size
alloc = Allocator(n_pages)
prefix = RadixCache(page_size) if page_size > 1 else None
if prefix is not None:
alloc.on_evict = prefix.evict
self._strategy = PagedStrategy(
alloc, prefix, page_size, self._req_pool, device
)
@property
def strategy(self) -> AllocationStrategy:
return self._strategy
@property
def req_pool(self) -> ReqToTokenPool:
return self._req_pool
def bind_tasks(
self,
req_indices: List[int],
seq_lens: List[int],
workspace: InferenceWorkspace,
device: Optional[torch.device] = None,
start_pos: Optional[int] = None,
incremental: bool = False,
) -> KVCache:
"""Assemble the ``KVCache`` metadata for a batch of tasks.
Args:
req_indices: request slot indices (from ``ReqToTokenPool``).
seq_lens: current sequence length per task.
workspace: pre-allocated fixed-shape buffers (CUDA-graph safe).
start_pos: if set, produce **prefill** cache (full q_len range).
If ``None``, produce **decode** cache (last position).
incremental: if ``True``, reuse workspace state from previous step
by incrementing counters in-place (decode hot path).
Returns:
``KVCache`` dataclass with the correct output shapes for the
attention backend (prefill: ``[B, q_len]``, decode: ``[B, 1]``).
"""
if device is None:
device = workspace.device
b = len(req_indices)
rpi_buf = workspace.req_pool_indices
sl_buf = workspace.seq_lens
kvp_buf = workspace.kv_indptr
inc_buf = workspace.inc
ocl_buf = workspace.out_cache_loc
if incremental:
sl_buf[:b] += 1
kvp_buf[: b + 1] += inc_buf[: b + 1]
else:
rpi_buf[:b].copy_(
torch.tensor(req_indices, dtype=torch.int32, device=device)
)
sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device))
kvp_buf[: b + 1].zero_()
kvp_buf[1 : b + 1] = sl_buf[:b].cumsum(0).to(torch.int32)
req_pool_indices = rpi_buf[:b]
seq_lens_t = sl_buf[:b]
kv_indptr = kvp_buf[: b + 1]
if start_pos is not None:
# Packed prefill concatenates each request's query tokens.
q_lens = [seq_len - start_pos for seq_len in seq_lens]
if any(q_len <= 0 for q_len in q_lens):
raise ValueError("prefill sequence lengths must exceed start_pos")
out_cache_loc = torch.cat(
[
self._req_pool.req_to_token[
req_pool_indices[i], start_pos : seq_lens[i]
]
for i in range(b)
]
)
workspace.qo_indptr[: b + 1].zero_()
workspace.qo_indptr[1 : b + 1].copy_(
torch.tensor(q_lens, dtype=torch.int32, device=device).cumsum(0)
)
qo_indptr = workspace.qo_indptr[: b + 1]
tile_batches = []
tile_indices = []
for batch, q_len in enumerate(q_lens):
n_tiles = (q_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS
tile_batches.extend([batch] * n_tiles)
tile_indices.extend(range(n_tiles))
n_tiles = len(tile_batches)
workspace.q_tile_to_batch[:n_tiles].copy_(
torch.tensor(tile_batches, dtype=torch.int32, device=device)
)
workspace.q_tile_to_index[:n_tiles].copy_(
torch.tensor(tile_indices, dtype=torch.int32, device=device)
)
q_tile_to_batch = workspace.q_tile_to_batch[:n_tiles]
q_tile_to_index = workspace.q_tile_to_index[:n_tiles]
decode_o_part = decode_ml_part = decode_out = None
else:
# ---- decode: out_cache_loc is a single column (last position) ----
write_pos = seq_lens_t - 1
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
ocl_buf[:b].copy_(loc)
out_cache_loc = ocl_buf[:b].reshape(-1)
workspace.qo_indptr[: b + 1].copy_(inc_buf[: b + 1])
qo_indptr = workspace.qo_indptr[: b + 1]
q_tile_to_batch = q_tile_to_index = None
decode_o_part = getattr(workspace, "decode_o_part", None)
decode_ml_part = getattr(workspace, "decode_ml_part", None)
decode_out = getattr(workspace, "decode_out", None)
return KVCache(
k_buffer=self._storage.k_buffer,
v_buffer=self._storage.v_buffer,
req_to_token=self._req_pool.req_to_token,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens_t,
out_cache_loc=out_cache_loc,
max_len=max(seq_lens),
kv_indptr=kv_indptr,
qo_indptr=qo_indptr,
q_tile_to_batch=q_tile_to_batch,
q_tile_to_index=q_tile_to_index,
decode_o_part=decode_o_part,
decode_ml_part=decode_ml_part,
decode_out=decode_out,
)
class TaskCacheManager:
"""Task ↔ KV slot lifecycle manager.
Sole owner of ``task_id → TaskCacheState``. Delegates physical slot
allocation to the strategy (via ``pool.strategy``) and KV bind to
``pool.bind_tasks()``.
Usage::
pool = PagePool(...)
mgr = TaskCacheManager(pool)
mgr.task_alloc("req_1", [101, 202, 303])
...
kv = mgr.bind(["req_1"], workspace)
"""
def __init__(self, pool: PagePool):
self._pool = pool
self._strategy = pool.strategy
self._req_pool = pool.req_pool
self._max_seq_len = pool.max_seq_len
self._states: Dict[str, TaskCacheState] = {}
self._bind_state: Optional[_BindState] = None
self._bind_was_steady = False
# -- public task lifecycle --
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
self._bind_state = None
req_slots = self._req_pool.alloc(1)
if req_slots is None:
return False
state = TaskCacheState(req_idx=req_slots[0])
self._states[task_id] = state
if not self._strategy.alloc(state, prompt_ids):
self._rollback(state, task_id)
return False
self._strategy.write_indices(state, prompt_ids)
state.length = len(prompt_ids)
return True
def task_free(self, task_id: str):
self._bind_state = None
state = self._states.pop(task_id, None)
if state is None:
return
self._strategy.free(state)
self._req_pool.free([state.req_idx])
def task_extend(self, task_id: str, pos: int) -> bool:
state = self._states.get(task_id)
if state is None or pos >= self._max_seq_len:
return False
if not self._strategy.extend(state, pos):
return False
state.length = pos + 1
return True
def task_cached(self, task_id: str) -> int:
state = self._states.get(task_id)
return state.cached if state is not None else 0
def task_record_hashes(
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
):
state = self._states.get(task_id)
if state is not None:
self._strategy.record_hashes(state, prompt_ids, start_logical_page)
@staticmethod
def task_cacheable_ids(task_id: str, prompt_ids: List[int], output_ids: List[int]):
return list(prompt_ids) + list(output_ids[:-1])
# -- bind (assemble KVCache for the model forward) --
def bind(
self,
task_ids: List[str],
workspace: InferenceWorkspace,
device: Optional[torch.device] = None,
start_pos: Optional[int] = None,
) -> KVCache:
"""Build ``KVCache`` for an ordered list of task IDs."""
states = [self._states[tid] for tid in task_ids]
req_indices = [s.req_idx for s in states]
seq_lens = [s.length for s in states]
sig = tuple(req_indices)
prev = self._bind_state
incremental = (
start_pos is None
and prev is not None
and _is_steady_increment(prev.sig, prev.seq_lens, sig, seq_lens)
)
self._bind_state = _BindState(sig, list(seq_lens))
self._bind_was_steady = incremental
return self._pool.bind_tasks(
req_indices,
seq_lens,
workspace,
device=device,
start_pos=start_pos,
incremental=incremental,
)
@property
def bind_was_steady(self) -> bool:
return self._bind_was_steady
# -- internals --
def _rollback(self, state: TaskCacheState, task_id: str):
self._strategy.free(state)
self._req_pool.free([state.req_idx])
self._states.pop(task_id, None)
+318
View File
@@ -0,0 +1,318 @@
"""KV cache allocation layer.
Encapsulates the physical slot allocation policy, isolated from GPU buffers
and task lifecycle management.
- ``TaskCacheState``: data contract between strategy and manager (per-task slot state)
- ``Allocator``: bitmask-based page allocator with LRU eviction
- ``RadixCache``: page-granular prefix index (exact token match)
- ``AllocationStrategy``: ABC for physical slot allocation
- ``ContiguousStrategy``: statically partitioned, no dynamic allocation
- ``PagedStrategy``: dynamic paged allocation from a shared pool
"""
import threading
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Callable, Dict, List, Optional, OrderedDict
from astrai.inference.cache.buffer import ReqToTokenPool
# ---- data contract: per-task slot state ----
@dataclass
class TaskCacheState:
"""Per-task cache allocation state.
Co-locates all task-owned cache metadata so the alloc/free/extend
lifecycle is atomic. Owned by ``TaskCacheManager``, consumed by
every ``AllocationStrategy`` method.
"""
req_idx: int
length: int = 0
cached: int = 0
pages: List[int] = field(default_factory=list)
# ---- allocation primitives ----
class Allocator:
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
def __init__(self, n_pages: int):
self._free_mask = (1 << n_pages) - 1
self._refs: List[int] = [0] * n_pages
self._lru: OrderedDict[int, None] = OrderedDict()
self.on_evict: Optional[Callable[[int], None]] = None
self._lock = threading.Lock()
def alloc(self) -> int:
with self._lock:
if self._free_mask:
lsb = self._free_mask & -self._free_mask
idx = lsb.bit_length() - 1
self._free_mask ^= lsb
self._refs[idx] = 1
return idx
if self._lru:
idx, _ = self._lru.popitem(last=False)
if self.on_evict:
self.on_evict(idx)
self._refs[idx] = 1
self._free_mask &= ~(1 << idx)
return idx
return -1
def free(self, idx: int, keep_cached: bool = False):
with self._lock:
self._refs[idx] -= 1
if self._refs[idx] == 0:
if keep_cached:
self._lru[idx] = None
else:
self._free_mask |= 1 << idx
def inc_ref(self, idx: int):
with self._lock:
self._refs[idx] += 1
self._lru.pop(idx, None)
def ref_count(self, idx: int) -> int:
with self._lock:
return self._refs[idx]
def touch(self, idx: int):
with self._lock:
if idx in self._lru:
self._lru.move_to_end(idx)
class RadixNode:
"""A page-aligned edge in the CPU-side prefix radix trie."""
__slots__ = ("parent", "children", "page_idx", "tokens", "lock_ref")
def __init__(self, parent=None, tokens=(), page_idx=None):
self.parent = parent
self.children: Dict[tuple, "RadixNode"] = {}
self.page_idx = page_idx
self.tokens = tuple(tokens)
self.lock_ref = 0
class RadixCache:
"""Page-granular radix prefix index with exact token matching."""
def __init__(self, page_size: int):
self._page_size = page_size
self._root = RadixNode()
self._page_to_node: Dict[int, RadixNode] = {}
self._lock = threading.Lock()
def evict(self, idx: int):
with self._lock:
node = self._page_to_node.pop(idx, None)
if node is None:
return
node.page_idx = None
parent = node.parent
if parent is not None:
parent.children.pop(node.tokens, None)
def has_page(self, idx: int) -> bool:
with self._lock:
return idx in self._page_to_node
def lookup(self, token_ids: List[int]) -> List[int]:
with self._lock:
full_pages = len(token_ids) // self._page_size
hits: List[int] = []
node = self._root
for i in range(full_pages):
start = i * self._page_size
page_tokens = tuple(token_ids[start : start + self._page_size])
child = node.children.get(page_tokens)
if child is None or child.page_idx is None:
break
hits.append(child.page_idx)
node = child
return hits
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
with self._lock:
full_pages = len(token_ids) // self._page_size
if logical_page_idx >= full_pages:
return
old = self._page_to_node.pop(page_idx, None)
if old is not None and old.parent is not None:
old.parent.children.pop(old.tokens, None)
node = self._root
for i in range(logical_page_idx + 1):
start = i * self._page_size
page_tokens = tuple(token_ids[start : start + self._page_size])
child = node.children.get(page_tokens)
if child is None:
child = RadixNode(node, page_tokens)
node.children[page_tokens] = child
node = child
if node.page_idx is not None and node.page_idx != page_idx:
replaced = node.page_idx
self._page_to_node.pop(replaced, None)
node.page_idx = page_idx
self._page_to_node[page_idx] = node
def release(self, pages: List[int]) -> None:
with self._lock:
for page_idx in pages:
node = self._page_to_node.get(page_idx)
if node is not None and node.lock_ref:
node.lock_ref -= 1
class AllocationStrategy(ABC):
"""Physical slot allocation policy.
Subclasses implement the actual allocation semantics. This ABC declares
the contract; there are no default implementations.
"""
@abstractmethod
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool: ...
@abstractmethod
def free(self, state: TaskCacheState) -> None: ...
@abstractmethod
def extend(self, state: TaskCacheState, pos: int) -> bool: ...
@abstractmethod
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None: ...
@abstractmethod
def record_hashes(
self,
state: TaskCacheState,
prompt_ids: List[int],
start: int,
) -> None: ...
class ContiguousStrategy(AllocationStrategy):
"""Static contiguous allocation: slots are pre-assigned at pool init.
No dynamic allocation or prefix caching. All operations are no-ops
because ``ReqToTokenPool`` is pre-filled with contiguous ranges.
"""
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
return True
def free(self, state: TaskCacheState) -> None:
pass
def extend(self, state: TaskCacheState, pos: int) -> bool:
return True
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
pass
def record_hashes(
self,
state: TaskCacheState,
prompt_ids: List[int],
start: int,
) -> None:
pass
class PagedStrategy(AllocationStrategy):
"""Dynamic paged allocation from a shared bitmask pool.
``page_size`` is a parameter, not a separate strategy: at ``page_size=1``
each allocated page *is* one token slot (``page * 1 + 0``), and prefix
caching is simply disabled (``prefix=None``). The unified page formula
``pages[page_idx] * page_size + offset`` holds for both.
"""
def __init__(
self,
alloc: Allocator,
prefix: Optional[RadixCache],
page_size: int,
req_pool: ReqToTokenPool,
device,
):
self._alloc = alloc
self._prefix = prefix
self._page_size = page_size
self._req_pool = req_pool
self._device = device
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
if self._prefix is not None:
hits = self._prefix.lookup(prompt_ids)
state.cached = len(hits) * self._page_size
for p in hits:
self._alloc.inc_ref(p)
state.pages = list(hits)
remaining = len(prompt_ids) - state.cached
if remaining <= 0:
return True
n_new = (remaining + self._page_size - 1) // self._page_size
for _ in range(n_new):
p = self._alloc.alloc()
if p < 0:
return False
state.pages.append(p)
return True
def free(self, state: TaskCacheState) -> None:
if self._prefix is not None:
for p in state.pages:
keep = self._prefix.has_page(p)
self._alloc.free(p, keep_cached=keep)
if not keep:
self._prefix.evict(p)
else:
for p in state.pages:
self._alloc.free(p)
def extend(self, state: TaskCacheState, pos: int) -> bool:
page_idx = pos // self._page_size
if page_idx >= len(state.pages):
p = self._alloc.alloc()
if p < 0:
return False
state.pages.append(p)
offset = pos % self._page_size
self._req_pool.req_to_token[state.req_idx, pos] = (
state.pages[page_idx] * self._page_size + offset
)
return True
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
total = len(prompt_ids)
for pos in range(total):
page_idx = pos // self._page_size
offset = pos % self._page_size
if page_idx < len(state.pages):
self._req_pool.req_to_token[state.req_idx, pos] = (
state.pages[page_idx] * self._page_size + offset
)
def record_hashes(
self,
state: TaskCacheState,
prompt_ids: List[int],
start: int,
) -> None:
if self._prefix is None:
return
full = len(prompt_ids) // self._page_size
for i in range(start, min(full, len(state.pages))):
self._prefix.record(state.pages[i], prompt_ids, i)
-30
View File
@@ -1,30 +0,0 @@
"""Inference core: cache, executor, scheduler, task management."""
from astrai.inference.core.cache import (
Allocator,
KVCache,
KVStorage,
PagePool,
RadixCache,
ReqToTokenPool,
page_hash,
)
from astrai.inference.core.executor import Executor
from astrai.inference.core.scheduler import InferenceScheduler
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
__all__ = [
"Allocator",
"KVCache",
"KVStorage",
"PagePool",
"RadixCache",
"ReqToTokenPool",
"page_hash",
"Executor",
"InferenceScheduler",
"STOP",
"Task",
"TaskManager",
"TaskStatus",
]
-638
View File
@@ -1,638 +0,0 @@
"""KV cache architecture: three-layer separation (SGLang-inspired).
Layer 1 KVStorage: flat token-level K/V buffers [n_layers, size, H, D]
Layer 2 ReqToTokenPool: index table [req_idx, pos] physical token slot
Layer 3 Allocator: slot/page allocation with ref-counting and LRU
PagePool orchestrates all three plus RadixCache (prefix addressing).
KVCache is a pure dataclass passed to the model for direct buffer access.
Two modes:
- contiguous (default): pre-allocated per-request blocks, no dynamic alloc
- paged: shared pool with on-demand allocation, prefix caching support
"""
import threading
from collections import OrderedDict
from dataclasses import dataclass
from typing import Callable, Dict, List, Optional
import torch
from torch import Tensor
from astrai.inference.core.workspace import InferenceWorkspace
def page_hash(
token_ids: List[int], page_idx: int, page_size: int, parent_hash: int = 0
) -> int:
start = page_idx * page_size
end = min(start + page_size, len(token_ids))
h = parent_hash
for i in range(start, end):
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
return h
class Allocator:
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
def __init__(self, n_pages: int):
self._free_mask = (1 << n_pages) - 1
self._refs: List[int] = [0] * n_pages
self._lru: OrderedDict[int, None] = OrderedDict()
self.on_evict: Optional[Callable[[int], None]] = None
self._lock = threading.Lock()
def alloc(self) -> int:
with self._lock:
if self._free_mask:
lsb = self._free_mask & -self._free_mask
idx = lsb.bit_length() - 1
self._free_mask ^= lsb
self._refs[idx] = 1
return idx
if self._lru:
idx, _ = self._lru.popitem(last=False)
if self.on_evict:
self.on_evict(idx)
self._refs[idx] = 1
self._free_mask &= ~(1 << idx)
return idx
return -1
def free(self, idx: int, keep_cached: bool = False):
with self._lock:
self._refs[idx] -= 1
if self._refs[idx] == 0:
if keep_cached:
self._lru[idx] = None
else:
self._free_mask |= 1 << idx
def inc_ref(self, idx: int):
with self._lock:
self._refs[idx] += 1
self._lru.pop(idx, None)
def ref_count(self, idx: int) -> int:
with self._lock:
return self._refs[idx]
def touch(self, idx: int):
with self._lock:
if idx in self._lru:
self._lru.move_to_end(idx)
class RadixNode:
"""A page-aligned edge in the CPU-side prefix radix."""
__slots__ = ("parent", "children", "page_idx", "tokens", "lock_ref")
def __init__(self, parent=None, tokens=(), page_idx=None):
self.parent = parent
self.children: Dict[tuple, "RadixNode"] = {}
self.page_idx = page_idx
self.tokens = tuple(tokens)
self.lock_ref = 0
class RadixCache:
"""Page-granular radix prefix index with exact token matching."""
def __init__(self, page_size: int):
self._page_size = page_size
self._root = RadixNode()
self._page_to_node: Dict[int, RadixNode] = {}
# Retained as an introspection-compatible map; matching never relies on
# this lossy value.
self._page_to_hash: Dict[int, int] = {}
self._lock = threading.Lock()
def evict(self, idx: int):
with self._lock:
node = self._page_to_node.pop(idx, None)
self._page_to_hash.pop(idx, None)
if node is None:
return
node.page_idx = None
parent = node.parent
if parent is not None:
parent.children.pop(node.tokens, None)
def has_page(self, idx: int) -> bool:
with self._lock:
return idx in self._page_to_node
def lookup(self, token_ids: List[int]) -> List[int]:
with self._lock:
full_pages = len(token_ids) // self._page_size
hits: List[int] = []
node = self._root
for i in range(full_pages):
start = i * self._page_size
page_tokens = tuple(token_ids[start : start + self._page_size])
child = node.children.get(page_tokens)
if child is None or child.page_idx is None:
break
hits.append(child.page_idx)
node = child
return hits
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
with self._lock:
full_pages = len(token_ids) // self._page_size
if logical_page_idx >= full_pages:
return
old = self._page_to_node.pop(page_idx, None)
self._page_to_hash.pop(page_idx, None)
if old is not None and old.parent is not None:
old.parent.children.pop(old.tokens, None)
node = self._root
for i in range(logical_page_idx + 1):
start = i * self._page_size
page_tokens = tuple(token_ids[start : start + self._page_size])
child = node.children.get(page_tokens)
if child is None:
child = RadixNode(node, page_tokens)
node.children[page_tokens] = child
node = child
if node.page_idx is not None and node.page_idx != page_idx:
replaced = node.page_idx
self._page_to_node.pop(replaced, None)
self._page_to_hash.pop(replaced, None)
node.page_idx = page_idx
self._page_to_node[page_idx] = node
self._page_to_hash[page_idx] = page_hash(
token_ids, logical_page_idx, self._page_size
)
def release(self, pages: List[int]) -> None:
with self._lock:
for page_idx in pages:
node = self._page_to_node.get(page_idx)
if node is not None and node.lock_ref:
node.lock_ref -= 1
class ReqToTokenPool:
"""Maps [req_idx, pos] -> physical token slot in KV storage.
Each row is one request; each column is a sequence position. The value
at [req_idx, pos] is the flat index into the KV storage buffers.
"""
def __init__(self, size: int, max_context_len: int, device: torch.device):
self.size = size
self.max_context_len = max_context_len
self.req_to_token = torch.zeros(
(size, max_context_len), dtype=torch.long, device=device
)
self.free_slots = list(range(size))
self._lock = threading.Lock()
def alloc(self, num_reqs: int) -> Optional[List[int]]:
with self._lock:
if num_reqs > len(self.free_slots):
return None
slots = self.free_slots[:num_reqs]
self.free_slots = self.free_slots[num_reqs:]
return slots
def free(self, req_indices: List[int]):
with self._lock:
self.free_slots.extend(req_indices)
def write(self, indices, values):
self.req_to_token[indices] = values
class KVStorage:
"""Token-level KV cache storage.
Buffers: [n_layers, size, n_kv_heads, head_dim]. Each token occupies
one slot indexed by ReqToTokenPool.
"""
def __init__(
self,
size: int,
n_layers: int,
n_kv_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.size = size
self.k_buffer = torch.empty(
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
)
self.v_buffer = torch.empty(
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
)
def get_key_buffer(self, layer_id: int) -> Tensor:
return self.k_buffer[layer_id]
def get_value_buffer(self, layer_id: int) -> Tensor:
return self.v_buffer[layer_id]
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
self.k_buffer[layer_id, loc] = k
self.v_buffer[layer_id, loc] = v
@dataclass
class KVCache:
"""Pure data struct passed to model for KV cache I/O.
The attention layer does raw buffer indexing no methods, no abstraction.
Attributes:
k_buffer: [n_layers, size, n_kv_heads, head_dim]
v_buffer: [n_layers, size, n_kv_heads, head_dim]
req_to_token: [num_reqs, max_ctx_len] index table
req_pool_indices: [batch_size] row indices into req_to_token
seq_lens: [batch_size] per-request total sequence lengths
out_cache_loc: [batch, new_seq_len] or [batch, 1] write indices
max_len: max(seq_lens) as Python int avoids GPU sync in decode
kv_indptr: [batch+1] int32 prefix sum of seq_lens, precomputed once
per step so the attention backend avoids rebuilding it per layer.
qo_indptr: [batch+1] int32 prefill qo prefix-sum (None in decode)
decode_o_part: split-KV o partial workspace (mirrors FlashInfer)
decode_ml_part: split-KV m/l partial workspace (mirrors FlashInfer)
decode_out: pre-allocated decode output buffer (graph-safe)
"""
k_buffer: Tensor
v_buffer: Tensor
req_to_token: Tensor
req_pool_indices: Tensor
seq_lens: Tensor
out_cache_loc: Tensor
max_len: int = 0
kv_indptr: Optional[Tensor] = None
qo_indptr: Optional[Tensor] = None
decode_o_part: Optional[Tensor] = None
decode_ml_part: Optional[Tensor] = None
decode_out: Optional[Tensor] = None
class PagePool:
"""Top-level KV cache manager.
Combines KVStorage + ReqToTokenPool + Allocator + RadixCache.
Args:
n_layers: Number of transformer layers.
n_kv_heads: Number of KV attention heads.
head_dim: Dimension per head.
max_batch_size: Maximum concurrent requests.
max_seq_len: Maximum sequence length per request.
device, dtype: Tensor device and dtype.
page_size: Page size for paged mode (1 = token-level).
n_tokens: Total token slots for paged mode. None = contiguous mode
(pre-allocates max_batch_size * max_seq_len).
"""
def __init__(
self,
n_layers: int,
n_kv_heads: int,
head_dim: int,
max_batch_size: int,
max_seq_len: int,
device: torch.device,
dtype: torch.dtype,
page_size: int = 1,
n_tokens: Optional[int] = None,
):
self.page_size = page_size
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.device = device
self.dtype = dtype
self.n_layers = n_layers
self.n_kv_heads = n_kv_heads
self.head_dim = head_dim
self.contiguous = n_tokens is None
if self.contiguous:
self.n_tokens = max_batch_size * max_seq_len
else:
self.n_tokens = n_tokens
self._storage = KVStorage(
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
)
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
if self.contiguous:
for i in range(max_batch_size):
self._req_pool.req_to_token[i] = torch.arange(
i * max_seq_len, (i + 1) * max_seq_len, device=device
)
self._alloc: Optional[Allocator] = None
self._prefix: Optional[RadixCache] = None
else:
n_pages = self.n_tokens // page_size
self._alloc = Allocator(n_pages)
self._prefix = RadixCache(page_size) if page_size > 1 else None
if self._prefix is not None:
self._alloc.on_evict = self._prefix.evict
self._task_req: Dict[str, int] = {}
self._task_len: Dict[int, int] = {}
self._task_cached: Dict[str, int] = {}
self._task_slots: Dict[str, List[int]] = {}
self._task_pages: Dict[str, List[int]] = {}
self._lock = threading.Lock()
# Steady-state decode validation state: the ordered task set and its
# Python seq_lens mirror. When the same set advances every sequence
# by exactly one token per step, bind_tasks updates the stable
# buffers in-place (+=1 / +=inc) instead of re-cumsumming. Any
# task-set change is a miss and rebuilds.
self._bind_sig: Optional[tuple] = None
self._bind_seq_lens: Optional[List[int]] = None
# ---- task lifecycle ----
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
req_slots = self._req_pool.alloc(1)
if req_slots is None:
return False
req_idx = req_slots[0]
self._task_req[task_id] = req_idx
if self.contiguous:
self._task_len[req_idx] = len(prompt_ids)
self._task_cached[task_id] = 0
return True
n_tokens_needed = len(prompt_ids)
cached = 0
if self._prefix is not None:
hits = self._prefix.lookup(prompt_ids)
cached = len(hits) * self.page_size
for p in hits:
self._alloc.inc_ref(p)
self._task_pages[task_id] = list(hits)
self._task_slots[task_id] = []
else:
self._task_pages[task_id] = []
self._task_slots[task_id] = []
remaining = n_tokens_needed - cached
if remaining > 0:
if self.page_size == 1:
slots = self._alloc_tokens(remaining)
if slots is None:
for p in self._task_pages[task_id]:
self._alloc.free(p)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
self._task_slots[task_id] = slots
else:
n_new_pages = (remaining + self.page_size - 1) // self.page_size
new_pages = []
for _ in range(n_new_pages):
p = self._alloc.alloc()
if p < 0:
for hp in self._task_pages[task_id]:
self._alloc.free(hp)
for np_ in new_pages:
self._alloc.free(np_)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
new_pages.append(p)
self._task_pages[task_id].extend(new_pages)
self._write_req_to_token(task_id, prompt_ids, cached)
self._task_len[req_idx] = len(prompt_ids)
self._task_cached[task_id] = cached
return True
def task_free(self, task_id: str):
req_idx = self._task_req.pop(task_id, None)
if req_idx is None:
return
self._task_len.pop(req_idx, None)
self._task_cached.pop(task_id, None)
if not self.contiguous:
if self._prefix is not None:
for p in self._task_pages.get(task_id, []):
keep = self._prefix.has_page(p)
self._alloc.free(p, keep_cached=keep)
if not keep:
self._prefix.evict(p)
else:
for p in self._task_pages.get(task_id, []):
self._alloc.free(p)
self._task_pages.pop(task_id, None)
self._task_slots.pop(task_id, None)
self._req_pool.free([req_idx])
def task_extend(self, task_id: str, pos: int) -> bool:
req_idx = self._task_req.get(task_id)
if req_idx is None or pos >= self.max_seq_len:
return False
# Paged mode must also claim a physical slot for the new token;
# contiguous mode's block is pre-allocated so this is a no-op.
if not self.contiguous and not self._extend_slot(task_id, req_idx, pos):
return False
self._task_len[req_idx] = pos + 1
return True
def _extend_slot(self, task_id: str, req_idx: int, pos: int) -> bool:
"""Allocate the physical slot for one extended token (paged mode)."""
if self.page_size == 1:
slots = self._alloc_tokens(1)
if slots is None:
return False
self._task_slots.setdefault(task_id, []).extend(slots)
self._req_pool.req_to_token[req_idx, pos] = slots[0]
return True
page_idx = pos // self.page_size
existing = self._task_pages.get(task_id, [])
if page_idx >= len(existing):
p = self._alloc.alloc()
if p < 0:
return False
existing.append(p)
self._task_pages[task_id] = existing
page_offset = pos % self.page_size
page = existing[page_idx]
token_slot = page * self.page_size + page_offset
self._req_pool.req_to_token[req_idx, pos] = token_slot
return True
def task_cached(self, task_id: str) -> int:
return self._task_cached.get(task_id, 0)
def task_record_hashes(
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
):
if self._prefix is None or self.contiguous:
return
pages = self._task_pages.get(task_id, [])
full_pages = len(prompt_ids) // self.page_size
for i in range(start_logical_page, min(full_pages, len(pages))):
self._prefix.record(pages[i], prompt_ids, i)
def task_cacheable_ids(
self, task_id: str, prompt_ids: List[int], output_ids: List[int]
):
"""Return the sequence whose KV entries are already materialized.
The first sampled output is produced by prompt prefill, and the last
sampled output has not been decoded into KV yet. Therefore the cache
can safely retain the prompt plus every output except the last one.
"""
return list(prompt_ids) + list(output_ids[:-1])
# ---- bind for forward ----
def bind_tasks(
self,
task_ids: List[str],
workspace: InferenceWorkspace,
device: Optional[torch.device] = None,
start_pos: Optional[int] = None,
) -> KVCache:
if device is None:
device = workspace.device
req_indices = [self._task_req[tid] for tid in task_ids]
# Per-request lengths come from the pool's own tracking (task_alloc
# sets len(prompt_ids); task_extend sets pos+1), so callers need not
# pass them.
seq_lens = [self._task_len[req_idx] for req_idx in req_indices]
b = len(task_ids)
sig = tuple(task_ids)
# Write into the caller's workspace buffers (fixed addresses, sized
# to max_batch/max_seq at init) — the sole owner of the per-step
# KV bind tensors.
rpi_buf = workspace.req_pool_indices
sl_buf = workspace.seq_lens
kvp_buf = workspace.kv_indptr
inc_buf = workspace.inc
ocl_buf = workspace.out_cache_loc
incremental = (
start_pos is None
and self._bind_sig is not None
and self._bind_sig == sig
and self._bind_seq_lens is not None
and len(self._bind_seq_lens) == b
and all(s == p + 1 for s, p in zip(seq_lens, self._bind_seq_lens))
)
if incremental:
# Steady-state decode: advance the stable buffers in-place.
# Normal-mode buffers keep ``+=`` legal regardless of whether
# this runs inside ``torch.inference_mode()``.
sl_buf[:b] += 1
kvp_buf[: b + 1] += inc_buf[: b + 1]
req_pool_indices = rpi_buf[:b]
seq_lens_t = sl_buf[:b]
kv_indptr = kvp_buf[: b + 1]
else:
# Cold path: fill the stable buffers from fresh host tensors.
rpi_buf[:b].copy_(
torch.tensor(req_indices, dtype=torch.long, device=device)
)
sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device))
kvp_buf[: b + 1].zero_()
kvp_buf[1 : b + 1] = sl_buf[:b].cumsum(0).to(torch.int32)
req_pool_indices = rpi_buf[:b]
seq_lens_t = sl_buf[:b]
kv_indptr = kvp_buf[: b + 1]
self._bind_sig = sig
self._bind_seq_lens = list(seq_lens)
if start_pos is not None:
seq_len = seq_lens[0]
out_cache_loc = self._req_pool.req_to_token[
req_pool_indices, start_pos:seq_len
]
# Ragged query segmentation for the prefill kernel, computed once
# (was rebuilt per layer in CudaBackend.fwd_prefill).
q_len = seq_len - start_pos
workspace.qo_indptr[: b + 1].copy_(
torch.arange(b + 1, dtype=torch.int32, device=device) * q_len
)
qo_indptr = workspace.qo_indptr[: b + 1]
decode_o_part, decode_ml_part = None, None
decode_out = None
else:
write_pos = seq_lens_t - 1
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
ocl_buf[:b].copy_(loc)
out_cache_loc = ocl_buf[:b]
qo_indptr = None
decode_o_part = getattr(workspace, "decode_o_part", None)
decode_ml_part = getattr(workspace, "decode_ml_part", None)
decode_out = getattr(workspace, "decode_out", None)
return KVCache(
k_buffer=self._storage.k_buffer,
v_buffer=self._storage.v_buffer,
req_to_token=self._req_pool.req_to_token,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens_t,
out_cache_loc=out_cache_loc,
max_len=max(seq_lens),
kv_indptr=kv_indptr,
qo_indptr=qo_indptr,
decode_o_part=decode_o_part,
decode_ml_part=decode_ml_part,
decode_out=decode_out,
)
# ---- internals ----
def _alloc_tokens(self, n: int) -> Optional[List[int]]:
if self.page_size != 1:
raise RuntimeError("_alloc_tokens is for page_size=1 only")
slots = []
for _ in range(n):
p = self._alloc.alloc()
if p < 0:
for s in slots:
self._alloc.free(s)
return None
slots.append(p)
return slots
def _write_req_to_token(self, task_id: str, prompt_ids: List[int], cached: int):
req_idx = self._task_req[task_id]
total = len(prompt_ids)
if self.contiguous:
return
if self.page_size == 1:
slots = self._task_slots.get(task_id, [])
all_slots = slots[: total - cached]
if all_slots:
self._req_pool.req_to_token[req_idx, cached:total] = torch.tensor(
all_slots, dtype=torch.long, device=self.device
)
else:
pages = self._task_pages.get(task_id, [])
for pos in range(cached, total):
page_idx = pos // self.page_size
page_offset = pos % self.page_size
if page_idx < len(pages):
token_slot = pages[page_idx] * self.page_size + page_offset
self._req_pool.req_to_token[req_idx, pos] = token_slot
+31 -21
View File
@@ -8,9 +8,10 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
import torch import torch
import torch.nn as nn import torch.nn as nn
from astrai.inference.core.cache import PagePool from astrai.extension import ATTN_BACKEND, AttentionBackend, get_backend
from astrai.inference.core.scheduler import InferenceScheduler from astrai.inference.cache import PagePool
from astrai.inference.core.task import STOP from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import STOP
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
@@ -74,6 +75,8 @@ class InferenceEngine:
max_batch_size: int = 1, max_batch_size: int = 1,
max_seq_len: Optional[int] = None, max_seq_len: Optional[int] = None,
cache: Optional[PagePool] = None, cache: Optional[PagePool] = None,
enable_cuda_graph: bool = True,
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
): ):
self.model = model self.model = model
self.tokenizer = tokenizer self.tokenizer = tokenizer
@@ -83,6 +86,8 @@ class InferenceEngine:
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
max_seq_len=max_seq_len, max_seq_len=max_seq_len,
cache=cache, cache=cache,
enable_cuda_graph=enable_cuda_graph,
backend=backend,
) )
self.scheduler.start() self.scheduler.start()
@@ -151,9 +156,8 @@ class InferenceEngine:
async def _agen(): async def _agen():
loop = asyncio.get_event_loop() loop = asyncio.get_event_loop()
while True: while True:
try: token = await loop.run_in_executor(None, next, sync_gen, None)
token = await loop.run_in_executor(None, next, sync_gen) if token is None:
except StopIteration:
break break
yield token yield token
@@ -172,6 +176,7 @@ class InferenceEngine:
rep_window: int, rep_window: int,
) -> Union[Generator, str, List[str]]: ) -> Union[Generator, str, List[str]]:
n = len(prompts) n = len(prompts)
request_backend = get_backend(use_default=False)
result = GenerateResult(count=n) result = GenerateResult(count=n)
task_ids = [ task_ids = [
self.scheduler.add_task( self.scheduler.add_task(
@@ -182,6 +187,7 @@ class InferenceEngine:
top_k=top_k, top_k=top_k,
frequency_penalty=frequency_penalty, frequency_penalty=frequency_penalty,
rep_window=rep_window, rep_window=rep_window,
backend=request_backend,
stream_callback=lambda token, idx=i: result.append(token, idx), stream_callback=lambda token, idx=i: result.append(token, idx),
) )
for i, p in enumerate(prompts) for i, p in enumerate(prompts)
@@ -204,27 +210,31 @@ class InferenceEngine:
def gen(): def gen():
nonlocal remaining nonlocal remaining
try: while remaining > 0:
while remaining > 0: items = result.pop_all()
items = result.pop_all() for idx, token in items:
for idx, token in items: if token is STOP:
if token is STOP: if not finished[idx]:
if not finished[idx]: finished[idx] = True
finished[idx] = True remaining -= 1
remaining -= 1 else:
else: yield (idx, token) if is_batch else token
yield (idx, token) if is_batch else token if remaining > 0:
if remaining > 0: result.wait(timeout=0.05)
result.wait(timeout=0.05)
finally:
for tid in task_ids:
self.scheduler.remove_task(tid)
return gen() return gen()
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
return self.scheduler.get_stats() return self.scheduler.get_stats()
@property
def backend_name(self) -> str:
return self.scheduler.backend_name
@property
def cuda_graph_enabled(self) -> bool:
return self.scheduler.cuda_graph_enabled
def shutdown(self): def shutdown(self):
self.scheduler.stop() self.scheduler.stop()
if torch.cuda.is_available(): if torch.cuda.is_available():
+201
View File
@@ -0,0 +1,201 @@
"""Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting."""
import time
from collections import deque
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Any, Deque, Dict, Generator, List, Literal, Optional
@dataclass
class TaskTiming:
"""Timestamp snapshots and computed metrics for one generation task.
Created by :class:`MetricsCollector` at task-registration time;
updated via ``record`` / ``mark_finished``.
"""
task_id: str
arrival_time: float
prefill_start_time: Optional[float] = None
first_token_time: Optional[float] = None
finish_time: Optional[float] = None
input_tokens: int = 0
output_tokens: int = 0
_decode_steps: int = 0
_decode_total_s: float = 0.0
# derived metrics
@property
def queue_wait_ms(self) -> Optional[float]:
if self.prefill_start_time is not None:
return (self.prefill_start_time - self.arrival_time) * 1000
return None
@property
def ttft_ms(self) -> Optional[float]:
if self.first_token_time is not None:
return (self.first_token_time - self.arrival_time) * 1000
return None
@property
def prefill_tps(self) -> Optional[float]:
if self.prefill_start_time is not None and self.first_token_time is not None:
d = self.first_token_time - self.prefill_start_time
if d > 0 and self.input_tokens > 0:
return self.input_tokens / d
return None
@property
def decode_tps(self) -> Optional[float]:
if self.first_token_time is not None and self.finish_time is not None:
d = self.finish_time - self.first_token_time
dt = self.output_tokens - 1
if dt > 0 and d > 0:
return dt / d
return None
@property
def decode_avg_ms(self) -> Optional[float]:
if self._decode_steps > 0 and self._decode_total_s > 0:
return (self._decode_total_s / self._decode_steps) * 1000
return None
@property
def e2e_latency_ms(self) -> Optional[float]:
if self.finish_time is not None:
return (self.finish_time - self.arrival_time) * 1000
return None
@property
def total_tps(self) -> Optional[float]:
if self.finish_time is not None:
total = self.input_tokens + self.output_tokens
d = self.finish_time - self.arrival_time
if total > 0 and d > 0:
return total / d
return None
def to_dict(self) -> Dict[str, Any]:
return {
"task_id": self.task_id,
"input_tokens": self.input_tokens,
"output_tokens": self.output_tokens,
"queue_wait_ms": (
round(self.queue_wait_ms, 2) if self.queue_wait_ms is not None else None
),
"ttft_ms": (round(self.ttft_ms, 2) if self.ttft_ms is not None else None),
"prefill_tps": (
round(self.prefill_tps, 2) if self.prefill_tps is not None else None
),
"decode_tps": (
round(self.decode_tps, 2) if self.decode_tps is not None else None
),
"decode_avg_ms": (
round(self.decode_avg_ms, 2) if self.decode_avg_ms is not None else None
),
"total_tps": (
round(self.total_tps, 2) if self.total_tps is not None else None
),
"e2e_latency_ms": (
round(self.e2e_latency_ms, 2)
if self.e2e_latency_ms is not None
else None
),
}
class MetricsCollector:
"""Single-owner perf/stats hub for all generation tasks.
Usage::
metrics = MetricsCollector()
metrics.register(task_id, arrival_time)
with metrics.record(task_ids, "prefill"):
run_prefill(...)
metrics.mark_finished(task_id, input_tokens, output_tokens)
stats = metrics.get_stats()
"""
def __init__(self, max_recent: int = 128):
self._timings: Dict[str, TaskTiming] = {}
self._completed: Deque[TaskTiming] = deque(maxlen=max_recent)
self._ttft_ms_sum = 0.0
self._ttft_ms_count = 0
self._decode_tps_sum = 0.0
self._decode_tps_count = 0
self._e2e_ms_sum = 0.0
self._e2e_ms_count = 0
def register(self, task_id: str):
"""Create a timing record for a newly-created task."""
self._timings[task_id] = TaskTiming(task_id=task_id, arrival_time=time.time())
def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int):
"""Close timing for a finished/aborted task and move it to completed."""
timing = self._timings.pop(task_id, None)
if timing is None:
return
timing.finish_time = time.time()
timing.input_tokens = input_tokens
timing.output_tokens = output_tokens
self._completed.append(timing)
self._accumulate(timing)
# timing scopes
@contextmanager
def record(
self, task_ids: List[str], phase: Literal["prefill", "decode"]
) -> Generator[None, None, None]:
tic = time.time()
yield
toc = time.time()
dt = toc - tic
for tid in task_ids:
t = self._timings.get(tid)
if t is None:
continue
if phase == "prefill":
t.prefill_start_time = tic
t.first_token_time = toc
elif phase == "decode":
t._decode_steps += 1
t._decode_total_s += dt
# aggregate stats
def get_stats(self) -> Dict[str, Any]:
stats: Dict[str, Any] = {}
if self._ttft_ms_count > 0:
stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2)
if self._decode_tps_count > 0:
stats["avg_decode_tps"] = round(
self._decode_tps_sum / self._decode_tps_count, 2
)
if self._e2e_ms_count > 0:
stats["avg_e2e_latency_ms"] = round(
self._e2e_ms_sum / self._e2e_ms_count, 2
)
if self._completed:
stats["recent_tasks"] = [t.to_dict() for t in self._completed]
return stats
# internal
def _accumulate(self, t: TaskTiming):
if t.ttft_ms is not None:
self._ttft_ms_sum += t.ttft_ms
self._ttft_ms_count += 1
if t.decode_tps is not None:
self._decode_tps_sum += t.decode_tps
self._decode_tps_count += 1
if t.e2e_latency_ms is not None:
self._e2e_ms_sum += t.e2e_latency_ms
self._e2e_ms_count += 1
@@ -4,8 +4,7 @@
lazy singleton FastAPI instance. lazy singleton FastAPI instance.
""" """
from astrai.inference.api.protocol import GenContext, ProtocolHandler, StopChecker from astrai.inference.network.app import (
from astrai.inference.api.server import (
AnthropicMessage, AnthropicMessage,
ChatCompletionRequest, ChatCompletionRequest,
ChatMessage, ChatMessage,
@@ -15,7 +14,8 @@ from astrai.inference.api.server import (
get_app, get_app,
run_server, run_server,
) )
from astrai.inference.api.tool_parser import ( from astrai.inference.network.protocol import GenContext, ProtocolHandler, StopChecker
from astrai.inference.network.tool_parser import (
BaseToolParser, BaseToolParser,
SimpleJsonToolParser, SimpleJsonToolParser,
ToolParserFactory, ToolParserFactory,
@@ -6,13 +6,13 @@ from typing import Any, Dict, List, Tuple, Union
from pydantic import BaseModel from pydantic import BaseModel
from astrai.inference.api.protocol import ( from astrai.inference.engine import InferenceEngine
from astrai.inference.network.protocol import (
GenContext, GenContext,
ResponseBuilder, ResponseBuilder,
StopInfo, StopInfo,
sse_event, sse_event,
) )
from astrai.inference.engine import InferenceEngine
def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str: def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str:
@@ -18,10 +18,10 @@ import uvicorn
from fastapi import APIRouter, FastAPI, HTTPException from fastapi import APIRouter, FastAPI, HTTPException
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from astrai.inference.api.anthropic import AnthropicResponseBuilder
from astrai.inference.api.openai import OpenAIResponseBuilder
from astrai.inference.api.protocol import ProtocolHandler
from astrai.inference.engine import InferenceEngine from astrai.inference.engine import InferenceEngine
from astrai.inference.network.anthropic import AnthropicResponseBuilder
from astrai.inference.network.openai import OpenAIResponseBuilder
from astrai.inference.network.protocol import ProtocolHandler
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
@@ -7,14 +7,14 @@ from typing import Any, Dict, List, Optional, Tuple, Union
from pydantic import BaseModel from pydantic import BaseModel
from astrai.inference.api.protocol import ( from astrai.inference.engine import InferenceEngine
from astrai.inference.network.protocol import (
GenContext, GenContext,
ResponseBuilder, ResponseBuilder,
StopInfo, StopInfo,
sse_event, sse_event,
) )
from astrai.inference.api.tool_parser import BaseToolParser, ToolParserFactory from astrai.inference.network.tool_parser import BaseToolParser, ToolParserFactory
from astrai.inference.engine import InferenceEngine
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -181,12 +181,10 @@ class ProtocolHandler:
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str] self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
) -> Dict[str, Any]: ) -> Dict[str, Any]:
checker = StopChecker(stop_sequences) checker = StopChecker(stop_sequences)
chunks: List[str] = []
body = "" body = ""
matched = None matched = None
async for token in agen: async for token in agen:
chunks.append(token)
body += token body += token
matched = checker.check(body) matched = checker.check(body)
@@ -195,6 +193,5 @@ class ProtocolHandler:
ctx.completion_tokens += 1 ctx.completion_tokens += 1
content = "".join(chunks)
stop = StopInfo(matched=matched, body=body) stop = StopInfo(matched=matched, body=body)
return self.builder.format_response(ctx, content, stop) return self.builder.format_response(ctx, body, stop)
+25
View File
@@ -0,0 +1,25 @@
"""Execution primitives: forward passes, CUDA graphs, and sampling."""
from astrai.inference.runtime.executor import Executor
from astrai.inference.runtime.graph import CudaGraphContext
from astrai.inference.runtime.sample import (
BaseSamplingStrategy,
FrequencyPenaltyStrategy,
SamplingPipeline,
TemperatureStrategy,
TopKStrategy,
TopPStrategy,
sample,
)
__all__ = [
"Executor",
"CudaGraphContext",
"BaseSamplingStrategy",
"FrequencyPenaltyStrategy",
"SamplingPipeline",
"TemperatureStrategy",
"TopKStrategy",
"TopPStrategy",
"sample",
]
@@ -1,5 +1,4 @@
import logging import logging
import os
import time import time
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass from dataclasses import dataclass
@@ -8,34 +7,42 @@ from typing import List, Optional
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.extension.attention_backend import ( from astrai.extension.backend.attention import (
ATTN_BACKEND,
CudaBackend, CudaBackend,
attn_backend,
get_backend, get_backend,
) )
from astrai.inference.core.cache import PagePool from astrai.inference.cache import PagePool, TaskCacheManager
from astrai.inference.core.graph import CudaGraphContext from astrai.inference.runtime.graph import CudaGraphContext
from astrai.inference.core.task import Task from astrai.inference.runtime.sample import sample
from astrai.inference.core.workspace import InferenceWorkspace from astrai.inference.task import Task
from astrai.inference.sample import sample from astrai.inference.workspace import InferenceWorkspace
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_TIMED = os.environ.get("ASTRAI_TIMED", "") == "1"
@contextmanager @contextmanager
def timed(label: str, log: Optional[logging.Logger] = None): def timed(label: str, log: Optional[logging.Logger] = None):
"""Wall-clock debug timer, enabled via ``ASTRAI_TIMED=1``.""" """GPU-precise timer via CUDA events; falls back to perf_counter on CPU."""
if not _TIMED: log = log or logger
if not log.isEnabledFor(logging.DEBUG):
yield yield
return return
tic = time.perf_counter() use_cuda = torch.cuda.is_available()
if use_cuda:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
else:
tic = time.perf_counter()
yield yield
elapsed_ms = (time.perf_counter() - tic) * 1000 if use_cuda:
(log or logger).info("%s %.1fms", label, elapsed_ms) end.record()
torch.cuda.synchronize()
elapsed_ms = start.elapsed_time(end)
else:
elapsed_ms = (time.perf_counter() - tic) * 1000
log.debug("%s %.2fms", label, elapsed_ms)
@dataclass @dataclass
@@ -54,6 +61,23 @@ class SamplingBatchInfo:
has_freq: bool # any frequency_penalty != 0 (avoids per-step GPU .any()) has_freq: bool # any frequency_penalty != 0 (avoids per-step GPU .any())
@dataclass
class DecodeSteadyState:
"""Cached decode metadata for the steady-state case.
When the same ordered task set decodes one token per step, sampling
params and task signature are reused; only positions advance by 1.
``last_tokens`` keeps that step's sampled ids on-device so the next
step with an unchanged signature can fill ``input_ids`` via a
device-to-device copy.
"""
task_sig: tuple
positions: list[int]
sampling_info: SamplingBatchInfo
last_tokens: Optional[Tensor] = None
def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo: def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
pin = str(device).startswith("cuda") pin = str(device).startswith("cuda")
freq_penalties = torch.tensor( freq_penalties = torch.tensor(
@@ -77,12 +101,37 @@ def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
def _warmup_cuda_graphs( def _warmup_cuda_graphs(
model: AutoModel, model: AutoModel,
pool: PagePool, pool: PagePool,
task_cache: TaskCacheManager,
ws: InferenceWorkspace, ws: InferenceWorkspace,
gctx: "CudaGraphContext", gctx: CudaGraphContext,
max_batch_size: int, max_batch_size: int,
prompt_len: int = 1, prompt_len: int = 1,
device: Optional[str] = None, device: Optional[str] = None,
): ):
dev = device or next(model.parameters()).device
# Prefill warmup: cuBLAS auto-tunes for the actual prompt-length tensor
# shapes on first call (F.linear is the dominant cost). This also warms
# up the CUDA context (driver init) and compiles the graph-capture trace
# that follows. Custom .so kernels do NOT need this — they are pre-built.
warmup_len = 64
tid = "_warmup_prefill"
if task_cache.task_alloc(tid, list(range(warmup_len))):
with (
torch.inference_mode(),
timed("warmup prefill", logger),
):
kv = task_cache.bind([tid], ws, start_pos=0)
ids_in = torch.arange(warmup_len, device=dev)
pos_in = ids_in
model(
ids_in,
kv_cache=kv,
position_ids=pos_in,
fwd="prefill",
)
task_cache.task_free(tid)
batch_sizes = [1] batch_sizes = [1]
n = 2 n = 2
while n <= max_batch_size: while n <= max_batch_size:
@@ -91,60 +140,41 @@ def _warmup_cuda_graphs(
if max_batch_size not in batch_sizes: if max_batch_size not in batch_sizes:
batch_sizes.append(max_batch_size) batch_sizes.append(max_batch_size)
dev = device or next(model.parameters()).device
for b in batch_sizes: for b in batch_sizes:
task_ids = [f"_gr_{b}_{i}" for i in range(b)] task_ids = [f"_warmup_decode_{b}_{i}" for i in range(b)]
prompt_tokens = [list(range(prompt_len)) for _ in range(b)] prompt_tokens = [list(range(prompt_len)) for _ in range(b)]
alloc_ok = True alloc_ok = True
for tid, pt in zip(task_ids, prompt_tokens): for tid, pt in zip(task_ids, prompt_tokens):
if not pool.task_alloc(tid, pt): if not task_cache.task_alloc(tid, pt):
alloc_ok = False alloc_ok = False
break break
if not alloc_ok: if not alloc_ok:
for tid in task_ids: for tid in task_ids:
pool.task_free(tid) task_cache.task_free(tid)
continue continue
with ( with (
torch.inference_mode(), torch.inference_mode(),
attn_backend(ATTN_BACKEND.CUDA),
timed(f"warmup prefill b={b}", logger),
):
kv_cache = pool.bind_tasks(task_ids, ws, start_pos=0)
ids_in = torch.tensor(prompt_tokens, dtype=torch.long, device=dev)
pos_in = torch.arange(prompt_len, device=dev).unsqueeze(0).expand(b, -1)
model(
ids_in,
input_mask=pos_in.unsqueeze(-1) >= torch.arange(prompt_len, device=dev),
kv_cache=kv_cache,
position_ids=pos_in,
)
with (
torch.inference_mode(),
attn_backend(ATTN_BACKEND.CUDA),
timed(f"warmup decode b={b}", logger), timed(f"warmup decode b={b}", logger),
): ):
for step in range(2): for step in range(2):
seq_pos = prompt_len + step seq_pos = step
ws.position_ids[:b] = seq_pos ws.position_ids[:b] = seq_pos
for tid in task_ids: for tid in task_ids:
pool.task_extend(tid, seq_pos) task_cache.task_extend(tid, seq_pos)
kv = pool.bind_tasks(task_ids, ws) kv = task_cache.bind(task_ids, ws)
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
ids_buf = ws.fill_input_ids([step] * b) ids_buf = ws.fill_input_ids([step] * b)
gctx.forward( gctx.forward(
model, model,
key=(b,), key=(b,),
input_ids=ids_buf.unsqueeze(1), input_ids=ids_buf,
input_mask=input_mask,
kv_cache=kv, kv_cache=kv,
position_ids=ws.position_ids[:b].unsqueeze(1), position_ids=ws.position_ids[:b],
fwd="decode",
) )
for tid in task_ids: for tid in task_ids:
pool.task_free(tid) task_cache.task_free(tid)
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -154,22 +184,22 @@ class Executor:
def __init__( def __init__(
self, self,
model: AutoModel, model: AutoModel,
tokenizer: AutoTokenizer,
kv_cache: PagePool, kv_cache: PagePool,
task_cache: TaskCacheManager,
device: Optional[str] = None, device: Optional[str] = None,
dtype: Optional[torch.dtype] = None, dtype: Optional[torch.dtype] = None,
enable_cuda_graph: bool = True,
): ):
self.model = model self.model = model
self.tokenizer = tokenizer
self.kv_cache = kv_cache self.kv_cache = kv_cache
self.task_cache = task_cache
self.device = device or next(model.parameters()).device self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype self.dtype = dtype or next(model.parameters()).dtype
# Per-step decode cache for the steady-state case where the same # Per-step decode cache for the steady-state case (same ordered
# ordered task set decodes one token per step. Sampling params are # task set decodes one token per step). Sampling params stay
# constant across steps; position_ids grows by exactly 1. Single-slot: # constant; only positions advance.
# any task-set change is a cache miss. self._decode_cache: Optional[DecodeSteadyState] = None
self._decode_cache: Optional[tuple] = None
# Pre-allocated fixed-shape buffers for the decode hot path # Pre-allocated fixed-shape buffers for the decode hot path
# (input_ids, decode mask, KV bind metadata). Eagerly sized at init # (input_ids, decode mask, KV bind metadata). Eagerly sized at init
@@ -178,8 +208,10 @@ class Executor:
config = model.config config = model.config
max_q_heads = config.num_attention_heads max_q_heads = config.num_attention_heads
head_dim = config.hidden_size // config.num_attention_heads head_dim = config.hidden_size // config.num_attention_heads
self._head_dim = head_dim backend = get_backend()
self._graph_supported = CudaBackend.supports(head_dim=head_dim) self._graph_supported = backend.supports_graph() and (
CudaBackend.available() and head_dim in CudaBackend.HEAD_DIMS
)
self._workspace = InferenceWorkspace( self._workspace = InferenceWorkspace(
max_batch_size=kv_cache.max_batch_size, max_batch_size=kv_cache.max_batch_size,
max_seq_len=kv_cache.max_seq_len, max_seq_len=kv_cache.max_seq_len,
@@ -193,7 +225,8 @@ class Executor:
# Enabled at init-time via _warmup_cuda_graphs for CudaBackend # Enabled at init-time via _warmup_cuda_graphs for CudaBackend
# on supported head_dims; left disabled otherwise. # on supported head_dims; left disabled otherwise.
self._graph_ctx = CudaGraphContext() self._graph_ctx = CudaGraphContext()
self._try_enable_cuda_graph() if enable_cuda_graph:
self._try_enable_cuda_graph()
def _try_enable_cuda_graph(self): def _try_enable_cuda_graph(self):
if not self._graph_supported: if not self._graph_supported:
@@ -203,12 +236,17 @@ class Executor:
_warmup_cuda_graphs( _warmup_cuda_graphs(
self.model, self.model,
self.kv_cache, self.kv_cache,
self.task_cache,
self._workspace, self._workspace,
self._graph_ctx, self._graph_ctx,
max_batch_size=self.kv_cache.max_batch_size, max_batch_size=self.kv_cache.max_batch_size,
device=self.device, device=self.device,
) )
@property
def cuda_graph_enabled(self) -> bool:
return self._graph_ctx.enabled and self._graph_supported
def _sample_logits( def _sample_logits(
self, self,
logits: Tensor, logits: Tensor,
@@ -216,6 +254,13 @@ class Executor:
return_logprobs: bool = False, return_logprobs: bool = False,
info: Optional[SamplingBatchInfo] = None, info: Optional[SamplingBatchInfo] = None,
): ):
"""Sample from ``logits`` and return ``(host_payload, tokens)``.
``host_payload`` is the scheduler-facing list (token ids, or
``(token_id, logprob)`` tuples with ``return_logprobs``);
``tokens`` is the ``[B]`` device tensor that produced it, kept
for the steady-state decode fast path.
"""
info = info or _build_sampling_batch_info(tasks, self.device) info = info or _build_sampling_batch_info(tasks, self.device)
if info.has_freq: if info.has_freq:
history_lists = [ history_lists = [
@@ -250,14 +295,14 @@ class Executor:
return_logprobs=return_logprobs, return_logprobs=return_logprobs,
) )
if not return_logprobs: if not return_logprobs:
return result.tolist() return result.tolist(), result
tokens, logprobs = result tokens, logprobs = result
tokens_list = tokens.tolist() tokens_list = tokens.tolist()
logprobs_list = logprobs.tolist() logprobs_list = logprobs.tolist()
for task, logprob in zip(tasks, logprobs_list): for task, logprob in zip(tasks, logprobs_list):
task.output_logprobs.append(float(logprob)) task.output_logprobs.append(float(logprob))
return list(zip(tokens_list, logprobs_list)) return list(zip(tokens_list, logprobs_list)), tokens
def execute_prefill( def execute_prefill(
self, self,
@@ -273,20 +318,15 @@ class Executor:
batch_sz = len(tasks) batch_sz = len(tasks)
input_ids = torch.tensor( input_ids = torch.tensor(
[t.prompt_ids[start_pos:prompt_len] for t in tasks], [token for t in tasks for token in t.prompt_ids[start_pos:prompt_len]],
dtype=torch.long, dtype=torch.long,
device=self.device, device=self.device,
) )
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
position_ids = ( position_ids = torch.arange(
torch.arange(start_pos, prompt_len, dtype=torch.long, device=self.device) start_pos, prompt_len, dtype=torch.long, device=self.device
.unsqueeze(0) ).repeat(batch_sz)
.expand(batch_sz, -1)
)
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_len, device=self.device
)
with ( with (
torch.inference_mode(), torch.inference_mode(),
@@ -294,17 +334,21 @@ class Executor:
): ):
outputs = self.model( outputs = self.model(
input_ids, input_ids,
input_mask=input_mask,
position_ids=position_ids, position_ids=position_ids,
kv_cache=self.kv_cache.bind_tasks( kv_cache=self.task_cache.bind(
task_ids, task_ids,
self._workspace, self._workspace,
start_pos=start_pos, start_pos=start_pos,
), ),
fwd="prefill",
) )
logits = outputs["logits"][:, -1, :] q_len = prompt_len - start_pos
logits = outputs["logits"][
torch.arange(1, batch_sz + 1, device=self.device) * q_len - 1
]
return tasks, self._sample_logits(logits, tasks, return_logprobs) step_out, _ = self._sample_logits(logits, tasks, return_logprobs)
return tasks, step_out
def execute_decode( def execute_decode(
self, tasks: List[Task], return_logprobs: bool = False self, tasks: List[Task], return_logprobs: bool = False
@@ -328,37 +372,39 @@ class Executor:
b = len(tasks) b = len(tasks)
ws = self._workspace ws = self._workspace
task_ids = [t.task_id for t in tasks]
cur_positions = [t.next_pos for t in tasks]
task_sig = tuple(task_ids)
# ---- pre-replay: update input buffers in-place ---- # ---- pre-replay: update input buffers in-place ----
input_ids = ws.fill_input_ids( # When the previous decode step ran this same ordered task set, its
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks] # sampled tokens are still on-device and map 1:1 onto the current
) # slots — fill input ids device-to-device. inference_mode guards
# the read because the source was produced under sampling's
task_ids = [t.task_id for t in tasks] # inference-mode context.
cur_positions = [t.next_pos for t in tasks]
sig = tuple(task_ids)
cached = self._decode_cache cached = self._decode_cache
if ( sig_match = cached is not None and cached.task_sig == task_sig
cached is not None if sig_match and cached.last_tokens is not None:
and cached[0] == sig with torch.inference_mode():
and cur_positions == [p + 1 for p in cached[1]] input_ids = ws.fill_input_ids_from_device(cached.last_tokens)
): else:
info = cached[2] input_ids = ws.fill_input_ids(
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
)
kv_cache = self.task_cache.bind(task_ids, ws)
reuse_decode_state = self.task_cache.bind_was_steady and sig_match
if reuse_decode_state:
info = self._decode_cache.sampling_info
ws.position_ids[:b] += 1 ws.position_ids[:b] += 1
self._decode_cache = (sig, cur_positions, info)
else: else:
info = _build_sampling_batch_info(tasks, self.device) info = _build_sampling_batch_info(tasks, self.device)
ws.position_ids[:b].copy_( ws.position_ids[:b].copy_(
torch.tensor(cur_positions, dtype=torch.long, device=self.device) torch.tensor(cur_positions, dtype=torch.long, device=self.device)
) )
self._decode_cache = (sig, cur_positions, info) self._decode_cache = DecodeSteadyState(task_sig, cur_positions, info)
total_len = max(cur_positions) + 1
input_mask = ws.decode_mask(ws.position_ids[:b], total_len)
kv_cache = self.kv_cache.bind_tasks(task_ids, ws)
# ---- forward (graph replay or live run + capture) ---- # ---- forward (graph replay or live run + capture) ----
@@ -368,9 +414,6 @@ class Executor:
and get_backend().supports_graph() and get_backend().supports_graph()
) )
key = (b,) key = (b,)
if use_graph:
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
with ( with (
torch.inference_mode(), torch.inference_mode(),
timed(f"execute_decode forward b={b}", logger), timed(f"execute_decode forward b={b}", logger),
@@ -379,18 +422,22 @@ class Executor:
outputs = self._graph_ctx.forward( outputs = self._graph_ctx.forward(
self.model, self.model,
key=key, key=key,
input_ids=input_ids.unsqueeze(1), input_ids=input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=ws.position_ids[:b].unsqueeze(1), position_ids=ws.position_ids[:b],
fwd="decode",
) )
else: else:
outputs = self.model( outputs = self.model(
input_ids.unsqueeze(1), input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=ws.position_ids[:b].unsqueeze(1), position_ids=ws.position_ids[:b],
fwd="decode",
) )
logits = outputs["logits"][:, -1, :] logits = outputs["logits"]
return self._sample_logits(logits, tasks, return_logprobs, info=info) step_out, tokens_dev = self._sample_logits(
logits, tasks, return_logprobs, info=info
)
self._decode_cache.last_tokens = tokens_dev
return step_out
@@ -363,20 +363,6 @@ def sample(
``True`` a ``(token_ids, chosen_logprobs)`` tuple where ``True`` a ``(token_ids, chosen_logprobs)`` tuple where
``chosen_logprobs`` has shape ``[batch]``. ``chosen_logprobs`` has shape ``[batch]``.
""" """
greedy = (
bool((temperature == 0).all())
if isinstance(temperature, Tensor)
else temperature == 0
)
if greedy:
tokens = logits.argmax(dim=-1)
if not return_logprobs:
return tokens
log_probs = torch.log_softmax(logits.float(), dim=-1)
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
has_freq = ( has_freq = (
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any()) (isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
if isinstance(frequency_penalty, Tensor) if isinstance(frequency_penalty, Tensor)
@@ -1,13 +1,21 @@
import logging import logging
import threading import threading
import uuid import uuid
from typing import Any, Dict, List, Optional, Tuple from contextlib import nullcontext
from typing import Any, Dict, List, Optional, Tuple, Union
import torch import torch
from astrai.inference.core.cache import PagePool from astrai.extension import (
from astrai.inference.core.executor import Executor ATTN_BACKEND,
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus AttentionBackend,
attn_backend,
get_backend,
)
from astrai.inference.cache import PagePool, TaskCacheManager
from astrai.inference.metrics import MetricsCollector
from astrai.inference.runtime.executor import Executor
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer
@@ -26,6 +34,8 @@ class InferenceScheduler:
device: Optional[str] = None, device: Optional[str] = None,
dtype: Optional[torch.dtype] = None, dtype: Optional[torch.dtype] = None,
cache: Optional[PagePool] = None, cache: Optional[PagePool] = None,
enable_cuda_graph: bool = True,
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
): ):
config = model.config config = model.config
@@ -56,19 +66,34 @@ class InferenceScheduler:
dtype=self.dtype, dtype=self.dtype,
) )
self._metrics = MetricsCollector()
self._task_cache = TaskCacheManager(self._cache)
self._task_mgr = TaskManager( self._task_mgr = TaskManager(
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
max_seq_len=self.max_seq_len, max_seq_len=self.max_seq_len,
metrics=self._metrics,
) )
self._executor = Executor( if backend is None:
model=model, self._backend = None
tokenizer=tokenizer, active_backend = get_backend()
kv_cache=self._cache, else:
device=self.device, active_backend = backend
dtype=self.dtype, with attn_backend(active_backend):
) if backend is not None:
self._backend = get_backend()
self._backend_name = type(get_backend()).__name__
self._executor = Executor(
model=model,
kv_cache=self._cache,
task_cache=self._task_cache,
device=self.device,
dtype=self.dtype,
enable_cuda_graph=enable_cuda_graph,
)
self._stop_event = threading.Event() self._stop_event = threading.Event()
self._loop_thread: Optional[threading.Thread] = None self._loop_thread: Optional[threading.Thread] = None
@@ -78,11 +103,31 @@ class InferenceScheduler:
def remove_task(self, task_id: str): def remove_task(self, task_id: str):
for task in self._task_mgr.remove_task(task_id): for task in self._task_mgr.remove_task(task_id):
self._cache.task_free(task.task_id) self._task_cache.task_free(task.task_id)
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
return self._task_mgr.get_stats() return self._task_mgr.get_stats()
@property
def backend_name(self) -> str:
return self._backend_name
@property
def cuda_graph_enabled(self) -> bool:
return self._executor.cuda_graph_enabled
def _backend_context(self):
if self._backend is None:
return nullcontext()
return attn_backend(self._backend)
@staticmethod
def _task_backend_groups(tasks: List[Task]):
groups = {}
for task in tasks:
groups.setdefault(task.backend, (task.backend, []))[1].append(task)
return groups.values()
def _step( def _step(
self, tasks: List[Task], return_logprobs: bool = False self, tasks: List[Task], return_logprobs: bool = False
) -> Tuple[List[Task], List[Task]]: ) -> Tuple[List[Task], List[Task]]:
@@ -106,32 +151,45 @@ class InferenceScheduler:
already appended to ``output_ids``) and tasks that hit the already appended to ``output_ids``) and tasks that hit the
sequence cap and were marked ``ABORTED``. sequence cap and were marked ``ABORTED``.
""" """
cache = self._cache to_prefill = [t for t in tasks if not t.prefill_done and t.prompt_ids]
to_prefill = [t for t in tasks if t.output_tokens == 0 and t.prompt_ids]
prefilled_ids = set() prefilled_ids = set()
produced: List[Task] = [] produced: List[Task] = []
if to_prefill: if to_prefill:
for t in to_prefill: for t in to_prefill:
t.input_tokens = len(t.prompt_ids) t.input_tokens = len(t.prompt_ids)
groups: Dict[Tuple[int, int], List[Task]] = {} groups: Dict[Tuple[int, int, Optional[AttentionBackend]], List[Task]] = {}
for t in to_prefill: for t in to_prefill:
start_pos = min(cache.task_cached(t.task_id), len(t.prompt_ids) - 1) start_pos = min(
groups.setdefault((len(t.prompt_ids), start_pos), []).append(t) self._task_cache.task_cached(t.task_id), len(t.prompt_ids) - 1
for (prompt_len, start_pos), group in groups.items():
prefilled, step_out = self._executor.execute_prefill(
group, prompt_len, start_pos, return_logprobs=return_logprobs
) )
groups.setdefault((len(t.prompt_ids), start_pos, t.backend), []).append(
t
)
for (prompt_len, start_pos, _), group in groups.items():
backend = group[0].backend
backend_context = (
attn_backend(backend) if backend is not None else nullcontext()
)
with (
backend_context,
self._metrics.record([t.task_id for t in group], "prefill"),
):
prefilled, step_out = self._executor.execute_prefill(
group, prompt_len, start_pos, return_logprobs=return_logprobs
)
for t, out in zip(prefilled, step_out): for t, out in zip(prefilled, step_out):
t.output_ids.append(out[0] if return_logprobs else out) t.output_ids.append(out[0] if return_logprobs else out)
t.output_tokens += 1 t.output_tokens += 1
t.mark_prefill_done()
prefilled_ids.add(t.task_id) prefilled_ids.add(t.task_id)
produced.append(t) produced.append(t)
start_logical_page = start_pos // getattr(cache, "page_size", 64)
start_logical_page = start_pos // self._cache.page_size
for t in group: for t in group:
cache.task_record_hashes( self._task_cache.task_record_hashes(
t.task_id, t.prompt_ids, start_logical_page t.task_id, t.prompt_ids, start_logical_page
) )
@@ -140,82 +198,84 @@ class InferenceScheduler:
for t in tasks: for t in tasks:
if t.task_id in prefilled_ids: if t.task_id in prefilled_ids:
continue continue
if cache.task_extend(t.task_id, t.next_pos): if self._task_cache.task_extend(t.task_id, t.next_pos):
decoded.append(t) decoded.append(t)
else: else:
t.status = TaskStatus.ABORTED t.status = TaskStatus.ABORTED
aborted.append(t) aborted.append(t)
if decoded: for backend, group in self._task_backend_groups(decoded):
step_out = self._executor.execute_decode( backend_context = (
decoded, return_logprobs=return_logprobs attn_backend(backend) if backend is not None else nullcontext()
) )
for t, out in zip(decoded, step_out): with (
backend_context,
self._metrics.record([t.task_id for t in group], "decode"),
):
step_out = self._executor.execute_decode(
group, return_logprobs=return_logprobs
)
for t, out in zip(group, step_out):
t.output_ids.append(out[0] if return_logprobs else out) t.output_ids.append(out[0] if return_logprobs else out)
t.output_tokens += 1 t.output_tokens += 1
t.advance_kv()
produced.append(t) produced.append(t)
return produced, aborted return produced, aborted
def _run_generation_loop(self): def _run_generation_loop(self):
stop_ids = self._task_mgr.tokenizer.stop_ids stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache
try: try:
while not self._stop_event.is_set(): with self._backend_context():
finished = self._task_mgr.remove_finished_tasks(stop_ids) while not self._stop_event.is_set():
for task in finished: finished = self._task_mgr.remove_finished_tasks(stop_ids)
if task.status == TaskStatus.FINISHED: for task in finished:
cache.task_record_hashes( if task.status == TaskStatus.FINISHED:
task.task_id, self._task_cache.task_record_hashes(
cache.task_cacheable_ids( task.task_id,
task.task_id, task.prompt_ids, task.output_ids self._task_cache.task_cacheable_ids(
), task.task_id, task.prompt_ids, task.output_ids
) ),
cache.task_free(task.task_id) )
self._task_cache.task_free(task.task_id)
active = self._task_mgr.get_active_tasks() active = self._task_mgr.get_active_tasks()
available = self._task_mgr.max_batch_size - len(active) available = self._task_mgr.max_batch_size - len(active)
if available > 0: if available > 0:
candidates = self._task_mgr.pull_candidates(available) candidates = self._task_mgr.pull_candidates(available)
failed = [] failed = []
for task in candidates: for task in candidates:
if cache.task_alloc(task.task_id, task.prompt_ids): if self._task_cache.task_alloc(
self._task_mgr.activate(task) task.task_id, task.prompt_ids
else: ):
failed.append(task) self._task_mgr.activate(task)
if failed: else:
self._task_mgr.return_to_waiting(failed) failed.append(task)
if failed:
self._task_mgr.return_to_waiting(failed)
if not self._task_mgr.has_work(): if not self._task_mgr.has_work():
self._task_mgr.wait_for_tasks(timeout=1.0) self._task_mgr.wait_for_tasks(timeout=1.0)
continue continue
active = self._task_mgr.get_active_tasks() active = self._task_mgr.get_active_tasks()
decoded, aborted = self._step(active) decoded, aborted = self._step(active)
for t in aborted: for t in aborted:
self._task_mgr.invoke_callback(t.task_id, STOP)
for t in decoded:
new_text = t.decode_new_token(self._task_mgr.tokenizer)
if new_text:
self._task_mgr.invoke_callback(t.task_id, new_text)
if t.is_finished(stop_ids):
remaining = t.flush_remaining(self._task_mgr.tokenizer)
if remaining:
self._task_mgr.invoke_callback(t.task_id, remaining)
self._task_mgr.invoke_callback(t.task_id, STOP) self._task_mgr.invoke_callback(t.task_id, STOP)
for t in decoded:
new_text = t.decode_new_token(self._task_mgr.tokenizer)
if new_text:
self._task_mgr.invoke_callback(t.task_id, new_text)
if t.is_finished(stop_ids):
self._task_mgr.invoke_callback(t.task_id, STOP)
except Exception as e: except Exception as e:
self._stop_event.set() self._stop_event.set()
logger.error(f"Scheduler loop crashed: {e}", exc_info=True) logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
for task in self._task_mgr.get_active_tasks(): self._abort_and_clear(free_waiting=False)
self._task_mgr.invoke_callback(task.task_id, STOP)
cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_mgr.clear_queues()
def start(self): def start(self):
if self._loop_thread is not None and self._loop_thread.is_alive(): if self._loop_thread is not None and self._loop_thread.is_alive():
@@ -231,16 +291,21 @@ class InferenceScheduler:
if self._loop_thread is not None: if self._loop_thread is not None:
self._loop_thread.join(timeout=2.0) self._loop_thread.join(timeout=2.0)
self._loop_thread = None self._loop_thread = None
for task in self._task_mgr.get_active_tasks(): self._abort_and_clear(free_waiting=True)
self._task_mgr.invoke_callback(task.task_id, STOP)
self._cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._cache.task_free(task.task_id)
self._task_mgr.clear_queues()
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
def _abort_and_clear(self, free_waiting: bool):
"""Invoke STOP callbacks, release cache slots, and clear task queues."""
for task in self._task_mgr.get_active_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
if free_waiting:
self._task_cache.task_free(task.task_id)
self._task_mgr.clear_queues()
def run_batch( def run_batch(
self, self,
prompt_ids_list: List[List[int]], prompt_ids_list: List[List[int]],
@@ -275,8 +340,8 @@ class InferenceScheduler:
``List[Tuple[List[int], List[float]]]``. ``List[Tuple[List[int], List[float]]]``.
""" """
stop_ids = self._task_mgr.tokenizer.stop_ids stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache
seq_cap = self.max_seq_len seq_cap = self.max_seq_len
request_backend = get_backend(use_default=False)
tasks: List[Task] = [] tasks: List[Task] = []
for ids in prompt_ids_list: for ids in prompt_ids_list:
@@ -300,23 +365,29 @@ class InferenceScheduler:
top_k=top_k, top_k=top_k,
frequency_penalty=frequency_penalty, frequency_penalty=frequency_penalty,
rep_window=rep_window, rep_window=rep_window,
backend=request_backend,
) )
if not cache.task_alloc(task.task_id, task.prompt_ids): if not self._task_cache.task_alloc(task.task_id, task.prompt_ids):
tasks.append(None) tasks.append(None)
continue continue
task.input_tokens = len(task.prompt_ids) task.input_tokens = len(task.prompt_ids)
self._metrics.register(task.task_id)
tasks.append(task) tasks.append(task)
try: try:
live = [t for t in tasks if t is not None] live = [t for t in tasks if t is not None]
while live: with self._backend_context():
decoded, _ = self._step(live, return_logprobs=return_logprobs) while live:
live = [t for t in decoded if not t.is_finished(stop_ids)] decoded, _ = self._step(live, return_logprobs=return_logprobs)
live = [t for t in decoded if not t.is_finished(stop_ids)]
finally: finally:
for t in tasks: for t in tasks:
if t is not None: if t is not None:
cache.task_free(t.task_id) self._metrics.mark_finished(
t.task_id, t.input_tokens, t.output_tokens
)
self._task_cache.task_free(t.task_id)
results: List[Any] = [] results: List[Any] = []
for t in tasks: for t in tasks:
@@ -1,16 +1,17 @@
import logging
import threading import threading
import time import time
import uuid import uuid
from collections import deque from collections import deque
from enum import Enum from enum import Enum
from typing import Any, Callable, Deque, Dict, List, Optional from typing import TYPE_CHECKING, Any, Callable, Deque, Dict, List, Optional
from tokenizers.decoders import DecodeStream from tokenizers.decoders import DecodeStream
from astrai.inference.metrics import MetricsCollector
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__) if TYPE_CHECKING:
from astrai.extension import AttentionBackend
STOP = object() STOP = object()
@@ -64,6 +65,7 @@ class Task:
top_k: int = 50, top_k: int = 50,
frequency_penalty: float = 0.0, frequency_penalty: float = 0.0,
rep_window: int = 64, rep_window: int = 64,
backend: Optional["AttentionBackend"] = None,
): ):
self.task_id = task_id self.task_id = task_id
self.prompt_ids = prompt_ids self.prompt_ids = prompt_ids
@@ -73,16 +75,25 @@ class Task:
self.top_k = top_k self.top_k = top_k
self.frequency_penalty = frequency_penalty self.frequency_penalty = frequency_penalty
self.rep_window = rep_window self.rep_window = rep_window
self.backend = backend
self.status = TaskStatus.PENDING self.status = TaskStatus.PENDING
self.output_ids: List[int] = [] self.output_ids: List[int] = []
self.output_logprobs: List[float] = [] self.output_logprobs: List[float] = []
self.input_tokens: int = 0 self.input_tokens: int = 0
self.output_tokens: int = 0 self.output_tokens: int = 0
self.arrival_time = time.time() self._kv_len: int = 0
self.finish_time: Optional[float] = None
self._decoder: Optional[StreamDecoder] = None self._decoder: Optional[StreamDecoder] = None
def mark_prefill_done(self):
"""Prompt KV is materialized by prefill; first output sampled but
not yet written to KV."""
self._kv_len = self.input_tokens
def advance_kv(self):
"""One more position written to KV (after a decode forward)."""
self._kv_len += 1
def decode_new_token(self, tokenizer: AutoTokenizer) -> str: def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
"""Decode the last appended output token, buffering incomplete """Decode the last appended output token, buffering incomplete
multi-byte sequences across calls. multi-byte sequences across calls.
@@ -93,20 +104,15 @@ class Task:
self._decoder = StreamDecoder(tokenizer) self._decoder = StreamDecoder(tokenizer)
return self._decoder.push(self.output_ids[-1]) return self._decoder.push(self.output_ids[-1])
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
"""Emit any text still buffered in the decoder.
With the Rust-native DecodeStream, the stream is always in a
correct state any completed text was already emitted by the
last ``push``. A trailing incomplete multi-byte sequence has no
valid text to emit, so this is a no-op.
"""
return ""
@property @property
def next_pos(self) -> int: def next_pos(self) -> int:
# The first output is sampled from prefill and enters KV on the next step. """KV position where the next decode step will write."""
return self.input_tokens + max(0, len(self.output_ids) - 1) return self._kv_len
@property
def prefill_done(self) -> bool:
"""True when all prompt KV entries are materialized."""
return self._kv_len >= self.input_tokens > 0
def is_finished(self, stop_ids: List[int]) -> bool: def is_finished(self, stop_ids: List[int]) -> bool:
if self.max_tokens is not None and self.output_tokens >= self.max_tokens: if self.max_tokens is not None and self.output_tokens >= self.max_tokens:
@@ -124,6 +130,7 @@ class TaskManager:
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
max_batch_size: int = 16, max_batch_size: int = 16,
max_seq_len: int = 8192, max_seq_len: int = 8192,
metrics: Optional["MetricsCollector"] = None,
): ):
self.tokenizer = tokenizer self.tokenizer = tokenizer
self.max_batch_size = max_batch_size self.max_batch_size = max_batch_size
@@ -139,6 +146,8 @@ class TaskManager:
self._total_tasks = 0 self._total_tasks = 0
self._total_tokens = 0 self._total_tokens = 0
self._metrics = metrics
def add_task( def add_task(
self, self,
prompt: str, prompt: str,
@@ -148,6 +157,7 @@ class TaskManager:
top_k: int = 50, top_k: int = 50,
frequency_penalty: float = 0.0, frequency_penalty: float = 0.0,
rep_window: int = 64, rep_window: int = 64,
backend: Optional["AttentionBackend"] = None,
stream_callback: Optional[Callable[[str], None]] = None, stream_callback: Optional[Callable[[str], None]] = None,
) -> str: ) -> str:
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}" task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
@@ -155,11 +165,6 @@ class TaskManager:
if len(prompt_ids) > self.max_seq_len: if len(prompt_ids) > self.max_seq_len:
prompt_ids = prompt_ids[-self.max_seq_len :] prompt_ids = prompt_ids[-self.max_seq_len :]
if len(prompt_ids) > self.max_seq_len:
if stream_callback:
stream_callback(STOP)
return task_id
if max_tokens is None: if max_tokens is None:
max_tokens = self.max_seq_len - len(prompt_ids) max_tokens = self.max_seq_len - len(prompt_ids)
else: else:
@@ -174,6 +179,7 @@ class TaskManager:
top_k=top_k, top_k=top_k,
frequency_penalty=frequency_penalty, frequency_penalty=frequency_penalty,
rep_window=rep_window, rep_window=rep_window,
backend=backend,
) )
with self._lock: with self._lock:
@@ -182,6 +188,9 @@ class TaskManager:
if stream_callback: if stream_callback:
self._callbacks[task_id] = stream_callback self._callbacks[task_id] = stream_callback
if self._metrics is not None:
self._metrics.register(task_id)
self._task_event.set() self._task_event.set()
return task_id return task_id
@@ -201,26 +210,33 @@ class TaskManager:
cb(token) cb(token)
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
return { stats: Dict[str, Any] = {
"total_tasks": self._total_tasks, "total_tasks": self._total_tasks,
"total_tokens": self._total_tokens, "total_tokens": self._total_tokens,
"active_tasks": len(self.active_tasks), "active_tasks": len(self.active_tasks),
"waiting_queue": len(self.waiting_queue), "waiting_queue": len(self.waiting_queue),
} }
if self._metrics is not None:
stats.update(self._metrics.get_stats())
return stats
def remove_finished_tasks(self, stop_ids: List[int]) -> List[Task]: def remove_finished_tasks(self, stop_ids: List[int]) -> List[Task]:
with self._lock: with self._lock:
finished = [] finished = []
for task in self.active_tasks: for task in self.active_tasks:
if task.status == TaskStatus.ABORTED: if task.status == TaskStatus.ABORTED:
task.finish_time = time.time()
finished.append(task) finished.append(task)
elif task.is_finished(stop_ids): elif task.is_finished(stop_ids):
task.status = TaskStatus.FINISHED task.status = TaskStatus.FINISHED
task.finish_time = time.time()
finished.append(task) finished.append(task)
self._total_tokens += task.output_tokens self._total_tokens += task.output_tokens
if self._metrics is not None:
for task in finished:
self._metrics.mark_finished(
task.task_id, task.input_tokens, task.output_tokens
)
self.active_tasks = [ self.active_tasks = [
t t
for t in self.active_tasks for t in self.active_tasks
@@ -10,6 +10,7 @@ import torch
from torch import Tensor from torch import Tensor
_MAX_SPLITS = 32 _MAX_SPLITS = 32
Q_TILE_ROWS = 64
class InferenceWorkspace: class InferenceWorkspace:
@@ -74,7 +75,7 @@ class InferenceWorkspace:
# when the Executor passes this workspace). Stable addresses make the # when the Executor passes this workspace). Stable addresses make the
# decode forward CUDA-graph capturable. # decode forward CUDA-graph capturable.
self.req_pool_indices = torch.empty( self.req_pool_indices = torch.empty(
(max_batch_size,), dtype=torch.long, device=device (max_batch_size,), dtype=torch.int32, device=device
) )
self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device) self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device)
self.kv_indptr = torch.empty( self.kv_indptr = torch.empty(
@@ -83,9 +84,16 @@ class InferenceWorkspace:
self.qo_indptr = torch.empty( self.qo_indptr = torch.empty(
(max_batch_size + 1,), dtype=torch.int32, device=device (max_batch_size + 1,), dtype=torch.int32, device=device
) )
max_q_tiles = max_batch_size * ((max_seq_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS)
self.q_tile_to_batch = torch.empty(
(max_q_tiles,), dtype=torch.int32, device=device
)
self.q_tile_to_index = torch.empty(
(max_q_tiles,), dtype=torch.int32, device=device
)
self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device) self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device)
self.out_cache_loc = torch.empty( self.out_cache_loc = torch.empty(
(max_batch_size, 1), dtype=torch.long, device=device (max_batch_size, 1), dtype=torch.int32, device=device
) )
# Per-step position IDs (must be at a fixed address for CUDA-graph capture). # Per-step position IDs (must be at a fixed address for CUDA-graph capture).
@@ -116,13 +124,6 @@ class InferenceWorkspace:
device=device, device=device,
) )
def decode_buffers(self, batch: int, q_heads: int):
"""Return ``(o_part, ml_part)`` view sliced to live dimensions."""
return (
self.decode_o_part[:batch, :q_heads],
self.decode_ml_part[:batch, :q_heads],
)
def fill_input_ids(self, ids: "list[int]") -> Tensor: def fill_input_ids(self, ids: "list[int]") -> Tensor:
"""Write ``ids`` into the device buffer and return ``[B]``. """Write ``ids`` into the device buffer and return ``[B]``.
@@ -138,6 +139,18 @@ class InferenceWorkspace:
self.input_ids[:b].copy_(pin[:b]) self.input_ids[:b].copy_(pin[:b])
return self.input_ids[:b] return self.input_ids[:b]
def fill_input_ids_from_device(self, tokens: Tensor) -> Tensor:
"""Copy device-resident ``[B]`` token ids into the device buffer.
Steady-state decode fast path: when the executor's cached task
signature still matches, the previous step's sampled tokens map
1:1 onto the current slots, so the ids transfer device-to-device
instead of round-tripping through the host staging buffers.
"""
b = tokens.size(0)
self.input_ids[:b].copy_(tokens)
return self.input_ids[:b]
def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor: def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor:
"""Return the ``[B, 1, total_len]`` validity mask for this step. """Return the ``[B, 1, total_len]`` validity mask for this step.
+37
View File
@@ -0,0 +1,37 @@
import logging
import os
from astrai.parallel.setup import get_rank, get_world_size
class _DistributedContextFilter(logging.Filter):
def filter(self, record: logging.LogRecord) -> bool:
record.rank = str(get_rank())
record.world_size = str(get_world_size())
return True
def setup_logging(level: str = "INFO"):
"""Attach a StreamHandler to the ``astrai`` logger (idempotent).
Call once per process at the top of CLI scripts.
Set ``ASTR_LOG_LEVEL`` env var to override the default level.
Level names: ``DEBUG``, ``INFO``, ``WARNING``, ``ERROR``, ``CRITICAL``.
``DEBUG`` enables per-step prefill/decode timing logs
(:func:`astrai.inference.runtime.executor.timed`).
"""
logger = logging.getLogger("astrai")
if logger.handlers:
return
level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper()
logger.setLevel(getattr(logging, level_name, logging.INFO))
handler = logging.StreamHandler()
handler.addFilter(_DistributedContextFilter())
handler.setFormatter(
logging.Formatter(
"%(asctime)s | %(levelname)-8s | rank=%(rank)2s/%(world_size)-2s | %(name)-32s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
)
logger.addHandler(handler)
+41 -7
View File
@@ -4,13 +4,21 @@ AutoModel base class for model loading and saving.
from contextlib import contextmanager from contextlib import contextmanager
from pathlib import Path from pathlib import Path
from typing import Self, Union from typing import Union
import torch.nn as nn import torch.nn as nn
from astrai.config.model_config import BaseModelConfig, ConfigFactory from astrai.config.model_config import BaseModelConfig, ConfigFactory
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.serialization import load_model_config, load_model_weights, save_model from astrai.serialization import (
HF_MODEL_TYPES,
adapt_config,
convert_hf_weights,
load_model_config,
load_model_weights,
looks_like_hf_state_dict,
save_model,
)
@contextmanager @contextmanager
@@ -57,7 +65,25 @@ class AutoModel(nn.Module):
path: Union[str, Path], path: Union[str, Path],
disable_random_init: bool = True, disable_random_init: bool = True,
strict: bool = True, strict: bool = True,
weights_format: str = "auto",
) -> nn.Module: ) -> nn.Module:
"""Load a model directory.
Args:
path: Directory containing ``config.json`` and optionally
``model.safetensors``.
disable_random_init: Replace parameter initializers with no-ops
while building the model.
strict: Passed to ``load_state_dict``.
weights_format: ``"auto"`` detects HuggingFace checkpoints
(LLaMA-style keys and ``model_type``) and converts them;
``"astrai"`` skips conversion; ``"hf"`` forces it.
"""
if weights_format not in ("auto", "astrai", "hf"):
raise ValueError(
f"weights_format must be one of 'auto', 'astrai', 'hf', "
f"got {weights_format!r}"
)
model_path = Path(path) model_path = Path(path)
@@ -66,6 +92,12 @@ class AutoModel(nn.Module):
raise FileNotFoundError(f"Config file not found: {config_path}") raise FileNotFoundError(f"Config file not found: {config_path}")
raw = load_model_config(str(model_path)) raw = load_model_config(str(model_path))
is_hf_config = weights_format == "hf" or (
weights_format == "auto" and raw.get("model_type") in HF_MODEL_TYPES
)
if is_hf_config:
raw = adapt_config(raw)
config = ConfigFactory.load(raw) config = ConfigFactory.load(raw)
model_type = config.model_type or "autoregressive_lm" model_type = config.model_type or "autoregressive_lm"
@@ -75,8 +107,14 @@ class AutoModel(nn.Module):
model = actual_cls(config) model = actual_cls(config)
weights_path = model_path / "model.safetensors" weights_path = model_path / "model.safetensors"
if weights_path.exists(): index_path = model_path / "model.safetensors.index.json"
if weights_path.exists() or index_path.exists():
state_dict = load_model_weights(str(model_path)) state_dict = load_model_weights(str(model_path))
is_hf_weights = is_hf_config or (
weights_format == "auto" and looks_like_hf_state_dict(state_dict)
)
if is_hf_weights:
state_dict = convert_hf_weights(state_dict, config)
model.load_state_dict(state_dict, strict=strict) model.load_state_dict(state_dict, strict=strict)
return model return model
@@ -90,7 +128,3 @@ class AutoModel(nn.Module):
state_dict=self.state_dict(), state_dict=self.state_dict(),
save_directory=str(save_directory), save_directory=str(save_directory),
) )
def to(self, *args, **kwargs) -> Self:
"""Move model to device/dtype."""
return super().to(*args, **kwargs)
+1 -1
View File
@@ -1,4 +1,4 @@
from astrai.extension.rotary_backend import apply_rotary_emb from astrai.extension.backend.rotary import apply_rotary_emb
from astrai.model.components.attention import GQA, MLA from astrai.model.components.attention import GQA, MLA
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding from astrai.model.components.embedding import Embedding
+13 -12
View File
@@ -5,10 +5,9 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.extension import attention from astrai.extension.backend import apply_rotary_emb, attention
from astrai.extension.rotary_backend import apply_rotary_emb
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.inference.core.cache import KVCache from astrai.inference.cache import KVCache
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
@@ -56,9 +55,7 @@ class GQA(nn.Module):
self.gate = Linear(dim, dim) self.gate = Linear(dim, dim)
def _split_heads(self, x: Tensor, n_heads) -> Tensor: def _split_heads(self, x: Tensor, n_heads) -> Tensor:
batch_size, seq_len, _ = x.shape return x.reshape(*x.shape[:-1], n_heads, self.head_dim)
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
return x
def forward( def forward(
self, self,
@@ -67,6 +64,7 @@ class GQA(nn.Module):
attn_mask: Tensor = None, attn_mask: Tensor = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor: ) -> Tensor:
q = self._split_heads(self.q_proj(x), self.n_heads) q = self._split_heads(self.q_proj(x), self.n_heads)
k = self._split_heads(self.k_proj(x), self.n_kv_heads) k = self._split_heads(self.k_proj(x), self.n_kv_heads)
@@ -76,7 +74,9 @@ class GQA(nn.Module):
if self.use_qk_norm: if self.use_qk_norm:
q, k = self.q_norm(q), self.k_norm(k) q, k = self.q_norm(q), self.k_norm(k)
sdqa_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal) sdqa_out = attention(
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
).reshape(*x.shape[:-1], self.dim)
if self.use_gated_attention: if self.use_gated_attention:
sdqa_out = sdqa_out * F.sigmoid(self.gate(x)) sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
@@ -141,17 +141,16 @@ class MLA(nn.Module):
attn_mask: Tensor = None, attn_mask: Tensor = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor: ) -> Tensor:
bsz, seq_len, _ = x.size()
q = self.q_proj(x) q = self.q_proj(x)
q = q.view(bsz, seq_len, self.n_heads, self.head_dim) q = q.reshape(*x.shape[:-1], self.n_heads, self.head_dim)
kv_compressed = self.kv_a_proj(x) kv_compressed = self.kv_a_proj(x)
kv_compressed = self.kv_norm(kv_compressed) kv_compressed = self.kv_norm(kv_compressed)
kv = self.kv_b_proj(kv_compressed) kv = self.kv_b_proj(kv_compressed)
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1) kv = kv.reshape(*x.shape[:-1], self.n_kv_heads, -1)
k_nope, k_rope, v = torch.split( k_nope, k_rope, v = torch.split(
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1 kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
@@ -171,7 +170,9 @@ class MLA(nn.Module):
q = self.q_norm(q) q = self.q_norm(q)
k = self.k_norm(k) k = self.k_norm(k)
attn_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal) attn_out = attention(
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
).reshape(*x.shape[:-1], self.dim)
if self.use_gated_attention: if self.use_gated_attention:
attn_out = attn_out * F.sigmoid(self.gate(x)) attn_out = attn_out * F.sigmoid(self.gate(x))
+3 -1
View File
@@ -4,7 +4,7 @@ from typing import Optional, TypedDict
import torch.nn as nn import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.inference.core.cache import KVCache from astrai.inference.cache import KVCache
from astrai.model.components.attention import AttnFactory from astrai.model.components.attention import AttnFactory
from astrai.model.components.mlp import FFNFactory, RouterStats from astrai.model.components.mlp import FFNFactory, RouterStats
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
@@ -54,6 +54,7 @@ class DecoderBlock(nn.Module):
attention_mask: Optional[Tensor] = None, attention_mask: Optional[Tensor] = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> DecoderOutput: ) -> DecoderOutput:
attn_output = self.attention( attn_output = self.attention(
self.input_norm(x), self.input_norm(x),
@@ -61,6 +62,7 @@ class DecoderBlock(nn.Module):
attention_mask, attention_mask,
kv_cache, kv_cache,
is_causal, is_causal,
fwd,
) )
x = attn_output + x x = attn_output + x
normalized = self.post_attention_norm(x) normalized = self.post_attention_norm(x)
+4 -9
View File
@@ -29,12 +29,6 @@ class FFNOutput(TypedDict):
router_stats: Optional[RouterStats] router_stats: Optional[RouterStats]
class RoutedOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
router_stats: Optional[RouterStats]
@FFNFactory.register("mlp") @FFNFactory.register("mlp")
class MLP(nn.Module): class MLP(nn.Module):
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02): def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
@@ -100,13 +94,14 @@ class DeepSeekMoE(nn.Module):
def forward(self, x: Tensor) -> FFNOutput: def forward(self, x: Tensor) -> FFNOutput:
include_aux_loss = self.training and torch.is_grad_enabled() include_aux_loss = self.training and torch.is_grad_enabled()
bsz, seq_len, dim = x.shape shape = x.shape
dim = shape[-1]
x_flat = x.view(-1, dim) x_flat = x.view(-1, dim)
shared_out = self._shared_forward(x_flat) shared_out = self._shared_forward(x_flat)
routed_output = self._routed_forward(x_flat, include_aux_loss) routed_output = self._routed_forward(x_flat, include_aux_loss)
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim) out = (shared_out + routed_output["hidden_states"]).view(shape)
return { return {
"hidden_states": out, "hidden_states": out,
"aux_loss": routed_output["aux_loss"], "aux_loss": routed_output["aux_loss"],
@@ -121,7 +116,7 @@ class DeepSeekMoE(nn.Module):
/ self.n_shared_experts / self.n_shared_experts
) )
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput: def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> FFNOutput:
N, D = x.shape N, D = x.shape
K = self.n_activated_experts K = self.n_activated_experts
E = self.n_routed_experts E = self.n_routed_experts
+8 -5
View File
@@ -65,9 +65,12 @@ class RotaryEmbedding(nn.Module):
[batch, seq_len, dim/2, 2] (f32) [cos, sin] pairs. [batch, seq_len, dim/2, 2] (f32) [cos, sin] pairs.
""" """
if position_ids is None: if position_ids is None:
position_ids = ( if x.ndim == 2:
torch.arange(x.size(1), device=x.device) position_ids = torch.arange(x.size(0), device=x.device)
.unsqueeze(0) else:
.expand(x.size(0), -1) position_ids = (
) torch.arange(x.size(1), device=x.device)
.unsqueeze(0)
.expand(x.size(0), -1)
)
return self.freqs_cis[position_ids].float() return self.freqs_cis[position_ids].float()
+15 -2
View File
@@ -5,7 +5,7 @@ import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.config.model_config import AutoRegressiveLMConfig from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.inference.core.cache import KVCache from astrai.inference.cache import KVCache
from astrai.model.automodel import AutoModel, ModelFactory from astrai.model.automodel import AutoModel, ModelFactory
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding from astrai.model.components.embedding import Embedding
@@ -105,8 +105,20 @@ class AutoRegressiveLM(AutoModel):
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
kv_cache: Optional[KVCache] = None, kv_cache: Optional[KVCache] = None,
position_ids: Optional[Tensor] = None, position_ids: Optional[Tensor] = None,
fwd: Optional[str] = None,
) -> Dict[str, Tensor]: ) -> Dict[str, Tensor]:
assert input_ids.ndim == 2 if fwd is None:
if input_ids.ndim != 2:
raise ValueError("training input_ids must be [batch, seq_len]")
if kv_cache is not None:
raise ValueError("training forward does not accept a KV cache")
elif fwd in ("prefill", "decode"):
if input_ids.ndim != 1:
raise ValueError("inference input_ids must be packed [tokens]")
if kv_cache is None:
raise ValueError("inference forward requires a KV cache")
else:
raise ValueError(f"unsupported forward mode: {fwd}")
x = self.embed_tokens(input_ids) x = self.embed_tokens(input_ids)
rotary_emb = self.rotary_embedding(x, position_ids) rotary_emb = self.rotary_embedding(x, position_ids)
@@ -122,6 +134,7 @@ class AutoRegressiveLM(AutoModel):
attn_mask, attn_mask,
kv_cache, kv_cache,
use_sdpa_causal_mask, use_sdpa_causal_mask,
fwd,
) )
x = layer_output["hidden_states"] x = layer_output["hidden_states"]
stats = layer_output.get("router_stats") stats = layer_output.get("router_stats")
+9 -15
View File
@@ -30,15 +30,13 @@ def get_current_device():
def get_world_size() -> int: def get_world_size() -> int:
if dist.is_available() and dist.is_initialized(): if dist.is_available() and dist.is_initialized():
return dist.get_world_size() return dist.get_world_size()
else: return int(os.environ.get("WORLD_SIZE", "1"))
return 1
def get_rank() -> int: def get_rank() -> int:
if dist.is_available() and dist.is_initialized(): if dist.is_available() and dist.is_initialized():
return dist.get_rank() return dist.get_rank()
else: return int(os.environ.get("RANK", "0"))
return 0
@contextmanager @contextmanager
@@ -247,18 +245,15 @@ class LocalStrategy(LaunchStrategy):
ctx.join() ctx.join()
def _detect_launcher() -> str: def _is_external_launcher() -> bool:
"""Detect the distributed launcher from environment. """Whether an external launcher (torchrun/elastic/manual env) started us."""
Returns one of: "torchelastic", "torchrun", "external", "local".
"""
if dist.is_torchelastic_launched(): if dist.is_torchelastic_launched():
return "torchelastic" return True
if "LOCAL_WORLD_SIZE" in os.environ: if "LOCAL_WORLD_SIZE" in os.environ:
return "torchrun" return True
if "RANK" in os.environ and "WORLD_SIZE" in os.environ: if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
return "external" return True
return "local" return False
def spawn_parallel_fn( def spawn_parallel_fn(
@@ -273,8 +268,7 @@ def spawn_parallel_fn(
): ):
if master_port is None: if master_port is None:
master_port = find_free_port() master_port = find_free_port()
launcher = _detect_launcher() if _is_external_launcher():
if launcher in ("torchelastic", "torchrun", "external"):
strategy = TorchrunStrategy( strategy = TorchrunStrategy(
world_size, backend, master_addr, master_port, device_type, start_method world_size, backend, master_addr, master_port, device_type, start_method
) )
+12
View File
@@ -22,9 +22,21 @@ from astrai.serialization.dataset import (
load_bin_offsets, load_bin_offsets,
save_bin, save_bin,
) )
from astrai.serialization.hf_adapter import (
HF_MODEL_TYPES,
adapt_config,
convert_hf_config,
convert_hf_weights,
looks_like_hf_state_dict,
)
__all__ = [ __all__ = [
"Checkpoint", "Checkpoint",
"HF_MODEL_TYPES",
"adapt_config",
"convert_hf_config",
"convert_hf_weights",
"looks_like_hf_state_dict",
"load_json", "load_json",
"load_model_config", "load_model_config",
"load_model_weights", "load_model_weights",
+31 -23
View File
@@ -5,7 +5,7 @@ import json
import time import time
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Optional, Union from typing import Any, Callable, Dict, Optional, Union
import safetensors.torch as st import safetensors.torch as st
import torch import torch
@@ -22,39 +22,31 @@ def save_safetensors(state_dict: dict, path: Union[str, Path]):
st.save_file(state_dict, str(path)) st.save_file(state_dict, str(path))
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict: def _broadcast_load(loader: Callable[[], dict], broadcast: bool) -> dict:
"""Load on rank 0 and broadcast the object to all ranks."""
if not broadcast or not dist.is_initialized(): if not broadcast or not dist.is_initialized():
return st.load_file(str(path)) return loader()
rank = get_rank() rank = get_rank()
if rank == 0: if rank == 0:
state_dict = st.load_file(str(path)) data = loader()
else: else:
state_dict = {} data = {}
tmp = [state_dict] tmp = [data]
dist.broadcast_object_list(tmp, src=0) dist.broadcast_object_list(tmp, src=0)
return tmp[0] return tmp[0]
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
return _broadcast_load(lambda: st.load_file(str(path)), broadcast)
def save_json(data: dict, path: Union[str, Path]): def save_json(data: dict, path: Union[str, Path]):
with open(str(path), "w") as f: with open(str(path), "w") as f:
json.dump(data, f, indent=2) json.dump(data, f, indent=2)
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict: def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
if not broadcast or not dist.is_initialized(): return _broadcast_load(lambda: json.loads(Path(path).read_text()), broadcast)
with open(str(path), "r") as f:
return json.load(f)
rank = get_rank()
if rank == 0:
with open(str(path), "r") as f:
data = json.load(f)
else:
data = {}
tmp = [data]
dist.broadcast_object_list(tmp, src=0)
return tmp[0]
def save_torch(obj: Any, path: Union[str, Path]): def save_torch(obj: Any, path: Union[str, Path]):
@@ -99,7 +91,21 @@ def load_model_config(save_directory: str) -> dict:
def load_model_weights(save_directory: str) -> dict: def load_model_weights(save_directory: str) -> dict:
return load_state_dict(Path(save_directory) / _WEIGHTS_FILE) save_path = Path(save_directory)
weights_file = save_path / _WEIGHTS_FILE
if weights_file.exists():
return load_state_dict(weights_file)
index_path = save_path / "model.safetensors.index.json"
if index_path.exists():
index = load_json(index_path)
weight_map = index.get("weight_map", {})
state_dict = {}
for shard in sorted(set(weight_map.values())):
state_dict.update(load_state_dict(save_path / shard))
return state_dict
raise FileNotFoundError(f"No model weights found in {save_directory}")
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict: def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
@@ -190,8 +196,10 @@ class Checkpoint:
if meta_path.exists(): if meta_path.exists():
return cls.load(save_dir, broadcast=broadcast) return cls.load(save_dir, broadcast=broadcast)
if weights_path.exists(): weights_path = save_path / _WEIGHTS_FILE
state_dict = load_state_dict(weights_path, broadcast=broadcast) index_path = save_path / "model.safetensors.index.json"
if weights_path.exists() or index_path.exists():
state_dict = load_model_weights(save_dir)
config = {} config = {}
config_path = save_path / _CONFIG_FILE config_path = save_path / _CONFIG_FILE
if config_path.exists(): if config_path.exists():
+271
View File
@@ -0,0 +1,271 @@
"""HuggingFace checkpoint adaptation for LLaMA-style decoder models.
AstrAI stores weights with its own key names (``layers.<i>.input_norm``,
``layers.<i>.mlp.gate``), while HuggingFace decoder-only checkpoints use
``model.layers.<i>.input_layernorm`` / ``model.layers.<i>.mlp.gate_proj``.
This module translates HF configs and state dicts so external checkpoints
can be loaded directly.
Supported families (LLaMA layout, dense and MoE):
- dense FFN: llama, mistral, qwen2, gemma, gemma2, phi3
- MoE FFN (Mixtral / Qwen2-MoE / DeepSeek-V3 layout): router
``mlp.gate``, routed experts ``mlp.experts.<j>``, shared experts
``mlp.shared_experts.<j>``
Not supported:
- MLA attention (DeepSeek-V2/V3 ``kv_a_proj_with_mqa``) uses a different
KV factorization and cannot be converted numerically.
- Attention/MLP bias (``attention_bias`` / ``mlp_bias``) AstrAI
projections are bias-free.
"""
import logging
import re
from typing import Any, Dict, Mapping
import torch
from astrai.config.base import BaseConfig
logger = logging.getLogger(__name__)
HF_MODEL_TYPES = frozenset(
{
"llama",
"mistral",
"mixtral",
"qwen2",
"qwen2_moe",
"gemma",
"gemma2",
"phi3",
}
)
_EMBED = re.compile(r"^model\.embed_tokens\.weight$")
_ATTN = re.compile(r"^model\.layers\.(\d+)\.self_attn\.(q|k|v|o)_proj\.(weight|bias)$")
_Q_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.q_norm\.weight$")
_K_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.k_norm\.weight$")
_INPUT_NORM = re.compile(r"^model\.layers\.(\d+)\.input_layernorm\.weight$")
_POST_NORM = re.compile(r"^model\.layers\.(\d+)\.post_attention_layernorm\.weight$")
_FINAL_NORM = re.compile(r"^model\.norm\.weight$")
_LM_HEAD = re.compile(r"^lm_head\.weight$")
_DENSE_MLP = re.compile(
r"^model\.layers\.(\d+)\.mlp\.(gate|up|down)_proj\.(weight|bias)$"
)
_MOE_ROUTER = re.compile(r"^model\.layers\.(\d+)\.mlp\.gate\.weight$")
_MOE_EXPERTS = re.compile(
r"^model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.(gate|up|down)_proj\.(weight|bias)$"
)
_MOE_SHARED = re.compile(
r"^model\.layers\.(\d+)\.mlp\.shared_expert(?:s)?\.(\d+)\."
r"(gate|up|down)_proj\.(weight|bias)$"
)
_ASTR_PREFIXES = ("embed_tokens.", "layers.", "norm.", "lm_head.")
def looks_like_hf_state_dict(state_dict: Mapping[str, Any]) -> bool:
"""Return True if *state_dict* uses HuggingFace key names."""
return any(
key.startswith("model.")
or "self_attn." in key
or "input_layernorm" in key
or "mlp.experts." in key
for key in state_dict
)
def _is_dense_mlp_layer(config: BaseConfig, layer_id: int) -> bool:
"""Return whether a layer uses dense MLP instead of routed experts."""
if getattr(config, "ffn_type", "mlp") != "moe":
return True
mlp_only = getattr(config, "mlp_only_layers", None) or []
if layer_id in mlp_only:
return True
step = getattr(config, "decoder_sparse_step", 1) or 1
return step > 1 and (layer_id + 1) % step != 0
def adapt_config(raw: Dict[str, Any]) -> Dict[str, Any]:
"""Translate *raw* for AstrAI if it looks like an HF model config."""
if raw.get("model_type") in HF_MODEL_TYPES:
return convert_hf_config(raw)
return raw
def convert_hf_config(raw: Dict[str, Any]) -> Dict[str, Any]:
"""Convert an HF LLaMA-style config dict to AstrAI field names."""
if raw.get("attention_bias") or raw.get("mlp_bias"):
raise NotImplementedError(
"attention_bias / mlp_bias checkpoints are not supported; "
"AstrAI projections are bias-free"
)
cfg: Dict[str, Any] = {}
for key in (
"vocab_size",
"hidden_size",
"num_hidden_layers",
"intermediate_size",
"rms_norm_eps",
"tie_word_embeddings",
"max_position_embeddings",
"rope_theta",
"rope_scaling",
"num_attention_heads",
"num_key_value_heads",
"use_qk_norm",
"use_gated_attention",
"kv_lora_rank",
"qk_nope_head_dim",
"qk_rope_head_dim",
"moe_intermediate_size",
"shared_expert_intermediate_size",
"topk_method",
"norm_topk_prob",
"moe_aux_loss_coef",
"decoder_sparse_step",
"mlp_only_layers",
"neftune_alpha",
):
if key in raw:
cfg[key] = raw[key]
if "qk_norm" in raw and "use_qk_norm" not in cfg:
cfg["use_qk_norm"] = raw["qk_norm"]
if (
raw.get("model_type") in ("gemma", "gemma2")
and "use_qk_norm" not in cfg
and "qk_norm" not in raw
):
# Gemma/Gemma2 always apply RMSNorm to Q and K before attention.
cfg["use_qk_norm"] = True
n_heads = raw.get("num_attention_heads")
if cfg.get("num_key_value_heads") is None and n_heads is not None:
cfg["num_key_value_heads"] = n_heads
if raw.get("head_dim") is not None and n_heads and raw.get("hidden_size"):
expected = raw["hidden_size"] // n_heads
if raw["head_dim"] != expected:
raise NotImplementedError(
f"HF head_dim={raw['head_dim']} differs from the computed "
f"head dim {expected}; AstrAI derives head_dim from "
"hidden_size / num_attention_heads"
)
if "kv_lora_rank" in raw:
cfg["attn_type"] = "mla"
n_experts = raw.get("num_local_experts") or raw.get("n_routed_experts")
if n_experts:
cfg["ffn_type"] = "moe"
cfg["n_routed_experts"] = n_experts
if "num_experts_per_tok" in raw:
cfg["n_activated_experts"] = raw["num_experts_per_tok"]
if "n_activated_experts" in raw:
cfg["n_activated_experts"] = raw["n_activated_experts"]
if "n_shared_experts" in raw:
cfg["n_shared_experts"] = raw["n_shared_experts"]
else:
# Mixtral has no shared experts; AstrAI defaults to one.
cfg["n_shared_experts"] = 0
if cfg.get("moe_intermediate_size") is None and "intermediate_size" in raw:
# MoE configs store the per-expert FFN size in intermediate_size.
cfg["moe_intermediate_size"] = raw["intermediate_size"]
first_k_dense = raw.get("first_k_dense_replace")
if isinstance(first_k_dense, int) and first_k_dense > 0:
cfg["mlp_only_layers"] = list(range(first_k_dense))
cfg["decoder_sparse_step"] = 1
cfg["model_type"] = "autoregressive_lm"
return cfg
def convert_hf_weights(
state_dict: Mapping[str, Any],
config: BaseConfig,
) -> Dict[str, torch.Tensor]:
"""Rename HF state dict keys to AstrAI names.
Keys that are already AstrAI-style pass through unchanged; unmapped
HF keys are dropped with a warning. Use with ``strict=True`` to fail
loudly when the checkpoint does not match the config.
"""
if getattr(config, "attn_type", "gqa") == "mla":
if any("kv_a_proj_with_mqa" in key for key in state_dict):
raise NotImplementedError(
"MLA attention (DeepSeek-V2/V3 kv_a_proj_with_mqa) uses a "
"different KV factorization and cannot be converted"
)
ffn_type = getattr(config, "ffn_type", "mlp")
converted: Dict[str, torch.Tensor] = {}
skipped: list[str] = []
for key, tensor in state_dict.items():
if key.startswith(_ASTR_PREFIXES):
converted[key] = tensor
continue
new_key = None
if ffn_type == "moe":
m = _MOE_ROUTER.match(key)
if m:
new_key = f"layers.{m.group(1)}.mlp.router.weight"
else:
m = _MOE_EXPERTS.match(key)
if m:
new_key = (
f"layers.{m.group(1)}.mlp.routed_experts.{m.group(2)}."
f"{m.group(3)}.{m.group(4)}"
)
else:
m = _MOE_SHARED.match(key)
if m:
new_key = (
f"layers.{m.group(1)}.mlp.shared_experts.{m.group(2)}."
f"{m.group(3)}.{m.group(4)}"
)
if new_key is None:
m = _DENSE_MLP.match(key)
if m and _is_dense_mlp_layer(config, int(m.group(1))):
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
else:
m = _DENSE_MLP.match(key)
if m:
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
if new_key is None:
m = _ATTN.match(key)
if m:
new_key = (
f"layers.{m.group(1)}.attention.{m.group(2)}_proj.{m.group(3)}"
)
elif (m := _Q_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.attention.q_norm.weight"
elif (m := _K_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.attention.k_norm.weight"
elif (m := _INPUT_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.input_norm.weight"
elif (m := _POST_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.post_attention_norm.weight"
elif (m := _EMBED.match(key)) is not None:
new_key = "embed_tokens.weight"
elif (m := _FINAL_NORM.match(key)) is not None:
new_key = "norm.weight"
elif (m := _LM_HEAD.match(key)) is not None:
new_key = "lm_head.weight"
if new_key is None:
skipped.append(key)
else:
converted[new_key] = tensor
if skipped:
logger.warning(
"Dropped %d unmapped HuggingFace weight key(s): %s",
len(skipped),
", ".join(sorted(skipped)[:10]),
)
return converted
+7 -19
View File
@@ -1,3 +1,4 @@
import math
from typing import Dict from typing import Dict
import torch import torch
@@ -27,6 +28,8 @@ class GradSNRTracker:
SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2) SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2)
The reported value is the power ratio in decibels: ``10 * log10(SNR)``.
The tracker accumulates per-parameter EMA moments across optimizer steps. The tracker accumulates per-parameter EMA moments across optimizer steps.
Call ``update`` after backward (before ``optimizer.step``) and read Call ``update`` after backward (before ``optimizer.step``) and read
``snr`` to get the aggregate SNR across all parameters. ``snr`` to get the aggregate SNR across all parameters.
@@ -64,7 +67,8 @@ class GradSNRTracker:
noise = (v - m.pow(2)).clamp(min=0).sum().item() noise = (v - m.pow(2)).clamp(min=0).sum().item()
total_signal += signal total_signal += signal
total_noise += noise total_noise += noise
return total_signal / (total_noise + self.eps) snr = total_signal / (total_noise + self.eps)
return 10.0 * math.log10(max(snr, self.eps))
def ctx_get_loss(ctx): def ctx_get_loss(ctx):
@@ -90,21 +94,5 @@ def ctx_get_grad_snr(ctx):
return tracker.snr return tracker.snr
def ctx_get_moe_aux_loss(ctx): def ctx_get_moe_metric(ctx, key):
return ctx.strategy._moe_metrics.get("aux_loss") return ctx.strategy._moe_metrics.get(key)
def ctx_get_router_entropy(ctx):
return ctx.strategy._moe_metrics.get("router_entropy")
def ctx_get_dead_expert_fraction(ctx):
return ctx.strategy._moe_metrics.get("dead_expert_fraction")
def ctx_get_load_imbalance_mean(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_mean")
def ctx_get_load_imbalance_max(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_max")
+3 -3
View File
@@ -6,7 +6,7 @@ Provides:
- :class:`BaseRewardModel` pluggable reward interface - :class:`BaseRewardModel` pluggable reward interface
- :class:`RolloutGenerator` KV-cache-backed generation of grouped - :class:`RolloutGenerator` KV-cache-backed generation of grouped
responses + decoding (no reward); delegates the generation loop to responses + decoding (no reward); delegates the generation loop to
:class:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch` :class:`~astrai.inference.scheduler.InferenceScheduler.run_batch`
so rollout and the production inference server share one code path so rollout and the production inference server share one code path
- :class:`RolloutRunner` orchestrates generation + scoring with a - :class:`RolloutRunner` orchestrates generation + scoring with a
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)`` step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
@@ -20,7 +20,7 @@ from typing import Dict, List, Optional, Tuple
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.inference.core.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
@dataclass(kw_only=True) @dataclass(kw_only=True)
@@ -101,7 +101,7 @@ class RolloutGenerator:
"""Pure generation + decoding for a group of responses per prompt. """Pure generation + decoding for a group of responses per prompt.
Delegates the prefill/decode loop to Delegates the prefill/decode loop to
:meth:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`, :meth:`~astrai.inference.scheduler.InferenceScheduler.run_batch`,
which uses a real KV cache (no O() recompute). Has no dependency which uses a real KV cache (no O() recompute). Has no dependency
on any reward model; can be reused in isolation for offline on any reward model; can be reused in isolation for offline
generation, qualitative sampling, or eval pipelines. generation, qualitative sampling, or eval pipelines.
+1 -7
View File
@@ -2,7 +2,7 @@
import math import math
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Dict, List from typing import List
from torch.optim.lr_scheduler import LRScheduler from torch.optim.lr_scheduler import LRScheduler
@@ -20,12 +20,6 @@ class BaseScheduler(LRScheduler, ABC):
"""Calculate the current learning rate.""" """Calculate the current learning rate."""
raise NotImplementedError raise NotImplementedError
def state_dict(self) -> Dict[str, Any]:
return super().state_dict()
def load_state_dict(self, state_dict: Dict[str, Any]):
super().load_state_dict(state_dict)
class SchedulerFactory(BaseFactory["BaseScheduler"]): class SchedulerFactory(BaseFactory["BaseScheduler"]):
"""Factory class for creating learning rate schedulers. """Factory class for creating learning rate schedulers.
+3 -16
View File
@@ -1,6 +1,6 @@
"""Training strategy implementations with factory pattern.""" """Training strategy implementations with factory pattern."""
from abc import ABC, abstractmethod from abc import ABC
from typing import Callable, Dict, List, Optional, TypedDict, Union from typing import Callable, Dict, List, Optional, TypedDict, Union
import torch import torch
@@ -184,10 +184,9 @@ class BaseStrategy(ABC):
self.executor = kwargs.pop("executor", None) self.executor = kwargs.pop("executor", None)
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01) self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
self._moe_metrics: Dict[str, float] = {} self._moe_metrics: Dict[str, float] = {}
self.extra_kwargs = kwargs self.strategy_kwargs = kwargs
self._rollout_runner = None self._rollout_runner = None
@abstractmethod
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
"""Compute loss for the given batch. """Compute loss for the given batch.
@@ -197,7 +196,7 @@ class BaseStrategy(ABC):
Returns: Returns:
Computed loss tensor Computed loss tensor
""" """
raise NotImplementedError return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
return self._normalize_output(self.compute_loss(batch)) return self._normalize_output(self.compute_loss(batch))
@@ -328,9 +327,6 @@ class SEQStrategy(BaseStrategy):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
self.label_smoothing = label_smoothing self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
input_ids, target_ids = batch["input_ids"], batch["target_ids"] input_ids, target_ids = batch["input_ids"], batch["target_ids"]
@@ -369,9 +365,6 @@ class SFTStrategy(BaseStrategy):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
self.label_smoothing = label_smoothing self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
input_ids, target_ids, position_ids, loss_mask = ( input_ids, target_ids, position_ids, loss_mask = (
@@ -426,9 +419,6 @@ class DPOStrategy(BaseStrategy):
self.beta = beta self.beta = beta
self.reduction = reduction self.reduction = reduction
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"] chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
@@ -553,9 +543,6 @@ class GRPOStrategy(BaseStrategy):
if state_dict is not None: if state_dict is not None:
self.old_model.load_state_dict(state_dict) self.old_model.load_state_dict(state_dict)
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
prompts = batch["prompts"] prompts = batch["prompts"]
+13 -10
View File
@@ -3,6 +3,7 @@ import logging
import os import os
import sys import sys
import time import time
from functools import partial
from pathlib import Path from pathlib import Path
from typing import IO, Callable, List, Optional, Protocol, runtime_checkable from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
@@ -17,15 +18,11 @@ from astrai.parallel import only_on_rank
from astrai.parallel.setup import get_current_device from astrai.parallel.setup import get_current_device
from astrai.serialization import Checkpoint from astrai.serialization import Checkpoint
from astrai.trainer.metric_util import ( from astrai.trainer.metric_util import (
ctx_get_dead_expert_fraction,
ctx_get_grad_norm, ctx_get_grad_norm,
ctx_get_grad_snr, ctx_get_grad_snr,
ctx_get_load_imbalance_max,
ctx_get_load_imbalance_mean,
ctx_get_loss, ctx_get_loss,
ctx_get_lr, ctx_get_lr,
ctx_get_moe_aux_loss, ctx_get_moe_metric,
ctx_get_router_entropy,
ctx_get_val_loss, ctx_get_val_loss,
) )
from astrai.trainer.train_context import TrainContext from astrai.trainer.train_context import TrainContext
@@ -119,6 +116,8 @@ class GradientCheckpointingCallback(TrainCallback):
del module._original_forward del module._original_forward
def on_train_begin(self, context: TrainContext): def on_train_begin(self, context: TrainContext):
if not self.modules:
return
context.model.apply(self._enable) context.model.apply(self._enable)
logger.info("Gradient checkpointing enabled") logger.info("Gradient checkpointing enabled")
@@ -262,11 +261,15 @@ class MetricCallback(TrainCallback):
"val_loss": ctx_get_val_loss, "val_loss": ctx_get_val_loss,
"grad_norm": ctx_get_grad_norm, "grad_norm": ctx_get_grad_norm,
"grad_snr": ctx_get_grad_snr, "grad_snr": ctx_get_grad_snr,
"moe_aux_loss": ctx_get_moe_aux_loss, "moe_aux_loss": partial(ctx_get_moe_metric, key="aux_loss"),
"router_entropy": ctx_get_router_entropy, "router_entropy": partial(ctx_get_moe_metric, key="router_entropy"),
"dead_expert_fraction": ctx_get_dead_expert_fraction, "dead_expert_fraction": partial(
"load_imbalance_mean": ctx_get_load_imbalance_mean, ctx_get_moe_metric, key="dead_expert_fraction"
"load_imbalance_max": ctx_get_load_imbalance_max, ),
"load_imbalance_mean": partial(
ctx_get_moe_metric, key="load_imbalance_mean"
),
"load_imbalance_max": partial(ctx_get_moe_metric, key="load_imbalance_max"),
} }
def _metrics(self, context: TrainContext, names): def _metrics(self, context: TrainContext, names):
+196 -152
View File
@@ -8,14 +8,21 @@ import torch
import torch.nn as nn import torch.nn as nn
from torch.utils.data import DataLoader, random_split from torch.utils.data import DataLoader, random_split
from astrai.config.model_config import ConfigFactory
from astrai.config.train_config import TrainConfig from astrai.config.train_config import TrainConfig
from astrai.dataset import RDSampler from astrai.dataset import RDSampler
from astrai.inference.core.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.model.components.lora import inject_lora from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory, create_ref_model from astrai.parallel.executor import BaseExecutor, ExecutorFactory, create_ref_model
from astrai.parallel.setup import get_current_device, get_rank, get_world_size from astrai.parallel.setup import get_current_device, get_rank, get_world_size
from astrai.protocols import OptimizerProtocol, SchedulerProtocol from astrai.protocols import OptimizerProtocol, SchedulerProtocol
from astrai.serialization import Checkpoint, load_json from astrai.serialization import (
Checkpoint,
adapt_config,
convert_hf_weights,
load_json,
looks_like_hf_state_dict,
)
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
from astrai.trainer.metric_util import GradSNRTracker from astrai.trainer.metric_util import GradSNRTracker
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
@@ -66,6 +73,15 @@ class TrainContext:
) )
@dataclass
class _PreloadedState:
model_config: dict = field(default_factory=dict)
state_dict: Optional[dict] = None
epoch: int = 0
consumed_samples: int = 0
checkpoint: Optional[Checkpoint] = None
class TrainContextBuilder: class TrainContextBuilder:
def __init__( def __init__(
self, self,
@@ -81,213 +97,241 @@ class TrainContextBuilder:
return self return self
def build(self) -> TrainContext: def build(self) -> TrainContext:
cfg = self.config # Resolve persisted state.
device = get_current_device() preloaded_state = self._load_preloaded_state()
executor = ExecutorFactory.create( # Build the core training components and restore their persisted state.
executor = self._create_executor()
context = self._create_context(preloaded_state, executor)
self._prepare_model(context, executor, preloaded_state)
self._restore_optimizer_state(context)
# Resolve datasets.
train_dataset, val_dataset = self._get_datasets()
self._create_dataloaders(context, train_dataset, val_dataset)
# Strategies depend on the prepared model; online rollout depends on both.
strategy_kwargs = self._create_strategy(context, executor)
self._configure_rollout(context, strategy_kwargs)
return context
def _create_executor(self) -> BaseExecutor:
cfg = self.config
return ExecutorFactory.create(
cfg.parallel_mode, cfg.parallel_mode,
grad_accum_steps=cfg.grad_accum_steps, grad_accum_steps=cfg.grad_accum_steps,
**cfg.executor_kwargs, **cfg.executor_kwargs,
) )
model_config = {} def _load_preloaded_state(self) -> _PreloadedState:
cfg = self.config
state = _PreloadedState(
epoch=cfg.start_epoch,
consumed_samples=cfg.start_samples * get_world_size(),
)
if self._param_path: if self._param_path:
config_path = Path(self._param_path) / "config.json" config_path = Path(self._param_path) / "config.json"
if config_path.exists(): if config_path.exists():
model_config = load_json(config_path) state.model_config = adapt_config(load_json(config_path))
preloaded_state_dict = None
preloaded_epoch = cfg.start_epoch
preloaded_consumed = cfg.start_samples * get_world_size()
preloaded_checkpoint = None
if self._param_path:
checkpoint = Checkpoint.load_any(self._param_path) checkpoint = Checkpoint.load_any(self._param_path)
if checkpoint is not None: if checkpoint is not None:
preloaded_state_dict = checkpoint.state_dict
if checkpoint.config: if checkpoint.config:
model_config = checkpoint.config checkpoint.config = adapt_config(checkpoint.config)
if checkpoint.state_dict and looks_like_hf_state_dict(
checkpoint.state_dict
):
checkpoint.state_dict = convert_hf_weights(
checkpoint.state_dict,
ConfigFactory.load(checkpoint.config or state.model_config),
)
state.state_dict = checkpoint.state_dict
state.model_config = checkpoint.config or state.model_config
if self._resume: if self._resume:
preloaded_epoch = checkpoint.epoch state.epoch = checkpoint.epoch
per_step = ( per_step = (
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
) )
preloaded_consumed = ( state.consumed_samples = (
checkpoint.consumed_samples // per_step checkpoint.consumed_samples // per_step * per_step
) * per_step )
preloaded_checkpoint = checkpoint state.checkpoint = checkpoint
if not state.model_config:
model = cfg.model_fn()
if hasattr(model, "config"):
state.model_config = model.config.to_dict()
return state
if not model_config and hasattr(cfg.model_fn(), "config"): def _create_context(
model_config = cfg.model_fn().config.to_dict() self, state: _PreloadedState, executor: BaseExecutor
) -> TrainContext:
return TrainContext(
world_size=get_world_size(),
rank=get_rank(),
config=self.config,
model_config=state.model_config,
executor=executor,
epoch=state.epoch,
consumed_samples=state.consumed_samples,
checkpoint=state.checkpoint,
)
def _before_wrap(m): def _prepare_model(
m = m.to(device=device) self, context: TrainContext, executor: BaseExecutor, state: _PreloadedState
) -> None:
cfg = self.config
device = get_current_device()
def before_wrap(model):
model = model.to(device=device)
if cfg.lora is not None: if cfg.lora is not None:
inject_lora( inject_lora(
m, model,
r=cfg.lora.r, r=cfg.lora.r,
alpha=cfg.lora.alpha, alpha=cfg.lora.alpha,
target_modules=set(cfg.lora.target_modules), target_modules=set(cfg.lora.target_modules),
) )
if preloaded_state_dict is not None: if state.state_dict is not None:
m.load_state_dict(preloaded_state_dict, strict=False) model.load_state_dict(state.state_dict, strict=False)
return m return model
def _after_wrap(m): def after_wrap(model):
if cfg.compile_mode is not None: if cfg.compile_mode is not None:
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode) logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
m = torch.compile(m, mode=cfg.compile_mode) model = torch.compile(model, mode=cfg.compile_mode)
return m return model
context = TrainContext(
world_size=get_world_size(),
rank=get_rank(),
config=cfg,
model_config=model_config,
executor=executor,
epoch=preloaded_epoch,
consumed_samples=preloaded_consumed,
checkpoint=preloaded_checkpoint,
)
context.model, context.optimizer, context.scheduler = executor.prepare( context.model, context.optimizer, context.scheduler = executor.prepare(
cfg.model_fn, cfg.model_fn,
cfg.optimizer_fn, cfg.optimizer_fn,
cfg.scheduler_fn, cfg.scheduler_fn,
before_wrap=_before_wrap, before_wrap=before_wrap,
after_wrap=_after_wrap, after_wrap=after_wrap,
) )
train_dataset = cfg.dataset def _get_datasets(self):
val_dataset = cfg.val_dataset cfg = self.config
if cfg.val_dataset is not None or cfg.val_split is None:
return cfg.dataset, cfg.val_dataset
n_val = max(1, int(len(cfg.dataset) * cfg.val_split))
generator = torch.Generator().manual_seed(cfg.random_seed)
return random_split(
cfg.dataset, [len(cfg.dataset) - n_val, n_val], generator=generator
)
if val_dataset is None and cfg.val_split is not None: def _create_dataloaders(
n_total = len(cfg.dataset) self, context: TrainContext, train_dataset, val_dataset
n_val = max(1, int(n_total * cfg.val_split)) ) -> None:
n_train = n_total - n_val sampler_offset = context.consumed_samples // context.world_size
generator = torch.Generator().manual_seed(cfg.random_seed) if self._resume and sampler_offset > 0:
train_dataset, val_dataset = random_split( samples_per_replica = (
cfg.dataset, [n_train, n_val], generator=generator len(train_dataset) + context.world_size - 1
) // context.world_size
if samples_per_replica > 0:
context.epoch = sampler_offset // samples_per_replica
context.dataloader = self._create_dataloader(
train_dataset, context.epoch, sampler_offset
)
if val_dataset is not None:
context.val_dataloader = self._create_dataloader(
val_dataset, 0, 0, shuffle=False
) )
sampler_offset = context.consumed_samples // context.world_size def _create_dataloader(
self, dataset, epoch: int, start_iter: int, shuffle: bool = True
if self._resume and sampler_offset > 0: ):
offset = context.world_size - 1 cfg = self.config
num_samples_per_replica = (
len(train_dataset) + offset
) // context.world_size
if num_samples_per_replica > 0:
context.epoch = sampler_offset // num_samples_per_replica
sampler = RDSampler( sampler = RDSampler(
data_source=train_dataset, dataset,
start_epoch=context.epoch, start_epoch=epoch,
start_iter=sampler_offset, start_iter=start_iter,
seed=cfg.random_seed, seed=cfg.random_seed,
shuffle=shuffle,
) )
context.dataloader = DataLoader( loader_kwargs = dict(
train_dataset, dataset=dataset,
batch_size=cfg.batch_per_device, batch_size=cfg.batch_per_device,
sampler=sampler, sampler=sampler,
num_workers=cfg.num_workers, num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory, pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor,
collate_fn=cfg.collate_fn, collate_fn=cfg.collate_fn,
) )
# PyTorch rejects prefetch_factor/persistent_workers when workers=0.
if val_dataset is not None: if cfg.num_workers > 0:
val_sampler = RDSampler( loader_kwargs["persistent_workers"] = cfg.persistent_workers
data_source=val_dataset, if cfg.prefetch_factor is not None:
start_epoch=0, loader_kwargs["prefetch_factor"] = cfg.prefetch_factor
start_iter=0, return DataLoader(
seed=cfg.random_seed, **loader_kwargs,
shuffle=False,
)
context.val_dataloader = DataLoader(
val_dataset,
batch_size=cfg.batch_per_device,
sampler=val_sampler,
num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor,
collate_fn=cfg.collate_fn,
)
if context.checkpoint and context.checkpoint.extra:
extra = context.checkpoint.extra
for name in ("optimizer", "scheduler"):
if name in extra:
obj = getattr(context, name, None)
if obj is not None:
obj.load_state_dict(extra[name])
strategy_kwargs = dict(cfg.extra_kwargs)
strategy_kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
needs_ref = cfg.strategy in (
"dpo",
"grpo",
"online_grpo",
"online_dpo",
) )
needs_old = cfg.strategy in ("grpo", "online_grpo")
if needs_ref: def _restore_optimizer_state(self, context: TrainContext) -> None:
strategy_kwargs["ref_model"] = create_ref_model( if context.checkpoint and context.checkpoint.extra:
cfg.model_fn, executor=executor, model=context.model, device=device for name in ("optimizer", "scheduler"):
if (
name in context.checkpoint.extra
and getattr(context, name, None) is not None
):
getattr(context, name).load_state_dict(
context.checkpoint.extra[name]
)
def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict:
cfg = self.config
kwargs = dict(cfg.strategy_kwargs)
kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"):
kwargs["ref_model"] = create_ref_model(
cfg.model_fn,
executor=executor,
model=context.model,
device=get_current_device(),
) )
if cfg.strategy in ("grpo", "online_grpo"):
if needs_old: kwargs["old_model"] = create_ref_model(
strategy_kwargs["old_model"] = create_ref_model( cfg.model_fn,
cfg.model_fn, executor=executor, model=context.model, device=device executor=executor,
model=context.model,
device=get_current_device(),
) )
context.strategy = StrategyFactory.create( context.strategy = StrategyFactory.create(
cfg.strategy, cfg.strategy,
model=context.model, model=context.model,
device=device, device=get_current_device(),
executor=executor, executor=executor,
**strategy_kwargs, **kwargs,
) )
return kwargs
# Enable online rollout when the train_type is an ``online_*`` variant. def _configure_rollout(self, context: TrainContext, strategy_kwargs: dict) -> None:
is_online = cfg.strategy.startswith("online_") cfg = self.config
if is_online: if not cfg.strategy.startswith("online_"):
if not context.strategy.supports_online(): return
raise ValueError( if not context.strategy.supports_online():
f"Strategy '{cfg.strategy}' does not support online rollout" raise ValueError(
) f"Strategy '{cfg.strategy}' does not support online rollout"
if cfg.reward_model_fn is None:
raise ValueError("reward_model_fn is required for online RL strategies")
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
reward_model = cfg.reward_model_fn()
group_size = strategy_kwargs.get("group_size", 1)
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
max_seq_len = getattr(context.model.config, "max_position_embeddings", None)
scheduler = InferenceScheduler(
model=context.model,
tokenizer=tokenizer,
max_batch_size=rollout_batch_size,
max_seq_len=max_seq_len,
) )
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
generator = RolloutGenerator( group_size = strategy_kwargs.get("group_size", 1)
scheduler=scheduler, scheduler = InferenceScheduler(
tokenizer=tokenizer, model=context.model,
max_tokens=cfg.rollout_max_tokens, tokenizer=tokenizer,
group_size=group_size, max_batch_size=group_size * max(1, cfg.batch_per_device),
temperature=cfg.rollout_temperature, max_seq_len=getattr(context.model.config, "max_position_embeddings", None),
top_k=cfg.rollout_top_k, )
top_p=cfg.rollout_top_p, generator = RolloutGenerator(
) scheduler=scheduler,
runner = RolloutRunner( tokenizer=tokenizer,
max_tokens=cfg.rollout_max_tokens,
group_size=group_size,
temperature=cfg.rollout_temperature,
top_k=cfg.rollout_top_k,
top_p=cfg.rollout_top_p,
)
context.strategy.set_rollout_runner(
RolloutRunner(
generator=generator, generator=generator,
reward_model=reward_model, reward_model=cfg.reward_model_fn(),
rollout_interval=cfg.rollout_interval, rollout_interval=cfg.rollout_interval,
) )
context.strategy.set_rollout_runner(runner) )
return context
+37 -3
View File
@@ -48,10 +48,44 @@ set(TORCH_LIBS
set(CMAKE_CUDA_ARCHITECTURES "${ASTRAI_CUDA_ARCH}") set(CMAKE_CUDA_ARCHITECTURES "${ASTRAI_CUDA_ARCH}")
set(KERNELS attn_decode attn_prefill attn_paged_decode attn_paged_prefill rotary_emb) # Kernel registry parallel lists of module names (.so / pybind names,
# globally unique across families) and their per-family source paths under
# kernels/. `loader.py` auto-discovers the .so files in astrai/extension/lib/,
# so this CMake registry is the single place to register a new kernel.
#
# FP8 MMA instructions require sm_89+. Keep the target out of the build on
# older architectures instead of instantiating templates that cannot compile.
# The remaining kernels are still useful on sm_80+ (including sm_86).
set(KERNEL_NAMES
attn_decode
attn_prefill
attn_paged_decode
attn_paged_prefill
rotary_emb
)
set(KERNEL_SRCS
attention/decode.cu
attention/prefill.cu
attention/paged_decode.cu
attention/paged_prefill.cu
rotary/rotary_emb.cu
)
foreach(name ${KERNELS}) if(ASTRAI_CUDA_ARCH GREATER_EQUAL 89)
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${name}.cu") list(APPEND KERNEL_NAMES fp8_ops)
list(APPEND KERNEL_SRCS fp8/ops.cu)
else()
message(WARNING
"FP8 operator disabled: ASTRAI_CUDA_ARCH=${ASTRAI_CUDA_ARCH} "
"requires compute capability 89 or newer")
endif()
list(LENGTH KERNEL_NAMES _kernel_count)
math(EXPR _kernel_last "${_kernel_count} - 1")
foreach(i RANGE ${_kernel_last})
list(GET KERNEL_NAMES ${i} name)
list(GET KERNEL_SRCS ${i} src)
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${src}")
target_compile_definitions(${name} PRIVATE TORCH_EXTENSION_NAME=${name}) target_compile_definitions(${name} PRIVATE TORCH_EXTENSION_NAME=${name})
+1 -1
View File
@@ -1,2 +1,2 @@
# Source directory for CUDA kernels — build-time only. # Source directory for CUDA kernels — build-time only.
# Compiled .so files live in astrAI/_ext/. # Compiled .so files live in astrai/extension/lib/ (see csrc/CMakeLists.txt).
+100
View File
@@ -0,0 +1,100 @@
#pragma once
// Pure POD header
namespace astrai {
namespace attention {
// Tensor layout for Q/K/V tensors passed to attention kernels.
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
enum TensorLayout : int {
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
};
// Split-KV workspace cap: max decode splits per (batch, q_head).
constexpr int MAX_SPLITS = 32;
// Paged-prefill host Q-tile granularity in q rows: one q_tile_to_index unit
// covers this many query rows of one request. Must match Q_TILE_ROWS in
// astrai/inference/workspace.py, which builds the device-side tile maps.
constexpr int HOST_Q_TILE_ROWS = 64;
// Unified attention params covering BOTH addressing modes:
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
// Each kernel selects the addressing via a KVSource policy (see
// layout_policies.cuh); a given call only touches the fields of one mode, so
// this is a POD shared by both paths rather than two parallel structs that
// drift out of sync.
//
// Pointer/flag members carry default member initializers: the pointers gate
// optional paths via null checks (new_k_ptr, mask, o_part, ...), so a stack
// `AttentionParams<T> p;` left partially packed must never see garbage
// non-null pointers or a garbage use_mask/causal_offset — that class of bug
// reads through wild addresses. NSDMI keeps the struct an aggregate (C++17)
// and trivially copyable, so `= {}`, memcpy-style packing and by-value kernel
// params all behave exactly as before.
template<typename T, typename AT = float>
struct AttentionParams {
// Shape
int batch;
int q_head;
int kv_head;
int head_dim;
int q_len; // Per-request in contiguous mode; total_q in paged mode.
int kv_len; // Contiguous mode; paged mode uses kv_indptr.
// Attention behavior
float scale;
// -1 = non-causal; >=0 = absolute position of first Q token
int causal_offset = -1;
int use_mask = 0;
// pointers
const T* __restrict__ q_ptr = nullptr;
const T* __restrict__ k_ptr = nullptr;
const T* __restrict__ v_ptr = nullptr;
const T* __restrict__ new_k_ptr = nullptr;
const T* __restrict__ new_v_ptr = nullptr;
T* __restrict__ o_ptr = nullptr;
const bool* __restrict__ mask = nullptr;
// strides
int q_b_stride;
int q_h_stride;
int q_l_stride;
int q_d_stride;
int kv_b_stride;
int kv_h_stride;
int kv_l_stride;
int kv_d_stride;
int new_kv_b_stride;
int new_kv_h_stride;
int mask_b_stride;
int mask_h_stride;
int mask_l_stride;
// Paged K/V addressing
const int* __restrict__ req_to_token = nullptr; // [num_reqs, max_context_len]
const int* __restrict__ req_pool_indices = nullptr; // [batch]
const int* __restrict__ kv_indptr = nullptr; // [batch + 1]
const int* __restrict__ qo_indptr = nullptr; // [batch + 1] or nullptr for decode
const int* __restrict__ q_tile_to_batch = nullptr; // [num_q_tiles], prefill only
const int* __restrict__ q_tile_to_index = nullptr; // [num_q_tiles], prefill only
int num_q_tiles;
int max_context_len; // req_to_token stride (dim 1)
// Decode split-KV workspace
int num_splits;
AT* __restrict__ o_part = nullptr;
AT* __restrict__ ml_part = nullptr;
};
} // namespace attention
} // namespace astrai
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh" #include "dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_decode( torch::Tensor attn_decode(
torch::Tensor q, torch::Tensor q,
@@ -22,15 +24,22 @@ torch::Tensor attn_decode(
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O; auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr(); p.o_ptr = (bf16*)O_view.data_ptr();
if (o_part_buf.has_value() && ml_part_buf.has_value() if (o_part_buf.has_value() && ml_part_buf.has_value()
&& o_part_buf->defined() && ml_part_buf->defined()) { && o_part_buf->defined() && ml_part_buf->defined()) {
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32"); TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32"); TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim; int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
int64_t ml_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * 2;
TORCH_CHECK(o_part_buf->numel() >= o_needed, TORCH_CHECK(o_part_buf->numel() >= o_needed,
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel()); "o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
TORCH_CHECK(ml_part_buf->numel() >= ml_needed,
"ml_part_buf too small: need ", ml_needed, " got ", ml_part_buf->numel());
TORCH_CHECK(o_part_buf->is_cuda() && ml_part_buf->is_cuda(),
"split buffers must be CUDA tensors");
TORCH_CHECK(o_part_buf->is_contiguous() && ml_part_buf->is_contiguous(),
"split buffers must be contiguous");
p.o_part = (float*)o_part_buf->data_ptr(); p.o_part = (float*)o_part_buf->data_ptr();
p.ml_part = (float*)ml_part_buf->data_ptr(); p.ml_part = (float*)ml_part_buf->data_ptr();
} else { } else {
@@ -1,9 +1,13 @@
#pragma once #pragma once
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include <float.h> #include <float.h>
#include "attn_common.h" #include "common.h"
#include "attn_kv_source.cuh" #include "layout_policies.cuh"
#include "attn_warp_utils.cuh" #include "../common/reduce.cuh"
namespace astrai {
namespace attention {
constexpr int DC_CHUNK = 64; constexpr int DC_CHUNK = 64;
// Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified // Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified
@@ -27,9 +31,9 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// Q: [batch, q_head, q_len=1, head_dim] — stride-based // Q: [batch, q_head, q_len=1, head_dim] — stride-based
float q_reg[8]; float q_reg[8];
int q_off = KV::q_decode_base(p, batch, q_head) int q_off = KV::q_decode_base(p, batch, q_head)
+ lane * hd_per_thread * p.q_stride_d; + lane * hd_per_thread * p.q_d_stride;
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); q_reg[i] = __bfloat162float(p.q_ptr[q_off + i * p.q_d_stride]);
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride; int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
@@ -57,7 +61,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int s = i / p.head_dim; int s = i / p.head_dim;
int d_dim = i % p.head_dim; int d_dim = i % p.head_dim;
int kc = chunk_start + s; int kc = chunk_start + s;
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); KVAddr a = KV::template decode_addr<1>(
p, kctx, batch, kv_head, kc, d_dim, true, true);
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f; k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f; v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
} }
@@ -138,6 +143,9 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
} }
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_stride_d; int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
p.o[o_off] = __float2bfloat16(acc * inv); p.o_ptr[o_off] = __float2bfloat16(acc * inv);
} }
} // namespace attention
} // namespace astrai
@@ -1,10 +1,12 @@
#pragma once #pragma once
#include <cfloat> #include <cfloat>
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "common.h"
#include "attn_kv_source.cuh" #include "layout_policies.cuh"
#include "attn_mma_utils.cuh" #include "mma_utils.cuh"
#include "attn_warp_utils.cuh"
namespace astrai {
namespace attention {
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified // Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified
// across contiguous and paged (SGLang flat-pool) K/V via the KV template // across contiguous and paged (SGLang flat-pool) K/V via the KV template
@@ -48,7 +50,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
const int qrb = gid + 8; const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G; const bool va = qra < G, vb = qrb < G;
unsigned Qa[Traits::KD][4]; unsigned Qa[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d, load_q_mma_frags<Traits::KD>(p.q_ptr + q_base, p.q_h_stride, p.q_d_stride,
qra, qrb, va, vb, tid4, Qa); qra, qrb, va, vb, tid4, Qa);
float Oacc[Traits::DN8][4]; float Oacc[Traits::DN8][4];
@@ -73,12 +75,15 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < seq_len; bool valid = kc < seq_len;
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); // All GQA passes consume new K/V directly. Only the first pass
// persists it, so no cross-block synchronization is required.
KVAddr a = KV::template decode_addr<Traits::VEC>(
p, kctx, batch, kv_head, kc, d, valid, pass == 0);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid); astrai::cp_async_16(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid); astrai::cp_async_16(&dV[off], a.v, a.valid);
} }
cp_async_commit(); astrai::cp_async_commit_group();
}; };
// ---- Multi-stage cp.async pipeline ---- // ---- Multi-stage cp.async pipeline ----
@@ -107,10 +112,11 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
int maxc = IsCausal ? KV::decode_attend_len(p, batch) : seq_len; int maxc = IsCausal ? KV::decode_attend_len(p, batch) : seq_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc, mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
0, 0, 0, 0,
p.mask_b_stride, 0, 0, p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
batch, 0, batch, q_head0 + gid, q_head0 + gid + 8,
p.mask, p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane); va, vb,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc); mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
}; };
@@ -121,7 +127,10 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
load_tile(ti_begin + i, i); load_tile(ti_begin + i, i);
for (int it = 0; it < ntiles; it++) { for (int it = 0; it < ntiles; it++) {
cp_async_wait_group<STAGES - 1>(); if (it + 1 == ntiles)
astrai::cp_async_wait_group<0>();
else
astrai::cp_async_wait_group<STAGES - 1>();
__syncwarp(); __syncwarp();
process_tile(it, it & (STAGES - 1)); process_tile(it, it & (STAGES - 1));
__syncwarp(); __syncwarp();
@@ -132,7 +141,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
// Fewer tiles than stages: load all, wait for all, process. // Fewer tiles than stages: load all, wait for all, process.
for (int i = 0; i < ntiles; i++) for (int i = 0; i < ntiles; i++)
load_tile(ti_begin + i, i); load_tile(ti_begin + i, i);
cp_async_wait_group<0>(); astrai::cp_async_wait_all();
__syncwarp(); __syncwarp();
for (int it = 0; it < ntiles; it++) for (int it = 0; it < ntiles; it++)
process_tile(it, it); process_tile(it, it);
@@ -174,3 +183,6 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
} }
} }
} }
} // namespace attention
} // namespace astrai
@@ -3,22 +3,24 @@
// No torch dependency; pure CUDA. // No torch dependency; pure CUDA.
// //
// The paged and contiguous kernels are unified by the KVSource policy // The paged and contiguous kernels are unified by the KVSource policy
// (ContigKV / PagedKV from attn_kv_source.cuh), so each launcher struct // (ContigKV / PagedKV from layout_policies.cuh), so each launcher struct
// below is templated on KV and the paged dispatch is just the same launcher // below is templated on KV and the paged dispatch is just the same launcher
// instantiated with PagedKV. Only the grid/split math differs, and that is // instantiated with PagedKV. Only the grid/split math differs, and that is
// covered by KV::host_q_len / KV::host_kv_len. // covered by KV::host_q_len / KV::host_kv_len.
#include <cuda_runtime.h> #include <cuda_runtime.h>
#include <algorithm> #include <algorithm>
#include "attn_warp_utils.cuh" #include "layout_policies.cuh"
#include "attn_kv_source.cuh" #include "prefill_split_q.cuh"
#include "attn_prefill_split_q.cuh" #include "decode_split_kv.cuh"
#include "attn_decode_split_kv.cuh"
#ifndef ASTRAI_NO_MMA #ifndef ASTRAI_NO_MMA
#include "attn_prefill_split_q_mma.cuh" #include "prefill_split_q_mma.cuh"
#include "attn_decode_split_kv_mma.cuh" #include "decode_split_kv_mma.cuh"
#endif #endif
namespace astrai {
namespace attention {
// Split-KV: compute number of splits to fill all SMs for small-batch decode. // Split-KV: compute number of splits to fill all SMs for small-batch decode.
// Caps splits so each split processes at least `min_tiles_per_split` tiles, // Caps splits so each split processes at least `min_tiles_per_split` tiles,
// avoiding excessive loop/prologue overhead when tiles are small. // avoiding excessive loop/prologue overhead when tiles are small.
@@ -59,32 +61,62 @@ inline int compute_num_splits(int base_blocks, int tiles_total,
// ====================================================================== // ======================================================================
#ifndef ASTRAI_NO_MMA #ifndef ASTRAI_NO_MMA
template <typename KV> template <int BC_>
struct PrefillKernelConfig {
static constexpr int BC = BC_;
static constexpr int WARPS = 4;
static constexpr int STAGES = 2;
};
// Compile-time configuration map shared by contiguous and paged prefill.
// Unsupported head dimensions intentionally have no mapping.
template <int HEAD_DIM, bool IsCausal>
struct PrefillConfigMap;
template <> struct PrefillConfigMap<32, false> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<32, true> : PrefillKernelConfig<64> {};
template <> struct PrefillConfigMap<64, false> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<64, true> : PrefillKernelConfig<64> {};
template <> struct PrefillConfigMap<128, false> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<128, true> : PrefillKernelConfig<32> {};
template <> struct PrefillConfigMap<256, false> : PrefillKernelConfig<16> {};
template <> struct PrefillConfigMap<256, true> : PrefillKernelConfig<16> {};
template <typename QSchedule, typename KV>
struct PrefillLauncherMMA { struct PrefillLauncherMMA {
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) { static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
constexpr int WARPS = 4; using Config = PrefillConfigMap<HEAD_DIM, IsCausal>;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16; using Traits = KernelTraits<HEAD_DIM, Config::BC, Config::WARPS, Config::STAGES>;
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>; // GQA head packing: HB = min(G, WARPS) q-heads of one kv-head group
int q_len = KV::host_q_len(p); // share each block's K/V stream (~HB× less global K/V traffic).
dim3 grid((q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), // Each head gets WPH = WARPS/HB 16-row chunks per block, so per-head
p.q_head, p.batch); // rows drop from 64 to BR*WPH while total mma work per K/V byte is
// unchanged. G=1 (MHA) reproduces the historical grid exactly.
const int G = p.q_head / p.kv_head;
const int HB = std::min(G, Config::WARPS);
const int WPH = Config::WARPS / HB;
constexpr int BR = Traits::BR;
dim3 grid(QSchedule::packed_grid_x(p, BR * WPH),
p.kv_head * ((G + HB - 1) / HB),
QSchedule::host_grid_batch(p));
dim3 block(Traits::NUM_THREADS); dim3 block(Traits::NUM_THREADS);
attn_prefill_split_q_mma_kernel<Traits, KV, IsCausal, HasMask> attn_prefill_split_q_mma_kernel<Traits, QSchedule, KV, IsCausal, HasMask>
<<<grid, block, 0, stream>>>(p); <<<grid, block, 0, stream>>>(p);
} }
}; };
#endif #endif
template <typename KV> template <typename QSchedule, typename KV>
struct PrefillLauncherScalar { struct PrefillLauncherScalar {
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) { static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
constexpr int G = 8, ROWS = 32, P_BC = 32; constexpr int G = (HEAD_DIM == 32) ? 4 : 8, ROWS = 64, P_BC = 32;
int q_len = KV::host_q_len(p); dim3 grid(QSchedule::host_q_blocks(p, ROWS), p.q_head,
dim3 grid((q_len + ROWS - 1) / ROWS, p.q_head, p.batch); QSchedule::host_grid_batch(p));
dim3 block(G, ROWS); dim3 block(G, ROWS);
attn_prefill_split_q_kernel_t<HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask> attn_prefill_split_q_kernel_t<HEAD_DIM, QSchedule, KV, G, ROWS, P_BC,
IsCausal, HasMask>
<<<grid, block, 0, stream>>>(p); <<<grid, block, 0, stream>>>(p);
} }
}; };
@@ -95,12 +127,14 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t strea
bool has_mask = (p.use_mask && p.mask); bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA #ifndef ASTRAI_NO_MMA
using Launcher = PrefillLauncherMMA<DenseQSchedule, ContigKV>;
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherMMA<ContigKV>::template launch, Launcher::template launch,
HEAD_DIM, p, stream); HEAD_DIM, p, stream);
#else #else
using Launcher = PrefillLauncherScalar<DenseQSchedule, ContigKV>;
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherScalar<ContigKV>::template launch, Launcher::template launch,
HEAD_DIM, p, stream); HEAD_DIM, p, stream);
#endif #endif
} }
@@ -111,12 +145,14 @@ static inline void dispatch_paged_prefill(AttentionParams<bf16>& p, cudaStream_t
bool has_mask = (p.use_mask && p.mask); bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA #ifndef ASTRAI_NO_MMA
using Launcher = PrefillLauncherMMA<PackedQSchedule, PagedKV>;
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherMMA<PagedKV>::template launch, Launcher::template launch,
HEAD_DIM, p, stream); HEAD_DIM, p, stream);
#else #else
using Launcher = PrefillLauncherScalar<PackedQSchedule, PagedKV>;
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherScalar<PagedKV>::template launch, Launcher::template launch,
HEAD_DIM, p, stream); HEAD_DIM, p, stream);
#endif #endif
} }
@@ -133,7 +169,7 @@ static inline void dispatch_paged_prefill(AttentionParams<bf16>& p, cudaStream_t
template <typename KV> template <typename KV>
struct DecodeLauncherMMA { struct DecodeLauncherMMA {
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) { static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
int G = p.q_head / p.kv_head; int G = p.q_head / p.kv_head;
constexpr int MAX_G = 16; constexpr int MAX_G = 16;
int num_passes = (G + MAX_G - 1) / MAX_G; int num_passes = (G + MAX_G - 1) / MAX_G;
@@ -153,14 +189,19 @@ struct DecodeLauncherMMA {
template <typename KV> template <typename KV>
struct DecodeLauncherScalar { struct DecodeLauncherScalar {
template <int HEAD_DIM, bool IsCausal, bool HasMask> template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) { static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
int kv_len = KV::host_kv_len(p); int kv_len = KV::host_kv_len(p);
int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK; int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = 2 * DC_CHUNK * p.head_dim * sizeof(bf16); size_t smem = 2 * DC_CHUNK * p.head_dim * sizeof(bf16);
int group_size = p.q_head / p.kv_head;
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, g); dim3 block(32, g);
cudaFuncSetAttribute(
attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
smem);
attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask> attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask>
<<<grid, block, smem, stream>>>(p); <<<grid, block, smem, stream>>>(p);
} }
@@ -170,16 +211,15 @@ template <int HEAD_DIM>
static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) { static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
bool is_causal = (p.causal_offset >= 0); bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask); bool has_mask = (p.use_mask && p.mask);
int group_size = p.q_head / p.kv_head;
#ifndef ASTRAI_NO_MMA #ifndef ASTRAI_NO_MMA
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherMMA<ContigKV>::template launch, DecodeLauncherMMA<ContigKV>::template launch,
HEAD_DIM, p, group_size, stream); HEAD_DIM, p, stream);
#else #else
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherScalar<ContigKV>::template launch, DecodeLauncherScalar<ContigKV>::template launch,
HEAD_DIM, p, group_size, stream); HEAD_DIM, p, stream);
#endif #endif
attn_decode_combine_kernel<ContigKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p); attn_decode_combine_kernel<ContigKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
@@ -189,17 +229,19 @@ template <int HEAD_DIM>
static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t stream) { static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
bool is_causal = (p.causal_offset >= 0); bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask); bool has_mask = (p.use_mask && p.mask);
int group_size = p.q_head / p.kv_head;
#ifndef ASTRAI_NO_MMA #ifndef ASTRAI_NO_MMA
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherMMA<PagedKV>::template launch, DecodeLauncherMMA<PagedKV>::template launch,
HEAD_DIM, p, group_size, stream); HEAD_DIM, p, stream);
#else #else
DISPATCH_CAUSAL_MASK(is_causal, has_mask, DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherScalar<PagedKV>::template launch, DecodeLauncherScalar<PagedKV>::template launch,
HEAD_DIM, p, group_size, stream); HEAD_DIM, p, stream);
#endif #endif
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p); attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
} }
} // namespace attention
} // namespace astrai
@@ -2,10 +2,7 @@
#include <float.h> #include <float.h>
#include <torch/extension.h> #include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h> #include <c10/cuda/CUDAGuard.h>
#include "attn_common.h" #include "common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax. // Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
// Usage: DISPATCH_HEAD_DIM(hd, fn, args...) // Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
@@ -21,6 +18,11 @@ using bf16 = __nv_bfloat16;
" (supported: 32, 64, 128, 256)"); \ " (supported: 32, 64, 128, 256)"); \
} }
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
// The split kernel unconditionally writes every (batch, q_head, split) slot it // The split kernel unconditionally writes every (batch, q_head, split) slot it
// owns — including empty split ranges, which store m = -FLT_MAX so the combine // owns — including empty split ranges, which store m = -FLT_MAX so the combine
// skips them. Allocators are therefore left uninitialized (torch::empty); the // skips them. Allocators are therefore left uninitialized (torch::empty); the
@@ -42,10 +44,10 @@ inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
p.q_head = (int)q.size(1); p.q_head = (int)q.size(1);
p.q_len = (int)q.size(2); p.q_len = (int)q.size(2);
p.head_dim = (int)q.size(3); p.head_dim = (int)q.size(3);
p.q_stride_b = (int)q.stride(0); p.q_b_stride = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1); p.q_h_stride = (int)q.stride(1);
p.q_stride_l = (int)q.stride(2); p.q_l_stride = (int)q.stride(2);
p.q_stride_d = (int)q.stride(3); p.q_d_stride = (int)q.stride(3);
} }
// ---- Shared mask packing ---- // ---- Shared mask packing ----
@@ -63,17 +65,17 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
if (m.dim() == 2) { if (m.dim() == 2) {
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
} else if (m.dim() == 3) { } else if (m.dim() == 3) {
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch"); TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1); p.mask_l_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
} else if (m.dim() == 4) { } else if (m.dim() == 4) {
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch"); TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch");
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1); p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2); p.mask_l_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
} else { } else {
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D"); TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
} }
@@ -82,7 +84,7 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
p.mask = nullptr; p.mask = nullptr;
p.mask_b_stride = 0; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
} }
} }
@@ -106,28 +108,33 @@ inline void attn_pack_params(
TORCH_CHECK(v.dtype() == torch::kBFloat16); TORCH_CHECK(v.dtype() == torch::kBFloat16);
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes"); TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D"); TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
extract_q_dims_and_strides(q, layout, p); extract_q_dims_and_strides(q, layout, p);
if (layout == BLHD) k = k.transpose(1, 2), v = v.transpose(1, 2); if (layout == BLHD) k = k.transpose(1, 2), v = v.transpose(1, 2);
p.kv_head = (int)k.size(1); p.kv_head = (int)k.size(1);
p.kv_len = (int)k.size(2); p.kv_len = (int)k.size(2);
TORCH_CHECK(p.q_head % p.kv_head == 0,
"q_head must be divisible by kv_head");
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q"); TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
TORCH_CHECK(q.stride(3) == 1 && k.stride(3) == 1 && v.stride(3) == 1,
"Q/K/V head_dim must be contiguous");
p.kv_stride_b = (int)k.stride(0); p.kv_b_stride = (int)k.stride(0);
p.kv_stride_h = (int)k.stride(1); p.kv_h_stride = (int)k.stride(1);
p.kv_stride_l = (int)k.stride(2); p.kv_l_stride = (int)k.stride(2);
p.kv_stride_d = (int)k.stride(3); p.kv_d_stride = (int)k.stride(3);
p.causal_offset = (int)causal_offset; p.causal_offset = (int)causal_offset;
p.use_mask = mask.has_value() ? 1 : 0; p.use_mask = mask.has_value() ? 1 : 0;
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim); p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.q = (const T*)q.data_ptr(); p.q_ptr = (const T*)q.data_ptr();
p.k = (const T*)k.data_ptr(); p.k_ptr = (const T*)k.data_ptr();
p.v = (const T*)v.data_ptr(); p.v_ptr = (const T*)v.data_ptr();
p.o = nullptr; p.new_k_ptr = nullptr;
p.new_v_ptr = nullptr;
p.o_ptr = nullptr;
p.o_part = nullptr; p.o_part = nullptr;
p.ml_part = nullptr; p.ml_part = nullptr;
@@ -145,7 +152,8 @@ inline void attn_pack_paged_decode_params(
torch::Tensor req_to_token, torch::Tensor req_to_token,
torch::Tensor req_pool_indices, torch::Tensor req_pool_indices,
torch::Tensor kv_indptr, torch::Tensor kv_indptr,
int64_t max_seq_len, const c10::optional<torch::Tensor>& new_k,
const c10::optional<torch::Tensor>& new_v,
c10::optional<torch::Tensor> mask, c10::optional<torch::Tensor> mask,
int64_t causal_offset, int64_t causal_offset,
double scale, double scale,
@@ -158,8 +166,9 @@ inline void attn_pack_paged_decode_params(
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16"); TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16"); TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16"); TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64"); TORCH_CHECK(req_to_token.dtype() == torch::kInt32, "req_to_token must be int32");
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64"); TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
"req_pool_indices must be int32");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32"); TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match"); TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]"); TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
@@ -170,24 +179,50 @@ inline void attn_pack_paged_decode_params(
p.head_dim = (int)q.size(2); p.head_dim = (int)q.size(2);
p.kv_head = (int)k_cache.size(1); p.kv_head = (int)k_cache.size(1);
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch"); TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
TORCH_CHECK(q.stride(2) == 1 && k_cache.stride(2) == 1 && v_cache.stride(2) == 1,
"Q/K/V head_dim must be contiguous");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32"); TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head"); TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
p.q_stride_l = (int)q.stride(0); p.q_l_stride = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1); p.q_h_stride = (int)q.stride(1);
p.q_stride_d = (int)q.stride(2); p.q_d_stride = (int)q.stride(2);
p.k_cache = (const T*)k_cache.data_ptr(); p.k_ptr = (const T*)k_cache.data_ptr();
p.v_cache = (const T*)v_cache.data_ptr(); p.v_ptr = (const T*)v_cache.data_ptr();
p.q = (const T*)q.data_ptr(); p.q_ptr = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int64_t>(); p.req_to_token = req_to_token.data_ptr<int>();
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>(); p.req_pool_indices = req_pool_indices.data_ptr<int>();
p.kv_indptr = kv_indptr.data_ptr<int>(); p.kv_indptr = kv_indptr.data_ptr<int>();
p.qo_indptr = nullptr; p.qo_indptr = nullptr;
p.max_context_len = (int)req_to_token.size(1); p.max_context_len = (int)req_to_token.size(1);
p.max_seq_len = (int)max_seq_len;
p.total_q = p.batch; // decode: 1 Q token per request TORCH_CHECK(new_k.has_value() == new_v.has_value(),
p.max_q_len = 1; "new_k and new_v must be provided together");
if (new_k.has_value()) {
auto nk = new_k.value();
auto nv = new_v.value();
TORCH_CHECK(nk.is_cuda() && nv.is_cuda(), "new K/V must be CUDA tensors");
TORCH_CHECK(nk.dtype() == torch::kBFloat16 && nv.dtype() == torch::kBFloat16,
"new K/V must be bf16");
TORCH_CHECK(nk.dim() == 3 && nv.dim() == 3,
"new K/V must be 3D [batch, kv_head, head_dim]");
TORCH_CHECK(nk.sizes() == nv.sizes(), "new K and V must have identical shapes");
TORCH_CHECK(nk.strides() == nv.strides(),
"new K and V must have identical strides");
TORCH_CHECK(nk.size(0) == p.batch && nk.size(1) == p.kv_head
&& nk.size(2) == p.head_dim, "new K/V shape mismatch");
TORCH_CHECK(nk.stride(2) == 1 && nv.stride(2) == 1,
"new K/V head_dim must be contiguous");
p.new_k_ptr = (const T*)nk.data_ptr();
p.new_v_ptr = (const T*)nv.data_ptr();
p.new_kv_b_stride = (int)nk.stride(0);
p.new_kv_h_stride = (int)nk.stride(1);
} else {
p.new_k_ptr = nullptr;
p.new_v_ptr = nullptr;
p.new_kv_b_stride = p.new_kv_h_stride = 0;
}
p.causal_offset = (int)causal_offset; p.causal_offset = (int)causal_offset;
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0; p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
@@ -199,16 +234,16 @@ inline void attn_pack_paged_decode_params(
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch"); TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
p.mask = m.data_ptr<bool>(); p.mask = m.data_ptr<bool>();
} else { } else {
p.mask = nullptr; p.mask = nullptr;
p.mask_b_stride = 0; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
} }
p.o = nullptr; p.o_ptr = nullptr;
p.o_part = nullptr; p.o_part = nullptr;
p.ml_part = nullptr; p.ml_part = nullptr;
} }
@@ -225,8 +260,9 @@ inline void attn_pack_paged_prefill_params(
torch::Tensor req_pool_indices, torch::Tensor req_pool_indices,
torch::Tensor kv_indptr, torch::Tensor kv_indptr,
torch::Tensor qo_indptr, torch::Tensor qo_indptr,
torch::Tensor q_tile_to_batch,
torch::Tensor q_tile_to_index,
c10::optional<torch::Tensor> mask, c10::optional<torch::Tensor> mask,
int64_t max_q_len,
int64_t causal_offset, int64_t causal_offset,
double scale, double scale,
AttentionParams<T>& p AttentionParams<T>& p
@@ -236,44 +272,57 @@ inline void attn_pack_paged_prefill_params(
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda()); TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda()); TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda());
TORCH_CHECK(kv_indptr.is_cuda() && qo_indptr.is_cuda()); TORCH_CHECK(kv_indptr.is_cuda() && qo_indptr.is_cuda());
TORCH_CHECK(q_tile_to_batch.is_cuda() && q_tile_to_index.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16"); TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16"); TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16"); TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64"); TORCH_CHECK(req_to_token.dtype() == torch::kInt32, "req_to_token must be int32");
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64"); TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
"req_pool_indices must be int32");
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32"); TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32"); TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32");
TORCH_CHECK(q_tile_to_batch.dtype() == torch::kInt32,
"q_tile_to_batch must be int32");
TORCH_CHECK(q_tile_to_index.dtype() == torch::kInt32,
"q_tile_to_index must be int32");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match"); TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]"); TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
TORCH_CHECK(q.dim() == 3, "q must be 3D [total_q, q_head, head_dim]"); TORCH_CHECK(q.dim() == 3, "q must be 3D [total_q, q_head, head_dim]");
p.q_head = (int)q.size(1); p.q_head = (int)q.size(1);
p.head_dim = (int)q.size(2); p.head_dim = (int)q.size(2);
p.q_len = (int)q.size(0);
p.kv_head = (int)k_cache.size(1); p.kv_head = (int)k_cache.size(1);
p.batch = (int)req_pool_indices.size(0); p.batch = (int)req_pool_indices.size(0);
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch"); TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
TORCH_CHECK(q.stride(2) == 1 && k_cache.stride(2) == 1 && v_cache.stride(2) == 1,
"Q/K/V head_dim must be contiguous");
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16"); TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head"); TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
TORCH_CHECK(kv_indptr.size(0) == p.batch + 1, "kv_indptr must be [batch+1]"); TORCH_CHECK(kv_indptr.size(0) == p.batch + 1, "kv_indptr must be [batch+1]");
TORCH_CHECK(qo_indptr.size(0) == p.batch + 1, "qo_indptr must be [batch+1]"); TORCH_CHECK(qo_indptr.size(0) == p.batch + 1, "qo_indptr must be [batch+1]");
TORCH_CHECK(q_tile_to_batch.dim() == 1 && q_tile_to_index.dim() == 1,
"Q tile mappings must be 1D");
TORCH_CHECK(q_tile_to_batch.size(0) == q_tile_to_index.size(0),
"Q tile mappings must have equal length");
p.q_stride_l = (int)q.stride(0); p.q_l_stride = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1); p.q_h_stride = (int)q.stride(1);
p.q_stride_d = (int)q.stride(2); p.q_d_stride = (int)q.stride(2);
p.k_cache = (const T*)k_cache.data_ptr(); p.k_ptr = (const T*)k_cache.data_ptr();
p.v_cache = (const T*)v_cache.data_ptr(); p.v_ptr = (const T*)v_cache.data_ptr();
p.q = (const T*)q.data_ptr(); p.new_k_ptr = nullptr;
p.req_to_token = req_to_token.data_ptr<int64_t>(); p.new_v_ptr = nullptr;
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>(); p.q_ptr = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int>();
p.req_pool_indices = req_pool_indices.data_ptr<int>();
p.kv_indptr = kv_indptr.data_ptr<int>(); p.kv_indptr = kv_indptr.data_ptr<int>();
p.qo_indptr = qo_indptr.data_ptr<int>(); p.qo_indptr = qo_indptr.data_ptr<int>();
p.q_tile_to_batch = q_tile_to_batch.data_ptr<int>();
p.q_tile_to_index = q_tile_to_index.data_ptr<int>();
p.num_q_tiles = (int)q_tile_to_batch.size(0);
p.max_context_len = (int)req_to_token.size(1); p.max_context_len = (int)req_to_token.size(1);
p.total_q = (int)q.size(0); // prefill: flattened Q across all requests
p.max_q_len = (int)max_q_len;
// max_seq_len is unused by the prefill path (decode uses it for split
// computation); fill with max_q_len only to keep the POD struct defined.
p.max_seq_len = p.max_q_len;
p.causal_offset = (int)causal_offset; p.causal_offset = (int)causal_offset;
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0; p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
@@ -285,14 +334,14 @@ inline void attn_pack_paged_prefill_params(
TORCH_CHECK(m.size(1) <= p.max_context_len, "mask kv_len mismatch"); TORCH_CHECK(m.size(1) <= p.max_context_len, "mask kv_len mismatch");
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
} else if (m.dim() == 4) { } else if (m.dim() == 4) {
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_head, "mask head mismatch"); TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_head, "mask head mismatch");
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.max_q_len, "mask q_len mismatch"); TORCH_CHECK(m.size(2) > 0 && m.size(2) <= p.q_len, "mask q_len mismatch");
TORCH_CHECK(m.size(3) <= p.max_context_len, "mask kv_len mismatch"); TORCH_CHECK(m.size(3) <= p.max_context_len, "mask kv_len mismatch");
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1); p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2); p.mask_l_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
} else { } else {
TORCH_CHECK(false, "mask must be 2D or 4D"); TORCH_CHECK(false, "mask must be 2D or 4D");
} }
@@ -301,11 +350,14 @@ inline void attn_pack_paged_prefill_params(
p.mask = nullptr; p.mask = nullptr;
p.mask_b_stride = 0; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
} }
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim); p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.o = nullptr; p.o_ptr = nullptr;
p.o_part = nullptr; p.o_part = nullptr;
p.ml_part = nullptr; p.ml_part = nullptr;
} }
} // namespace attention
} // namespace astrai
+292
View File
@@ -0,0 +1,292 @@
#pragma once
#include <cuda_bf16.h>
#include "common.h"
// ============================================================================
// Attention layout policies keep Q scheduling independent from K/V storage.
// DenseQSchedule / PackedQSchedule map blocks to Q tiles; ContigKV / PagedKV
// resolve logical K/V positions to physical addresses. This lets the shared
// kernels compose Q layout and K/V storage without coupling the two concerns.
//
// ContigKV: K/V are dense [batch, kv_head, kv_len, head_dim] tensors.
// Params fields used: k, v, kv_stride_*, kv_len, q_len,
// q_b_stride, causal_offset.
// PagedKV: K/V live in a flat pool [size, kv_head, head_dim] indexed via
// req_to_token. Params fields used: k_cache, v_cache,
// req_to_token, req_pool_indices, kv_indptr, qo_indptr,
// max_context_len, q_l_stride.
//
// Addressing state that is constant across a whole kernel invocation for one
// (batch, kv_head) pair is captured once by make_ctx<HEAD_DIM>() and passed
// to kv_addr, so the load loops never redo the hoistable base computation
// (e.g. the req_pool_indices global read) element-by-element.
// ============================================================================
#define HOST_FORCEINLINE static __host__ __forceinline__
#define DEVICE_FORCEINLINE static __device__ __forceinline__
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
// ============================================================================
// Q scheduling policies
//
// Map CUDA blocks to request-local Q tiles independently of K/V storage.
// Dense tensors encode the request in blockIdx.z; packed ragged tensors use
// a compact precomputed work map indexed by blockIdx.x.
// ============================================================================
struct DenseQSchedule {
HOST_FORCEINLINE int host_q_blocks(
const AttentionParams<bf16>& p, int rows) {
return (p.q_len + rows - 1) / rows;
}
HOST_FORCEINLINE int host_grid_batch(
const AttentionParams<bf16>& p) {
return p.batch;
}
DEVICE_FORCEINLINE void map_block(
const AttentionParams<bf16>&, int& batch, int& q_tile) {
batch = blockIdx.z;
q_tile = blockIdx.x;
}
// GQA-packed prefill mapping: HB q-heads of one kv-head group share a
// block's K/V stream, each head owning `rows` = BR*WPH consecutive q rows
// per block. Dense tensors tile q_len directly, one block per range.
HOST_FORCEINLINE int packed_grid_x(
const AttentionParams<bf16>& p, int rows) {
return (p.q_len + rows - 1) / rows;
}
DEVICE_FORCEINLINE void map_packed_block(
const AttentionParams<bf16>&, int rows, int& batch, int& row_base) {
batch = blockIdx.z;
row_base = blockIdx.x * rows;
}
DEVICE_FORCEINLINE int q_len(
const AttentionParams<bf16>& p, int) {
return p.q_len;
}
DEVICE_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_b_stride + q_head * p.q_h_stride;
}
};
struct PackedQSchedule {
HOST_FORCEINLINE int host_q_blocks(
const AttentionParams<bf16>& p, int) {
return p.num_q_tiles;
}
HOST_FORCEINLINE int host_grid_batch(
const AttentionParams<bf16>&) {
return 1;
}
DEVICE_FORCEINLINE void map_block(
const AttentionParams<bf16>& p, int& batch, int& q_tile) {
batch = p.q_tile_to_batch[blockIdx.x];
q_tile = p.q_tile_to_index[blockIdx.x];
}
// GQA-packed prefill mapping: the host tile maps are built in
// HOST_Q_TILE_ROWS granularity, so each host tile splits into
// HOST_Q_TILE_ROWS / rows packed blocks along blockIdx.x.
HOST_FORCEINLINE int packed_grid_x(
const AttentionParams<bf16>& p, int rows) {
return p.num_q_tiles * (HOST_Q_TILE_ROWS / rows);
}
DEVICE_FORCEINLINE void map_packed_block(
const AttentionParams<bf16>& p, int rows, int& batch, int& row_base) {
const int hb = HOST_Q_TILE_ROWS / rows;
const int host_tile = blockIdx.x / hb;
batch = p.q_tile_to_batch[host_tile];
row_base = p.q_tile_to_index[host_tile] * HOST_Q_TILE_ROWS
+ (blockIdx.x - host_tile * hb) * rows;
}
DEVICE_FORCEINLINE int q_len(
const AttentionParams<bf16>& p, int batch) {
return p.qo_indptr[batch + 1] - p.qo_indptr[batch];
}
DEVICE_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return p.qo_indptr[batch] * p.q_l_stride + q_head * p.q_h_stride;
}
};
// Hoisted per-(batch, kv_head) addressing context.
struct KVContext {
int kv_base; // contig: batch*kv_b_stride + kv_head*kv_h_stride
int req_idx; // paged: req_pool_indices[batch]
int64_t rtt_stride; // paged: max_context_len
int64_t pool_stride; // paged: kv_head * HEAD_DIM
int64_t head_off; // paged: kv_head * HEAD_DIM
};
// Per-element K/V global addresses for one (kc, d) position of a K/V tile.
// The pointers are ALWAYS the computed addresses (never nullptr) — callers
// gate on `valid` (cp.async src_size=0, or a guarded scalar deref). `valid`
// starts as "within the request's seq_len"; the paged policy further degrades
// it when req_to_token maps the position to a negative slot (empty padding).
// This matches the original hand-rolled load loops, where the address was
// always formed and the predicate decided whether anything was read.
struct KVAddr {
const void* k;
const void* v;
bool valid;
};
// ---- Contiguous K/V ----
struct ContigKV {
static constexpr bool kPaged = false;
HOST_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
return p.kv_len;
}
// decode: same offset (q_len == 1, so there is no row stride component)
DEVICE_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_b_stride + q_head * p.q_h_stride;
}
DEVICE_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int) {
return p.kv_len;
}
DEVICE_FORCEINLINE int causal_offset(
const AttentionParams<bf16>& p, int, int) {
return p.causal_offset;
}
// decode: exclusive bound of the single query's attend range
DEVICE_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int) {
return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1);
}
template <int HEAD_DIM>
DEVICE_FORCEINLINE KVContext make_ctx(
const AttentionParams<bf16>& p, int batch, int kv_head) {
KVContext c = {};
c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride;
return c;
}
DEVICE_FORCEINLINE int resolve_token(
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
return valid ? kc : -1;
}
DEVICE_FORCEINLINE KVAddr kv_addr_from_token(
const AttentionParams<bf16>& p, const KVContext& c, int token, int d) {
const bool valid = token >= 0;
const int safe_token = valid ? token : 0;
const int64_t gmem_off = (int64_t)c.kv_base
+ (int64_t)safe_token * p.kv_l_stride
+ (int64_t)d * p.kv_d_stride;
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
}
template <int VEC>
DEVICE_FORCEINLINE KVAddr decode_addr(
const AttentionParams<bf16>& p, const KVContext& c,
int, int, int kc, int d, bool valid, bool) {
int token = resolve_token(p, c, kc, valid);
return kv_addr_from_token(p, c, token, d);
}
};
// ---- Paged (SGLang-style flat pool) K/V ----
struct PagedKV {
static constexpr bool kPaged = true;
HOST_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
return p.max_context_len;
}
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
DEVICE_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_l_stride + q_head * p.q_h_stride;
}
DEVICE_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
return p.kv_indptr[batch + 1] - p.kv_indptr[batch];
}
DEVICE_FORCEINLINE int causal_offset(
const AttentionParams<bf16>& p, int batch, int q_len) {
return kv_len(p, batch) - q_len;
}
// decode: the query is the last token, so [0, seq_len) IS its causal range
DEVICE_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
return kv_len(p, batch);
}
template <int HEAD_DIM>
DEVICE_FORCEINLINE KVContext make_ctx(
const AttentionParams<bf16>& p, int batch, int kv_head) {
KVContext c = {};
c.req_idx = p.req_pool_indices[batch];
c.rtt_stride = (int64_t)p.max_context_len;
c.pool_stride = (int64_t)p.kv_head * HEAD_DIM;
c.head_off = (int64_t)kv_head * HEAD_DIM;
return c;
}
DEVICE_FORCEINLINE int resolve_token(
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
return valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : -1;
}
DEVICE_FORCEINLINE KVAddr kv_addr_from_token(
const AttentionParams<bf16>& p, const KVContext& c, int slot, int d) {
const bool valid = slot >= 0;
const int safe_slot = valid ? slot : 0;
const int64_t gmem_off = (int64_t)safe_slot * c.pool_stride + c.head_off + d;
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
}
DEVICE_FORCEINLINE KVAddr new_kv_addr(
const AttentionParams<bf16>& p, int batch, int kv_head, int d) {
const int64_t off = (int64_t)batch * p.new_kv_b_stride
+ (int64_t)kv_head * p.new_kv_h_stride + d;
return {&p.new_k_ptr[off], &p.new_v_ptr[off], true};
}
DEVICE_FORCEINLINE void store_new_kv(
const AttentionParams<bf16>& p, const KVContext& c,
int seq_len, int d, const KVAddr& src) {
int slot = resolve_token(p, c, seq_len - 1, true);
const int64_t off = (int64_t)slot * c.pool_stride + c.head_off + d;
const_cast<bf16*>(p.k_ptr)[off] = *reinterpret_cast<const bf16*>(src.k);
const_cast<bf16*>(p.v_ptr)[off] = *reinterpret_cast<const bf16*>(src.v);
}
template <int VEC>
DEVICE_FORCEINLINE KVAddr decode_addr(
const AttentionParams<bf16>& p, const KVContext& c,
int batch, int kv_head, int kc, int d, bool valid, bool persist) {
if (p.new_k_ptr && valid && kc == kv_len(p, batch) - 1) {
KVAddr src = new_kv_addr(p, batch, kv_head, d);
if (persist) {
#pragma unroll
for (int j = 0; j < VEC; j++) {
KVAddr value = new_kv_addr(p, batch, kv_head, d + j);
store_new_kv(p, c, kc + 1, d + j, value);
}
}
return src;
}
int token = resolve_token(p, c, kc, valid);
return kv_addr_from_token(p, c, token, d);
}
};
} // namespace attention
} // namespace astrai
@@ -3,12 +3,18 @@
#include <cuda_fp16.h> #include <cuda_fp16.h>
#include <cuda_runtime.h> #include <cuda_runtime.h>
#include "../common/cp_async.cuh"
#include "../common/mma.cuh"
// Predicated cp.async (4-operand form) requires CUDA 11.2+. // Predicated cp.async (4-operand form) requires CUDA 11.2+.
// bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA). // bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA).
#if CUDART_VERSION < 11020 #if CUDART_VERSION < 11020
#error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)." #error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)."
#endif #endif
namespace astrai {
namespace attention {
// ============================================================================ // ============================================================================
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle. // KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
// //
@@ -24,10 +30,10 @@ struct KernelTraits {
static constexpr int BR = 16; // Q rows per warp (mma M=16) static constexpr int BR = 16; // Q rows per warp (mma M=16)
// Derived: mma.sync.m16n8k16 tile counts // Derived: mma tile counts from the shared mma_shape (m16n8k16 for bf16)
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides static constexpr int KD = HEAD_DIM / astrai::mma_shape<bf16>::k; // Q/K k-slides
static constexpr int NC8 = BC / 8; // S n-tiles (N=8) static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
static constexpr int KT2 = BC / 16; // P k-tiles (K=16) static constexpr int KT2 = BC / astrai::mma_shape<bf16>::k; // P k-tiles (K=16)
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8) static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
static constexpr int LD = HEAD_DIM; // smem leading dim static constexpr int LD = HEAD_DIM; // smem leading dim
@@ -43,16 +49,7 @@ struct KernelTraits {
// ---- PTX wrappers ---- // ---- PTX wrappers ----
using bf16 = __nv_bfloat16; using bf16 = __nv_bfloat16;
// bf16 mma.sync lives in the shared astrai::mma_sync template (common/mma.cuh).
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
const unsigned* b, const float* c) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
}
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high) // read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
__device__ __forceinline__ unsigned ld2(const bf16* p) { __device__ __forceinline__ unsigned ld2(const bf16* p) {
@@ -73,68 +70,24 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
return *reinterpret_cast<unsigned*>(&v); return *reinterpret_cast<unsigned*>(&v);
} }
// ldmatrix: cooperatively load mma fragments from smem (one instruction per // ldmatrix lives in the shared template (common/mma.cuh):
// 16x16 / 16x8 tile) with the exact register layout mma expects. // `astrai::ldmatrix_x2<bf16>` / `<bf16, /*Trans=*/true>` load the K/V
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) { // fragments with the exact register layout mma expects.
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(a));
}
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity. // XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) { __device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7); return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
} }
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly. // cp.async primitives live in the shared template (common/cp_async.cuh):
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) { // `astrai::cp_async_16` (predicated), `astrai::cp_async_commit_group`,
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr); // `astrai::cp_async_wait_group<N>` / `_wait_all` stage the K/V tiles.
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
:: "r"(smem_addr), "l"(gmem_ptr));
}
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// src_size=0 → no bytes read from src, so out-of-bounds src address is safe.
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr,
bool pred) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
int src_size = pred ? 16 : 0;
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
}
__device__ __forceinline__ void cp_async_commit() {
asm volatile("cp.async.commit_group;");
}
__device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;");
}
template <int N>
__device__ __forceinline__ void cp_async_wait_group() {
asm volatile("cp.async.wait_group %0;" :: "n"(N));
}
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
// Q-load: load query rows directly from global memory into mma A-operand // Q-load: load query rows directly from global memory into mma A-operand
// register layout. One call replaces ~15 duplicated lines in each MMA kernel. // register layout. One call replaces ~15 duplicated lines in each MMA kernel.
// stride_row is p.q_stride_h for decode (q_len=1, G heads) or // stride_row is p.q_h_stride for decode (q_len=1, G heads) or
// p.q_stride_l for prefill (multi-q rows). // p.q_l_stride for prefill (multi-q rows).
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
template <int KD> template <int KD>
__device__ inline void load_q_mma_frags( __device__ inline void load_q_mma_frags(
@@ -180,9 +133,9 @@ __device__ inline void mma_compute_scores(
#pragma unroll #pragma unroll
for (int kt = 0; kt < Traits::KD; kt++) { for (int kt = 0; kt < Traits::KD; kt++) {
unsigned b[2]; unsigned b[2];
ldmatrix_x2(b, &sK[krow_l * Traits::LD astrai::ldmatrix_x2<bf16>(b, &sK[krow_l * Traits::LD
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]); + swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]); astrai::mma_sync<bf16>(Sacc[n8], Qa[kt], b, Sacc[n8]);
} }
} }
} }
@@ -198,9 +151,10 @@ __device__ inline void mma_softmax_tile(
int kv0, int kv0,
int maxc0, int maxc1, int maxc0, int maxc1,
int qrow0, int qrow1, int qrow0, int qrow1,
int mask_b_stride, int mask_h_stride, int mask_q_stride, int mask_b_stride, int mask_h_stride, int mask_l_stride,
int mask_batch, int mask_head, int mask_batch, int mask_head0, int mask_head1,
const bool* __restrict__ mask, const bool* __restrict__ mask,
bool valid0, bool valid1,
float Sacc[Traits::NC8][4], float Sacc[Traits::NC8][4],
float Oacc[Traits::DN8][4], float Oacc[Traits::DN8][4],
float& m0, float& m1, float& m0, float& m1,
@@ -210,16 +164,16 @@ __device__ inline void mma_softmax_tile(
int tid4 = lane & 3; int tid4 = lane & 3;
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX; float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
int mask_base0 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow0 * mask_q_stride; int mask_base0 = mask_batch * mask_b_stride + mask_head0 * mask_h_stride + qrow0 * mask_l_stride;
int mask_base1 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow1 * mask_q_stride; int mask_base1 = mask_batch * mask_b_stride + mask_head1 * mask_h_stride + qrow1 * mask_l_stride;
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < Traits::NC8; n8++) { for (int n8 = 0; n8 < Traits::NC8; n8++) {
int cc = kv0 + n8 * 8 + 2 * tid4; int cc = kv0 + n8 * 8 + 2 * tid4;
int c1 = cc + 1; int c1 = cc + 1;
bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]); bool b0 = !valid0 || (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]); bool b1 = !valid0 || (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]); bool b2 = !valid1 || (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]); bool b3 = !valid1 || (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
float s0 = b0 ? -FLT_MAX : Sacc[n8][0]; float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
float s1 = b1 ? -FLT_MAX : Sacc[n8][1]; float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
float s2 = b2 ? -FLT_MAX : Sacc[n8][2]; float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
@@ -289,9 +243,12 @@ __device__ inline void mma_pv_accumulate(
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
unsigned b[2]; unsigned b[2];
ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD astrai::ldmatrix_x2<bf16, true>(b, &sV[vrow_l * Traits::LD
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]); + swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]); astrai::mma_sync<bf16>(Oacc[dn8], Pa, b, Oacc[dn8]);
} }
} }
} }
} // namespace attention
} // namespace astrai
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh" #include "dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_paged_decode( torch::Tensor attn_paged_decode(
torch::Tensor q, torch::Tensor q,
@@ -8,7 +10,8 @@ torch::Tensor attn_paged_decode(
torch::Tensor req_to_token, torch::Tensor req_to_token,
torch::Tensor req_pool_indices, torch::Tensor req_pool_indices,
torch::Tensor kv_indptr, torch::Tensor kv_indptr,
int64_t max_seq_len, c10::optional<torch::Tensor> new_k,
c10::optional<torch::Tensor> new_v,
c10::optional<torch::Tensor> mask, c10::optional<torch::Tensor> mask,
int64_t causal_offset, int64_t causal_offset,
double scale, double scale,
@@ -22,29 +25,39 @@ torch::Tensor attn_paged_decode(
AttentionParams<bf16> p; AttentionParams<bf16> p;
attn_pack_paged_decode_params(q, k_cache, v_cache, attn_pack_paged_decode_params(q, k_cache, v_cache,
req_to_token, req_pool_indices, kv_indptr, req_to_token, req_pool_indices, kv_indptr,
max_seq_len, mask, causal_offset, scale, p); new_k, new_v,
mask, causal_offset, scale, p);
torch::Tensor O; torch::Tensor O;
if (out_buf.has_value() && out_buf->defined()) { if (out_buf.has_value() && out_buf->defined()) {
TORCH_CHECK(out_buf->dtype() == q.dtype(), "out_buf dtype must match q"); TORCH_CHECK(out_buf->dtype() == q.dtype(), "out_buf dtype must match q");
TORCH_CHECK(out_buf->is_cuda() && out_buf->is_contiguous(),
"out_buf must be a contiguous CUDA tensor");
TORCH_CHECK(out_buf->size(0) >= q.size(0), "out_buf batch too small"); TORCH_CHECK(out_buf->size(0) >= q.size(0), "out_buf batch too small");
TORCH_CHECK(out_buf->size(1) >= q.size(1), "out_buf heads too small"); TORCH_CHECK(out_buf->size(1) == q.size(1), "out_buf heads must match q");
TORCH_CHECK(out_buf->size(2) >= q.size(2), "out_buf head_dim too small"); TORCH_CHECK(out_buf->size(2) == q.size(2), "out_buf head_dim must match q");
O = out_buf.value().slice(0, 0, q.size(0)) TORCH_CHECK(q.is_contiguous(),
.slice(1, 0, q.size(1)) "q must be contiguous when out_buf is provided");
.slice(2, 0, q.size(2)); O = out_buf.value().slice(0, 0, q.size(0));
} else { } else {
O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options()); O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
} }
p.o = (bf16*)O.data_ptr(); p.o_ptr = (bf16*)O.data_ptr();
if (o_part_buf.has_value() && ml_part_buf.has_value() if (o_part_buf.has_value() && ml_part_buf.has_value()
&& o_part_buf->defined() && ml_part_buf->defined()) { && o_part_buf->defined() && ml_part_buf->defined()) {
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32"); TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32"); TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim; int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
int64_t ml_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * 2;
TORCH_CHECK(o_part_buf->numel() >= o_needed, TORCH_CHECK(o_part_buf->numel() >= o_needed,
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel()); "o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
TORCH_CHECK(ml_part_buf->numel() >= ml_needed,
"ml_part_buf too small: need ", ml_needed, " got ", ml_part_buf->numel());
TORCH_CHECK(o_part_buf->is_cuda() && ml_part_buf->is_cuda(),
"split buffers must be CUDA tensors");
TORCH_CHECK(o_part_buf->is_contiguous() && ml_part_buf->is_contiguous(),
"split buffers must be contiguous");
p.o_part = (float*)o_part_buf->data_ptr(); p.o_part = (float*)o_part_buf->data_ptr();
p.ml_part = (float*)ml_part_buf->data_ptr(); p.ml_part = (float*)ml_part_buf->data_ptr();
} else { } else {
@@ -63,7 +76,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("req_to_token"), py::arg("req_to_token"),
py::arg("req_pool_indices"), py::arg("req_pool_indices"),
py::arg("kv_indptr"), py::arg("kv_indptr"),
py::arg("max_seq_len"), py::arg("new_k") = py::none(),
py::arg("new_v") = py::none(),
py::arg("mask") = py::none(), py::arg("mask") = py::none(),
py::arg("causal_offset") = -1, py::arg("causal_offset") = -1,
py::arg("scale") = 0.0, py::arg("scale") = 0.0,
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh" #include "dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_paged_prefill( torch::Tensor attn_paged_prefill(
torch::Tensor q, torch::Tensor q,
@@ -9,8 +11,9 @@ torch::Tensor attn_paged_prefill(
torch::Tensor req_pool_indices, torch::Tensor req_pool_indices,
torch::Tensor kv_indptr, torch::Tensor kv_indptr,
torch::Tensor qo_indptr, torch::Tensor qo_indptr,
torch::Tensor q_tile_to_batch,
torch::Tensor q_tile_to_index,
c10::optional<torch::Tensor> mask, c10::optional<torch::Tensor> mask,
int64_t max_q_len,
int64_t causal_offset, int64_t causal_offset,
double scale double scale
) { ) {
@@ -19,12 +22,13 @@ torch::Tensor attn_paged_prefill(
AttentionParams<bf16> p; AttentionParams<bf16> p;
attn_pack_paged_prefill_params(q, k_cache, v_cache, attn_pack_paged_prefill_params(q, k_cache, v_cache,
req_to_token, req_pool_indices, req_to_token, req_pool_indices,
kv_indptr, qo_indptr, mask, kv_indptr, qo_indptr,
max_q_len, causal_offset, scale, p); q_tile_to_batch, q_tile_to_index, mask,
causal_offset, scale, p);
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options()); auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
p.o = (bf16*)O.data_ptr(); p.o_ptr = (bf16*)O.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream); DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream);
C10_CUDA_CHECK(cudaGetLastError()); C10_CUDA_CHECK(cudaGetLastError());
@@ -40,8 +44,9 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("req_pool_indices"), py::arg("req_pool_indices"),
py::arg("kv_indptr"), py::arg("kv_indptr"),
py::arg("qo_indptr"), py::arg("qo_indptr"),
py::arg("q_tile_to_batch"),
py::arg("q_tile_to_index"),
py::arg("mask") = py::none(), py::arg("mask") = py::none(),
py::arg("max_q_len"),
py::arg("causal_offset") = -1, py::arg("causal_offset") = -1,
py::arg("scale") = 0.0, py::arg("scale") = 0.0,
"SGLang-style paged prefill: flat KV pool + ragged batch."); "SGLang-style paged prefill: flat KV pool + ragged batch.");
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh" #include "dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_prefill( torch::Tensor attn_prefill(
torch::Tensor q, torch::Tensor q,
@@ -19,7 +21,7 @@ torch::Tensor attn_prefill(
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O; auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr(); p.o_ptr = (bf16*)O_view.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream); DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream);
C10_CUDA_CHECK(cudaGetLastError()); C10_CUDA_CHECK(cudaGetLastError());
@@ -1,8 +1,12 @@
#pragma once #pragma once
#include <cfloat> #include <cfloat>
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "common.h"
#include "attn_kv_source.cuh" #include "layout_policies.cuh"
#include "../common/reduce.cuh"
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16; using bf16 = __nv_bfloat16;
@@ -11,14 +15,7 @@ using bf16 = __nv_bfloat16;
// compile-time bools — the compiler eliminates dead branches. // compile-time bools — the compiler eliminates dead branches.
// Unified across contiguous and paged (SGLang flat-pool) K/V via KV. // Unified across contiguous and paged (SGLang flat-pool) K/V via KV.
// Templated on <HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>. // Templated on <HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>.
// group_reduce_sum<G> lives in common/reduce.cuh (astrai::).
template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
#pragma unroll
for (int o = G / 2; o > 0; o >>= 1)
v += __shfl_xor_sync(mask, v, o);
return v;
}
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4 // load 8 contiguous bf16 from (16-byte aligned) smem as one float4
__device__ __forceinline__ void ld8(const bf16* p, float* o) { __device__ __forceinline__ void ld8(const bf16* p, float* o) {
@@ -32,21 +29,23 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) {
} }
} }
template <int HEAD_DIM, typename KV, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask> template <int HEAD_DIM, typename QSchedule, typename KV, int G, int ROWS, int P_BC,
bool IsCausal, bool HasMask>
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) { __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
constexpr int DPT = HEAD_DIM / G; constexpr int DPT = HEAD_DIM / G;
int q_tile = blockIdx.x; int batch, q_tile;
QSchedule::map_block(p, batch, q_tile);
int q_head = blockIdx.y; int q_head = blockIdx.y;
int batch = blockIdx.z;
int gpos = threadIdx.x; // 0..G-1 (which d-chunk) int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
int row = threadIdx.y; // 0..ROWS-1 int row = threadIdx.y; // 0..ROWS-1
int q_row = q_tile * ROWS + row; int q_row = q_tile * ROWS + row;
// Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr). // Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
const int seq_len = KV::kv_len(p, batch); const int seq_len = KV::kv_len(p, batch);
const int q_len = KV::q_len(p, batch); const int q_len = QSchedule::q_len(p, batch);
const int causal_off = KV::causal_offset(p, batch); const int causal_off = KV::causal_offset(p, batch, q_len);
const int kv_head = q_head / (p.q_head / p.kv_head); const int kv_head = q_head / (p.q_head / p.kv_head);
const KVContext kctx = KV::template make_ctx<HEAD_DIM>(p, batch, kv_head); const KVContext kctx = KV::template make_ctx<HEAD_DIM>(p, batch, kv_head);
@@ -54,13 +53,13 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM]; __shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
// Q: stride-based load [batch, q_head, q_len, head_dim] // Q: stride-based load [batch, q_head, q_len, head_dim]
const int q_base = KV::q_base(p, batch, q_head); const int q_base = QSchedule::q_base(p, batch, q_head);
float qreg[DPT]; float qreg[DPT];
if (q_row < q_len) { if (q_row < q_len) {
int q_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; int q_off = q_base + q_row * p.q_l_stride + gpos * DPT * p.q_d_stride;
#pragma unroll #pragma unroll
for (int i = 0; i < DPT; i++) for (int i = 0; i < DPT; i++)
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); qreg[i] = __bfloat162float(p.q_ptr[q_off + i * p.q_d_stride]);
} }
float m = -FLT_MAX, l = 0.0f; float m = -FLT_MAX, l = 0.0f;
@@ -88,7 +87,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
int s = i / HEAD_DIM; int s = i / HEAD_DIM;
int d_dim = i % HEAD_DIM; int d_dim = i % HEAD_DIM;
int kc = kv0 + s; int kc = kv0 + s;
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); int token = KV::resolve_token(p, kctx, kc, true);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f; sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f; sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
} }
@@ -105,7 +105,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
} }
} }
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride; int mask_row_base = mask_batch_base + q_row * p.mask_l_stride;
for (int s = 0; s < lim; s++) { for (int s = 0; s < lim; s++) {
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT; const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
float part = 0.0f; float part = 0.0f;
@@ -145,10 +145,13 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
} }
if (q_row < q_len) { if (q_row < q_len) {
int o_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; int o_off = q_base + q_row * p.q_l_stride + gpos * DPT * p.q_d_stride;
float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f; float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
#pragma unroll #pragma unroll
for (int i = 0; i < DPT; i++) for (int i = 0; i < DPT; i++)
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl); p.o_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl);
} }
} }
} // namespace attention
} // namespace astrai
@@ -1,37 +1,60 @@
#pragma once #pragma once
#include <cfloat> #include <cfloat>
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "common.h"
#include "attn_kv_source.cuh" #include "layout_policies.cuh"
#include "attn_mma_utils.cuh" #include "mma_utils.cuh"
namespace astrai {
namespace attention {
// Tensor-core prefill flash attention (raw mma.sync PTX), unified across // Tensor-core prefill flash attention (raw mma.sync PTX), unified across
// contiguous and paged (SGLang flat-pool) K/V via the KV template parameter. // contiguous and paged (SGLang flat-pool) K/V via the KV template parameter.
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor // One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
// cores via mma.sync.m16n8k16 (f32 accumulate). // cores via mma.sync.m16n8k16 (f32 accumulate).
// //
// GQA head packing (FA2/FA3-style): HB = min(G, WARPS) query heads of one
// kv-head group share a block's K/V tiles, so each K/V element is read from
// global memory once per block instead of once per q head (~HB× less K/V
// traffic). WARPS = WPH × HB: warp w handles head slot w/WPH, chunk w%WPH;
// all warps of a block cover the same token range, keeping the causal sweep
// end block-uniform. G=1 (MHA) degenerates to the unpadded layout.
//
// KV = ContigKV (dense [batch, kv_head, kv_len, head_dim]) or PagedKV // KV = ContigKV (dense [batch, kv_head, kv_len, head_dim]) or PagedKV
// (flat pool + req_to_token, ragged batches via qo_indptr/kv_indptr). // (flat pool + req_to_token, ragged batches via qo_indptr/kv_indptr).
// IsCausal and HasMask are compile-time bools — the compiler eliminates all // IsCausal and HasMask are compile-time bools — the compiler eliminates all
// dead branches in the inner compute loop (FA2-style). // dead branches in the inner compute loop (FA2-style).
// //
// Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>. // Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
template <typename Traits, typename KV, bool IsCausal, bool HasMask> template <typename Traits, typename QSchedule, typename KV, bool IsCausal, bool HasMask>
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) { __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
const int warp = threadIdx.x / 32; const int warp = threadIdx.x / 32;
const int lane = threadIdx.x % 32; const int lane = threadIdx.x % 32;
const int gid = lane >> 2; // 0..7 const int gid = lane >> 2; // 0..7
const int tid4 = lane & 3; // 0..3 const int tid4 = lane & 3; // 0..3
const int q_head = blockIdx.y; const int G = p.q_head / p.kv_head;
const int batch = blockIdx.z; const int HB = min(G, Traits::WARPS); // q heads packed per block
const int kv_head = q_head / (p.q_head / p.kv_head); const int WPH = Traits::WARPS / HB; // 16-row chunks per head
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR; const int BPG = (G + HB - 1) / HB; // blocks per GQA group
const int chunk = warp % WPH;
int batch, row_base;
QSchedule::map_packed_block(p, Traits::BR * WPH, batch, row_base);
const int kv_head = blockIdx.y / BPG;
const int slot = blockIdx.y - kv_head * BPG;
const int head_idx = slot * HB + warp / WPH;
// G % HB tail blocks have idle head slots: clamp to the last head so all
// warps do valid work (cp.async + __syncthreads stay block-uniform) and
// just skip the O store via `active`.
const bool active = head_idx < G;
const int q_head = kv_head * G + min(head_idx, G - 1);
const int qrow0 = row_base + chunk * Traits::BR;
// Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr). // Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
const int seq_len = KV::kv_len(p, batch); const int seq_len = KV::kv_len(p, batch);
const int q_len = KV::q_len(p, batch); const int q_len = QSchedule::q_len(p, batch);
const int causal_off = KV::causal_offset(p, batch); const int causal_off = KV::causal_offset(p, batch, q_len);
const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head); const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head);
// Static shared memory: double-buffered K/V (no sQ — Q goes direct // Static shared memory: double-buffered K/V (no sQ — Q goes direct
@@ -40,12 +63,12 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD]; __shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// Load Q fragments straight from global into mma A-operand layout. // Load Q fragments straight from global into mma A-operand layout.
const int q_base = KV::q_base(p, batch, q_head); const int q_base = QSchedule::q_base(p, batch, q_head);
const int qra = qrow0 + gid; const int qra = qrow0 + gid;
const int qrb = qrow0 + gid + 8; const int qrb = qrow0 + gid + 8;
const bool va = qra < q_len, vb = qrb < q_len; const bool va = qra < q_len, vb = qrb < q_len;
unsigned Qa[Traits::KD][4]; unsigned Qa[Traits::KD][4];
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d, load_q_mma_frags<Traits::KD>(p.q_ptr + q_base, p.q_l_stride, p.q_d_stride,
qra, qrb, va, vb, tid4, Qa); qra, qrb, va, vb, tid4, Qa);
float Oacc[Traits::DN8][4]; float Oacc[Traits::DN8][4];
@@ -58,11 +81,11 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
const int qr0 = qrow0 + gid; const int qr0 = qrow0 + gid;
const int qr1 = qrow0 + gid + 8; const int qr1 = qrow0 + gid + 8;
// Causal tile-skip bounds (dead code when IsCausal == false) // Causal tile-skip bounds (dead code when IsCausal == false).
// max_kv is per-warp (its own 16 rows); block_max_kv is the last row of
// the whole block's range and must be uniform for the shared sweep loop.
const int max_kv = qrow0 + Traits::BR - 1 + causal_off; const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
const int block_max_kv = const int block_max_kv = row_base + WPH * Traits::BR - 1 + causal_off;
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
+ causal_off;
int t_end = tiles - 1; int t_end = tiles - 1;
if constexpr (IsCausal) { if constexpr (IsCausal) {
@@ -81,12 +104,13 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < seq_len; bool valid = kc < seq_len;
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); int token = KV::resolve_token(p, kctx, kc, valid);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid); astrai::cp_async_16(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid); astrai::cp_async_16(&dV[off], a.v, a.valid);
} }
cp_async_commit(); astrai::cp_async_commit_group();
}; };
// ---- Prologue: issue first tile load ---- // ---- Prologue: issue first tile load ----
@@ -96,7 +120,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int buf = ti & 1; int buf = ti & 1;
// Wait for current tile, then publish cross-warp + guard buffer reuse. // Wait for current tile, then publish cross-warp + guard buffer reuse.
cp_async_wait_group<0>(); astrai::cp_async_wait_group<0>();
__syncthreads(); __syncthreads();
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1); if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
@@ -122,9 +146,10 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
: seq_len; : seq_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1, mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
qr0, qr1, qr0, qr1,
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride, p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
batch, q_head, batch, q_head, q_head,
p.mask, p.mask,
va, vb,
Sacc, Oacc, m0, m1, l0, l1, lane); Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc); mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
@@ -134,21 +159,24 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
// ---- write output: packed bf16x2 stores ---- // ---- write output: packed bf16x2 stores ----
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f; float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f; float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
const int o_base = KV::q_base(p, batch, q_head); const int o_base = QSchedule::q_base(p, batch, q_head);
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4; int d = dn8 * 8 + 2 * tid4;
if (qr0 < q_len) { if (active && qr0 < q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0, __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
Oacc[dn8][1] * rl0); Oacc[dn8][1] * rl0);
*reinterpret_cast<__nv_bfloat162*>( *reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v; &p.o_ptr[o_base + qr0 * p.q_l_stride + d * p.q_d_stride]) = v;
} }
if (qr1 < q_len) { if (active && qr1 < q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1, __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
Oacc[dn8][3] * rl1); Oacc[dn8][3] * rl1);
*reinterpret_cast<__nv_bfloat162*>( *reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v; &p.o_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v;
} }
} }
} }
} // namespace attention
} // namespace astrai
-66
View File
@@ -1,66 +0,0 @@
#pragma once
// Tensor layout for Q/K/V tensors passed to attention kernels.
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
enum TensorLayout : int {
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
};
// Unified attention params covering BOTH addressing modes:
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
// Each kernel selects the addressing via a KVSource policy (see
// attn_kv_source.cuh); a given call only touches the fields of one mode, so
// this is a POD shared by both paths rather than two parallel structs that
// drift out of sync.
template<typename T, typename AT = float>
struct AttentionParams {
// ---- shared across all paths ----
int batch;
int q_head;
int kv_head;
int head_dim;
int use_mask;
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
int num_splits;
float scale;
// Q strides (element offsets for each dim — layout-agnostic)
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
// Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
int mask_b_stride; // batch stride
int mask_h_stride; // head stride (0 = broadcast across heads)
int mask_q_stride; // q stride (0 = all q rows share)
const bool* __restrict__ mask;
const T* __restrict__ q;
T* __restrict__ o;
AT* __restrict__ o_part;
AT* __restrict__ ml_part;
// ---- contiguous K/V mode ----
int q_len;
int kv_len;
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
const T* __restrict__ k;
const T* __restrict__ v;
// ---- paged (SGLang flat pool) mode ----
const T* __restrict__ k_cache;
const T* __restrict__ v_cache;
// Indexing
const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len]
const int64_t* __restrict__ req_pool_indices; // [batch]
const int* __restrict__ kv_indptr; // [batch+1]
const int* __restrict__ qo_indptr; // [batch+1] or nullptr (decode)
int max_context_len; // req_to_token stride (dim 1)
int max_seq_len; // max per-request seq_len (host-side, for split computation)
int total_q; // total Q tokens across all requests (host-side, for grid)
int max_q_len; // max per-request q_len (host-side, for prefill grid)
};
-159
View File
@@ -1,159 +0,0 @@
#pragma once
#include <cuda_bf16.h>
#include "attn_common.h"
// ============================================================================
// KVSource policies — the single dimension along which the paged and
// non-paged attention kernels differ. Each kernel is templated on one of
// these (ContigKV / PagedKV) and stays fully generic: the policy owns every
// place where "where does K/V live" and "what is this request's seq_len"
// are answered. All methods are __host__ __device__ so the same policy
// serves both the device kernels (addressing, seq_len) and the host-side
// launchers (grid / split computation).
//
// ContigKV: K/V are dense [batch, kv_head, kv_len, head_dim] tensors.
// Params fields used: k, v, kv_stride_*, kv_len, q_len,
// q_stride_b, causal_offset.
// PagedKV: K/V live in a flat pool [size, kv_head, head_dim] indexed via
// req_to_token. Params fields used: k_cache, v_cache,
// req_to_token, req_pool_indices, kv_indptr, qo_indptr,
// max_context_len, q_stride_l.
//
// Addressing state that is constant across a whole kernel invocation for one
// (batch, kv_head) pair is captured once by make_ctx<HEAD_DIM>() and passed
// to kv_addr, so the load loops never redo the hoistable base computation
// (e.g. the req_pool_indices global read) element-by-element.
// ============================================================================
// Every policy method is static + callable from both host and device code.
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
using bf16 = __nv_bfloat16;
// Hoisted per-(batch, kv_head) addressing context.
struct KVContext {
int kv_base; // contig: batch*kv_stride_b + kv_head*kv_stride_h
int64_t req_idx; // paged: req_pool_indices[batch]
int64_t rtt_stride; // paged: max_context_len
int64_t pool_stride; // paged: kv_head * HEAD_DIM
int64_t head_off; // paged: kv_head * HEAD_DIM
};
// Per-element K/V global addresses for one (kc, d) position of a K/V tile.
// The pointers are ALWAYS the computed addresses (never nullptr) — callers
// gate on `valid` (cp.async src_size=0, or a guarded scalar deref). `valid`
// starts as "within the request's seq_len"; the paged policy further degrades
// it when req_to_token maps the position to a negative slot (empty padding).
// This matches the original hand-rolled load loops, where the address was
// always formed and the predicate decided whether anything was read.
struct KVAddr {
const void* k;
const void* v;
bool valid;
};
// ---- Contiguous K/V ----
struct ContigKV {
static constexpr bool kPaged = false;
// host-side length hooks (grid + split computation in the launchers)
HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams<bf16>& p) {
return p.q_len;
}
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
return p.kv_len;
}
// prefill: element offset of the request's Q rows (kernel adds qrow*q_stride_l)
HOST_DEV_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_b + q_head * p.q_stride_h;
}
// decode: same offset (q_len == 1, so there is no row stride component)
HOST_DEV_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_b + q_head * p.q_stride_h;
}
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
return p.kv_len;
}
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
return p.q_len;
}
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
return p.causal_offset;
}
// decode: exclusive bound of the single query's attend range
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1);
}
template <int HEAD_DIM>
HOST_DEV_FORCEINLINE KVContext make_ctx(
const AttentionParams<bf16>& p, int batch, int kv_head) {
KVContext c = {};
c.kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
return c;
}
HOST_DEV_FORCEINLINE KVAddr kv_addr(
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
const int g_off = c.kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
return {&p.k[g_off], &p.v[g_off], valid};
}
};
// ---- Paged (SGLang-style flat pool) K/V ----
struct PagedKV {
static constexpr bool kPaged = true;
HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams<bf16>& p) {
return p.max_q_len;
}
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
return p.max_context_len;
}
// prefill: Q rows start at qo_indptr[batch] (ragged batch base)
HOST_DEV_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return p.qo_indptr[batch] * p.q_stride_l + q_head * p.q_stride_h;
}
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
HOST_DEV_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_l + q_head * p.q_stride_h;
}
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
return p.kv_indptr[batch + 1] - p.kv_indptr[batch];
}
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
return p.qo_indptr[batch + 1] - p.qo_indptr[batch];
}
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
return kv_len(p, batch) - q_len(p, batch);
}
// decode: the query is the last token, so [0, seq_len) IS its causal range
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
return kv_len(p, batch);
}
template <int HEAD_DIM>
HOST_DEV_FORCEINLINE KVContext make_ctx(
const AttentionParams<bf16>& p, int batch, int kv_head) {
KVContext c = {};
c.req_idx = p.req_pool_indices[batch];
c.rtt_stride = (int64_t)p.max_context_len;
c.pool_stride = (int64_t)p.kv_head * HEAD_DIM;
c.head_off = (int64_t)kv_head * HEAD_DIM;
return c;
}
HOST_DEV_FORCEINLINE KVAddr kv_addr(
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
const int64_t slot = valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : 0;
const bool ok = valid && (slot >= 0);
const int64_t gmem_off = slot * c.pool_stride + c.head_off + d;
return {&p.k_cache[gmem_off], &p.v_cache[gmem_off], ok};
}
};
-13
View File
@@ -1,13 +0,0 @@
#pragma once
#include <cuda_bf16.h>
using bf16 = __nv_bfloat16;
static constexpr int MAX_SPLITS = 32;
__device__ inline float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
+82
View File
@@ -0,0 +1,82 @@
// Shared cp.async primitives — pure CUDA, no torch.
//
// One header for the async-copy pipeline used by both the attention kernels
// (predicated 16-byte K/V tile staging) and the fp8 GEMM (predicated operand
// staging + the fixed-depth wait_group). The emitter is split from its
// policies: cp_async_16_raw owns the single PTX site, and each wrapper states
// one destination contract (generic pointer vs loop-carried shared offset)
// and one predication contract (unconditional vs zero-fill-when-false), so
// call sites never pass a dead `true` predicate or re-convert a carried
// offset. PTX requires wait_group's operand to be an immediate, hence the
// template form below.
#pragma once
#include <cuda_runtime.h>
namespace astrai {
// Raw emitter: read src_size bytes (<= 16) from gmem into the shared
// offset. src_size = 0 reads nothing, so a predicated-off call zero-fills
// its destination without touching the (possibly out-of-range) source.
// BypassL1 selects .cg (L2 only, default) vs .ca (L1 + L2).
template <bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16_raw(unsigned smem_addr,
const void* gmem_ptr,
int src_size) {
if constexpr (BypassL1) {
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
} else {
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
}
}
// Unconditional 16-byte copy to a generic shared pointer.
// `T` is the smem element type; only the destination pointer's type matters.
template <typename T, bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16(T* smem_ptr,
const void* gmem_ptr) {
cp_async_16_raw<BypassL1>(__cvta_generic_to_shared(smem_ptr), gmem_ptr,
16);
}
// Predicated: full copy when `pred`, zero-fill otherwise.
template <typename T, bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16(T* smem_ptr, const void* gmem_ptr,
bool pred) {
cp_async_16_raw<BypassL1>(__cvta_generic_to_shared(smem_ptr), gmem_ptr,
pred ? 16 : 0);
}
// Predicated raw-offset form: the destination is an already-converted
// shared-memory offset (e.g. a loop-carried swizzled stage address), so
// steady-state prefetch sites issue one LDGSTS straight from the register.
template <bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16(unsigned smem_addr,
const void* gmem_ptr, bool pred) {
cp_async_16_raw<BypassL1>(smem_addr, gmem_ptr, pred ? 16 : 0);
}
// Commit all outstanding cp.async ops of this thread as one group.
__device__ __forceinline__ void cp_async_commit_group() {
asm volatile("cp.async.commit_group;");
}
// Wait for every committed group (pipeline drain).
__device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;");
}
// Wait until at most KeepGroups committed groups are still in flight.
// PTX requires an immediate operand; keep it as a template argument so the
// stage policy stays compile-time configurable.
template <int KeepGroups>
__device__ __forceinline__ void cp_async_wait_group() {
static_assert(KeepGroups >= 0 && KeepGroups <= 7,
"cp.async.wait_group supports immediates in [0, 7]");
asm volatile("cp.async.wait_group %0;" :: "n"(KeepGroups));
}
} // namespace astrai
+23
View File
@@ -0,0 +1,23 @@
// Pure-CUDA device helpers shared across kernel families (no torch).
//
// Family-local headers under kernels/<family>/ own their POD params and
// strategy traits; anything cross-cutting (compute-capability checks, device
// constants) lives here.
#pragma once
namespace astrai {
// Compute-capability comparison: is the device at least (major, minor)?
inline bool sm_at_least(int device_major, int device_minor, int major,
int minor) {
return device_major > major ||
(device_major == major && device_minor >= minor);
}
// FP8 tensor-core MMA (`mma.sync.aligned.m16n8k32` with fp8 inputs) exists on
// Ada (sm_89) and Hopper (sm_90+); sm_80 has no fp8 instructions.
inline constexpr int kMinSmForFp8Major = 8;
inline constexpr int kMinSmForFp8Minor = 9;
} // namespace astrai
+165
View File
@@ -0,0 +1,165 @@
// Shared mma.sync wrappers — pure CUDA, no torch.
//
// One template for every tensor-core MMA used by the kernel families. The
// instruction shape follows from the input element type:
// __nv_bfloat16 -> mma.sync.aligned.m16n8k16 (sm_80+), A = 4x b32, B = 2x b32
// __nv_fp8_e4m3/e5m2 -> mma.sync.aligned.m16n8k32 (sm_89+), A = 4x b32, B = 2x b32
// All variants accumulate into fp32: d = a*b + c, with the PTX mnemonic and
// the K dimension differing per type. `d` may alias `c` (in-place accumulate,
// as the FP8 GEMM does).
#pragma once
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <type_traits>
#define DEVICE_FORCEINLINE static __device__ __forceinline__
namespace astrai {
// Compute capability of the current compilation pass: 0 in the host pass,
// the numeric CC (e.g. 890) in device passes where __CUDA_ARCH__ is defined.
// Defined() cannot appear in expressions, so this macro lets mma_sync use
// the arch in a static_assert instead of per-branch #if guards.
#ifndef __CUDA_ARCH__
#define ASTRAI_DEVICE_ARCH 0
#else
#define ASTRAI_DEVICE_ARCH __CUDA_ARCH__
#endif
// Compile-time shape of the MMA instruction for an input element type.
// `min_arch` is the numeric compute capability the instruction requires —
// the single place that encodes the hardware floor for each type.
template <typename InT>
struct mma_shape {
static constexpr int k = 16; // m16n8k16
static constexpr int a_regs = 4; // A fragment: 4x b32
static constexpr int b_regs = 2; // B fragment: 2x b32
static constexpr int min_arch = 800; // bf16 mma.sync, sm_80+
};
template <>
struct mma_shape<__nv_fp8_e4m3> {
static constexpr int k = 32; // m16n8k32
static constexpr int a_regs = 4;
static constexpr int b_regs = 2;
static constexpr int min_arch = 890; // fp8 mma.sync, sm_89+ (Ada/Hopper)
};
template <>
struct mma_shape<__nv_fp8_e5m2> {
static constexpr int k = 32;
static constexpr int a_regs = 4;
static constexpr int b_regs = 2;
static constexpr int min_arch = 890;
};
// d[4] = a[4] x b[2] + c[4], row-major A, col-major B, fp32 accumulator.
// The PTX mnemonic is selected from InT. Building for a compute capability
// below `mma_shape<InT>::min_arch` is a **compile error** — the instruction
// does not exist there, and a silent no-op would produce wrong results.
template <typename InT>
DEVICE_FORCEINLINE void mma_sync(float d[4], const unsigned a[4],
const unsigned b[2],
const float c[4]) {
static_assert(ASTRAI_DEVICE_ARCH == 0 ||
ASTRAI_DEVICE_ARCH >= mma_shape<InT>::min_arch,
"mma_sync: this MMA shape requires a newer compute "
"capability than the build target");
if constexpr (std::is_same_v<InT, __nv_bfloat16>) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
} else if constexpr (std::is_same_v<InT, __nv_fp8_e5m2>) {
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e5m2.e5m2.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
} else {
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
}
}
#undef ASTRAI_DEVICE_ARCH
// ---------------------------------------------------------------------------
// ldmatrix — cooperatively load 8x8 b16 matrices from smem into registers.
//
// The instruction is identical for every 16-bit-storage element type: bf16
// maps 1:1 onto b16 slots; fp8 is stored packed two-per-slot (see
// fp8/gemm.cuh), so one b16 slot holds two fp8 values. `T` is the element
// type and only serves as a semantic tag.
//
// x2 (single address): matrix0 = p (8 rows), matrix1 = p + 8*16 bytes
// x4: four matrices at p, +128, +256, +384 bytes
// Trans: transpose variant (V fragments of attention)
//
// ldmatrix takes a *single* smem address per thread, but the addresses of
// the 32 lanes are *not* all the same: lane i supplies the start address of
// matrix-row i (modulo 8) for matrix (i/8) — lanes 0-7 feed matrix 0's rows,
// lanes 8-15 matrix 1's rows (x2/x4), lanes 16-23 / 24-31 matrix 2 / 3's rows
// (x4 only; their addresses are ignored by x2). Each matrix is 8 rows x 16
// bytes, and consecutive matrices of one instruction are contiguous at
// 128-byte strides. fp8 fragment layouts in fp8/gemm.cuh are arranged around
// this constraint.
// ---------------------------------------------------------------------------
template <typename T, bool Trans = false>
DEVICE_FORCEINLINE void ldmatrix_x2(unsigned r[2], const T* p) {
const unsigned a = __cvta_generic_to_shared(p);
if constexpr (Trans) {
asm volatile(
"ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
} else {
asm volatile(
"ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
}
// Four matrices at p, p+128, p+256, p+384 bytes (16-byte row stride).
template <typename T>
DEVICE_FORCEINLINE void ldmatrix_x4(unsigned r[4], const T* p) {
const unsigned a = __cvta_generic_to_shared(p);
asm volatile(
"ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(a));
}
// Per-lane-address variants: the caller supplies a raw shared-memory address
// per lane instead of one common pointer. Use when the fragment tiles are
// XOR-swizzled per 16B chunk so each lane must compute its own row and chunk
// address (see fp8/gemm.cuh's frag_addr + lane selectors for the m16n8k32
// operand layouts).
DEVICE_FORCEINLINE void ldmatrix_x2_lane(unsigned r[2],
unsigned addr) {
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(addr));
}
DEVICE_FORCEINLINE void ldmatrix_x4_lane(unsigned r[4],
unsigned addr) {
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(addr));
}
} // namespace astrai
+48
View File
@@ -0,0 +1,48 @@
// Shared warp/block reduction + atomic helpers — pure CUDA, no torch.
//
// Extracted from the attention and fp8 families so both share one
// implementation: warp_reduce_sum (decode scalar kernel), warp_reduce_max +
// atomic_max_float (fp8 quantize amax), group_reduce_sum<G> (prefill scalar
// kernel).
#pragma once
namespace astrai {
// Full-warp butterfly sum reduction (32 lanes).
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
// Full-warp butterfly max reduction (32 lanes).
__device__ __forceinline__ float warp_reduce_max(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
value = fmaxf(value, __shfl_xor_sync(0xffffffffu, value, offset));
return value;
}
// Sub-warp group reduction over G consecutive lanes (G a power of two).
// `mask` is the full participating-lane mask of the group (see the
// prefill scalar kernel's gmask computation).
template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
#pragma unroll
for (int o = G / 2; o > 0; o >>= 1)
v += __shfl_xor_sync(mask, v, o);
return v;
}
// Unsigned-bit-pattern atomicMax for non-negative floats; a null
// destination disables the update (kernels with optional amax slots).
__device__ __forceinline__ void atomic_max_float(float* destination,
float value) {
if (destination)
atomicMax(reinterpret_cast<unsigned*>(destination),
__float_as_uint(value));
}
} // namespace astrai
+130
View File
@@ -0,0 +1,130 @@
#pragma once
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <cstdint>
// Pure POD/traits header — no .cuh/CUDA-kernel includes; raw __nv_* type
// spellings only.
namespace astrai {
namespace fp8 {
// Compile-time FP8 format: E4M3 (forward, max 448) or E5M2 (gradients,
// max 57344).
enum class FP8Format : int {
E4M3 = 0,
E5M2 = 1,
};
// Operand storage tags (CUTLASS-style) relative to the canonical matrices
// A [M][K] / B [K][N]: A RowMajor = [M][K] (default), A ColMajor = [K][M],
// B RowMajor = [K][N], B ColMajor = [N][K] (the nn.Linear weight). Selection
// is by type at compile time (see gemm.cuh's stage loads).
struct RowMajor {};
struct ColMajor {};
// Compile-time tile configuration, mirroring KernelTraits in the attention
// kernels: CTA tile, warp tile (WarpM x WarpN — e.g. 64x32 on the 128x128
// CTA, 32x32 on the 64x64 small CTA) and cp.async pipeline depth.
template <FP8Format Fmt, int BlockM, int BlockN, int K, int Stages,
int WarpM = 64, int WarpN = 32>
struct Fp8GemmTraits {
static constexpr FP8Format kFormat = Fmt;
static constexpr int kBlockM = BlockM;
static constexpr int kBlockN = BlockN;
static constexpr int kK = K;
static constexpr int kStages = Stages;
static constexpr int kWarpM = WarpM;
static constexpr int kWarpN = WarpN;
static constexpr bool kIsE5M2 = (Fmt == FP8Format::E5M2);
static constexpr __nv_fp8_interpretation_t kNvFormat =
kIsE5M2 ? __NV_E5M2 : __NV_E4M3;
static constexpr float kFp8Max = kIsE5M2 ? 57344.0f : 448.0f;
// Derived geometry: warp tiles tile the CTA. The smem budget is
// layout-aware, so it lives in Fp8GemmSmem (gemm.cuh).
static constexpr int kWarpsM = BlockM / WarpM;
static constexpr int kWarpsN = BlockN / WarpN;
static constexpr int kCtaThreads = kWarpsM * kWarpsN * 32;
static_assert(kWarpsM * WarpM == BlockM && kWarpsN * WarpN == BlockN,
"warp tiles must exactly tile the CTA");
static_assert(WarpM % 16 == 0 && WarpN % 8 == 0,
"warp tile must be a multiple of the m16n8 MMA shape");
};
// Quantize output orientation: RowMajor = x8 only; Transposed = the
// [cols][rows] x8T only; Dual = both from a single read. Transposed/Dual
// produce K-contiguous operands so crosswise consumers (backward
// grad_x / grad_w) route through the NT fast path.
enum class QuantLayout : int {
RowMajor = 0,
Transposed = 1,
Dual = 2,
};
// Quantize-kernel parameter POD: float input -> FP8 with fused amax.
struct FP8QuantizeParams {
const void* __restrict__ input_ptr = nullptr;
void* __restrict__ output_ptr = nullptr;
void* __restrict__ output_transposed_ptr = nullptr; // [cols][rows]
QuantLayout out_layout = QuantLayout::RowMajor;
const float* __restrict__ scale = nullptr; // device multiplier
float* __restrict__ amax = nullptr; // raw-domain max out
// Optional delayed-scaling ring fold: when fold_ring is set, the kernel's
// last-finishing block folds the final amax into hist[hist_idx], reduces
// the window and publishes the next scale — replacing the host-side
// update chain. amax then points at a persistent self-cleaning slot
// (zeroed by the same last block) inside the caller's ring state.
bool fold_ring = false;
float* __restrict__ hist = nullptr; // [hist_len] amax history window
float* __restrict__ scale_out = nullptr;
unsigned int* __restrict__ done = nullptr; // block-completion counter
int hist_len = 0;
int hist_idx = 0;
float fp8_max = 448.0f; // scale = max(hist) / fp8_max / pow2_margin
float pow2_margin = 1.0f;
// Element count (elementwise kernel); the tiled kernel views the same
// buffer as [rows][cols] row-major.
int total = 0;
int rows = 0;
int cols = 0;
};
// Unified GEMM parameter POD, mirroring AttentionParams: one struct flows
// through the kernels; each kernel touches only the fields it needs.
struct FP8Params {
// FP8 operands + output; scales are quantization steps (device
// scalars). Optional bf16 bias fuses into the epilogue (fp32 add before
// the single bf16 rounding); null disables.
const void* __restrict__ a_ptr = nullptr;
const void* __restrict__ b_ptr = nullptr;
const void* __restrict__ bias_ptr = nullptr;
void* __restrict__ out_ptr = nullptr;
const float* __restrict__ scale = nullptr;
// NN-swap mode (canonicalize_gemm): the kernel computes the transposed
// problem and the epilogue scatters D[row][col] to out[col * p.m + row]
// in the caller's [M][N] buffer. Zero in the plain orientation.
int out_transposed = 0;
int m, n, k; // int covers LLM shapes; kernels promote to int64
// Batched (bmm) geometry: grid.z steps these element strides (0
// broadcasts the operand across batches).
int batch = 1;
int64_t a_batch_stride = 0;
int64_t b_batch_stride = 0;
int64_t out_batch_stride = 0;
// Physical leading dims (row strides) of A and B; the binding packs
// them so the kernel reads each buffer naturally or transposed per the
// LayoutA/LayoutB tags.
int a_ld, b_ld;
};
} // namespace fp8
} // namespace astrai
+271
View File
@@ -0,0 +1,271 @@
#pragma once
// FP8 GEMM umbrella: the kernel orchestrator and the host-side launch
// planning. Device layers live in gemm/ (policy / load / scheduler /
// mainloop / epilogue) — pure CUDA, no torch; launchers are plain functions
// shared by the torch binding and the C tests. Layout tags and the NN swap
// semantics are documented in common.h and the design notes
// (docs/developer/cuda_kernels.md).
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <type_traits>
#include "../common/cp_async.cuh"
#include "common.h"
#include "gemm/epilogue.cuh"
#include "gemm/load.cuh"
#include "gemm/mainloop.cuh"
#include "gemm/policy.cuh"
#include "gemm/scheduler.cuh"
namespace astrai {
namespace fp8 {
template <typename Policy>
__global__ void __launch_bounds__(Policy::kCtaThreads, Policy::kMinCtas)
fp8_gemm_kernel(FP8Params p) {
using Traits = typename Policy::Traits;
using Mainloop = Fp8CollectiveMainloop<Policy>;
using Epilogue = Fp8CollectiveEpilogue<Policy>;
// Stages live in dynamic shared memory so deep pipelines (> 48KB
// static limit) opt in via cudaFuncSetAttribute in the launcher.
extern __shared__ __align__(16) char fp8_gemm_smem[];
// Batch slice (grid.z): broadcast operands carry a 0 stride, so the
// same pointer serves every batch.
using T8 = typename Mainloop::T8;
const T8* a = reinterpret_cast<const T8*>(p.a_ptr) +
(int64_t)blockIdx.z * p.a_batch_stride;
const T8* b = reinterpret_cast<const T8*>(p.b_ptr) +
(int64_t)blockIdx.z * p.b_batch_stride;
auto* out_bf16 = reinterpret_cast<__nv_bfloat16*>(p.out_ptr) +
(int64_t)blockIdx.z * p.out_batch_stride;
static_assert(Mainloop::kBlockM * Mainloop::kBlockN * 2 <=
Mainloop::kARing * Mainloop::kBlockM * Mainloop::kK +
Mainloop::kBRing * Mainloop::kBlockN * Mainloop::kK,
"output tile must fit the reclaimed operand smem");
const int2 bn = Fp8GemmTileScheduler<Policy::kGroupRaster>::tile(blockIdx, gridDim);
Mainloop mainloop(fp8_gemm_smem, a, b, p.m, p.n, p.k, p.a_ld, p.b_ld,
threadIdx.x, bn);
float acc[Mainloop::kNt][Mainloop::kMt][4] = {}; // [nt][mt][acc]
mainloop.prologue();
mainloop.accumulate(acc);
// Drain the pipeline before the epilogue reclaims the operand rings.
astrai::cp_async_wait_all();
Epilogue(fp8_gemm_smem, p, bn.x, bn.y, threadIdx.x).run(acc, out_bf16);
}
// ---------------------------------------------------------------------------
// Launchers — pure CUDA (no torch), usable from the binding and pure C tests.
// ---------------------------------------------------------------------------
// SM count of the current device (cached per device; benign init race —
// every writer stores the same value).
inline int device_sm_count() {
static int cached[64] = {};
int dev = 0;
cudaGetDevice(&dev);
const bool cacheable = dev >= 0 && dev < 64;
int sms = cacheable ? cached[dev] : 0;
if (!sms) {
cudaDeviceGetAttribute(&sms, cudaDevAttrMultiProcessorCount, dev);
sms = sms > 0 ? sms : 1;
if (cacheable) cached[dev] = sms;
}
return sms;
}
// Launch one kernel instantiation with its shared-memory budget: budgets
// beyond the 48KB static limit opt in once per instantiation via
// cudaFuncSetAttribute. Templated on the kernel *value* (auto NTTP) so
// every instantiation owns its own armed flag — same-signature kernels
// must not share it. A failed opt-in arms nothing, so the launch below
// fails loudly through the caller's error checks.
template <auto Kernel, typename... Args>
void launch_with_smem(int smem_bytes, dim3 grid, dim3 block,
cudaStream_t stream, Args... args) {
if (smem_bytes > 48 * 1024) {
static bool armed = false; // per instantiation
if (!armed) {
const cudaError_t err = cudaFuncSetAttribute(
Kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
smem_bytes);
armed = (err == cudaSuccess);
}
}
Kernel<<<grid, block, smem_bytes, stream>>>(args...);
}
// Padding-driven small-CTA rule: m or n <= 64 wastes half a 128-row CTA's
// MMA work, and a non-128-divisible shape drags its edge tiles through the
// predicated generic path — when 64 divides both dims, the 64x64 CTA tiles
// exactly and wins that band.
inline bool small_cta_padding(int64_t m, int64_t n) {
if (m <= 64 || n <= 64) return true;
const bool big_div = (m % 128 == 0) && (n % 128 == 0);
const bool small_div = (m % 64 == 0) && (n % 64 == 0);
return !big_div && small_div;
}
// Launch configuration — a pure function of the problem (unit-testable
// without a GPU). Raster order is not a plan field: every canonical layout
// runs grouped raster; the plain-raster knob stays available through
// launch_plan's GroupRaster parameter for experiments.
struct Fp8GemmPlan {
enum class Cta { kSmall64, kNarrow128x64, kBig128 };
Cta cta;
bool small_s3; // kSmall64 only: cp.async pipeline depth (2 vs 3 stages)
};
// crosswise_ops counts the operands taking the direct crosswise load
// (A ColMajor / B RowMajor storage): 0 = dual-congruous NT, 1 = TN and the
// NN swap, 2 = TT. The layout shifts the crossovers (measured tables in
// the design notes): the small CTA hides the crosswise LDG+PRMT latency
// far better, while the big CTA's operand reuse buys back load bandwidth
// the crosswise path does not traffic in.
inline Fp8GemmPlan plan_gemm(const FP8Params& p, int crosswise_ops = 0) {
const int64_t sm = device_sm_count();
const int64_t tiles_128 =
(int64_t)p.batch * ((p.m + 127) / 128) * ((p.n + 127) / 128);
const auto small = [&](bool s3) {
return Fp8GemmPlan{Fp8GemmPlan::Cta::kSmall64, s3};
};
const auto big = [] {
return Fp8GemmPlan{Fp8GemmPlan::Cta::kBig128, false};
};
const auto narrow = [] {
return Fp8GemmPlan{Fp8GemmPlan::Cta::kNarrow128x64, false};
};
// Padding rules first: predication waste beats any wave-fill effect.
if (small_cta_padding(p.m, p.n)) return small(crosswise_ops > 0);
if (crosswise_ops > 0) {
// Crosswise ladder (L20 measured): the small s3 CTA holds ~3/4 of
// the big CTA's per-SM throughput but tiles 4x finer, so it owns
// the whole sub-wave band and past it; the big CTA takes over once
// its grid fills ~1.5 waves.
if (tiles_128 >= sm * 3 / 2) return big();
return small(true);
}
if (tiles_128 >= sm) {
// Wave band: pick by the wave-quantization cost ceil(tiles/sm) *
// T_tile. The narrow tile carries half the big tile's MMA work at
// ~94% of its per-SM efficiency (T_narrow ~= 0.53 * T_big,
// integer-scaled by 100 below) — reproduces every measured
// crossover.
const int64_t tiles_narrow =
(int64_t)p.batch * ((p.m + 127) / 128) * ((p.n + 63) / 64);
const auto waves = [sm](int64_t tiles) { return (tiles + sm - 1) / sm; };
if (waves(tiles_narrow) * 53 < waves(tiles_128) * 100) return narrow();
return big();
}
// Sub-wave band: the narrow CTA fills the wave with N-tiles at full
// warp depth once its grid passes ~3/8 of a wave; below that the plain
// 64x64 CTA's extra parallelism wins, and past ~5/8 of a wave of
// 128x128 tiles the big CTA's operand reuse wins instead.
if (tiles_128 >= sm * 5 / 8) return big();
const int64_t tiles_narrow =
(int64_t)p.batch * ((p.m + 127) / 128) * ((p.n + 63) / 64);
if (tiles_narrow >= sm * 3 / 8) return narrow();
// Full-ring small CTAs: the 24KB s2 variant keeps 4 CTAs/SM while the
// whole grid stays resident; past that the 32KB s3 variant's deeper
// pipeline wins on multi-wave grids.
const int64_t tiles_64 =
(int64_t)p.batch * ((p.m + 63) / 64) * ((p.n + 63) / 64);
return small(tiles_64 > sm * 3);
}
// Grid + launch for one concrete Policy — the only place a GEMM kernel goes
// to the wire.
template <typename Policy>
void launch_policy(const FP8Params& p, cudaStream_t stream) {
using Traits = typename Policy::Traits;
dim3 grid((p.n + Traits::kBlockN - 1) / Traits::kBlockN,
(p.m + Traits::kBlockM - 1) / Traits::kBlockM, p.batch);
launch_with_smem<fp8_gemm_kernel<Policy>>(
Policy::kSmemBytes, grid, dim3(Traits::kCtaThreads), stream, p);
}
// Plan -> Policy: the production-tuned configs. Big CTA: 128x128 of 8 warps
// x 64x32, kK=64, 2-stage full ring, fast loop only for dual-congruous
// layouts. Narrow: 128x64. Small CTA: 64x64 of 4 warps x 32x32, kK=64,
// kFastLoop always on.
template <FP8Format Fmt, typename LayoutA, typename LayoutB, int GroupRaster>
void launch_plan(const FP8Params& p, const Fp8GemmPlan& plan,
cudaStream_t stream) {
constexpr bool kBigFast = !std::is_same_v<LayoutA, ColMajor> &&
!std::is_same_v<LayoutB, RowMajor>;
switch (plan.cta) {
case Fp8GemmPlan::Cta::kBig128: {
using Policy =
Fp8GemmPolicy<Fmt, 128, 128, LayoutA, LayoutB, 64, 32, 64, 2,
GroupRaster, false, kBigFast>;
launch_policy<Policy>(p, stream);
break;
}
case Fp8GemmPlan::Cta::kNarrow128x64: {
using Policy =
Fp8GemmPolicy<Fmt, 128, 64, LayoutA, LayoutB, 32, 32, 64, 2,
GroupRaster, false, true>;
launch_policy<Policy>(p, stream);
break;
}
case Fp8GemmPlan::Cta::kSmall64: {
if (plan.small_s3) {
using Policy = Fp8GemmPolicy<Fmt, 64, 64, LayoutA, LayoutB, 32, 32,
64, 3, GroupRaster, false, true>;
launch_policy<Policy>(p, stream);
} else {
using Policy = Fp8GemmPolicy<Fmt, 64, 64, LayoutA, LayoutB, 32, 32,
64, 2, GroupRaster, false, true>;
launch_policy<Policy>(p, stream);
}
break;
}
}
}
// Pure problem rewrite: the dual-N-contiguous problem (trans_a/trans_b both
// false) has no dedicated instantiation — it runs as its transpose
// E[N][M] = B^T @ A^T (CUTLASS-sm90's is_swapAB) over swapped operands,
// with p.out_transposed making the epilogue scatter into the caller's
// [M][N] row-major buffer. The rewritten trans flags become the layout tags
// the launcher instantiates; the NN path pays a scalar-store scatter, which
// its rare usage makes the right trade.
inline void canonicalize_gemm(FP8Params& p, bool& trans_a, bool& trans_b) {
if (!trans_a && !trans_b) {
FP8Params s = p; // E = B^T * A^T: swap roles, M <-> N
s.m = p.n;
s.n = p.m;
s.a_ptr = p.b_ptr;
s.b_ptr = p.a_ptr;
s.a_ld = p.b_ld;
s.b_ld = p.a_ld;
s.a_batch_stride = p.b_batch_stride;
s.b_batch_stride = p.a_batch_stride;
s.out_transposed = 1;
p = s;
trans_a = trans_b = true;
}
}
// Entry point: canonicalize the problem, plan the launch, wire the layout
// tags through.
template <FP8Format Fmt>
void gemm(FP8Params p, cudaStream_t stream, bool trans_a, bool trans_b) {
canonicalize_gemm(p, trans_a, trans_b);
// Crosswise operand count for the plan: transposed-A storage (ColMajor)
// and plain-B storage (RowMajor) both take the direct crosswise load.
const int crosswise = (trans_a ? 1 : 0) + (trans_b ? 0 : 1);
const Fp8GemmPlan plan = plan_gemm(p, crosswise);
if (trans_a && trans_b)
launch_plan<Fmt, ColMajor, ColMajor, 8>(p, plan, stream);
else if (trans_b)
launch_plan<Fmt, RowMajor, ColMajor, 8>(p, plan, stream);
else
launch_plan<Fmt, ColMajor, RowMajor, 8>(p, plan, stream);
}
} // namespace fp8
} // namespace astrai
+180
View File
@@ -0,0 +1,180 @@
#pragma once
// Collective epilogue: fused bias, the bf16 scatter of the fp32 accumulators
// through the reclaimed operand shared memory, and the coalesced copy-out.
#include "../common.h"
#include "policy.cuh"
namespace astrai {
namespace fp8 {
template <typename Policy>
struct Fp8CollectiveEpilogue {
using Traits = typename Policy::Traits;
static constexpr bool kStreamOut = Policy::kStreamOut;
static constexpr int kBlockM = Traits::kBlockM;
static constexpr int kBlockN = Traits::kBlockN;
static constexpr int kMt = Traits::kWarpM / 16;
static constexpr int kNt = Traits::kWarpN / 8;
__nv_bfloat16* const tile_out;
const float output_scale;
const __nv_bfloat16* const bias;
const int64_t m, n;
const bool t_out;
const int row_elems, row_chunks;
const int warp_m, warp_n, group, thread_in_group;
const int64_t block_m, block_n;
__device__ Fp8CollectiveEpilogue(char* smem, const FP8Params& p,
int64_t block_m, int64_t block_n, int tid)
: tile_out(reinterpret_cast<__nv_bfloat16*>(smem)),
output_scale(*p.scale),
bias(reinterpret_cast<const __nv_bfloat16*>(p.bias_ptr)),
m(p.m), n(p.n), t_out(p.out_transposed != 0),
row_elems(t_out ? kBlockM : kBlockN),
row_chunks(row_elems / 8),
warp_m((tid >> 5) / Traits::kWarpsN),
warp_n((tid >> 5) % Traits::kWarpsN),
group((tid & 31) >> 2),
thread_in_group(tid & 3),
block_m(block_m), block_n(block_n) {}
// Swizzled address of one 16B chunk (row r, chunk c) of the staged
// tile. Plain orientation: kBlockM rows of kBlockN elems; out-
// transposed (swap dispatch): rows and row length trade places. Both
// row-chunk counts are powers of two, keeping the XOR swizzle
// well-defined.
__device__ __forceinline__ __nv_bfloat16* out_chunk(int r, int c) const {
return tile_out + (size_t)r * row_elems +
((c ^ (r & (row_chunks - 1))) * 8);
}
__device__ __forceinline__ __nv_bfloat16* out_elem(int r, int v) const {
return out_chunk(r, v >> 3) + (v & 7);
}
// Scatter the accumulators into the staging tile: the operand rings are
// dead once the mainloop ends, so their space stages the bf16 output
// tile. Threads scatter (STS.32 of bf16x2 pairs), a barrier makes the
// tile coherent, then the whole CTA copies it out in fully-coalesced
// 16B chunks. The 16B-chunk XOR swizzle keeps both the scatter and the
// gather conflict-free.
__device__ __forceinline__ void stage(float acc[kNt][kMt][4]) const {
// Fused bias: added to the fp32 accumulator before the single bf16
// rounding. The per-lane loads are L1 broadcasts; rows past the
// edge skip the load (their smem slots never copy out). Under
// out_transposed the bias indexes D-cols = the kernel's rows.
const int local_col0 = warp_n * Traits::kWarpN + thread_in_group * 2;
const int64_t bias_col0 = block_n * kBlockN;
const int64_t bias_row0 = block_m * kBlockM;
if (!t_out) {
#pragma unroll
for (int nt = 0; nt < kNt; ++nt) {
const int col = local_col0 + nt * 8;
const int64_t gcol = bias_col0 + col;
const float b0 =
bias && gcol < n ? __bfloat162float(bias[gcol]) : 0.0f;
const float b1 =
bias && gcol + 1 < n ? __bfloat162float(bias[gcol + 1])
: 0.0f;
#pragma unroll
for (int mt = 0; mt < kMt; ++mt) {
const int r0 = warp_m * Traits::kWarpM + group + mt * 16;
const float* tile_acc = acc[nt][mt];
// Two bf16x2 stores per accumulator tile: rows g and
// g+8 of the m16n8 output, columns tig*2/tig*2+1 inside
// one 16B chunk.
const int off = col & 7; // element offset in the chunk
*reinterpret_cast<__nv_bfloat162*>(
out_chunk(r0, col >> 3) + off) =
__floats2bfloat162_rn(tile_acc[0] * output_scale + b0,
tile_acc[1] * output_scale + b1);
*reinterpret_cast<__nv_bfloat162*>(
out_chunk(r0 + 8, col >> 3) + off) =
__floats2bfloat162_rn(tile_acc[2] * output_scale + b0,
tile_acc[3] * output_scale + b1);
}
}
} else {
// Transposed scatter: accumulator (kernel row r0, col) is
// D[col0_global + col][row0_global + r0], staged at T[col][r0].
// The acc pair spans two staged rows, so these are scalar
// stores (the swap path is the rare NN layout). OOB elements
// store dead lanes of the tile, never copied out.
#pragma unroll
for (int nt = 0; nt < kNt; ++nt) {
const int col = local_col0 + nt * 8;
#pragma unroll
for (int mt = 0; mt < kMt; ++mt) {
const int r0 = warp_m * Traits::kWarpM + group + mt * 16;
const int64_t grow = bias_row0 + r0;
const float b =
bias && grow < m ? __bfloat162float(bias[grow]) : 0.0f;
const float* tile_acc = acc[nt][mt];
*out_elem(col, r0) =
__float2bfloat16(tile_acc[0] * output_scale + b);
*out_elem(col + 1, r0) =
__float2bfloat16(tile_acc[1] * output_scale + b);
*out_elem(col, r0 + 8) =
__float2bfloat16(tile_acc[2] * output_scale + b);
*out_elem(col + 1, r0 + 8) =
__float2bfloat16(tile_acc[3] * output_scale + b);
}
}
}
}
// Coalesced copy-out: thread -> one 16B chunk; consecutive threads walk
// a row so each global transaction covers a full 128B line. Under the
// swap the staged rows are D-rows counted from block_n's stripe while
// the row length is kernel m', so row/stride flip to the swapped dims.
__device__ __forceinline__ void store(__nv_bfloat16* out_bf16) const {
constexpr int kTotalChunks =
kBlockM * (kBlockN / 8); // == kBlockN * (kBlockM/8)
const int64_t row0_global = block_m * kBlockM;
const int64_t col0_global = block_n * kBlockN;
for (int idx = threadIdx.x; idx < kTotalChunks; idx += kCtaThreads) {
const int r = idx / row_chunks;
const int c = idx % row_chunks;
const uint4 v = *reinterpret_cast<const uint4*>(out_chunk(r, c));
const int64_t row = t_out ? (int64_t)block_n * kBlockN + r
: row0_global + r;
const int64_t col = t_out ? row0_global + (int64_t)c * 8
: col0_global + (int64_t)c * 8;
const int64_t rows_total = t_out ? n : m;
const int64_t row_stride = t_out ? m : n;
if (row >= rows_total) break; // rows are consecutive: nothing left
auto* dst = out_bf16 + row * row_stride + col;
if (col + 8 <= row_stride &&
(reinterpret_cast<uintptr_t>(dst) & 15) == 0) {
if constexpr (kStreamOut) {
// Evict-first streaming store knob: neutral on L20
// squares, -3..4% on rects; kept for other SKUs.
__stcs(reinterpret_cast<uint4*>(dst), v);
} else {
*reinterpret_cast<uint4*>(dst) = v;
}
} else {
// Row-edge chunk or an odd-stride row base: spill the
// elements that survive the row edge.
const __nv_bfloat16* elems =
reinterpret_cast<const __nv_bfloat16*>(&v);
for (int e = 0; e < 8 && col + e < row_stride; ++e)
dst[e] = elems[e];
}
}
}
__device__ __forceinline__ void run(float acc[kNt][kMt][4],
__nv_bfloat16* out_bf16) {
stage(acc);
__syncthreads();
store(out_bf16);
}
private:
static constexpr int kCtaThreads = Traits::kCtaThreads;
};
} // namespace fp8
} // namespace astrai
+224
View File
@@ -0,0 +1,224 @@
#pragma once
// Operand loaders: swizzled shared-memory staging for congruous operands
// (cp.async, predicated and interior variants, plus the loop-carried
// prefetch state) and the direct LDG+PRMT path for crosswise operands.
// The staging invariants and the swizzle derivation live in
// docs/developer/cuda_kernels.md.
#include "../../common/cp_async.cuh"
#include "../common.h"
#include "policy.cuh"
namespace astrai {
namespace fp8 {
// log2 of a compile-time power of two (for the swizzle shifts).
template <int N, int Acc = 0>
struct log2_const : log2_const<(N >> 1), Acc + 1> {};
template <int Acc>
struct log2_const<1, Acc> {
static constexpr int value = Acc;
};
// Swizzled address inside a flat [rows * K] staging tile: the 16B chunk
// index is XORed with the row bits at [3, 3+log2(kChunks)) so a warp's
// ldmatrix fragment load (8 consecutive rows x 16B) hits all 32 banks
// exactly once; chunks stay contiguous, so cp.async staging is unaffected.
template <int K, typename T8>
__device__ __forceinline__ T8* tile_at(T8* tile, int row, int col) {
constexpr int kChunks = K / 16; // 16B chunks per row
static_assert(kChunks >= 1 && (kChunks & (kChunks - 1)) == 0,
"swizzle needs a power-of-two 16B-chunk count");
constexpr int kShift = 3 - log2_const<kChunks>::value;
return tile + row * K +
((((col >> 4) ^ ((row >> kShift) & (kChunks - 1))) << 4) + (col & 15));
}
// Stage-load a CONGRUOUS operand (contract-contiguous storage — the only
// cp.async-able shape) into the flat [rows * K] swizzled tile. kInterior
// drops all predication: valid only for a fully interior CTA (whole rows,
// 16B-aligned base|ld, k_base + K <= contract); the address math then folds
// to one immediate XOR per chunk (see the design notes). Crosswise operands
// go through load_crosswise_direct instead.
template <typename T8, int K, int RowsTile, int kThreads,
bool kInterior = false>
__device__ __forceinline__ void
load_operand_tile(T8* tile, const T8* __restrict__ operand, int64_t rows,
int64_t contract, int64_t ld, int tid, int64_t k_base,
int64_t block_row) {
constexpr int kChunks = K / 16;
static_assert(RowsTile * kChunks % kThreads == 0,
"tile chunks must divide evenly across threads");
constexpr int kCpt = RowsTile * kChunks / kThreads; // chunks per thread
constexpr int kCpr = kChunks / kCpt; // chunks per row slice
const int r = tid / kCpr;
const int c0 = (tid % kCpr) * kCpt * 16;
if constexpr (kInterior) {
const char* src = reinterpret_cast<const char*>(
operand + (block_row + r) * ld + k_base + c0);
const uintptr_t dst =
reinterpret_cast<uintptr_t>(tile_at<K>(tile, r, c0));
#pragma unroll
for (int j = 0; j < kCpt; ++j)
astrai::cp_async_16(reinterpret_cast<T8*>(dst ^ (j << 4)),
src + j * 16);
} else {
const int64_t row = block_row + r;
const bool row_ok = row < rows;
// k_base and every c are multiples of 16, so all chunks share the
// row base's alignment verdict.
const auto* src = operand + row * ld + k_base;
const bool chunk_aligned = (reinterpret_cast<uintptr_t>(src) & 15) == 0;
#pragma unroll
for (int j = 0; j < kCpt; ++j) {
const int c = c0 + j * 16;
T8* dst = tile_at<K>(tile, r, c);
if (row_ok && chunk_aligned && k_base + c + 15 < contract) {
astrai::cp_async_16(dst, src + c);
} else {
// Tail chunk / misaligned base / OOB row: scalar fill.
#pragma unroll
for (int i = 0; i < 16; ++i)
dst[i] =
row_ok && k_base + c + i < contract ? src[c + i] : T8(0.0f);
}
}
}
}
// Loop-carried prefetch state for one congruous operand ring: per-thread
// (r, c0) mapping with the swizzled stage destination and global source
// pointer carried across k-tiles, so each prefetch chunk is one LDGSTS
// issued straight from registers. The guard is a property of the operand's
// layout, so it lives in the type: the false specialization (crosswise
// operand) is an empty no-op.
template <bool kAsync, typename T8, int kK, int kRowsTile, int kThreads>
struct PrefetchCarry;
template <typename T8, int kK, int kRowsTile, int kThreads>
struct PrefetchCarry<true, T8, kK, kRowsTile, kThreads> {
static constexpr int kCpt = kRowsTile * (kK / 16) / kThreads;
static constexpr int kCpr = (kK / 16) / kCpt;
unsigned wr = 0; // current stage's swizzled destination offset
unsigned wr0 = 0; // slot-0 wrap base
unsigned wrEnd = 0; // one-past-the-ring sentinel
const char* src = nullptr; // current tile's global source bytes
__device__ __forceinline__ PrefetchCarry(
const T8* ring, int ringSlots, int stageElems, const T8* operand,
int64_t ld, int64_t blockRow, int tid, int firstTile) {
const int r = tid / kCpr;
const int c0 = (tid % kCpr) * kCpt * 16;
const T8* slot0 = ring + (firstTile % ringSlots) * stageElems;
const unsigned laneOff = static_cast<unsigned>(
(const char*)tile_at<kK>(slot0, r, c0) - (const char*)slot0);
const unsigned base = __cvta_generic_to_shared(ring) + laneOff;
wr = base + (unsigned)((firstTile % ringSlots) * stageElems);
wr0 = base;
wrEnd = base + (unsigned)(ringSlots * stageElems);
src = reinterpret_cast<const char*>(
operand + (blockRow + r) * ld + c0) +
(int64_t)firstTile * kK;
}
// Emit this thread's chunks for the current tile; pf false (loop tail)
// zero-fills into the slot compute(i-1) already released.
__device__ __forceinline__ void emit(bool pf) const {
#pragma unroll
for (int j = 0; j < kCpt; ++j)
astrai::cp_async_16(wr ^ (unsigned)(j << 4), src + j * 16, pf);
}
__device__ __forceinline__ void advance(int stageElems) {
wr += (unsigned)stageElems;
if (wr == wrEnd) wr = wr0;
src += kK;
}
};
template <typename T8, int kK, int kRowsTile, int kThreads>
struct PrefetchCarry<false, T8, kK, kRowsTile, kThreads> {
__device__ __forceinline__ PrefetchCarry(
const T8*, int, int, const T8*, int64_t, int64_t, int, int) {}
__device__ __forceinline__ void emit(bool) const {}
__device__ __forceinline__ void advance(int) {}
};
// Direct (synchronous) crosswise load into a canonical rotating stage:
// LDG.128 x4 (4 consecutive contract bytes x 16 rows) + in-register PRMT
// transpose + 16 STS.32. Crosswise operands cannot cp.async into the
// canonical tile (a 16B global run holds one contract byte for each of 16
// rows), so they take this path; a staged smem->smem variant measured
// 15-20% slower and was removed (see git history).
template <typename T8, int K, int RowsTile, int kThreads>
__device__ __forceinline__ void
load_crosswise_direct(T8* tile, const T8* __restrict__ operand, int64_t rows,
int64_t contract, int64_t ld, int tid, int64_t k_base,
int64_t block_row) {
constexpr int kQuads = K / 4; // 4-byte contract quads per tile
constexpr int kGroups = RowsTile / 16;
constexpr int kTChunks = kQuads * kGroups; // 64B chunks per tile
// r0 is a multiple of 16 and p*ld preserves alignment whenever ld has
// it, so every run of a chunk shares one alignment verdict.
const bool run_aligned =
((reinterpret_cast<uintptr_t>(operand) | ld) & 15) == 0;
for (int chunk = tid; chunk < kTChunks; chunk += kThreads) {
const int quad = chunk / kGroups;
const int rg = chunk % kGroups;
const int64_t r0 = block_row + rg * 16;
const bool rows_full = r0 + 15 < rows;
if (rows_full && run_aligned) {
const int64_t p0 = k_base + quad * 4;
uint4 v[4];
#pragma unroll
for (int s = 0; s < 4; ++s) {
// Contract tail: a run past k carries zero bytes; they flow
// through the PRMT transpose like any other value.
if (p0 + s < contract)
v[s] = *reinterpret_cast<const uint4*>(
operand + (p0 + s) * ld + r0);
else
v[s] = make_uint4(0u, 0u, 0u, 0u);
}
const unsigned* bytes = reinterpret_cast<const unsigned*>(v);
#pragma unroll
for (int i = 0; i < 16; ++i) {
// word i = row r0+i's quad: byte i of each of the four runs
// [v0.b(i), v1.b(i), v2.b(i), v3.b(i)].
const unsigned nib = i & 3;
const unsigned sel = nib | ((nib + 4) << 4);
const unsigned w01 =
__byte_perm(bytes[0 + (i >> 2)], bytes[4 + (i >> 2)], sel);
const unsigned w23 =
__byte_perm(bytes[8 + (i >> 2)], bytes[12 + (i >> 2)], sel);
*reinterpret_cast<unsigned*>(tile_at<K>(tile, rg * 16 + i,
quad * 4)) =
__byte_perm(w01, w23, 0x5410u);
}
} else {
// Row-tail or misaligned chunk: byte-granular gather with
// per-row predication; contract-tail columns zero-fill.
#pragma unroll
for (int s = 0; s < 4; ++s) {
const int col = quad * 4 + s;
if (k_base + col >= contract) {
#pragma unroll
for (int i = 0; i < 16; ++i)
*tile_at<K>(tile, rg * 16 + i, col) = T8(0.0f);
continue;
}
#pragma unroll
for (int i = 0; i < 16; ++i) {
const int64_t r_idx = r0 + i;
*tile_at<K>(tile, rg * 16 + i, col) =
r_idx < rows
? operand[(k_base + col) * ld + r_idx]
: T8(0.0f);
}
}
}
}
}
} // namespace fp8
} // namespace astrai
+336
View File
@@ -0,0 +1,336 @@
#pragma once
// Collective mainloop: shared-memory stage rings, the gmem->smem stage loads
// (congruous cp.async / crosswise LDG+PRMT), the per-lane ldmatrix fragment
// addressing and the software-pipelined mma.sync loop. The fragment
// addressing scheme and the fast-loop peel rationale live in
// docs/developer/cuda_kernels.md.
#include <type_traits>
#include "../../common/mma.cuh"
#include "../common.h"
#include "load.cuh"
#include "policy.cuh"
namespace astrai {
namespace fp8 {
template <typename Policy>
struct Fp8CollectiveMainloop {
using Traits = typename Policy::Traits;
using LayoutA = typename Policy::LayoutTagA;
using LayoutB = typename Policy::LayoutTagB;
using Smem = Fp8GemmSmem<Traits, LayoutA, LayoutB>;
static constexpr bool kFastLoop = Policy::kFastLoop;
using T8 = std::conditional_t<Traits::kIsE5M2, __nv_fp8_e5m2, __nv_fp8_e4m3>;
static constexpr int kBlockM = Traits::kBlockM;
static constexpr int kBlockN = Traits::kBlockN;
static constexpr int kK = Traits::kK;
static constexpr int kStages = Traits::kStages;
static constexpr int kCtaThreads = Traits::kCtaThreads;
static constexpr bool kDirectA = Smem::kDirectA;
static constexpr bool kDirectB = Smem::kDirectB;
static_assert(kStages >= 1 && kStages <= 8,
"FP8 GEMM stages must be in [1, 8]");
// CTA = (BlockM/WarpM) x (BlockN/WarpN) warps, each warp computing
// kMt x kNt m16n8k32 MMAs. Rings rotate kStages+1 buffers (see
// Fp8GemmSmem) — one __syncthreads per k-tile.
static constexpr int kMt = Traits::kWarpM / 16; // 16-row MMA tiles per warp
static constexpr int kNt = Traits::kWarpN / 8; // 8-col MMA tiles per warp
static constexpr int kSegs = kK / kMmaK; // mma-sized k segments per tile
static constexpr int kARing = Smem::kRingDepth;
static constexpr int kBRing = Smem::kRingDepth;
static constexpr int kAStageBytes = kBlockM * kK;
static constexpr int kBStageBytes = kBlockN * kK;
T8* const a_base;
T8* const b_base;
const T8* const a;
const T8* const b;
const int64_t m, n, k, a_ld, b_ld;
const int tid;
const int64_t block_m, block_n;
const int warp_m, warp_n;
const int a_row0; // + mt * 16 in the loop
const int b_row0; // + nt * 8
const int64_t tile_count;
// Interior-CTA peel (kFastLoop instantiations only): whole-CTA,
// 16B-aligned, K without tail — the mainloop then runs a compile-time
// specialized copy with no per-chunk predication (measured +4.5..10% on
// the issue-bound small CTA; the 128x128 CTA regressed, so only the
// small CTA opts in). The verdict is uniform per CTA.
const bool fast_cta;
__device__ Fp8CollectiveMainloop(char* smem, const T8* a, const T8* b,
int64_t m, int64_t n, int64_t k,
int64_t a_ld, int64_t b_ld, int tid,
int2 block)
: a_base(reinterpret_cast<T8*>(smem)),
b_base(reinterpret_cast<T8*>(smem + kARing * kAStageBytes)),
a(a), b(b), m(m), n(n), k(k), a_ld(a_ld), b_ld(b_ld), tid(tid),
block_m(block.x), block_n(block.y),
warp_m((tid >> 5) / Traits::kWarpsN),
warp_n((tid >> 5) % Traits::kWarpsN),
a_row0(warp_m * Traits::kWarpM),
b_row0(warp_n * Traits::kWarpN),
tile_count((k + kK - 1) / kK),
fast_cta(kFastLoop && !kDirectA && !kDirectB &&
((int64_t)block.x * kBlockM + kBlockM <= m) &&
((int64_t)block.y * kBlockN + kBlockN <= n) &&
((reinterpret_cast<uintptr_t>(a) | (uint64_t)a_ld) & 15) == 0 &&
((reinterpret_cast<uintptr_t>(b) | (uint64_t)b_ld) & 15) == 0 &&
(k % kK) == 0) {}
// Stage-slot helpers: rings rotate one slot per k-tile, so callers
// either compute the slot from the tile index (prologue, generic loop)
// or carry an advancing pointer (steady-state fast loop).
__device__ __forceinline__ T8* a_stage_of(int64_t tile) const {
return a_base + (size_t)(tile % kARing) * kAStageBytes;
}
__device__ __forceinline__ T8* b_stage_of(int64_t tile) const {
return b_base + (size_t)(tile % kBRing) * kBStageBytes;
}
// Asynchronous congruous loads for one k-tile: cp.async into the
// canonical rings; kFast selects the predication-free interior copy
// (fast_cta admits only congruous operands). Called after the
// post-compute barrier, alongside the commit.
template <bool kFast = false>
__device__ __forceinline__ void load_async(T8* a_stage, T8* b_stage,
int64_t k_base) const {
if constexpr (!kDirectA)
load_operand_tile<T8, kK, kBlockM, kCtaThreads, kFast>(
a_stage, a, m, k, a_ld, tid, k_base, block_m * kBlockM);
if constexpr (!kDirectB)
load_operand_tile<T8, kK, kBlockN, kCtaThreads, kFast>(
b_stage, b, n, k, b_ld, tid, k_base, block_n * kBlockN);
}
// Synchronous direct-crosswise loads for one k-tile. In the steady
// state this runs right after barrier 1, so the LDG latency and the
// PRMT transpose overlap the MMA phase instead of stalling the
// inter-barrier window.
__device__ __forceinline__ void load_direct(T8* a_stage, T8* b_stage,
int64_t k_base) const {
if constexpr (kDirectA)
load_crosswise_direct<T8, kK, kBlockM, kCtaThreads>(
a_stage, a, m, k, a_ld, tid, k_base, block_m * kBlockM);
if constexpr (kDirectB)
load_crosswise_direct<T8, kK, kBlockN, kCtaThreads>(
b_stage, b, n, k, b_ld, tid, k_base, block_n * kBlockN);
}
// Prime the pipeline: kStages committed groups, one per stage slot.
// The commit is unconditional — when K is shorter than the pipeline the
// skipped stages commit empty groups, so the group sequence stays
// tile-indexed and the steady-state wait count never needs a runtime
// dispatch.
__device__ __forceinline__ void prologue() const {
#pragma unroll
for (int stage = 0; stage < kStages; ++stage) {
if (stage < tile_count) {
if (fast_cta)
load_async<true>(a_stage_of(stage), b_stage_of(stage),
(int64_t)stage * kK);
else
load_async(a_stage_of(stage), b_stage_of(stage),
(int64_t)stage * kK);
load_direct(a_stage_of(stage), b_stage_of(stage),
(int64_t)stage * kK);
}
astrai::cp_async_commit_group();
}
}
// Steady-state mainloop, compile-time specialized on kFast: the fast
// copy runs predication-free loads with loop-carried read/write
// pointers; the generic copy keeps full predication. kFastLoop=false
// instantiates only the generic copy.
template <bool kFast>
__device__ __forceinline__ void run_loop(float acc[kNt][kMt][4]) const {
const int lane = tid & 31;
// Fast-path write carries: one per congruous operand (crosswise
// operands get the empty no-op type), targeting the first
// prefetched tile (kStages). Steady-state read carries: the LDSM
// base of the current k-tile's stage with the lane offset folded
// in, advanced one stage per iteration with an equality wrap —
// replaces the per-k-tile (tile % ring) * stage_bytes
// recomputation (a UIMAD.WIDE magic-division ladder in SASS).
PrefetchCarry<!kDirectA, T8, kK, kBlockM, kCtaThreads> carry_a(
a_base, kARing, kAStageBytes, a, a_ld, block_m * kBlockM, tid,
kStages);
PrefetchCarry<!kDirectB, T8, kK, kBlockN, kCtaThreads> carry_b(
b_base, kBRing, kBStageBytes, b, b_ld, block_n * kBlockN, tid,
kStages);
const unsigned a_rd0 = __cvta_generic_to_shared(a_base) + a_lane_off(lane);
const unsigned b_rd0 =
__cvta_generic_to_shared(b_base) +
(kPairB ? b4_lane_off(lane) : b_lane_off(lane));
const unsigned a_rd_end = a_rd0 + (unsigned)(kARing * kAStageBytes);
const unsigned b_rd_end = b_rd0 + (unsigned)(kBRing * kBStageBytes);
unsigned a_rd = a_rd0, b_rd = b_rd0;
for (int64_t tile_index = 0; tile_index < tile_count; ++tile_index) {
// In the steady state exactly kStages-1 younger groups are in flight
// when this fires; the tail's unconditional (possibly empty)
// commits keep that invariant true for every iteration.
const bool prefetch = tile_index + kStages < tile_count;
astrai::cp_async_wait_group<kStages - 1>();
// Barrier 1: every thread's cp.async for this stage is complete
// before any thread reads tiles written by other threads.
__syncthreads();
// Direct chunks for tile i+kStages: issue LDG+PRMT+STS now so the
// global-load latency hides behind the MMA phase below.
if (prefetch)
load_direct(a_stage_of(tile_index + kStages),
b_stage_of(tile_index + kStages),
(tile_index + kStages) * kK);
const unsigned a_addr = a_rd;
const unsigned b_addr = b_rd;
// Per-k_seg base pair (cuBLAS's scheme): seg s lives at the seg-0
// base XOR (s<<5) — one LOP3 per extra seg per k-tile, never per
// fragment. Every LDSM below addresses [base + immediate].
unsigned a_seg[kSegs], b_seg[kSegs];
#pragma unroll
for (int s = 0; s < kSegs; ++s) {
a_seg[s] = a_addr ^ (unsigned)(s * kSegXor);
b_seg[s] = b_addr ^ (unsigned)(s * kSegXor);
}
// kNt ldmatrix.x2 (B) + kMt ldmatrix.x4 (A) feed kMt*kNt*2 mma.sync
// per k_seg — 0.5 load instructions per MMA. B fragments
// double-buffer across k_segs; kPairB folds the two adjacent nt
// fragments of one pair into a single x4 (see b4_lane_off).
unsigned b_frag[2][kNt][2];
unsigned b_frag4[2][kNt / 2][4];
load_b_frags(b_frag[0][0], b_frag4[0][0], b_seg[0]);
#pragma unroll
for (int k_seg = 0; k_seg < kSegs; ++k_seg) {
const int bcur = k_seg & 1, bnext = bcur ^ 1;
if (k_seg + 1 < kSegs)
load_b_frags(b_frag[bnext][0], b_frag4[bnext][0],
b_seg[k_seg + 1]);
// Software-pipelined A fragments: the ldmatrix.x4 for row mt+1 is
// issued before the MMAs consuming row mt, so the LDS latency hides
// behind tensor-pipe work. Costs 4 extra registers.
unsigned a_frag[kMt + 1][4];
astrai::ldmatrix_x4_lane(a_frag[0], a_seg[k_seg]);
#pragma unroll
for (int mt = 0; mt < kMt; ++mt) {
if (mt + 1 < kMt)
astrai::ldmatrix_x4_lane(a_frag[mt + 1],
a_seg[k_seg] + (mt + 1) * kMtStep);
#pragma unroll
for (int nt = 0; nt < kNt; ++nt) {
const unsigned* bops =
kPairB ? (b_frag4[bcur][nt >> 1] + (nt & 1) * 2)
: b_frag[bcur][nt];
astrai::mma_sync<T8>(acc[nt][mt], a_frag[mt], bops,
acc[nt][mt]);
}
}
// Next tile's LDGSTS chunks inside the MMA phase: A's after the
// first k_seg's MMA batch, B's after the last.
if constexpr (kFast) {
if (k_seg == 0) carry_a.emit(prefetch);
if (k_seg == kSegs - 1) carry_b.emit(prefetch);
}
}
// Generic loop (no interleaved prefetch): the next tile's predicated
// loads run after the MMA phase.
if constexpr (!kFast) {
if (prefetch) {
load_async(a_stage_of(tile_index + kStages),
b_stage_of(tile_index + kStages),
(tile_index + kStages) * kK);
}
}
// Unconditional commit: empty in the tail, it pads the group
// sequence so the fixed wait above stays correct.
astrai::cp_async_commit_group();
a_rd += (unsigned)kAStageBytes;
if (a_rd == a_rd_end) a_rd = a_rd0;
b_rd += (unsigned)kBStageBytes;
if (b_rd == b_rd_end) b_rd = b_rd0;
if constexpr (kFast) {
carry_a.advance(kAStageBytes);
carry_b.advance(kBStageBytes);
}
}
}
__device__ __forceinline__ void accumulate(float acc[kNt][kMt][4]) const {
if constexpr (kFastLoop) {
if (fast_cta)
run_loop<true>(acc);
else
run_loop<false>(acc);
} else {
run_loop<false>(acc);
}
}
private:
// Per-lane ldmatrix fragment addressing (base-pair scheme, mirrored
// from the cuBLAS SASS; derivation in the design notes): one base
// register per operand per k_seg, every fragment offset an LDSM
// immediate — zero address arithmetic inside the MMA phase.
__device__ __forceinline__ unsigned a_lane_off(int lane) const {
const int r7 = lane & 7; // row within the 8-row matrix
const int rh8 = (lane >> 3) & 1; // +8 rows (A: lanes 8-15, 24-31)
const int rh16 = lane >> 4; // +1 chunk (A: lanes 16-31)
constexpr int kChunks = kK / 16;
constexpr int kShift = 3 - log2_const<kChunks>::value; // tile_at's shift
const unsigned lswz =
static_cast<unsigned>((r7 >> kShift) & (kChunks - 1));
// Stage-relative, loop-invariant per-lane base; A's fragment row
// carries the +8-row (rh8) and +1-chunk (rh16) halves.
return static_cast<unsigned>((a_row0 + rh8 * 8 + r7) * kK +
((rh16 ^ lswz) << 4));
}
__device__ __forceinline__ unsigned b_lane_off(int lane) const {
const int r7 = lane & 7;
const int rh8 = (lane >> 3) & 1; // +8 rows (B uses rh8 as its chunk half)
constexpr int kChunks = kK / 16;
constexpr int kShift = 3 - log2_const<kChunks>::value;
const unsigned lswz =
static_cast<unsigned>((r7 >> kShift) & (kChunks - 1));
return static_cast<unsigned>((b_row0 + r7) * kK + ((rh8 ^ lswz) << 4));
}
// x4-paired B loads: one ldmatrix.x4 feeds the two adjacent nt
// fragments. Lane contract: lanes 0-7 address rows n0..n7 chunk c,
// lanes 8-15 rows n0..n7 chunk c+1, lanes 16-23 rows n8..n15 chunk c,
// lanes 24-31 rows n8..n15 chunk c+1. The +8-row step never reaches
// the swizzle source bits for kK <= 64; kK=128 swizzles on row[2:0]
// where +8 flips bits, so that config keeps the x2 loads.
static constexpr unsigned kMtStep = 16 * kK; // bytes per m-tile row step
static constexpr unsigned kNtStep = 8 * kK; // bytes per n-tile row step
static constexpr unsigned kSegXor = 32; // chunk-index +2 per k_seg
static constexpr bool kPairB = kK / 16 <= 4;
static_assert(!kPairB || kNt % 2 == 0, "B pairing needs even kNt");
static constexpr unsigned kPairStep = 16 * kK; // bytes per nt-pair row step
__device__ __forceinline__ unsigned b4_lane_off(int lane) const {
return b_lane_off(lane) + (lane >> 4) * kPairStep / 2;
}
// One k_seg's B-fragment loads, shared by the initial fill and the
// double-buffer's next-seg fill. frag2/frag4 are the flat bases of one
// b_frag / b_frag4 buffer (the unused one is never touched).
__device__ __forceinline__ void
load_b_frags(unsigned* frag2, unsigned* frag4, unsigned seg_base) const {
#pragma unroll
for (int p = 0; p < kNt / 2; ++p) {
if constexpr (kPairB) {
astrai::ldmatrix_x4_lane(frag4 + p * 4,
seg_base + p * kPairStep);
} else {
astrai::ldmatrix_x2_lane(frag2 + p * 4,
seg_base + p * 2 * kNtStep);
astrai::ldmatrix_x2_lane(frag2 + p * 4 + 2,
seg_base + (p * 2 + 1) * kNtStep);
}
}
}
};
} // namespace fp8
} // namespace astrai
+53
View File
@@ -0,0 +1,53 @@
#pragma once
// Kernel policy layer: shared-memory budget, occupancy hint and the
// single Policy type the kernel and collectives take (CUTLASS-style
// consolidation of traits + layout tags + scheduling knobs).
#include <type_traits>
#include "../common.h"
namespace astrai {
namespace fp8 {
// m16n8k32 (see astrai::mma_shape<fp8 type>::k in common/mma.cuh)
constexpr int kMmaK = 32;
// Layout-aware shared-memory budget and occupancy hint. Every operand ring
// holds kStages+1 buffers: the load for tile i+kStages targets slot
// (i-1)%(kStages+1) — already consumed — so neither load path needs a
// post-compute barrier (one __syncthreads per k-tile; see the design notes
// in docs/developer/cuda_kernels.md). The 48KB static watermark picks the
// resident-CTA hint for __launch_bounds__.
template <typename Traits, typename LayoutA, typename LayoutB>
struct Fp8GemmSmem {
// Crosswise (direct-load) operands: A ColMajor storage, B RowMajor
// storage (B's tag is relative to the canonical [K][N]).
static constexpr bool kDirectA = std::is_same_v<LayoutA, ColMajor>;
static constexpr bool kDirectB = std::is_same_v<LayoutB, RowMajor>;
static constexpr int kRingDepth = Traits::kStages + 1;
static constexpr int kBytes =
kRingDepth * (Traits::kBlockM + Traits::kBlockN) * Traits::kK;
static constexpr int kMinCtas = kBytes <= 48 * 1024 ? 2 : 1;
};
template <FP8Format Fmt_, int BlockM_, int BlockN_, typename LayoutA_,
typename LayoutB_, int WarpM_, int WarpN_, int kK_, int Stages_,
int GroupRaster_, bool StreamOut_ = false, bool FastLoop_ = false>
struct Fp8GemmPolicy {
using Traits =
Fp8GemmTraits<Fmt_, BlockM_, BlockN_, kK_, Stages_, WarpM_, WarpN_>;
using LayoutTagA = LayoutA_;
using LayoutTagB = LayoutB_;
static constexpr int kGroupRaster = GroupRaster_;
static constexpr bool kStreamOut = StreamOut_;
static constexpr bool kFastLoop = FastLoop_;
using Smem = Fp8GemmSmem<Traits, LayoutA_, LayoutB_>;
// Flattened for __launch_bounds__, which takes no dependent type names.
static constexpr int kCtaThreads = Traits::kCtaThreads;
static constexpr int kMinCtas = Smem::kMinCtas;
static constexpr int kSmemBytes = Smem::kBytes;
};
} // namespace fp8
} // namespace astrai
+28
View File
@@ -0,0 +1,28 @@
#pragma once
// Tile scheduler: the linear CTA id maps to (block_m, block_n) in grouped
// (L2-friendly) raster — consecutive CTAs share one B column stripe — or
// plain N-fastest raster (kRasterGroup=0, the measured best for dX's
// crosswise-B layouts where grouping was neutral).
namespace astrai {
namespace fp8 {
template <int kRasterGroup>
struct Fp8GemmTileScheduler {
static __device__ int2 tile(const uint3& block, const dim3& blocks) {
if constexpr (kRasterGroup > 0) {
constexpr int kGroupM = kRasterGroup;
const int bid = int(block.y) * int(blocks.x) + int(block.x);
const int group_first_m = (bid / (kGroupM * int(blocks.x))) * kGroupM;
const int group_rows =
min(int(blocks.y) - group_first_m, kGroupM); // M-tail group is short
return int2{group_first_m + bid % group_rows,
(bid % (kGroupM * int(blocks.x))) / group_rows};
} else {
return int2{int(block.y), int(block.x)};
}
}
};
} // namespace fp8
} // namespace astrai
+307
View File
@@ -0,0 +1,307 @@
// CUDA bindings for the stateless FP8 quantize/GEMM primitives.
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <torch/extension.h>
#include <cstdint>
#include <mutex>
#include <unordered_map>
#include "../common/device.cuh"
#include "gemm.cuh"
#include "quantize.cuh"
using namespace astrai::fp8;
namespace {
void check_fp8_device(const torch::Tensor& tensor) {
static std::mutex mutex;
static std::unordered_map<int, bool> supported;
const int device = tensor.device().index();
{
std::lock_guard<std::mutex> lock(mutex);
auto it = supported.find(device);
if (it != supported.end()) {
TORCH_CHECK(it->second, "FP8 MMA requires compute capability 8.9+");
return;
}
}
const auto* properties = at::cuda::getDeviceProperties(device);
const bool ok = astrai::sm_at_least(
properties->major, properties->minor, astrai::kMinSmForFp8Major,
astrai::kMinSmForFp8Minor);
{
std::lock_guard<std::mutex> lock(mutex);
supported.emplace(device, ok);
}
TORCH_CHECK(ok, "FP8 MMA requires compute capability 8.9+");
}
void check_scale(const torch::Tensor& scale, const torch::Tensor& input) {
TORCH_CHECK(scale.is_cuda() && scale.device() == input.device() &&
scale.scalar_type() == torch::kFloat32 && scale.numel() == 1,
"scale must be a CUDA float32 scalar on the input device");
}
// Inner-layout resolution for one GEMM operand. The user flag names the
// math (0 = last two dims are [rows][contract], 1 = transposed); the
// storage may independently be a col-major view (.t() of a contiguous
// buffer), which folds into the returned dispatch flag at zero copy — the
// kernel's LayoutA/LayoutB tags cover both storages. m/n/k derive from the
// user flag only. Tensors whose inner dims are neither natural layout fall
// back to .contiguous().
bool resolve_operand(const torch::Tensor& t_in, bool flag, int64_t& ld,
int64_t& batch_stride, torch::Tensor& storage) {
torch::Tensor t = t_in;
bool col_major = false;
if (t.stride(-1) != 1) {
if (t.stride(-2) == 1) {
col_major = true;
} else {
t = t.contiguous();
}
}
storage = t;
ld = col_major ? t.stride(-1) : t.stride(-2);
batch_stride = t.dim() == 3 ? t.stride(0) : 0;
return flag ^ col_major;
}
// Dtype dispatch over the unified quantize launcher.
template <bool Tiled, FP8Format Fmt>
void launch_for_dtype(const torch::Tensor& x, const FP8QuantizeParams& p,
cudaStream_t stream) {
switch (x.scalar_type()) {
case torch::kHalf:
launch_fp8_quantize<Fmt, __half, Tiled>(p, stream);
break;
case torch::kFloat32:
launch_fp8_quantize<Fmt, float, Tiled>(p, stream);
break;
default:
launch_fp8_quantize<Fmt, __nv_bfloat16, Tiled>(p, stream);
}
}
template <bool Tiled>
void launch_quantize_for(const torch::Tensor& x, const FP8QuantizeParams& p,
bool e5m2, cudaStream_t stream) {
if (e5m2)
launch_for_dtype<Tiled, FP8Format::E5M2>(x, p, stream);
else
launch_for_dtype<Tiled, FP8Format::E4M3>(x, p, stream);
}
// Shared binding body for the two quantize entry points: RowMajor /
// Transposed (single output) serve quantize(), Dual (both orientations from
// one read) serves quantize_dual(). A ring tensor switches
// on the in-kernel delayed-scaling fold: state layout
// [hist n | scale | legacy | amax | done-as-int], and the returned amax is
// the (self-cleaned) persistent slot. Without it, amax is reduced into a
// fresh buffer armed by a driver memset — cheaper than the zeros() fill
// kernel.
py::object quantize_impl(torch::Tensor x, torch::Tensor scale, int64_t fmt,
QuantLayout layout, py::object ring, int64_t hist_idx,
double fp8_max, double pow2_margin) {
TORCH_CHECK(x.is_cuda(), "CUDA tensors required");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16 ||
x.scalar_type() == torch::kHalf ||
x.scalar_type() == torch::kFloat32,
"x must be bf16, fp16 or fp32");
TORCH_CHECK(fmt == static_cast<int64_t>(FP8Format::E4M3) ||
fmt == static_cast<int64_t>(FP8Format::E5M2),
"unsupported quantization type: expected E4M3 (0) or E5M2 (1)");
TORCH_CHECK(layout == QuantLayout::RowMajor || x.dim() >= 2,
"transposed quantize layouts need a 2D+ tensor");
check_scale(scale, x);
check_fp8_device(x);
const at::cuda::OptionalCUDAGuard guard(x.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto input = x.contiguous();
auto out_opts = input.options().dtype(
fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn);
torch::Tensor amax;
float *ring_hist = nullptr, *ring_scale_out = nullptr;
unsigned int* ring_done = nullptr;
int ring_len = 0;
if (!ring.is_none()) {
auto st = ring.cast<torch::Tensor>();
TORCH_CHECK(st.is_cuda() && st.dim() == 1 &&
st.scalar_type() == torch::kFloat32,
"ring state must be a 1D float32 CUDA tensor");
const int64_t n = st.numel() - 4;
TORCH_CHECK(n > 0 && hist_idx >= 0 && hist_idx < n,
"ring state too small or hist_idx out of range");
float* base = st.data_ptr<float>();
amax = st.narrow(0, n + 2, 1);
ring_hist = base;
ring_scale_out = base + n;
ring_done = reinterpret_cast<unsigned int*>(base + n + 3);
ring_len = static_cast<int>(n);
} else {
amax = torch::empty({1}, input.options().dtype(torch::kFloat32));
cudaMemsetAsync(amax.data_ptr(), 0, sizeof(float), stream.stream());
}
FP8QuantizeParams p;
p.input_ptr = input.data_ptr();
p.scale = scale.data_ptr<float>();
p.amax = amax.data_ptr<float>();
if (ring_hist) {
p.fold_ring = true;
p.hist = ring_hist;
p.scale_out = ring_scale_out;
p.done = ring_done;
p.hist_len = ring_len;
p.hist_idx = static_cast<int>(hist_idx);
p.fp8_max = static_cast<float>(fp8_max);
p.pow2_margin = static_cast<float>(pow2_margin);
}
p.total = static_cast<int>(input.numel());
p.out_layout = layout;
p.rows = static_cast<int>(input.size(-2));
p.cols = static_cast<int>(input.size(-1));
torch::Tensor output, output_t;
if (layout != QuantLayout::Transposed) {
output = torch::empty_like(input, out_opts);
p.output_ptr = output.data_ptr();
}
if (layout != QuantLayout::RowMajor) {
output_t = torch::empty({input.size(-1), input.size(-2)}, out_opts);
p.output_transposed_ptr = output_t.data_ptr();
}
const bool e5m2 = fmt == static_cast<int64_t>(FP8Format::E5M2);
if (layout == QuantLayout::RowMajor)
launch_quantize_for<false>(input, p, e5m2, stream.stream());
else
launch_quantize_for<true>(input, p, e5m2, stream.stream());
C10_CUDA_CHECK(cudaGetLastError());
if (layout == QuantLayout::Dual)
return py::make_tuple(output, output_t, amax);
return py::make_tuple(
layout == QuantLayout::Transposed ? output_t : output, amax);
}
} // namespace
// Single-orientation quantize binding: row-major x8, or its [cols][rows]
// transpose when transposed is set — the K-contiguous operand orientation
// NT GEMMs want. Returns (x8|x8T, amax).
py::object quantize(torch::Tensor x, torch::Tensor scale, int64_t fmt,
bool transposed, py::object ring, int64_t hist_idx,
double fp8_max, double pow2_margin) {
const QuantLayout layout =
transposed ? QuantLayout::Transposed : QuantLayout::RowMajor;
return quantize_impl(x, scale, fmt, layout, ring, hist_idx, fp8_max,
pow2_margin);
}
// Dual-orientation quantize binding: one read of x produces both the
// row-major x8 and its transpose (plus amax), for tensors consumed by GEMMs
// in both orientations (backward g). Returns (x8, x8T, amax).
py::object quantize_dual(torch::Tensor x, torch::Tensor scale, int64_t fmt,
py::object ring, int64_t hist_idx, double fp8_max,
double pow2_margin) {
return quantize_impl(x, scale, fmt, QuantLayout::Dual, ring, hist_idx,
fp8_max, pow2_margin);
}
torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor scale,
bool trans_a, bool trans_b, py::object bias) {
TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required");
TORCH_CHECK(a.scalar_type() == torch::kFloat8_e4m3fn ||
a.scalar_type() == torch::kFloat8_e5m2,
"a and b must be fp8");
TORCH_CHECK(a.scalar_type() == b.scalar_type(), "a and b must share format");
TORCH_CHECK((a.dim() == 2 || a.dim() == 3) &&
(b.dim() == 2 || b.dim() == 3),
"a and b must be 2D or 3D (batched)");
TORCH_CHECK(a.device() == b.device(), "a and b must share device");
// Python None and an omitted argument both mean "no bias" — an undefined
// tensor below. (py::isinstance<torch::Tensor> is false for real tensors
// here — torch's caster registers no pybind type info — so validate by
// attempting the cast itself.)
torch::Tensor bias_t;
if (!bias.is_none()) {
try {
bias_t = bias.cast<torch::Tensor>();
} catch (const py::cast_error&) {
TORCH_CHECK(false, "bias must be a torch.Tensor or None");
}
}
check_scale(scale, a);
check_fp8_device(a);
const at::cuda::OptionalCUDAGuard guard(a.device());
auto stream = at::cuda::getCurrentCUDAStream();
// Batched operands follow matmul broadcast rules: 2D acts as a batch
// of 1; a size-1 batch broadcasts across the other side (stride 0).
const int64_t batch_a = a.dim() == 3 ? a.size(0) : 1;
const int64_t batch_b = b.dim() == 3 ? b.size(0) : 1;
TORCH_CHECK(batch_a == batch_b || batch_a == 1 || batch_b == 1,
"batch dim mismatch (got ", batch_a, " and ", batch_b, ")");
const int64_t batch = std::max(batch_a, batch_b);
TORCH_CHECK(batch <= 65535, "batch dim exceeds the grid.z launch limit");
torch::Tensor a_st, b_st;
int64_t a_ld, b_ld, a_bstride, b_bstride;
const bool tag_a = resolve_operand(a, trans_a, a_ld, a_bstride, a_st);
const bool tag_b = resolve_operand(b, trans_b, b_ld, b_bstride, b_st);
// GEMM dims from the user flags; storage layout never swaps them.
const int64_t m = trans_a ? a.size(-1) : a.size(-2);
const int64_t k = trans_a ? a.size(-2) : a.size(-1);
const int64_t n = trans_b ? b.size(-2) : b.size(-1);
TORCH_CHECK(k == (trans_b ? b.size(-1) : b.size(-2)), "inner dim mismatch");
const bool batched_out = a.dim() == 3 || b.dim() == 3;
torch::Tensor output =
batched_out
? torch::empty({batch, m, n}, a.options().dtype(torch::kBFloat16))
: torch::empty({m, n}, a.options().dtype(torch::kBFloat16));
FP8Params p;
p.a_ptr = a_st.data_ptr();
p.b_ptr = b_st.data_ptr();
p.out_ptr = output.data_ptr();
p.scale = scale.data_ptr<float>();
p.m = static_cast<int>(m);
p.n = static_cast<int>(n);
p.k = static_cast<int>(k);
p.a_ld = static_cast<int>(a_ld);
p.b_ld = static_cast<int>(b_ld);
// Fused epilogue bias (bf16, broadcast over rows and batches). An
// undefined or 0-element tensor keeps the plain scaled output.
if (bias_t.defined() && bias_t.numel() > 0) {
TORCH_CHECK(bias_t.is_cuda() && bias_t.scalar_type() == torch::kBFloat16,
"fp8 gemm bias must be a CUDA bf16 tensor");
TORCH_CHECK(bias_t.dim() == 1 && bias_t.size(0) == n,
"fp8 gemm bias must be 1D of length n=", n);
TORCH_CHECK(bias_t.is_contiguous(), "fp8 gemm bias must be contiguous");
p.bias_ptr = bias_t.data_ptr();
}
p.batch = static_cast<int>(batch);
p.a_batch_stride = (batch_a == 1 && batch > 1) ? 0 : a_bstride;
p.b_batch_stride = (batch_b == 1 && batch > 1) ? 0 : b_bstride;
p.out_batch_stride = m * n;
if (a.scalar_type() == torch::kFloat8_e4m3fn)
gemm<FP8Format::E4M3>(p, stream.stream(), tag_a, tag_b);
else
gemm<FP8Format::E5M2>(p, stream.stream(), tag_a, tag_b);
C10_CUDA_CHECK(cudaGetLastError());
return output;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("quantize", &quantize, py::arg("x"), py::arg("scale"),
py::arg("fmt"), py::arg("transposed") = false,
py::arg("ring") = py::none(), py::arg("hist_idx") = 0,
py::arg("fp8_max") = 448.0, py::arg("pow2_margin") = 1.0);
m.def("quantize_dual", &quantize_dual, py::arg("x"), py::arg("scale"),
py::arg("fmt"), py::arg("ring") = py::none(),
py::arg("hist_idx") = 0, py::arg("fp8_max") = 448.0,
py::arg("pow2_margin") = 1.0);
m.def("mm_fp8", &mm_fp8, py::arg("a"), py::arg("b"), py::arg("scale"),
py::arg("trans_a") = false, py::arg("trans_b") = false,
py::arg("bias") = py::none());
}
+311
View File
@@ -0,0 +1,311 @@
#pragma once
// FP8 quantize device code — pure CUDA, no torch: kernels take the
// FP8QuantizeParams POD, format and input type ride on template parameters,
// and the launcher is shared by the torch binding and the C tests.
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <cstdint>
#include "common.h"
#include "../common/reduce.cuh"
namespace astrai {
namespace fp8 {
// Input element type traits: one element -> float, the unpack of one
// 16-byte load into kVecElems floats, and a native 2-element pair load.
template <typename InT>
struct quant_in_traits;
template <>
struct quant_in_traits<__nv_bfloat16> {
static constexpr int kVecElems = 8;
static __device__ __forceinline__ float to_float(__nv_bfloat16 v) {
return __bfloat162float(v);
}
static __device__ __forceinline__ void load_vec(const uint4& raw,
float* f) {
const __nv_bfloat162* b2 =
reinterpret_cast<const __nv_bfloat162*>(&raw);
#pragma unroll
for (int j = 0; j < 4; ++j) {
const float2 p = __bfloat1622float2(b2[j]);
f[2 * j] = p.x;
f[2 * j + 1] = p.y;
}
}
static __device__ __forceinline__ void load_pair(const __nv_bfloat16* p,
float* f) {
const float2 v = __bfloat1622float2(
*reinterpret_cast<const __nv_bfloat162*>(p));
f[0] = v.x;
f[1] = v.y;
}
};
template <>
struct quant_in_traits<__half> {
static constexpr int kVecElems = 8;
static __device__ __forceinline__ float to_float(__half v) {
return __half2float(v);
}
static __device__ __forceinline__ void load_vec(const uint4& raw,
float* f) {
const __half2* h2 = reinterpret_cast<const __half2*>(&raw);
#pragma unroll
for (int j = 0; j < 4; ++j) {
const float2 p = __half22float2(h2[j]);
f[2 * j] = p.x;
f[2 * j + 1] = p.y;
}
}
static __device__ __forceinline__ void load_pair(const __half* p,
float* f) {
const float2 v =
__half22float2(*reinterpret_cast<const __half2*>(p));
f[0] = v.x;
f[1] = v.y;
}
};
template <>
struct quant_in_traits<float> {
static constexpr int kVecElems = 4;
static __device__ __forceinline__ float to_float(float v) { return v; }
static __device__ __forceinline__ void load_vec(const uint4& raw,
float* f) {
const unsigned* w = reinterpret_cast<const unsigned*>(&raw);
#pragma unroll
for (int j = 0; j < 4; ++j) f[j] = __uint_as_float(w[j]);
}
static __device__ __forceinline__ void load_pair(const float* p,
float* f) {
f[0] = p[0];
f[1] = p[1];
}
};
// One float -> one fp8 byte (round-nearest-even + satfinite).
template <FP8Format Fmt>
__device__ __forceinline__ uint8_t cvt_fp8(float v) {
if constexpr (Fmt == FP8Format::E5M2)
return __nv_fp8_e5m2(v).__x;
else
return __nv_fp8_e4m3(v).__x;
}
// One float pair -> one packed fp8x2 word (round-nearest-even + satfinite).
template <FP8Format Fmt>
__device__ __forceinline__ unsigned cvt_fp8x2(float a, float b) {
constexpr __nv_fp8_interpretation_t kFmt =
Fmt == FP8Format::E5M2 ? __NV_E5M2 : __NV_E4M3;
return static_cast<unsigned>(__nv_cvt_float2_to_fp8x2(
make_float2(a, b), __NV_SATFINITE, kFmt));
}
// Block-wide amax reduce -> one atomic per block: warp-reduce, park one
// value per warp, thread 0 folds. kWarps must cover the block's warp count.
// With p.fold_ring, the last-finishing block additionally folds the final
// amax into the history window and publishes the next scale (atomicAdd
// ticket + fences), re-zeroing the amax slot and the counter for the next
// launch — the host-side delayed-scaling update chain disappears.
template <int kWarps>
__device__ __forceinline__ void publish_amax(const FP8QuantizeParams& p,
float v) {
v = warp_reduce_max(v);
__shared__ float slots[kWarps];
const int tid = threadIdx.y * blockDim.x + threadIdx.x;
if ((tid & 31) == 0) slots[tid >> 5] = v;
__syncthreads();
if (tid == 0) {
#pragma unroll
for (int w = 1; w < kWarps; ++w) v = fmaxf(v, slots[w]);
atomic_max_float(p.amax, v);
if (!p.fold_ring) return;
__threadfence();
const unsigned int ticket = atomicAdd(p.done, 1u);
__threadfence();
if (ticket != gridDim.x - 1u) return;
p.hist[p.hist_idx] = *p.amax;
float peak = p.hist[0];
for (int i = 1; i < p.hist_len; ++i) peak = fmaxf(peak, p.hist[i]);
*p.scale_out = fmaxf(peak / p.fp8_max / p.pow2_margin, 1e-12f);
*p.amax = 0.0f;
*p.done = 0u;
}
}
// Elementwise quantize kernel (QuantLayout::RowMajor): vectorized 16B loads
// -> fp8 stores, fused amax over raw values.
template <FP8Format Fmt, typename InT>
__global__ void fp8_quantize_kernel(FP8QuantizeParams p) {
const float mult = *p.scale;
const auto* x = static_cast<const InT*>(p.input_ptr);
uint8_t* x8 = static_cast<uint8_t*>(p.output_ptr);
float local_amax = 0.0f;
const int64_t stride = (int64_t)blockDim.x * gridDim.x;
// One 16B load -> kVecElems bytes per step. Torch allocations are >=16B
// aligned, so element 0 keeps the uint4 access natural; a misaligned
// base (odd storage offset view) falls to the scalar tail via
// total_vec = 0.
constexpr int kVecElems = quant_in_traits<InT>::kVecElems;
const bool aligned =
((reinterpret_cast<uintptr_t>(x) |
reinterpret_cast<uintptr_t>(x8)) &
15) == 0;
const int64_t total_vec = aligned ? p.total / kVecElems : 0;
const uint4* xv = reinterpret_cast<const uint4*>(x);
for (int64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < total_vec;
i += stride) {
float f[kVecElems];
quant_in_traits<InT>::load_vec(xv[i], f);
// One 32-bit word packs two fp8x2 pairs (4 elements).
unsigned packed[kVecElems / 4];
#pragma unroll
for (int j = 0; j < kVecElems / 4; ++j) {
local_amax = fmaxf(
local_amax,
fmaxf(fmaxf(fabsf(f[4 * j]), fabsf(f[4 * j + 1])),
fmaxf(fabsf(f[4 * j + 2]), fabsf(f[4 * j + 3]))));
const unsigned lo =
cvt_fp8x2<Fmt>(f[4 * j] * mult, f[4 * j + 1] * mult);
const unsigned hi =
cvt_fp8x2<Fmt>(f[4 * j + 2] * mult, f[4 * j + 3] * mult);
packed[j] = (lo & 0xffffu) | (hi << 16);
}
if constexpr (kVecElems == 8)
reinterpret_cast<uint2*>(x8)[i] = make_uint2(packed[0], packed[1]);
else
reinterpret_cast<unsigned*>(x8)[i] = packed[0];
}
// Scalar tail (and full fallback for misaligned bases).
for (int64_t i = total_vec * kVecElems + blockIdx.x * blockDim.x +
threadIdx.x;
i < p.total; i += stride) {
const float v = quant_in_traits<InT>::to_float(x[i]);
local_amax = fmaxf(local_amax, fabsf(v));
x8[i] = cvt_fp8<Fmt>(v * mult);
}
if (p.amax) publish_amax<8>(p, local_amax);
}
// Tiled transpose quantize (QuantLayout::Transposed/Dual): reads the
// [rows][cols] input
// once and writes the fp8 bytes transposed ([cols][rows], so the contract
// dim lands K-contiguous for NT GEMM operands) and, in mode 2, the row-major
// copy too. 64x32 tiles, one native pair load per row (a full 128B warp
// read); rows whose pair is unaligned or ragged (odd widths, misaligned
// bases) fall back to element loads in place. Staging goes through a byte
// tile whose pitch keeps the store stride coprime with the 32 banks.
// (+25-35% over the former 32x32 scalar kernel on sub-4M tensors; ~5%
// slower once DRAM-saturated — accepted for the single-kernel shape.)
template <FP8Format Fmt, typename InT>
__global__ void fp8_quantize_tiled_kernel(FP8QuantizeParams p) {
constexpr int kTileC = 64, kTileR = 32;
// 34B pitch: staging stride is 17 words (coprime with the 32 banks) so
// the pair-byte stores stay conflict-free, and the byte-wise consume
// reads still span distinct words.
__shared__ uint8_t tile[kTileC][kTileR + 2];
const float mult = *p.scale;
const auto* x = static_cast<const InT*>(p.input_ptr);
const int r0 = blockIdx.y * kTileR;
const int c0 = blockIdx.x * kTileC;
const int r = r0 + threadIdx.y * 4;
const int c = c0 + threadIdx.x * 2; // cols even => the pair is in-bounds
uint8_t q[4][2];
float local_amax = 0.0f;
// Vectorize the pair when both elements are in-bounds and the native
// 2-element load is aligned; odd widths, misaligned bases and ragged
// edges fall back to element loads row by row.
constexpr int kPairAlign = 2 * (int)sizeof(InT);
#pragma unroll
for (int j = 0; j < 4; ++j) {
q[j][0] = 0;
q[j][1] = 0;
if (r + j < p.rows && c < p.cols) {
const InT* a = x + (int64_t)(r + j) * p.cols + c;
if (c + 1 < p.cols &&
(reinterpret_cast<uintptr_t>(a) & (kPairAlign - 1)) == 0) {
float f[2];
quant_in_traits<InT>::load_pair(a, f);
#pragma unroll
for (int k = 0; k < 2; ++k) {
local_amax = fmaxf(local_amax, fabsf(f[k]));
q[j][k] = cvt_fp8<Fmt>(f[k] * mult);
}
} else {
const float v0 = quant_in_traits<InT>::to_float(a[0]);
local_amax = fmaxf(local_amax, fabsf(v0));
q[j][0] = cvt_fp8<Fmt>(v0 * mult);
if (c + 1 < p.cols) {
const float v1 = quant_in_traits<InT>::to_float(a[1]);
local_amax = fmaxf(local_amax, fabsf(v1));
q[j][1] = cvt_fp8<Fmt>(v1 * mult);
}
}
}
}
if (p.out_layout == QuantLayout::Dual) {
uint8_t* out = static_cast<uint8_t*>(p.output_ptr);
#pragma unroll
for (int j = 0; j < 4; ++j)
if (r + j < p.rows && c < p.cols) {
uint8_t* o = out + (int64_t)(r + j) * p.cols + c;
const int64_t off = (int64_t)(r + j) * p.cols + c;
if (c + 1 < p.cols && (off & 1) == 0)
*reinterpret_cast<unsigned short*>(o) =
(unsigned short)(q[j][0] | (q[j][1] << 8));
else {
o[0] = q[j][0];
if (c + 1 < p.cols) o[1] = q[j][1];
}
}
}
#pragma unroll
for (int j = 0; j < 4; ++j)
#pragma unroll
for (int k = 0; k < 2; ++k)
tile[threadIdx.x * 2 + k][threadIdx.y * 4 + j] = q[j][k];
__syncthreads();
// Transposed scatter: output element (c, r) lives at c * rows + r;
// threadIdx.x tracks r so each warp writes one contiguous run. tile is
// [col][row]; warp y walks 8 columns, threads read down one column.
uint8_t* out_t = static_cast<uint8_t*>(p.output_transposed_ptr);
#pragma unroll
for (int i = 0; i < 8; ++i) {
const int oc = c0 + threadIdx.y * 8 + i;
if (oc < p.cols && r0 + threadIdx.x < p.rows)
out_t[(int64_t)oc * p.rows + r0 + threadIdx.x] =
tile[threadIdx.y * 8 + i][threadIdx.x];
}
if (p.amax) publish_amax<8>(p, local_amax);
}
// Unified quantize launcher: Tiled selects the transpose kernel
// (QuantLayout::Transposed/Dual) over the vectorized elementwise one. The
// transpose kernel vectorizes
// pair loads in-kernel and falls back to scalar loads at unaligned/ragged
// rows, so the host side picks only the grid.
template <FP8Format Fmt, typename InT, bool Tiled = false>
void launch_fp8_quantize(const FP8QuantizeParams& p, cudaStream_t stream) {
if constexpr (Tiled) {
const dim3 grid((p.cols + 63) / 64, (p.rows + 31) / 32);
if (grid.x == 0 || grid.y == 0) return;
fp8_quantize_tiled_kernel<Fmt, InT><<<grid, dim3(32, 8), 0, stream>>>(p);
} else {
constexpr int kThreads = 256;
constexpr int kVecElems = quant_in_traits<InT>::kVecElems;
// Grid-stride loops: any grid >= 1 is correct; one block per 256
// vectors plus the tail block covers tiny and misaligned tensors.
const int64_t blocks = 1 + p.total / (kVecElems * kThreads);
fp8_quantize_kernel<Fmt, InT><<<blocks, kThreads, 0, stream>>>(p);
}
}
} // namespace fp8
} // namespace astrai
@@ -7,13 +7,12 @@ __global__ void rotary_emb_kernel(
const __nv_bfloat16* __restrict__ x, const __nv_bfloat16* __restrict__ x,
const float* __restrict__ freqs_cis, const float* __restrict__ freqs_cis,
__nv_bfloat16* __restrict__ out, __nv_bfloat16* __restrict__ out,
int batch, int n_tokens,
int seq_len,
int n_heads, int n_heads,
int head_dim int head_dim
) { ) {
const int half_dim = head_dim >> 1; const int half_dim = head_dim >> 1;
const int total = batch * seq_len * n_heads * half_dim; const int total = n_tokens * n_heads * half_dim;
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; for (int idx = blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx < total;
@@ -23,11 +22,10 @@ __global__ void rotary_emb_kernel(
int tmp = idx / half_dim; int tmp = idx / half_dim;
int head = tmp % n_heads; int head = tmp % n_heads;
tmp /= n_heads; tmp /= n_heads;
int seq = tmp % seq_len; int token = tmp;
int b = tmp / seq_len;
int x_offset = ((b * seq_len + seq) * n_heads + head) * head_dim + (pair << 1); int x_offset = (token * n_heads + head) * head_dim + (pair << 1);
int cs_offset = ((b * seq_len + seq) * half_dim + pair) * 2; int cs_offset = (token * half_dim + pair) * 2;
__nv_bfloat162 x_pair = *reinterpret_cast<const __nv_bfloat162*>(x + x_offset); __nv_bfloat162 x_pair = *reinterpret_cast<const __nv_bfloat162*>(x + x_offset);
float x_even = __bfloat162float(__low2bfloat16(x_pair)); float x_even = __bfloat162float(__low2bfloat16(x_pair));
@@ -54,27 +52,28 @@ torch::Tensor rotary_emb(
TORCH_CHECK(x.is_cuda(), "x must be on CUDA"); TORCH_CHECK(x.is_cuda(), "x must be on CUDA");
TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA"); TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16"); TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
TORCH_CHECK(x.dim() == 4, "x must be 4D [batch, seq_len, n_heads, head_dim]"); TORCH_CHECK(x.dim() == 3 || x.dim() == 4,
"x must be [tokens, n_heads, head_dim] or "
"[batch, seq_len, n_heads, head_dim]");
TORCH_CHECK(x.is_contiguous(), "x must be contiguous"); TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
TORCH_CHECK(freqs_cis.dim() == 4, "freqs_cis must be 4D [batch, seq_len, dim/2, 2]"); TORCH_CHECK(freqs_cis.dim() == x.dim(), "freqs_cis rank must match x rank");
TORCH_CHECK(freqs_cis.is_contiguous(), "freqs_cis must be contiguous"); TORCH_CHECK(freqs_cis.is_contiguous(), "freqs_cis must be contiguous");
TORCH_CHECK(freqs_cis.scalar_type() == torch::kFloat32, "freqs_cis must be f32"); TORCH_CHECK(freqs_cis.scalar_type() == torch::kFloat32, "freqs_cis must be f32");
int batch = x.size(0); int n_tokens = x.dim() == 3 ? x.size(0) : x.size(0) * x.size(1);
int seq_len = x.size(1); int n_heads = x.size(x.dim() - 2);
int n_heads = x.size(2); int head_dim = x.size(x.dim() - 1);
int head_dim = x.size(3);
TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even"); TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even");
TORCH_CHECK(freqs_cis.size(0) == batch, "freqs_cis batch mismatch"); TORCH_CHECK(freqs_cis.numel() == (int64_t)n_tokens * head_dim,
TORCH_CHECK(freqs_cis.size(1) == seq_len, "freqs_cis seq_len mismatch"); "freqs_cis token or rotary dimension mismatch");
TORCH_CHECK(freqs_cis.size(2) == head_dim / 2, "freqs_cis dim/2 mismatch"); TORCH_CHECK(freqs_cis.size(-2) == head_dim / 2, "freqs_cis dim/2 mismatch");
TORCH_CHECK(freqs_cis.size(3) == 2, "freqs_cis last dim must be 2 [cos, sin]"); TORCH_CHECK(freqs_cis.size(-1) == 2, "freqs_cis last dim must be 2 [cos, sin]");
auto out = torch::empty_like(x); auto out = torch::empty_like(x);
int half_dim = head_dim / 2; int half_dim = head_dim / 2;
int total = batch * seq_len * n_heads * half_dim; int total = n_tokens * n_heads * half_dim;
int block = 256; int block = 256;
int grid = std::min((total + block - 1) / block, 1024); int grid = std::min((total + block - 1) / block, 1024);
@@ -82,7 +81,7 @@ torch::Tensor rotary_emb(
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()), reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
freqs_cis.data_ptr<float>(), freqs_cis.data_ptr<float>(),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
batch, seq_len, n_heads, head_dim n_tokens, n_heads, head_dim
); );
C10_CUDA_CHECK(cudaGetLastError()); C10_CUDA_CHECK(cudaGetLastError());
@@ -93,6 +92,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("rotary_emb", &rotary_emb, m.def("rotary_emb", &rotary_emb,
py::arg("x"), py::arg("x"),
py::arg("freqs_cis"), py::arg("freqs_cis"),
"Fused rotary embedding (bf16 x, f32 freqs_cis [b,s,d/2,2], bf16 out)" "Fused rotary embedding for packed 3D or dense 4D tensors"
); );
} }
+147 -100
View File
@@ -7,7 +7,32 @@
#include <cstring> #include <cstring>
#include <vector> #include <vector>
#include "test_utils.cuh" #include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh" #include "../kernels/attention/dispatchers.cuh"
using namespace astrai::attention;
struct PagedDecodeDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_paged_decode<H>(p, 0); } };
struct PagedPrefillDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_paged_prefill<H>(p, 0); } };
static int make_q_tile_mapping(const std::vector<int>& q_lens,
int** d_batch, int** d_tile) {
constexpr int ROWS = 64;
std::vector<int> h_batch;
std::vector<int> h_tile;
for (int b = 0; b < (int)q_lens.size(); ++b) {
int n_tiles = (q_lens[b] + ROWS - 1) / ROWS;
for (int tile = 0; tile < n_tiles; ++tile) {
h_batch.push_back(b);
h_tile.push_back(tile);
}
}
size_t bytes = h_batch.size() * sizeof(int);
cudaMalloc(d_batch, bytes);
cudaMalloc(d_tile, bytes);
cudaMemcpy(*d_batch, h_batch.data(), bytes, cudaMemcpyHostToDevice);
cudaMemcpy(*d_tile, h_tile.data(), bytes, cudaMemcpyHostToDevice);
return (int)h_batch.size();
}
// ---- CPU reference: paged decode with variable seq_lens ---- // ---- CPU reference: paged decode with variable seq_lens ----
// Q: [B, Hq, D], K/V pool: [pool_size, Hkv, D] // Q: [B, Hq, D], K/V pool: [pool_size, Hkv, D]
@@ -15,7 +40,7 @@
// kv_indptr: [B+1]. mask: [B, max_seq_len] bool (True=keep) or NULL. // kv_indptr: [B+1]. mask: [B, max_seq_len] bool (True=keep) or NULL.
static void cpu_paged_decode_ref( static void cpu_paged_decode_ref(
const float* Q, const float* K_pool, const float* V_pool, const float* Q, const float* K_pool, const float* V_pool,
const int64_t* req_to_token, const int64_t* req_pool_indices, const int* req_to_token, const int* req_pool_indices,
const int* kv_indptr, const bool* mask, int mask_b_stride, const int* kv_indptr, const bool* mask, int mask_b_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len, int B, int Hq, int Hkv, int D, int max_ctx_len,
float* O) float* O)
@@ -24,7 +49,7 @@ static void cpu_paged_decode_ref(
int n_rep = Hq / Hkv; int n_rep = Hq / Hkv;
for (int b = 0; b < B; b++) { for (int b = 0; b < B; b++) {
int seq_len = kv_indptr[b + 1] - kv_indptr[b]; int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int64_t req_idx = req_pool_indices[b]; int req_idx = req_pool_indices[b];
#pragma omp parallel for schedule(dynamic) #pragma omp parallel for schedule(dynamic)
for (int h = 0; h < Hq; h++) { for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep; int kv_h = h / n_rep;
@@ -32,7 +57,7 @@ static void cpu_paged_decode_ref(
float accum[256] = {0.0f}; float accum[256] = {0.0f};
for (int kj = 0; kj < seq_len; kj++) { for (int kj = 0; kj < seq_len; kj++) {
if (mask && !mask[b * mask_b_stride + kj]) continue; if (mask && !mask[b * mask_b_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj]; int slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f; float dot = 0.0f;
for (int d = 0; d < D; d++) for (int d = 0; d < D; d++)
dot += Q[(b * Hq + h) * D + d] * dot += Q[(b * Hq + h) * D + d] *
@@ -63,9 +88,9 @@ static void cpu_paged_decode_ref(
// attention mask on top of the (unused) causal logic. // attention mask on top of the (unused) causal logic.
static void cpu_paged_prefill_ref( static void cpu_paged_prefill_ref(
const float* Q, const float* K_pool, const float* V_pool, const float* Q, const float* K_pool, const float* V_pool,
const int64_t* req_to_token, const int64_t* req_pool_indices, const int* req_to_token, const int* req_pool_indices,
const int* kv_indptr, const int* qo_indptr, const int* kv_indptr, const int* qo_indptr,
const bool* mask, int mask_q_stride, int mask_kv_stride, const bool* mask, int mask_l_stride, int mask_kv_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len, int causal, int B, int Hq, int Hkv, int D, int max_ctx_len, int causal,
float* O) float* O)
{ {
@@ -75,7 +100,7 @@ static void cpu_paged_prefill_ref(
int seq_len = kv_indptr[b + 1] - kv_indptr[b]; int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int q_len = qo_indptr[b + 1] - qo_indptr[b]; int q_len = qo_indptr[b + 1] - qo_indptr[b];
int causal_off = seq_len - q_len; int causal_off = seq_len - q_len;
int64_t req_idx = req_pool_indices[b]; int req_idx = req_pool_indices[b];
#pragma omp parallel for collapse(2) schedule(dynamic) #pragma omp parallel for collapse(2) schedule(dynamic)
for (int h = 0; h < Hq; h++) { for (int h = 0; h < Hq; h++) {
for (int qi = 0; qi < q_len; qi++) { for (int qi = 0; qi < q_len; qi++) {
@@ -84,9 +109,9 @@ static void cpu_paged_prefill_ref(
float accum[256] = {0.0f}; float accum[256] = {0.0f};
int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len; int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len;
for (int kj = 0; kj < lim; kj++) { for (int kj = 0; kj < lim; kj++) {
if (mask && !mask[b * mask_q_stride * mask_kv_stride if (mask && !mask[b * mask_l_stride * mask_kv_stride
+ qi * mask_kv_stride + kj]) continue; + qi * mask_kv_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj]; int slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f; float dot = 0.0f;
for (int d = 0; d < D; d++) for (int d = 0; d < D; d++)
dot += Q[(qo_indptr[b] + qi) * Hq * D + h * D + d] * dot += Q[(qo_indptr[b] + qi) * Hq * D + h * D + d] *
@@ -127,14 +152,15 @@ inline void print_paged_row(const char* cfg, float max_err, bool pass) {
// ====================================================================== // ======================================================================
template <int HEAD_DIM> template <int HEAD_DIM>
static int run_decode_test(int B, int Hq, int Hkv, int max_seq, static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
int causal, int seed) { int causal, int seed, int context_capacity = 0,
int fixed_seq_len = 0) {
// Variable seq_lens per request // Variable seq_lens per request
srand(seed); srand(seed);
std::vector<int> seq_lens(B); std::vector<int> seq_lens(B);
for (int b = 0; b < B; b++) for (int b = 0; b < B; b++)
seq_lens[b] = 8 + rand() % (max_seq - 8); seq_lens[b] = fixed_seq_len ? fixed_seq_len : 8 + rand() % (max_seq - 8);
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end()); int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
int max_ctx = max_sl + 16; int max_ctx = context_capacity ? context_capacity : max_sl + 16;
int pool_size = B * max_ctx; int pool_size = B * max_ctx;
int num_reqs = B + 4; int num_reqs = B + 4;
@@ -145,14 +171,14 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi; int *d_kvi;
float *d_op, *d_ml; float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
@@ -177,7 +203,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
// req_to_token: assign unique slots per request (scattered, not contiguous) // req_to_token: assign unique slots per request (scattered, not contiguous)
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -187,7 +213,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
// req_pool_indices: pick B random request rows // req_pool_indices: pick B random request rows
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -212,21 +238,21 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref); B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
// Kernel launch // Kernel launch
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B; p.head_dim = HEAD_DIM;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl; p.max_context_len = max_ctx;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0; p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0; p.mask = nullptr; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_q_stride = 0; p.mask_h_stride = 0; p.mask_l_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool; p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr; p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); }); dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p});
cudaDeviceSynchronize(); cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q); bf16* h_o_bf = (bf16*)malloc(sz_q);
@@ -234,7 +260,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float)); float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f; const float atol = 0.01f, rtol = 0.01f;
bool pass = true; bool pass = true;
float max_err = 0.0f; float max_err = 0.0f;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) { for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
@@ -274,15 +300,15 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_mask = (size_t)B * max_sl * sizeof(bool); size_t sz_mask = (size_t)B * max_sl * sizeof(bool);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi; int *d_kvi;
bool *d_mask; bool *d_mask;
float *d_op, *d_ml; float *d_op, *d_ml;
@@ -308,7 +334,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -317,7 +343,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
} }
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -347,21 +373,21 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
h_mask, max_sl, h_mask, max_sl,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref); B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B; p.head_dim = HEAD_DIM;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl; p.max_context_len = max_ctx;
p.causal_offset = -1; p.use_mask = 1; p.causal_offset = -1; p.use_mask = 1;
p.mask = d_mask; p.mask_b_stride = max_sl; p.mask = d_mask; p.mask_b_stride = max_sl;
p.mask_h_stride = 0; p.mask_q_stride = 0; p.mask_h_stride = 0; p.mask_l_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool; p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr; p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); }); dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p});
cudaDeviceSynchronize(); cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q); bf16* h_o_bf = (bf16*)malloc(sz_q);
@@ -369,7 +395,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float)); float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f; const float atol = 0.01f, rtol = 0.01f;
bool pass = true; bool pass = true;
float max_err = 0.0f; float max_err = 0.0f;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) { for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
@@ -413,13 +439,13 @@ static int run_prefill_test(int B, int Hq, int Hkv,
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int); size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi, *d_qoi; int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv); cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
@@ -442,7 +468,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -451,7 +477,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
} }
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -479,25 +505,28 @@ static int run_prefill_test(int B, int Hq, int Hkv,
nullptr, 0, 0, nullptr, 0, 0,
B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref); B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref);
int *d_qtb, *d_qti;
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
// Kernel launch // Kernel launch
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q; p.head_dim = HEAD_DIM;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl; p.max_context_len = max_ctx;
int max_ql = 0; p.q_len = total_q;
for (int b = 0; b < B; b++) max_ql = max(max_ql, q_lens[b]);
p.max_q_len = max_ql;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0; p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0; p.mask = nullptr; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_q_stride = 0; p.mask_h_stride = 0; p.mask_l_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool; p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; p.q_tile_to_batch = d_qtb; p.q_tile_to_index = d_qti;
p.num_q_tiles = num_q_tiles;
p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); }); dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
cudaDeviceSynchronize(); cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q); bf16* h_o_bf = (bf16*)malloc(sz_q);
@@ -505,7 +534,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float)); float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f; const float atol = 0.01f, rtol = 0.01f;
bool pass = true; bool pass = true;
float max_err = 0.0f; float max_err = 0.0f;
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) { for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
@@ -521,6 +550,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
free(h_o_ref); free(h_o_bf); free(h_o_got); free(h_o_ref); free(h_o_bf); free(h_o_got);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi); cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
cudaFree(d_qtb); cudaFree(d_qti);
return pass ? 0 : 1; return pass ? 0 : 1;
} }
@@ -544,14 +574,14 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int); size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
size_t sz_mask = (size_t)B * q_len * q_len * sizeof(bool); size_t sz_mask = (size_t)B * q_len * q_len * sizeof(bool);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi, *d_qoi; int *d_kvi, *d_qoi;
bool *d_mask; bool *d_mask;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
@@ -575,7 +605,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
int next_slot = 0; int next_slot = 0;
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) { for (int p = 0; p < max_ctx; p++) {
@@ -584,7 +614,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
} }
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
h_rpi[0] = 0; h_rpi[0] = 0;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
@@ -617,22 +647,28 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
h_mask, q_len, q_len, h_mask, q_len, q_len,
B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref); B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref);
AttentionParams<bf16> p; std::vector<int> q_lens(B, q_len);
int *d_qtb, *d_qti;
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q; p.head_dim = HEAD_DIM;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = q_len; p.max_context_len = max_ctx;
p.max_q_len = q_len; p.q_len = B * q_len;
p.causal_offset = -1; p.use_mask = 1; p.causal_offset = -1; p.use_mask = 1;
p.mask = d_mask; p.mask_b_stride = q_len * q_len; p.mask = d_mask; p.mask_b_stride = q_len * q_len;
p.mask_h_stride = 0; p.mask_q_stride = q_len; p.mask_h_stride = 0; p.mask_l_stride = q_len;
p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool; p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; p.q_tile_to_batch = d_qtb; p.q_tile_to_index = d_qti;
p.num_q_tiles = num_q_tiles;
p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); }); dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
cudaDeviceSynchronize(); cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q); bf16* h_o_bf = (bf16*)malloc(sz_q);
@@ -640,7 +676,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float)); float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]); for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f; const float atol = 0.01f, rtol = 0.01f;
bool pass = true; bool pass = true;
float max_err = 0.0f; float max_err = 0.0f;
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) { for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
@@ -657,6 +693,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi); cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
cudaFree(d_mask); cudaFree(d_mask);
cudaFree(d_qtb); cudaFree(d_qti);
return pass ? 0 : 1; return pass ? 0 : 1;
} }
@@ -665,20 +702,20 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
// ====================================================================== // ======================================================================
template <int HEAD_DIM> template <int HEAD_DIM>
static void bench_decode(int B, int Hq, int Hkv, int seq_len) { static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
int max_ctx = seq_len + 16; int max_ctx = max(16384, seq_len + 16);
int pool_size = B * max_ctx; int pool_size = B * (seq_len + 16);
int num_reqs = B; int num_reqs = B;
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float); size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float); size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi; int *d_kvi;
float *d_op, *d_ml; float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
@@ -694,12 +731,12 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size; h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi); int* h_kvi = (int*)malloc(sz_kvi);
@@ -707,21 +744,21 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len; for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice); cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B; p.head_dim = HEAD_DIM;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = seq_len; p.max_context_len = max_ctx;
p.causal_offset = 0; p.use_mask = 0; p.causal_offset = 0; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0; p.mask = nullptr; p.mask_b_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool; p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr; p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml;
auto launch = [&]() { auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); }); dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p});
}; };
// Decode: q_len=1, query is the last token → attends to all [0, seq_len). // Decode: q_len=1, query is the last token → attends to all [0, seq_len).
// FLOPs = 2 * (QK^T + PV) = 4 * B * Hq * seq_len * D. // FLOPs = 2 * (QK^T + PV) = 4 * B * Hq * seq_len * D.
@@ -747,13 +784,13 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16); size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16); size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t); size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int);
size_t sz_rpi = (size_t)B * sizeof(int64_t); size_t sz_rpi = (size_t)B * sizeof(int);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int); size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int); size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool; bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi; int *d_rtt, *d_rpi;
int *d_kvi, *d_qoi; int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q); cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv); cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
@@ -767,12 +804,12 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice); cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt); int* h_rtt = (int*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++) for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size; h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice); cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi); int* h_rpi = (int*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b; for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice); cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi); int* h_kvi = (int*)malloc(sz_kvi);
@@ -784,22 +821,28 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len; for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len;
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice); cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
AttentionParams<bf16> p; std::vector<int> q_lens(B, q_len);
int *d_qtb, *d_qti;
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q; p.head_dim = HEAD_DIM;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = kv_len; p.max_context_len = max_ctx;
p.total_q = total_q; p.max_q_len = q_len; p.q_len = B * q_len;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0; p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0; p.mask = nullptr; p.mask_b_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool; p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; p.q_tile_to_batch = d_qtb; p.q_tile_to_index = d_qti;
p.num_q_tiles = num_q_tiles;
p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
auto launch = [&]() { auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); }); dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
}; };
// FLOPs = 2 * (QK^T + PV) = 4 * effective_qk_pairs * Hq * D. // FLOPs = 2 * (QK^T + PV) = 4 * effective_qk_pairs * Hq * D.
// Non-causal: effective = q_len * kv_len. // Non-causal: effective = q_len * kv_len.
@@ -825,6 +868,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
free(tmp); free(h_rtt); free(h_rpi); free(h_kvi); free(h_qoi); free(tmp); free(h_rtt); free(h_rpi); free(h_kvi); free(h_qoi);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi); cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
cudaFree(d_qtb); cudaFree(d_qti);
} }
int main() { int main() {
@@ -844,6 +888,9 @@ int main() {
fail += run_decode_test<256>(1, 2, 1, 256, 0, 9); fail += run_decode_test<256>(1, 2, 1, 256, 0, 9);
fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10); fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10);
fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11); fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11);
// Production keeps a fixed 32768-wide request table. This forces 32
// splits, so seq_len > 512 gives each split multiple cp.async tiles.
fail += run_decode_test<64>(1, 24, 4, 1100, 0, 12, 32768, 1100);
// Decode with 2D mask (regression: mixed seq_lens + HasMask) // Decode with 2D mask (regression: mixed seq_lens + HasMask)
fail += run_decode_mask_test<128>(2, 8, 2, 256, 30); fail += run_decode_mask_test<128>(2, 8, 2, 256, 30);
@@ -928,9 +975,9 @@ int main() {
bench_decode<128>(1, 32, 4, 1024); bench_decode<128>(1, 32, 4, 1024);
bench_decode<128>(1, 32, 4, 2048); bench_decode<128>(1, 32, 4, 2048);
bench_decode<128>(1, 32, 4, 4096); bench_decode<128>(1, 32, 4, 4096);
bench_decode<128>(1, 32, 4, 16384);
bench_decode<128>(4, 32, 4, 2048); bench_decode<128>(4, 32, 4, 2048);
bench_decode<128>(16, 32, 4, 2048); bench_decode<128>(16, 32, 4, 2048);
bench_decode<128>(32, 32, 4, 1024);
printf("\n===== PAGED PREFILL BENCH =====\n"); printf("\n===== PAGED PREFILL BENCH =====\n");
print_bench_header(); print_bench_header();
+29 -16
View File
@@ -7,7 +7,12 @@ nvcc -I csrc -arch=sm_89 -O3 \
*/ */
#include "test_utils.cuh" #include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh" #include "../kernels/attention/dispatchers.cuh"
using namespace astrai::attention;
struct DecodeDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_decode<H>(p, 0); } };
struct PrefillDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_prefill<H>(p, 0); } };
// Split-K scratch (torch-free) // Split-K scratch (torch-free)
struct DecodeScratch { struct DecodeScratch {
@@ -30,8 +35,6 @@ static void free_scratch(DecodeScratch& sc) {
// ====================================================================== // ======================================================================
static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) { static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
int gs = Hq / Hk;
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D; size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV]; float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
for (size_t i=0;i<nQ;i++) hQ[i]=randf(); for (size_t i=0;i<nQ;i++) hQ[i]=randf();
@@ -55,19 +58,19 @@ static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice); cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice); cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D; p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
p.use_mask=0; p.causal_offset=causal?0:-1; p.use_mask=0; p.causal_offset=causal?0:-1;
p.scale=1.0f/sqrtf((float)D); p.scale=1.0f/sqrtf((float)D);
set_default_strides(p); set_default_strides(p);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO;
DecodeScratch sc; DecodeScratch sc;
setup_scratch(p, sc); setup_scratch(p, sc);
p.o_part = sc.o_part; p.ml_part = sc.ml_part; p.o_part = sc.o_part; p.ml_part = sc.ml_part;
double t0=now_ms(); double t0=now_ms();
dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p, 0); }); dispatch_by_head_dim(D, DecodeDispatch{p});
cudaDeviceSynchronize(); cudaDeviceSynchronize();
(void)t0; (void)t0;
cudaError_t err=cudaGetLastError(); cudaError_t err=cudaGetLastError();
@@ -117,7 +120,8 @@ static void bench_decode() {
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS); printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
print_bench_header(); print_bench_header();
for (int ci = 0; ci < 6; ci++) { int n = sizeof(cfgs) / sizeof(cfgs[0]);
for (int ci = 0; ci < n; ci++) {
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2]; int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
int sl = cfgs[ci][3], D = cfgs[ci][4]; int sl = cfgs[ci][3], D = cfgs[ci][4];
size_t nQ = (size_t)B * Hq * D; size_t nQ = (size_t)B * Hq * D;
@@ -135,18 +139,18 @@ static void bench_decode() {
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice); cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
delete[] tmp; delete[] tmp;
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl; p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1; p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
p.scale = 1.0f / sqrtf((float)D); p.scale = 1.0f / sqrtf((float)D);
set_default_strides(p); set_default_strides(p);
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO; p.q_ptr = dQ; p.k_ptr = dK; p.v_ptr = dV; p.mask = nullptr; p.o_ptr = dO;
DecodeScratch sc; DecodeScratch sc;
setup_scratch(p, sc); setup_scratch(p, sc);
p.o_part = sc.o_part; p.ml_part = sc.ml_part; p.o_part = sc.o_part; p.ml_part = sc.ml_part;
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p, 0); }); }; auto launch = [&]() { dispatch_by_head_dim(D, DecodeDispatch{p}); };
double flops = 4.0 * B * Hq * (double)sl * D; double flops = 4.0 * B * Hq * (double)sl * D;
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops); BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops);
@@ -182,15 +186,15 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]); for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice); cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
p.use_mask=0; p.causal_offset=causal?0:-1; p.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p); set_default_strides(p);
p.scale=1.0f/sqrtf((float)D); p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO;
double t0=now_ms(); double t0=now_ms();
dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p, 0); }); dispatch_by_head_dim(D, PrefillDispatch{p});
cudaDeviceSynchronize(); cudaDeviceSynchronize();
(void)t0; (void)t0;
cudaError_t err=cudaGetLastError(); cudaError_t err=cudaGetLastError();
@@ -228,6 +232,12 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
static void bench_prefill() { static void bench_prefill() {
const int cfgs[][7] = { const int cfgs[][7] = {
{1,32,4,1024,1024,32,0},
{1,32,4,1024,1024,32,1},
{1,32,4,4096,4096,32,1},
{1,32,4,1024,1024,64,0},
{1,32,4,1024,1024,64,1},
{1,32,4,4096,4096,64,1},
{1,32,4,512,512,128,0}, {1,32,4,512,512,128,0},
{1,32,4,1024,1024,128,0}, {1,32,4,1024,1024,128,0},
{1,32,4,2048,2048,128,0}, {1,32,4,2048,2048,128,0},
@@ -256,14 +266,14 @@ static void bench_prefill() {
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf()); for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice); cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
p.use_mask=0; p.causal_offset=causal?0:-1; p.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p); set_default_strides(p);
p.scale=1.0f/sqrtf((float)D); p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO; p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO;
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p, 0); }); }; auto launch = [&]() { dispatch_by_head_dim(D, PrefillDispatch{p}); };
for (int i=0;i<WARMUP;i++) launch(); for (int i=0;i<WARMUP;i++) launch();
cudaDeviceSynchronize(); cudaDeviceSynchronize();
cudaError_t err=cudaGetLastError(); cudaError_t err=cudaGetLastError();
@@ -322,7 +332,10 @@ int main() {
// ---- PREFILL ---- // ---- PREFILL ----
{ {
const int configs[][7] = { const int configs[][7] = {
{1,2,1,64,128,32,0}, // scalar fallback D=32
{1,4,2,256,256,32,1}, // causal D=32 dispatch
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal {1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,4,2,256,256,64,1}, // causal D=64 dispatch
{1,32,4,512,512,128,0}, // standard {1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium {1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal {1,4,2,256,256,128,1}, // causal
+344
View File
@@ -0,0 +1,344 @@
/*
FP8 family tests: single-warp MMA demo + full GEMM correctness.
Part 1 exercises one bf16 -> fp8 -> mma.sync m16n8k32 instruction pair
(sanity for astrai::mma_sync + the fragment layout contract).
Part 2 checks launch_fp8_gemm across all four operand layouts, both K
tiles, and ragged shapes against an fp32 CPU reference.
nvcc -I csrc -arch=sm_89 -std=c++17 -O3 csrc/tests/fp8_test.cu -o /tmp/fp8_test \
&& /tmp/fp8_test
*/
#include "test_utils.cuh"
#include <cuda_fp8.h>
#include <algorithm>
#include <cmath>
#include <cstdio>
#include <cstdlib>
#include <cuda_runtime.h>
#include <type_traits>
#include <vector>
#include "../kernels/common/mma.cuh"
#include "../kernels/fp8/gemm.cuh"
using namespace astrai::fp8;
// ---------------------------------------------------------------------------
// Part 1: single-kernel BF16 -> FP8 MMA -> BF16 demo (m16n8k32)
// ---------------------------------------------------------------------------
namespace {
constexpr int kMmaM = 16;
constexpr int kMmaN = 8;
constexpr int kMmaK = 32;
__device__ __forceinline__ unsigned pack_fp8x4(float x0, float x1, float x2,
float x3) {
__nv_fp8_e4m3 q0(x0);
__nv_fp8_e4m3 q1(x1);
__nv_fp8_e4m3 q2(x2);
__nv_fp8_e4m3 q3(x3);
return static_cast<unsigned>(q0.__x) |
(static_cast<unsigned>(q1.__x) << 8) |
(static_cast<unsigned>(q2.__x) << 16) |
(static_cast<unsigned>(q3.__x) << 24);
}
__device__ __forceinline__ unsigned load_quantize_fp8x4(
const bf16* src, float scale_inv) {
return pack_fp8x4(__bfloat162float(src[0]) * scale_inv,
__bfloat162float(src[1]) * scale_inv,
__bfloat162float(src[2]) * scale_inv,
__bfloat162float(src[3]) * scale_inv);
}
__global__ void fused_bf16_fp8_mma_kernel(
const bf16* __restrict__ a, const bf16* __restrict__ b,
bf16* __restrict__ out, float scale_a, float scale_b) {
const int lane = threadIdx.x;
const int group = lane >> 2;
const int thread_in_group = lane & 3;
const int k0 = thread_in_group * 4;
// PTX m16n8k32 A fragment: two rows, two 16-column K partitions.
unsigned a_frag[4];
a_frag[0] = load_quantize_fp8x4(&a[group * kMmaK + k0], 1.0f / scale_a);
a_frag[1] =
load_quantize_fp8x4(&a[(group + 8) * kMmaK + k0], 1.0f / scale_a);
a_frag[2] =
load_quantize_fp8x4(&a[group * kMmaK + k0 + 16], 1.0f / scale_a);
a_frag[3] = load_quantize_fp8x4(&a[(group + 8) * kMmaK + k0 + 16],
1.0f / scale_a);
// B is supplied as row-major [N,K], equivalent to the col-major [K,N]
// operand required by the MMA instruction.
unsigned b_frag[2];
b_frag[0] = load_quantize_fp8x4(&b[group * kMmaK + k0], 1.0f / scale_b);
b_frag[1] =
load_quantize_fp8x4(&b[group * kMmaK + k0 + 16], 1.0f / scale_b);
float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f};
astrai::mma_sync<__nv_fp8_e4m3>(acc, a_frag, b_frag, acc);
const int col = thread_in_group * 2;
const float output_scale = scale_a * scale_b;
*reinterpret_cast<__nv_bfloat162*>(&out[group * kMmaN + col]) =
__floats2bfloat162_rn(acc[0] * output_scale, acc[1] * output_scale);
*reinterpret_cast<__nv_bfloat162*>(&out[(group + 8) * kMmaN + col]) =
__floats2bfloat162_rn(acc[2] * output_scale, acc[3] * output_scale);
}
static float quantize_e4m3(float value) {
return static_cast<float>(__nv_fp8_e4m3(value));
}
static bool test_single_mma() {
srand(0);
std::vector<float> a(kMmaM * kMmaK), b(kMmaN * kMmaK),
reference(kMmaM * kMmaN, 0.0f);
std::vector<bf16> a_bf16(kMmaM * kMmaK), b_bf16(kMmaN * kMmaK),
output(kMmaM * kMmaN);
for (float& value : a) value = randf() * 4.0f;
for (float& value : b) value = randf() * 4.0f;
for (int i = 0; i < kMmaM * kMmaK; ++i) {
a_bf16[i] = f2bf(a[i]);
a[i] = bf2f(a_bf16[i]);
}
for (int i = 0; i < kMmaN * kMmaK; ++i) {
b_bf16[i] = f2bf(b[i]);
b[i] = bf2f(b_bf16[i]);
}
const float amax = *std::max_element(
a.begin(), a.end(),
[](float x, float y) { return fabsf(x) < fabsf(y); });
const float bmax = *std::max_element(
b.begin(), b.end(),
[](float x, float y) { return fabsf(x) < fabsf(y); });
const float scale_a = fabsf(amax) / 448.0f;
const float scale_b = fabsf(bmax) / 448.0f;
for (int row = 0; row < kMmaM; ++row) {
for (int col = 0; col < kMmaN; ++col) {
float sum = 0.0f;
for (int k = 0; k < kMmaK; ++k) {
float qa = quantize_e4m3(a[row * kMmaK + k] / scale_a);
float qb = quantize_e4m3(b[col * kMmaK + k] / scale_b);
sum = fmaf(qa, qb, sum);
}
reference[row * kMmaN + col] = sum * scale_a * scale_b;
}
}
bf16 *d_a, *d_b, *d_out;
CUDA_CHECK(cudaMalloc(&d_a, a_bf16.size() * sizeof(bf16)));
CUDA_CHECK(cudaMalloc(&d_b, b_bf16.size() * sizeof(bf16)));
CUDA_CHECK(cudaMalloc(&d_out, output.size() * sizeof(bf16)));
CUDA_CHECK(cudaMemcpy(d_a, a_bf16.data(), a_bf16.size() * sizeof(bf16),
cudaMemcpyHostToDevice));
CUDA_CHECK(cudaMemcpy(d_b, b_bf16.data(), b_bf16.size() * sizeof(bf16),
cudaMemcpyHostToDevice));
fused_bf16_fp8_mma_kernel<<<1, 32>>>(d_a, d_b, d_out, scale_a, scale_b);
CUDA_CHECK(cudaDeviceSynchronize());
CUDA_CHECK(cudaMemcpy(output.data(), d_out, output.size() * sizeof(bf16),
cudaMemcpyDeviceToHost));
float max_abs_error = 0.0f;
float max_rel_error = 0.0f;
for (int i = 0; i < kMmaM * kMmaN; ++i) {
float error = fabsf(bf2f(output[i]) - reference[i]);
max_abs_error = fmaxf(max_abs_error, error);
max_rel_error = fmaxf(
max_rel_error, error / fmaxf(fabsf(reference[i]), 1e-4f));
}
const bool pass = max_abs_error < 0.05f;
print_test_row("M=16 N=8 K=32 fused BF16->E4M3 MMA", max_abs_error,
max_rel_error, pass);
cudaFree(d_a);
cudaFree(d_b);
cudaFree(d_out);
return pass;
}
// ---------------------------------------------------------------------------
// Part 2: GEMM correctness — layouts x K-tiles vs fp32 CPU reference
// ---------------------------------------------------------------------------
// Naive fp32 reference on the GPU: same layout interpretation as the CPU
// loop it replaces (O(m*n) to check instead of O(m*n*k) to compute).
__global__ static void
naive_gemm_ref(const __nv_fp8_e4m3* a, const __nv_fp8_e4m3* b, float* out,
int m, int n, int k, int a_ld, int b_ld, int a_rm, int b_rm) {
const int i = blockIdx.y * 32 + threadIdx.y;
const int j = blockIdx.x * 32 + threadIdx.x;
if (i >= m || j >= n) return;
float acc = 0.f;
for (int kk = 0; kk < k; ++kk) {
float av = a_rm ? (float)a[i * a_ld + kk] : (float)a[kk * a_ld + i];
float bv = b_rm ? (float)b[kk * b_ld + j] : (float)b[j * b_ld + kk];
acc += av * bv;
}
out[i * n + j] = acc;
}
// Big-CTA policies for the direct-layout cases: kK/Stages vary per case;
// the fast interior loop follows the dual-congruous rule, grouped raster 8
// matches the production dispatch.
template <typename LA, typename LB>
constexpr bool kCaseFast =
!std::is_same_v<LA, ColMajor> && !std::is_same_v<LB, RowMajor>;
template <typename LA, typename LB, int kK, int Stages>
using CasePolicy =
Fp8GemmPolicy<FP8Format::E4M3, 128, 128, LA, LB, 64, 32, kK, Stages, 8,
false, kCaseFast<LA, LB>>;
template <typename LA, typename LB, int kK, int Stages>
static bool run_gemm_case(const float* ha, const float* hb, int m, int n,
int k, int a_ld, int b_ld, int dispatch = 0) {
__nv_fp8_e4m3 *da, *db;
__nv_bfloat16* dout;
float* dscale;
cudaMalloc(&da, (size_t)m * k);
cudaMalloc(&db, (size_t)n * k);
cudaMalloc(&dout, (size_t)m * n * 2);
cudaMalloc(&dscale, 4);
float one = 1.0f;
cudaMemcpy(dscale, &one, 4, cudaMemcpyHostToDevice);
// quantize inputs to e4m3 on host and upload byte-by-byte
std::vector<unsigned char> qa(m * k), qb(n * k);
for (int i = 0; i < m * k; ++i) {
__nv_fp8_e4m3 q(ha[i]);
qa[i] = *(unsigned char*)&q;
}
for (int i = 0; i < n * k; ++i) {
__nv_fp8_e4m3 q(hb[i]);
qb[i] = *(unsigned char*)&q;
}
cudaMemcpy(da, qa.data(), qa.size(), cudaMemcpyHostToDevice);
cudaMemcpy(db, qb.data(), qb.size(), cudaMemcpyHostToDevice);
FP8Params p = {};
p.a_ptr = da;
p.b_ptr = db;
p.out_ptr = dout;
p.scale = dscale;
p.m = m;
p.n = n;
p.k = k;
p.a_ld = a_ld;
p.b_ld = b_ld;
float* d_ref;
cudaMalloc(&d_ref, (size_t)m * n * 4);
naive_gemm_ref<<<dim3((n + 31) / 32, (m + 31) / 32), dim3(32, 32)>>>(
da, db, d_ref, m, n, k, a_ld, b_ld,
!std::is_same_v<LA, ColMajor>, !std::is_same_v<LB, ColMajor>);
std::vector<float> href((size_t)m * n);
cudaMemcpy(href.data(), d_ref, href.size() * 4, cudaMemcpyDeviceToHost);
cudaFree(d_ref);
if (dispatch == 1)
// Production route, NN: the dual-N-contiguous problem has no
// dedicated instantiation — canonicalize_gemm swaps to the
// transposed <ColMajor, ColMajor> kernel with its out-transposed
// epilogue (see gemm.cuh).
gemm<FP8Format::E4M3>(p, 0, false, false);
else if (dispatch == 2)
// Production route, NT: exercises plan_gemm's small/narrow/big
// selection for this shape.
gemm<FP8Format::E4M3>(p, 0, false, true);
else
launch_policy<CasePolicy<LA, LB, kK, Stages>>(p, 0);
cudaError_t e = cudaDeviceSynchronize();
if (e != cudaSuccess) {
printf(" CUDA err: %s\n", cudaGetErrorString(e));
return false;
}
std::vector<unsigned short> hb16(m * n);
cudaMemcpy(hb16.data(), dout, (size_t)m * n * 2, cudaMemcpyDeviceToHost);
const float tol = 0.06f;
double max_rel = 0;
bool ok = true;
for (int i = 0; i < m && ok; ++i) {
for (int j = 0; j < n && ok; ++j) {
const float ref = href[(size_t)i * n + j];
float got =
__bfloat162float(__ushort_as_bfloat16(hb16[i * n + j]));
float err = fabsf(got - ref);
float rel = err / fmaxf(fabsf(ref), 0.5f);
if (rel > max_rel) max_rel = rel;
if (err > tol * fmaxf(fabsf(ref), 1.0f)) ok = false;
}
}
printf(" max_rel=%.4f %s\n", max_rel, ok ? "PASS" : "FAIL");
cudaFree(da);
cudaFree(db);
cudaFree(dout);
cudaFree(dscale);
return ok;
}
static bool test_gemm() {
struct {
int m, n, k;
} cfgs[] = {
{128, 128, 128}, {256, 128, 256}, {128, 256, 64},
{100, 130, 96}, {64, 64, 160}, {300, 200, 320},
{2048, 256, 512}, {1024, 1024, 512},
};
bool all = true;
for (auto& c : cfgs) {
float* ha = new float[c.m * c.k];
float* hb_rowmajor = new float[c.k * c.n]; // [K][N] for B RowMajor
float* hb_colmajor = new float[c.n * c.k]; // [N][K] for B ColMajor
for (int i = 0; i < c.m * c.k; ++i) ha[i] = randf();
for (int i = 0; i < c.k * c.n; ++i) hb_rowmajor[i] = randf();
for (int i = 0; i < c.k * c.n; ++i)
hb_colmajor[i / c.k * c.k + i % c.k] = hb_rowmajor[i];
float* ha_t = new float[c.k * c.m]; // [K][M] for A ColMajor
for (int i = 0; i < c.m; ++i)
for (int p = 0; p < c.k; ++p) ha_t[p * c.m + i] = ha[i * c.k + p];
printf("%dx%dx%d:\n", c.m, c.n, c.k);
printf(" NT K32:");
all &= run_gemm_case<RowMajor, ColMajor, 32, 3>(ha, hb_colmajor, c.m,
c.n, c.k, c.k, c.k);
printf(" NT K64:");
all &= run_gemm_case<RowMajor, ColMajor, 64, 2>(ha, hb_colmajor, c.m,
c.n, c.k, c.k, c.k);
printf(" NN swap:");
all &= run_gemm_case<RowMajor, RowMajor, 64, 2>(
ha, hb_rowmajor, c.m, c.n, c.k, c.k, c.n, /*dispatch=*/1);
printf(" NT disp:");
all &= run_gemm_case<RowMajor, ColMajor, 64, 2>(
ha, hb_colmajor, c.m, c.n, c.k, c.k, c.k, /*dispatch=*/2);
printf(" TN K32:");
all &= run_gemm_case<ColMajor, ColMajor, 32, 3>(ha_t, hb_colmajor, c.m,
c.n, c.k, c.m, c.k);
printf(" TN K64:");
all &= run_gemm_case<ColMajor, ColMajor, 64, 2>(ha_t, hb_colmajor, c.m,
c.n, c.k, c.m, c.k);
printf(" TT K64:");
all &= run_gemm_case<ColMajor, RowMajor, 64, 2>(ha_t, hb_rowmajor, c.m,
c.n, c.k, c.m, c.n);
delete[] ha;
delete[] hb_rowmajor;
delete[] hb_colmajor;
delete[] ha_t;
}
return all;
}
} // namespace
int main() {
print_test_header();
bool ok = test_single_mma();
ok &= test_gemm();
printf(ok ? "All PASS\n" : "FAILURES\n");
return ok ? 0 : 1;
}
+14 -14
View File
@@ -107,29 +107,29 @@ void dispatch_by_head_dim(int head_dim, Fn&& fn) {
// Set default strides for contiguous b h l d layout on AttentionParams. // Set default strides for contiguous b h l d layout on AttentionParams.
template<typename P> template<typename P>
inline void set_default_strides(P& p) { inline void set_default_strides(P& p) {
p.q_stride_b = p.q_head * p.q_len * p.head_dim; p.q_b_stride = p.q_head * p.q_len * p.head_dim;
p.q_stride_h = p.q_len * p.head_dim; p.q_h_stride = p.q_len * p.head_dim;
p.q_stride_l = p.head_dim; p.q_l_stride = p.head_dim;
p.q_stride_d = 1; p.q_d_stride = 1;
p.kv_stride_b = p.kv_head * p.kv_len * p.head_dim; p.kv_b_stride = p.kv_head * p.kv_len * p.head_dim;
p.kv_stride_h = p.kv_len * p.head_dim; p.kv_h_stride = p.kv_len * p.head_dim;
p.kv_stride_l = p.head_dim; p.kv_l_stride = p.head_dim;
p.kv_stride_d = 1; p.kv_d_stride = 1;
p.mask_b_stride = p.kv_len; p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
} }
// Set default Q strides for a paged decode params struct. // Set default Q strides for a paged decode params struct.
template<typename P> template<typename P>
inline void set_default_paged_strides(P& p) { inline void set_default_paged_strides(P& p) {
p.q_stride_b = p.q_head * p.q_len * p.head_dim; p.q_b_stride = p.q_head * p.q_len * p.head_dim;
p.q_stride_h = p.q_len * p.head_dim; p.q_h_stride = p.q_len * p.head_dim;
p.q_stride_l = p.head_dim; p.q_l_stride = p.head_dim;
p.q_stride_d = 1; p.q_d_stride = 1;
p.mask_b_stride = p.kv_len; p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0; p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_l_stride = 0;
} }
// Generic CPU reference for multi-query / grouped-query attention. // Generic CPU reference for multi-query / grouped-query attention.
+56 -7
View File
@@ -1,22 +1,27 @@
services: services:
server: server:
image: astrai:latest
build: build:
context: . context: .
dockerfile: Dockerfile dockerfile: Dockerfile
args: args:
CUDA_TAG: ${CUDA_TAG:-cu128} CUDA_TAG: ${CUDA_TAG:-cu128}
user: "${UID:-1000}:${GID:-1000}" USER_UID: ${ASTRAI_UID:-1000}
USER_GID: ${ASTRAI_GID:-1000}
user: "${ASTRAI_UID:-1000}:${ASTRAI_GID:-1000}"
ports: ports:
- "8000:8000" - "${SERVE_PORT:-8000}:${SERVE_CONTAINER_PORT:-8000}"
volumes: volumes:
- ./params:/app/params:ro - ${SERVE_PARAM_DIR:-./params}:/app/params:ro
environment:
- CUDA_VISIBLE_DEVICES
command: python -m scripts.tools.server --port 8000 --device cuda command: python -m scripts.tools.server --port 8000 --device cuda
deploy: deploy:
resources: resources:
reservations: reservations:
devices: devices:
- driver: nvidia - driver: nvidia
count: 1 count: all
capabilities: [gpu] capabilities: [gpu]
healthcheck: healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8000/health"] test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
@@ -27,17 +32,20 @@ services:
restart: unless-stopped restart: unless-stopped
server-cpu: server-cpu:
image: astrai:latest
profiles: [cpu] profiles: [cpu]
build: build:
context: . context: .
dockerfile: Dockerfile dockerfile: Dockerfile
args: args:
CUDA_TAG: ${CUDA_TAG:-cu128} CUDA_TAG: ${CUDA_TAG:-cu128}
user: "${UID:-1000}:${GID:-1000}" USER_UID: ${ASTRAI_UID:-1000}
USER_GID: ${ASTRAI_GID:-1000}
user: "${ASTRAI_UID:-1000}:${ASTRAI_GID:-1000}"
ports: ports:
- "8000:8000" - "${SERVE_PORT:-8000}:${SERVE_CONTAINER_PORT:-8000}"
volumes: volumes:
- ./params:/app/params:ro - ${SERVE_PARAM_DIR:-./params}:/app/params:ro
command: python -m scripts.tools.server --port 8000 --device cpu command: python -m scripts.tools.server --port 8000 --device cpu
healthcheck: healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8000/health"] test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
@@ -46,3 +54,44 @@ services:
retries: 3 retries: 3
start_period: 120s start_period: 120s
restart: unless-stopped restart: unless-stopped
trainer:
image: astrai:latest
profiles: [train]
build:
context: .
dockerfile: Dockerfile
args:
CUDA_TAG: ${CUDA_TAG:-cu128}
USER_UID: ${ASTRAI_UID:-1000}
USER_GID: ${ASTRAI_GID:-1000}
init: true
user: "${ASTRAI_UID:-1000}:${ASTRAI_GID:-1000}"
volumes:
- ${TRAIN_DATA_DIR:-./data}:/data:ro
- ${TRAIN_MODEL_DIR:-./params}:/models/base:ro
- ${TRAIN_CHECKPOINT_DIR:-./checkpoints}:/checkpoints
environment:
- TRAIN_JOB_NAME=${TRAIN_JOB_NAME:-astrai-train}
- TRAIN_CONFIG=${TRAIN_CONFIG:-}
- BASE_MODEL=${BASE_MODEL:-/models/base}
- CHECKPOINT_ROOT=/checkpoints
- TRAIN_GPU_COUNT=${TRAIN_GPU_COUNT:-all}
- TRAIN_PARALLEL_MODE=${TRAIN_PARALLEL_MODE:-auto}
- CUDA_VISIBLE_DEVICES
entrypoint: ["bash", "/app/scripts/docker/train-entrypoint.sh"]
ipc: ${TRAIN_IPC_MODE:-host}
stop_grace_period: ${TRAIN_STOP_GRACE_PERIOD:-10m}
restart: "no"
logging:
driver: json-file
options:
max-size: ${TRAIN_LOG_MAX_SIZE:-100m}
max-file: ${TRAIN_LOG_MAX_FILES:-5}
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: all
capabilities: [gpu]
+7 -4
View File
@@ -57,7 +57,7 @@ AstrAI 是一个覆盖模型构建、训练、评测与部署的端到端 Transf
| **数据** | 声明式 JSON 预处理、可配置掩码与样本打包、二进制/JSONL 存储和流式数据集 | | **数据** | 声明式 JSON 预处理、可配置掩码与样本打包、二进制/JSONL 存储和流式数据集 |
| **推理** | 连续批处理、分页 KV Cache、Radix 前缀缓存、流式生成,以及 Torch/CUDA/FlashAttention 后端 | | **推理** | 连续批处理、分页 KV Cache、Radix 前缀缓存、流式生成,以及 Torch/CUDA/FlashAttention 后端 |
| **服务** | 基于 FastAPI 的 OpenAI 与 Anthropic 聊天补全协议,支持 SSE 流式输出和工具调用 | | **服务** | 基于 FastAPI 的 OpenAI 与 Anthropic 聊天补全协议,支持 SSE 流式输出和工具调用 |
| **评测** | Perplexity、MMLU、HumanEval、IFEval、IFDROUGE 评测工具 | | **评测** | Perplexity、MMLU、HumanEval、IFEval、IFDROUGE 和权重分析评测工具 |
| **扩展** | 基于工厂与注册表扩展模型、数据集、训练策略、回调、内核和协议组件 | | **扩展** | 基于工厂与注册表扩展模型、数据集、训练策略、回调、内核和协议组件 |
### 快速上手 ### 快速上手
@@ -71,8 +71,9 @@ AstrAI 需要 Python 3.12+,并精确固定 PyTorch 版本为 `2.11.0`。训练
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
pip install -e . # 纯 PyTorch(不含 CUDA 内核 pip install -e . # 检测到 nvcc + CUDA 时自动构建内核
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 可选:融合 CUDA 内核加速 # CSRC_KERNELS=false pip install -e . # 跳过内核(纯 PyTorch
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 强制构建融合 CUDA 内核
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff # pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff
``` ```
@@ -187,7 +188,7 @@ docker run --gpus all -it astrai:latest
# 运行推理服务 # 运行推理服务
docker run --gpus all -p 8000:8000 astrai:latest \ docker run --gpus all -p 8000:8000 astrai:latest \
python -m scripts.tools.server --port 8000 --device cuda python scripts/tools/server.py --port 8000 --device cuda
# 挂载数据卷 # 挂载数据卷
docker run --gpus all -v /path/to/data:/data -it astrai:latest docker run --gpus all -v /path/to/data:/data -it astrai:latest
@@ -242,6 +243,8 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
| [数据流程](./developer/dataflow.md) | 数据管道、存储后端与数据集架构 | | [数据流程](./developer/dataflow.md) | 数据管道、存储后端与数据集架构 |
| [内部实现](./developer/internals.md) | 训练原理:损失公式、回调生命周期、KV Cache | | [内部实现](./developer/internals.md) | 训练原理:损失公式、回调生命周期、KV Cache |
| [CUDA 内核](./developer/cuda_kernels.md) | 自定义 CUDA 注意力内核与基准测试 | | [CUDA 内核](./developer/cuda_kernels.md) | 自定义 CUDA 注意力内核与基准测试 |
| [Docker 服务部署](./developer/docker-serving.md) | YAML 驱动的容器化服务(`serve.yaml``serve.sh` |
| [Docker 训练部署](./developer/docker-training.md) | YAML 驱动的容器化训练(`train.yaml``train.sh` |
### 贡献 ### 贡献
+53 -27
View File
@@ -4,7 +4,7 @@
- [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces - [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces
- [Module Overview](#module-overview) — Component inventory per module - [Module Overview](#module-overview) — Component inventory per module
- [Design Patterns](#design-patterns) — 15 documented patterns with classes - [Design Patterns](#design-patterns) — 16 documented patterns with classes
- [Core Relationships](#core-relationships) — 11 key inter-component relationships - [Core Relationships](#core-relationships) — 11 key inter-component relationships
## Class Diagram ## Class Diagram
@@ -315,6 +315,7 @@ classDiagram
<<TypedDict>> <<TypedDict>>
+Tensor hidden_states +Tensor hidden_states
+Optional[Tensor] aux_loss +Optional[Tensor] aux_loss
+Optional[RouterStats] router_stats
} }
class GQA { class GQA {
@@ -361,6 +362,7 @@ classDiagram
<<TypedDict>> <<TypedDict>>
+Tensor hidden_states +Tensor hidden_states
+Optional[Tensor] aux_loss +Optional[Tensor] aux_loss
+Optional[RouterStats] router_stats
} }
class DeepSeekMoE { class DeepSeekMoE {
@@ -807,7 +809,6 @@ classDiagram
+AutoTokenizer tokenizer +AutoTokenizer tokenizer
+InferenceScheduler scheduler +InferenceScheduler scheduler
+generate(prompt, stream, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) Union[Generator, str, List[str]] +generate(prompt, stream, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) Union[Generator, str, List[str]]
+generate_with_request(request) Union[Generator, str, List[str]]
+generate_async(prompt, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) AsyncGenerator +generate_async(prompt, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) AsyncGenerator
+get_stats() Dict +get_stats() Dict
+shutdown() +shutdown()
@@ -815,8 +816,8 @@ classDiagram
class Executor { class Executor {
+AutoModel model +AutoModel model
+AutoTokenizer tokenizer
+PagePool kv_cache +PagePool kv_cache
+TaskCacheManager task_cache
+InferenceWorkspace _workspace +InferenceWorkspace _workspace
+Optional[str] device +Optional[str] device
+Optional[torch.dtype] dtype +Optional[torch.dtype] dtype
@@ -844,6 +845,7 @@ classDiagram
class InferenceScheduler { class InferenceScheduler {
+PagePool _cache +PagePool _cache
+TaskCacheManager _task_cache
+Executor _executor +Executor _executor
+TaskManager _task_mgr +TaskManager _task_mgr
+Event _stop_event +Event _stop_event
@@ -887,6 +889,24 @@ classDiagram
+release(pages) +release(pages)
} }
class AllocationStrategy {
<<abstract>>
+alloc(state, prompt_ids) bool
+free(state)
+extend(state, pos) bool
+write_indices(state, prompt_ids)
+record_hashes(state, prompt_ids, start_logical_page)
}
class ContiguousStrategy {
+write_indices(state, prompt_ids)
}
class PagedStrategy {
-Allocator _alloc
-RadixCache _prefix
}
class KVStorage { class KVStorage {
+int size +int size
+Tensor k_buffer +Tensor k_buffer
@@ -915,6 +935,9 @@ classDiagram
+int max_len +int max_len
+Optional[Tensor] kv_indptr +Optional[Tensor] kv_indptr
+Optional[Tensor] qo_indptr +Optional[Tensor] qo_indptr
+Optional[Tensor] decode_o_part
+Optional[Tensor] decode_ml_part
+Optional[Tensor] decode_out
} }
class PagePool { class PagePool {
@@ -922,14 +945,21 @@ classDiagram
+bool contiguous +bool contiguous
-KVStorage _storage -KVStorage _storage
-ReqToTokenPool _req_pool -ReqToTokenPool _req_pool
-Allocator _alloc -AllocationStrategy _strategy
-RadixCache _prefix +strategy AllocationStrategy
+req_pool ReqToTokenPool
+bind_tasks(req_indices, seq_lens, workspace, device, start_pos, incremental) KVCache
}
class TaskCacheManager {
-PagePool _pool
-Dict _states
+task_alloc(task_id, prompt_ids) bool +task_alloc(task_id, prompt_ids) bool
+task_free(task_id) +task_free(task_id)
+task_extend(task_id, pos) bool +task_extend(task_id, pos) bool
+task_cached(task_id) int +task_cached(task_id) int
+task_record_hashes(task_id, prompt_ids, start_logical_page) +task_record_hashes(task_id, prompt_ids, start_logical_page)
+bind_tasks(task_ids, workspace, device, start_pos) KVCache +bind(task_ids, workspace) KVCache
} }
class Task { class Task {
@@ -980,17 +1010,6 @@ classDiagram
+get_stats() Dict +get_stats() Dict
} }
class GenerationRequest {
+List[Dict] messages
+int top_k
+float top_p
+float temperature
+Optional[int] max_tokens
+float frequency_penalty
+int rep_window
+bool stream
}
class BaseSamplingStrategy { class BaseSamplingStrategy {
<<abstract>> <<abstract>>
+apply(logits, filter_value, input_ids, input_mask) Tensor +apply(logits, filter_value, input_ids, input_mask) Tensor
@@ -1323,17 +1342,22 @@ classDiagram
PositionIdStrategy <|-- DocResetPositionId PositionIdStrategy <|-- DocResetPositionId
PositionIdStrategy <|-- ContinuousPositionId PositionIdStrategy <|-- ContinuousPositionId
StoreWriter <|-- BinWriter StoreWriter <|-- BinWriter
AllocationStrategy <|-- ContiguousStrategy
AllocationStrategy <|-- PagedStrategy
RawRollout <|-- RolloutResult RawRollout <|-- RolloutResult
LaunchStrategy <|-- TorchrunStrategy LaunchStrategy <|-- TorchrunStrategy
LaunchStrategy <|-- LocalStrategy LaunchStrategy <|-- LocalStrategy
%% --- Composition (strong ownership, part destroyed with whole) --- %% --- Composition (strong ownership, part destroyed with whole) ---
PagePool *-- KVStorage PagePool *-- KVStorage
PagePool *-- ReqToTokenPool PagePool *-- ReqToTokenPool
PagePool *-- Allocator PagePool *-- AllocationStrategy
PagePool *-- RadixCache PagedStrategy *-- Allocator
PagedStrategy *-- RadixCache
TaskCacheManager o-- PagePool
RadixCache *-- RadixNode RadixCache *-- RadixNode
InferenceEngine *-- InferenceScheduler InferenceEngine *-- InferenceScheduler
InferenceScheduler *-- PagePool InferenceScheduler *-- PagePool
InferenceScheduler *-- TaskCacheManager
InferenceScheduler *-- Executor InferenceScheduler *-- Executor
Executor *-- InferenceWorkspace Executor *-- InferenceWorkspace
InferenceScheduler *-- TaskManager InferenceScheduler *-- TaskManager
@@ -1407,7 +1431,7 @@ classDiagram
CheckpointCallback ..> Checkpoint : creates CheckpointCallback ..> Checkpoint : creates
PagePool ..> KVCache : binds PagePool ..> KVCache : binds
PagePool ..> InferenceWorkspace : fills PagePool ..> InferenceWorkspace : fills
InferenceEngine ..> GenerationRequest : uses InferenceEngine ..> GenerateResult : uses
InferenceEngine ..> GenerateResult : creates InferenceEngine ..> GenerateResult : creates
OpenAIResponseBuilder ..> ChatCompletionRequest : receives OpenAIResponseBuilder ..> ChatCompletionRequest : receives
AnthropicResponseBuilder ..> MessagesRequest : receives AnthropicResponseBuilder ..> MessagesRequest : receives
@@ -1426,7 +1450,7 @@ classDiagram
Task --> TaskStatus Task --> TaskStatus
InferenceEngine --> AutoModel InferenceEngine --> AutoModel
Executor --> AutoModel Executor --> AutoModel
Executor --> AutoTokenizer Executor --> TaskCacheManager
TaskManager --> AutoTokenizer TaskManager --> AutoTokenizer
``` ```
@@ -1443,8 +1467,9 @@ classDiagram
| **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model | | **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template | | **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow | | **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, InferenceWorkspace, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, RadixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service | | **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, InferenceWorkspace, PagePool, TaskCacheManager, KVStorage, ReqToTokenPool, KVCache, Allocator, RadixCache, AllocationStrategy, ContiguousStrategy, PagedStrategy, Task, TaskManager, TaskStatus, StreamDecoder, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
| **astrai.extension** | AttentionBackend, TorchNativeBackend, CudaBackend, attn_backend, ATTN_BACKEND, attn_decode, attn_prefill, attn_paged_decode, attn_paged_prefill, rotary_emb, apply_rotary_emb, rotary_backend, is_available | CUDA attention + rotary kernels, backend abstraction, auto-dispatch | | **astrai.extension** | `backend` policy package, `ops` kernel-wrapper package, `fp8.py` FP8 strategy layer, AttentionBackend, TorchNativeBackend, CudaBackend, FlashAttnBackend, attention, attn_backend, ATTN_BACKEND, apply_rotary_emb, is_available | Stable API over attention/rotary/FP8 execution policy and optional CUDA kernels |
| **astrai.optim** | OptimizerFactory, MuonAdamW, NoraNadamW, ManoAdamW, composite_step/composite_zero_grad/composite_state_dict, partition_optimizer_parameters | Built-in optimizers (`muon_adamw` / `nora_nadamw` / `mano_adamw`) with shared composite-optimizer helpers |
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation | | **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation |
| **astrai.factory** | BaseFactory | Component registration | | **astrai.factory** | BaseFactory | Component registration |
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers | | **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
@@ -1462,12 +1487,13 @@ classDiagram
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring | | **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
| **Context** | `TrainContext` | Unified training state bag | | **Context** | `TrainContext` | Unified training state bag |
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction | | **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
| **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager | | **Strategy (Attention)** | `AttentionBackend`, `CudaBackend`, `FlashAttnBackend`, `TorchNativeBackend` | Attention computation backend switching via context manager |
| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback | | **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `backend/rotary.py`, `ops/rotary.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback |
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution | | **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
| **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support | | **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching | | **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
| **Model Registry** | `ModelFactory`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading | | **Model Registry** | `ModelFactory`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
| **Optimizer Routing** | `OptimizerFactory`, `MuonAdamW`, `NoraNadamW`, `ManoAdamW` | Route parameter groups (matrices vs. embeddings/heads/norms) through different optimizers |
## Core Relationships ## Core Relationships
@@ -1475,7 +1501,7 @@ classDiagram
2. **Training Flow**: `Trainer``TrainContextBuilder``TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution 2. **Training Flow**: `Trainer``TrainContextBuilder``TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type` 3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor` 4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels). Rotary embedding auto-dispatches to CUDA kernel when available (inference mode), else torch complex multiply (training). 5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. `astrai.extension.backend` owns attention/rotary dispatch, fallback, and KV cache policy; it calls the stateless compiled-kernel wrappers in `astrai.extension.ops`. Attention uses cuda > flash > torch priority unless explicitly selected by `ASTR_BACKEND` or `attn_backend()`. Rotary embedding auto-dispatches to the CUDA op when supported, else torch complex multiply.
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP 6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (`MmapStore`/`JsonlStore`) loads data with explicit `_length` and multi-segment `_data` 7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (`MmapStore`/`JsonlStore`) loads data with explicit `_length` and multi-segment `_data`
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata; `CheckpointCallback` performs rank-0 training saves, with extra state saved as `{key}.pt` 8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata; `CheckpointCallback` performs rank-0 training saves, with extra state saved as `{key}.pt`
@@ -1483,4 +1509,4 @@ classDiagram
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops 10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers 11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
> Document Update Time: 2026-08-02 > Document Update Time: 2026-08-29
+343 -44
View File
@@ -1,40 +1,140 @@
# CUDA Kernels # CUDA Kernels
AstrAI includes optional custom CUDA kernels for attention and rotary embedding. These are built when `nvcc` is available and CUDA is detected, and are dispatched via the `CudaBackend` attention backend or auto-dispatched for rotary. AstrAI includes optional custom CUDA kernels for attention, rotary embedding, and FP8 GEMM. These are built when `nvcc` is available and CUDA is detected, and are dispatched via the `CudaBackend` attention backend, auto-dispatched for rotary, or invoked through the FP8 linear primitives.
## Overview ## Overview
| Kernel | File | Description | | Kernel | File | Description |
|--------|------|-------------| |--------|------|-------------|
| `attn_decode` | `attn_decode.cu` | GQA decode attention (split-KV) | | `attn_decode` | `attention/decode.cu` | GQA decode attention (split-KV) |
| `attn_prefill` | `attn_prefill.cu` | GQA prefill attention (split-Q) | | `attn_prefill` | `attention/prefill.cu` | GQA prefill attention (split-Q) |
| `attn_paged_decode` | `attn_paged_decode.cu` | Paged KV cache decode attention | | `attn_paged_decode` | `attention/paged_decode.cu` | Paged KV cache decode attention |
| `attn_paged_prefill` | `attn_paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) | | `attn_paged_prefill` | `attention/paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) |
| `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) | | `rotary_emb` | `rotary/rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
| `fp8_ops` | `fp8/ops.cu` | FP8 quantization + tensor-core GEMM (sm_89+) |
Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist: Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist:
| Variant | File | Optimization | | Variant | File | Optimization |
|---------|------|--------------| |---------|------|--------------|
| Split-KV MMA decode | `attn_decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) | | Split-KV MMA decode | `attention/decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) |
| Split-Q MMA prefill | `attn_prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) | | Split-Q MMA prefill | `attention/prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) |
> The paged and non-paged paths are ONE kernel templated on a `KVSource` > The paged and non-paged paths share one kernel body. Prefill is templated on
> policy (`ContigKV` / `PagedKV` in `attn_kv_source.cuh`); there are no > an independent Q schedule (`DenseQSchedule` / `PackedQSchedule`) and KV
> separate `attn_paged_*.cuh` files anymore. > source (`ContigKV` / `PagedKV`); decode only needs the KV source. There are
> no separate `attn_paged_*.cuh` files.
### Rotary Embedding Kernel ### Rotary Embedding Kernel
The `rotary_emb` kernel (`csrc/kernels/rotary_emb.cu`) fuses cos/sin lookup and rotation into a single kernel: The `rotary_emb` kernel (`csrc/kernels/rotary/rotary_emb.cu`) fuses cos/sin lookup and rotation into a single kernel:
- One thread per (head, dim-pair), vectorized `__nv_bfloat162` load/store - One thread per (head, dim-pair), vectorized `__nv_bfloat162` load/store
- f32 cos/sin input, bf16 compute and output - f32 cos/sin input, bf16 compute and output
- 256-thread blocks, grid-stride loop - 256-thread blocks, grid-stride loop
- Auto-dispatched via `apply_rotary_emb` in `astrai/extension/rotary_backend.py` (CUDA when available + inference mode, else torch complex-multiply fallback) - Auto-dispatched via `apply_rotary_emb` in `astrai/extension/backend/rotary.py` (CUDA when available + inference mode, else torch complex-multiply fallback)
- No context-manager backend needed — rotary is backend-agnostic, both attention backends benefit - No context-manager backend needed — rotary is backend-agnostic, both attention backends benefit
Standalone benchmark vs torch complex-multiply (48 calls = 24 layers × q+k): 6-9x faster, max diff 0 (decode) to 3e-2 (large prefill, bf16). Standalone benchmark vs torch complex-multiply (48 calls = 24 layers × q+k): 6-9x faster, max diff 0 (decode) to 3e-2 (large prefill, bf16).
### FP8 GEMM / Linear Kernel
The `fp8_ops` family (`csrc/kernels/fp8/`) accelerates bf16 linear layers by
quantizing to FP8 and running tensor-core GEMMs (**requires sm_89+**; fp8
`mma.sync.m16n8k32` only exists on Ada/Hopper). Same three-layer style as
attention; the GEMM device code is split humming/CUTLASS-style into one
layered directory:
| File | Role |
|------|------|
| `fp8/common.h` | `FP8Format` enum (E4M3/E5M2), `Fp8GemmTraits<Fmt, BlockM, BlockN, K, Stages>`, `FP8Params` / `FP8QuantizeParams` PODs, layout tags — no torch |
| `fp8/quantize.cuh` | pure-CUDA device code: vectorized `fp8_quantize_kernel` + 32×32-tile transpose kernel (out_layout 0/1/2), `quant_in_traits<InT>` unpack — no torch |
| `fp8/gemm/policy.cuh` | smem budget / occupancy hint (`Fp8GemmSmem`) + `Fp8GemmPolicy` (traits + layouts + knobs — the kernel's single template parameter) |
| `fp8/gemm/load.cuh` | operand loaders: swizzle (`tile_at`), congruous cp.async (predicated + interior), `PrefetchCarry`, crosswise LDG+PRMT direct load |
| `fp8/gemm/scheduler.cuh` | CTA id → (block_m, block_n) grouped/plain raster |
| `fp8/gemm/mainloop.cuh` | `Fp8CollectiveMainloop`: stage rings, stage loads, fragment addressing, pipelined mma.sync loop |
| `fp8/gemm/epilogue.cuh` | `Fp8CollectiveEpilogue`: fused bias + bf16 smem scatter + coalesced copy-out |
| `fp8/gemm.cuh` | umbrella: `fp8_gemm_kernel<Policy>` orchestrator + host planning (`plan_gemm` / `launch_plan`; 64×64 / 128×64 / 128×128 CTA) + entry `gemm<Fmt>(params, stream, trans_a, trans_b)` = `canonicalize_gemm``plan_gemm``launch_plan` |
| `fp8/ops.cu` | binding only: `check_fp8_device` (sm_89+), param packing, launch dispatch, pybind → module `fp8_ops` |
Scale semantics: `quantize` takes the quantization *multiplier*; the
strategy layer passes `scale.reciprocal()` and the kernel multiplies by it.
`mm_fp8` takes the combined dequant scale (`sa * sb`). `amax` is always
returned in the original input domain.
Python layer (two levels): `astrai/extension/ops/fp8.py` provides stateless
primitives (`fp8_quantize` / `fp8_gemm`) via `torch.library.custom_op`, with
plain `quantize` / `mm_fp8` wrappers, and `astrai/extension/fp8.py` is the
strategy layer (`fp8_autocast`, delayed / dynamic scaling recipes,
`fp8_linear_forward/backward` wiring `aten::linear` on CUDA). See the FP8
section in `AGENTS.md` for full detail.
#### FP8 GEMM design notes
The load-bearing invariants behind the kernel code (all measurements on
L20/sm_89 unless noted):
**Swizzle.** Staging tiles are flat `[rows * kK]`; `tile_at` XORs the 16B
chunk index with row bits at `[3, 3+log2(kChunks))` so a warp's ldmatrix
fragment load (8 consecutive rows × 16B) hits all 32 banks exactly once
(the unswizzled row word-stride is `kK/4` words, so rows `r` and
`r + 8/kChunks` collide mod 32). Chunks stay contiguous, so cp.async
staging is unaffected.
**Fragment addressing (base-pair scheme).** One base register per operand
per k_seg, every fragment offset an LDSM immediate. The closure works
because the XOR swizzle's source bits come only from the lane's
row-within-matrix `r7`: the 8/16-row fragment steps never reach them, so
`addr(s, mt) = lane_base + mt*(16*kK) ^ (s<<5)` for A and
`addr(s, nt) = lane_base + nt*(8*kK) ^ (s<<5)` for B. This replaced
runtime offset tables that spilled at 131 registers (~55 of 146 hot-loop
instructions were address math; cuBLAS's inner loop has ~0). Steady-state
read pointers advance one stage per iteration with an equality wrap,
replacing the per-k-tile `(tile % ring) * stage_bytes` recomputation
(UIMAD.WIDE magic-division ladder).
**Pipeline depth and barriers.** Every operand ring holds `kStages+1`
buffers: the load for tile `i+kStages` targets slot `(i-1)%(kStages+1)`,
which compute(i-1) finished reading before this iteration's barrier — no
post-compute barrier, one `__syncthreads` per k-tile. Prologue and tail
commits are unconditional so the group sequence stays tile-indexed and the
fixed `wait_group<kStages-1>` is iteration-invariant (a runtime
wait-count dispatch ladder cost 16 instructions/k-tile). A lean
`kStages`-deep ring trading the barrier for a 4th resident CTA measured
+5..9% slower at 1280³ and was removed.
**Crosswise loads.** Crosswise operands (A `[K][M]` / B `[N][K]` storage)
cannot cp.async into the canonical tile; they take the direct LDG.128×4 +
in-register PRMT transpose + STS.32 path. A staged variant (cp.async into
K-major staging + per-tile smem→smem transpose) measured 15-20% slower
across every probed shape including DRAM-streaming B (git history 5745c2f).
**Fast-loop peel.** When both operands are congruous, the whole CTA is
interior, base|ld is 16B-aligned and K has no tail, the mainloop switches
to a predication-free copy with loop-carried prefetch state: +4.5..10% on
the issue-bound 64×64 CTA (256³..1024³), 3% on the 128×128 CTA, so only
the small CTA opts in.
**Launch planning crossovers** (L20, TFLOPS, big vs alternative):
crosswise problems keep the 64×64 s3 CTA below ~1.5 waves of 128×128
tiles (M=256: 129.7 vs 113.1; 1024³: 107.2 vs 94.8; the big CTA wins from
M=640/1536³ on). Dual-congruous wave band picks narrow vs big by
`ceil(tiles/sm) * T_tile` with `T_narrow ≈ 0.53 * T_big` (M=384: 134.3 vs
114.4 narrow wins; M=1024: 202.5 vs 178.8 big wins). Sub-wave: narrow
wins past ~3/8 of a wave (1024³ 174 vs 131T), the big CTA's operand reuse
wins past ~5/8 (forcing 64×64 there cost 2048³ 123→171T). Non-128-divisible
shapes with 64-divisibility take the 64×64 CTA (edge tiles otherwise drag
the single wave; 1088³: 76 vs 93T). Persistent schedules (static
round-robin and atomic ticket) both measured worse on L20 (4..8%; the
ticket variant recovers L2 locality but its loop-head barrier costs what
the CTA-restart overlap saves).
**NN swap.** The dual-N-contiguous problem runs as its transpose
`E = B^T @ A^T` over swapped operands with an out-transposed epilogue
scatter (CUTLASS-sm90 `is_swapAB`): one instantiation fewer per tile
config, at the cost of a scalar-store scatter on a path no LLM-linear
operand pair hits.
## Build System ## Build System
### Auto-detection ### Auto-detection
@@ -65,10 +165,19 @@ cmake --build build/cmake -j 16
### Architecture flags ### Architecture flags
`setup.py` passes the GPU compute capability to CMake via `ASTRAI_CUDA_ARCH` (default `89`, i.e. sm_89 / L20): `setup.py` passes the GPU compute capability to CMake via `ASTRAI_CUDA_ARCH`. When
unset, `setup.py` auto-detects the real GPU capability through
`torch.cuda.get_device_capability()`; the CMake fallback default is `80` (sm_80):
- **sm_80+** (Ampere and later): enables tensor-core MMA path (`mma.sync.m16n8k16.bf16`) - **sm_80+** (Ampere and later): enables the tensor-core MMA path
- **Below sm_80**: adds `-DASTRAI_NO_MMA` to disable the MMA path at compile time (`mma.sync.m16n8k16.bf16` for bf16 attention, `mma.sync.m16n8k32` for FP8).
- **sm_89+**: required for the FP8 family (`fp8_ops`) — FP8 tensor-core
instructions only exist on Ada/Hopper and newer. On older architectures,
CMake emits a warning and skips the `fp8_ops` target so the remaining CUDA
kernels still build successfully.
- **`-DASTRAI_NO_MMA`** is a manual escape hatch only — the build never defines
it automatically. To disable the MMA path, add it to `NVCC_FLAGS` yourself;
all supported build targets are sm_80+.
### Build configuration ### Build configuration
@@ -79,15 +188,128 @@ NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
--ptxas-options=-O3,-v --extra-device-vectorization --threads=16 --ptxas-options=-O3,-v --extra-device-vectorization --threads=16
``` ```
Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 module (one `.so` per kernel, named `<kernel>.cpython-*-x86_64-linux-gnu.so`). CMake builds all five kernel targets in parallel via `cmake --build -j N`. Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 module (one `.so` per kernel, named `<kernel>.cpython-*-x86_64-linux-gnu.so`). CMake builds all registered kernel targets in parallel via `cmake --build -j N` (the five base targets always; `fp8_ops` additionally on sm_89+). The target list is the **single source of truth**: `KERNEL_NAMES` and the parallel `KERNEL_SRCS` list in `csrc/CMakeLists.txt`; `astrai/extension/loader.py` auto-discovers the compiled `.so` files.
## Python Extension Architecture
The Python extension package separates low-level kernel bindings from execution
policy:
```text
astrai/extension/
├── __init__.py # Stable public API
├── loader.py # Optional compiled-module discovery and loading
├── ops/
│ ├── attention.py # Stateless attention kernel wrappers
│ ├── rotary.py # Stateless rotary kernel wrapper
│ └── fp8.py # Stateless FP8 primitives (custom_op)
├── fp8.py # FP8 strategy layer (fp8_autocast, recipes)
└── backend/
├── attention.py # Backend selection, KV cache I/O, and fallback
└── rotary.py # Per-call CUDA/torch rotary dispatch
```
The dependency direction is one-way:
```text
model / inference
|
v
extension public API
|
v
backend policy ---> ops wrappers ---> loader ---> compiled .so
|
+-----------> torch / flash-attn fallback
```
`ops` must not import `backend`. This keeps direct kernel bindings independent
of model, cache, fallback, and backend-selection policy.
### Ops Layer
`astrai.extension.ops` is the low-level boundary around compiled extensions:
- Wrappers are stateless and map Python arguments to pybind or
`torch.library.custom_op` calls.
- Wrappers validate kernel availability and raise `RuntimeError` when a
requested extension was not built.
- Wrappers do not choose another implementation, gather KV cache entries, or
decide whether an input is supported by a backend.
- Tests that specifically exercise a compiled kernel may import from
`astrai.extension.ops`.
For example, `attn_prefill(...)` means "run this CUDA kernel" rather than "run
attention using the best available implementation":
```python
from astrai.extension.ops import attn_prefill
output = attn_prefill(q, k, v, mask=mask, is_causal=True)
```
If the kernel is unavailable, this call fails. Callers that need fallback and
capability dispatch must use the public `attention(...)` entry point instead.
### Backend Layer
`astrai.extension.backend` owns execution policy:
- It selects CUDA, FlashAttention, or torch-native attention.
- It checks per-call constraints such as dtype, shape, head dimension, cache
availability, and installed optional dependencies.
- It owns KV cache writes and reads because those operations differ by backend.
- It provides torch fallbacks and raises when an explicitly requested backend
cannot handle a call.
- Rotary dispatch follows the same boundary without a backend class: the
policy layer chooses the fused op for supported inference calls and otherwise
uses the autograd-compatible torch implementation.
Normal model and inference code should import the stable API from
`astrai.extension`:
```python
from astrai.extension import ATTN_BACKEND, attention, attn_backend
output = attention(q, k, v, kv_cache=cache, layer_id=layer_id, fwd="decode")
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
output = attention(q, k, v)
```
The package root re-exports the supported high-level API and selected direct
kernel wrappers. Internal code should use `astrai.extension.backend` only when
it needs a backend type or policy implementation, and `astrai.extension.ops`
only when it deliberately requires one exact kernel.
### Placement Rules
When extending this package:
| Change | Location |
|--------|----------|
| Add a pybind call for a compiled kernel | `astrai/extension/ops/` |
| Add argument translation required by the compiled ABI | `astrai/extension/ops/` |
| Add capability checks or implementation selection | `astrai/extension/backend/` |
| Add a torch or third-party fallback | `astrai/extension/backend/` |
| Add attention KV cache behavior | `astrai/extension/backend/attention.py` |
| Expose a supported user-facing symbol | `astrai/extension/__init__.py` |
Imports belong at module scope. Optional dependencies such as `flash_attn` may
use a module-level guarded import. Type-only imports that would create a runtime
cycle belong under `TYPE_CHECKING`.
## Attention Backend ## Attention Backend
`astrai/extension/attention_backend.py` provides the backend abstraction: `astrai/extension/backend/attention.py` provides the backend abstraction:
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len - **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (default) - **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`). Default on GPU.
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`) - **`FlashAttnBackend`**: Optional flash-attn dispatch via `flash_attn_varlen_func` over gathered flat K/V.
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (always-available fallback)
Default priority: cuda > flash > torch. Set ``ASTR_BACKEND=cuda|torch_native|flash``
to override the default.
Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`): Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
@@ -98,11 +320,20 @@ with attn_backend(ATTN_BACKEND.CUDA):
engine.generate("hello") engine.generate("hello")
``` ```
`CudaBackend` falls back to `TorchNativeBackend` when a kernel is not available. The `attention(...)` policy entry point falls back to `FlashAttnBackend` (when
flash-attn is installed and supports the call) or `TorchNativeBackend` when the
automatically selected CUDA backend cannot handle an input. Resolution
precedence is: explicit `attn_backend(...)` context > `ASTR_BACKEND` env >
default. An explicit `attn_backend(...)` selection is strict and raises instead
of silently switching implementations; the env override (and the implicit
default) fall back to the first compatible backend when incapable. Training
calls (`fwd=None`, no KV cache) resolve by capability: the CUDA cache kernels
cannot run without a cache, so they fall back to flash (mask-free/causal calls
only) and finally to torch SDPA.
### Rotary Backend ### Rotary Backend
`astrai/extension/rotary_backend.py` provides `apply_rotary_emb(x, (cos, sin))` with auto-dispatch: `astrai/extension/backend/rotary.py` provides `apply_rotary_emb(x, (cos, sin))` with auto-dispatch:
- **CUDA path**: calls `rotary_emb` kernel directly when available, input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference) - **CUDA path**: calls `rotary_emb` kernel directly when available, input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference)
- **Torch fallback**: complex multiply (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd) or when kernel unavailable - **Torch fallback**: complex multiply (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd) or when kernel unavailable
@@ -111,9 +342,9 @@ No context-manager switching needed — the dispatch is automatic per call.
## Python Wrappers ## Python Wrappers
`astrai/extension/attention_ops.py` provides Python wrappers for each compiled attention kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions. `astrai/extension/ops/attention.py` provides Python wrappers for each compiled attention kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions.
`astrai/extension/rotary_ops.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `rotary_backend.py`. `astrai/extension/ops/rotary.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `backend/rotary.py`.
Interface (all functions): Interface (all functions):
``` ```
@@ -123,6 +354,54 @@ mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
Layout convention: all q/k/v are `[batch, seq_len, n_heads, head_dim]` (blhd). Scale is always `1/sqrt(head_dim)`. Layout convention: all q/k/v are `[batch, seq_len, n_heads, head_dim]` (blhd). Scale is always `1/sqrt(head_dim)`.
### Q Scheduling and KV Addressing
Prefill separates Q work scheduling from KV storage:
- `DenseQSchedule` maps a rectangular grid directly with
`batch = blockIdx.z` and `q_tile = blockIdx.x`.
- `PackedQSchedule` consumes a compact work map for a packed
`[total_q, q_heads, head_dim]` tensor.
- `ContigKV` and `PagedKV` only provide KV lengths and translate logical KV
positions into physical addresses. They do not schedule Q blocks.
For ragged Q lengths `[70, 10, 130]` and 64 rows per Q tile, cache binding
builds:
```text
qo_indptr = [0, 70, 80, 210]
q_tile_to_batch = [0, 0, 1, 2, 2, 2]
q_tile_to_index = [0, 1, 0, 0, 1, 2]
```
Paged prefill launches (MMA path, GQA head packing):
```text
grid.x = num_q_tiles * HB # HB = min(G, WARPS): q heads packed per block
grid.y = kv_heads * ceil(G / HB)
grid.z = 1
```
The tensor-core prefill kernel packs `HB = min(G, WARPS)` query heads of one
kv-head group into a block, so K/V tiles stream once per block instead of once
per q head (~HB× less global K/V traffic). Warp `w` handles head slot `w / WPH`
and 16-row chunk `w % WPH`, where `WPH = WARPS / HB`; `G = q_heads / kv_heads`
and `G = 1` (MHA) degenerates to the historical one-head-per-block layout.
Each host Q tile (64 rows, `Q_TILE_ROWS`) splits into `HB` packed blocks along
`grid.x`. Each block resolves its request and request-local row range in O(1):
```cpp
host_tile = blockIdx.x / HB;
batch = q_tile_to_batch[host_tile];
row_base = q_tile_to_index[host_tile] * 64 + (blockIdx.x % HB) * (64 / HB);
```
The kernel then uses `qo_indptr[batch]` for the packed Q base and adjacent
`qo_indptr` / `kv_indptr` entries for that request's Q and KV lengths. This
avoids the previous per-block linear scan over the batch, shared-memory
broadcast, mapping barrier, and upper-bound grid with potentially invalid
blocks.
## Standalone Testing ## Standalone Testing
Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment. Example: Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment. Example:
@@ -136,6 +415,7 @@ nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
Test files: Test files:
- `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks) - `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks)
- `attn_paged_test.cu` — paged decode/prefill kernels - `attn_paged_test.cu` — paged decode/prefill kernels
- `fp8_test.cu` — single-warp bf16→fp8→mma.sync sanity check + full FP8 GEMM correctness (sm_89)
## Benchmarks ## Benchmarks
@@ -158,29 +438,48 @@ nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
``` ```
csrc/ csrc/
├── CMakeLists.txt # CMake build: 5 kernel targets, torch/pybind11 linking ├── CMakeLists.txt # CMake build: kernel registry (KERNEL_NAMES / KERNEL_SRCS), torch/pybind11 linking
├── kernels/ ├── kernels/
│ ├── attn_common.h # Unified attention params (contig + paged modes) │ ├── common/ # cross-family pure-CUDA helpers (no torch)
│ ├── attn_decode.cu # Basic decode kernel (registered) │ ├── device.cuh # sm_at_least(), kMinSmForFp8* constants
│ ├── attn_prefill.cu # Basic prefill kernel (registered) │ ├── mma.cuh # shared mma_sync<InT> + mma_shape<InT> (bf16 m16n8k16 / fp8 m16n8k32) + ldmatrix_x2/x4<T>
│ ├── attn_paged_decode.cu # Paged decode kernel (registered) │ ├── cp_async.cuh # cp.async 16B primitives (predicated copy, commit/wait groups)
├── attn_paged_prefill.cu # Paged prefill kernel (registered) │ └── reduce.cuh # warp_reduce_max, atomic_max_float
│ ├── rotary_emb.cu # Fused rotary embedding kernel (registered) │ ├── attention/ # attention family (module names keep the attn_* prefix)
│ ├── attn_decode_split_kv.cuh # Split-KV variant (contig + paged via KVSource) │ ├── common.h # AttentionParams POD, TensorLayout enum (BHLD/BLHD)
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant (contig + paged) │ ├── warp_utils.cuh # warp reduction helpers
│ ├── attn_prefill_split_q.cuh # Split-Q variant (contig + paged via KVSource) │ ├── layout_policies.cuh # KV addressing policies: DenseQSchedule/PackedQSchedule, ContigKV/PagedKV
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant (contig + paged) │ ├── mma_utils.cuh # ldmatrix/pack helpers + online-softmax (bf16 mma via common/mma.cuh)
│ ├── attn_kv_source.cuh # KVSource policies (ContigKV / PagedKV) │ ├── entry_utils.cuh # torch binding helpers: DISPATCH_HEAD_DIM, pack_*_params
│ ├── attn_dispatchers.cuh # Kernel dispatch macros + KV-templated launchers │ ├── dispatchers.cuh # pure-CUDA launchers: dispatch_decode/prefill (+paged), split-K math
│ ├── attn_entry_utils.cuh # Entry point helpers │ ├── decode_split_kv.cuh # decode kernel, scalar (split-KV)
│ ├── attn_mma_utils.cuh # MMA utilities │ ├── decode_split_kv_mma.cuh # decode kernel, MMA + split-K
└── attn_warp_utils.cuh # Warp-level utilities │ ├── prefill_split_q.cuh # prefill kernel, scalar (split-Q)
│ │ ├── prefill_split_q_mma.cuh # prefill kernel, MMA (split-Q, GQA head packing, packed/ragged Q schedule)
│ │ ├── decode.cu # → module attn_decode
│ │ ├── prefill.cu # → module attn_prefill
│ │ ├── paged_decode.cu # → module attn_paged_decode
│ │ └── paged_prefill.cu # → module attn_paged_prefill
│ ├── rotary/
│ │ └── rotary_emb.cu # rotary embedding (kernel + binding in one file) → module rotary_emb
│ └── fp8/ # FP8 family (module name fp8_ops)
│ ├── common.h # FP8Format enum, Fp8GemmTraits, FP8Params / FP8QuantizeParams PODs, layout tags (no torch)
│ ├── quantize.cuh # quantize kernels: vectorized + 32×32-tile transpose (out_layout 0/1/2) (no torch)
│ ├── gemm.cuh # GEMM umbrella: kernel orchestrator + host launch planning (no torch)
│ ├── gemm/ # GEMM device layers (humming/CUTLASS-style split)
│ │ ├── policy.cuh # smem budget / occupancy hint + Fp8GemmPolicy
│ │ ├── load.cuh # operand loaders (swizzle, congruous cp.async, crosswise direct)
│ │ ├── scheduler.cuh # grouped/plain raster mapping
│ │ ├── mainloop.cuh # stage rings + pipelined mma.sync mainloop
│ │ └── epilogue.cuh # fused bias + bf16 scatter + copy-out
│ └── ops.cu # binding only: validation, param packing, launch dispatch, pybind
└── tests/ └── tests/
├── test_utils.cuh # Shared test utilities ├── test_utils.cuh # Shared test utilities (now_ms, f2bf, bf2f, randf)
├── attn_test.cu # Decode + prefill kernels ├── attn_test.cu # Decode + prefill kernels
── attn_paged_test.cu # Paged decode/prefill kernels ── attn_paged_test.cu # Paged decode/prefill kernels
└── fp8_test.cu # MMA demo + GEMM correctness across layouts/K tiles/ragged shapes
``` ```
Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files. Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.
> Document Update Time: 2026-07-31 > Document Update Time: 2026-08-29

Some files were not shown because too many files have changed in this diff Show More