78 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
130 changed files with 8993 additions and 3215 deletions
+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
+6 -2
View File
@@ -57,10 +57,14 @@ COPY docs/ ./docs/
COPY pyproject.toml . COPY pyproject.toml .
COPY README.md . COPY README.md .
# Create non-root user matching the host uid/gid (passed via build args) # Create non-root user matching the host uid/gid (passed via build args).
# 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_UID=1000
ARG USER_GID=1000 ARG USER_GID=1000
RUN groupadd -g "${USER_GID}" astrai \ 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 \ && useradd -m -u "${USER_UID}" -g astrai astrai \
&& chown -R astrai:astrai /app && chown -R astrai:astrai /app
ENV HOME=/home/astrai ENV HOME=/home/astrai
+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
+3 -8
View File
@@ -17,14 +17,9 @@ 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,
get_app,
run_server,
sample,
)
from astrai.logging import setup_logging from astrai.logging import setup_logging
from astrai.model import ( from astrai.model import (
AutoModel, AutoModel,
+14 -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"})
@@ -70,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]
@@ -125,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
-2
View File
@@ -14,7 +14,6 @@ from astrai.dataset.storage import (
Streamable, Streamable,
detect_format, detect_format,
) )
from astrai.dataset.streaming import StreamingSeqDataset
from astrai.serialization import ( from astrai.serialization import (
load_bin, load_bin,
save_bin, save_bin,
@@ -35,5 +34,4 @@ __all__ = [
"save_bin", "save_bin",
"load_bin", "load_bin",
"RDSampler", "RDSampler",
"StreamingSeqDataset",
] ]
+4 -7
View File
@@ -383,10 +383,10 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
transform = _build_jsonl_transform(load_path, tokenizer_path) transform = _build_jsonl_transform(load_path, tokenizer_path)
if transform is None: if transform is None:
raise FileNotFoundError( raise FileNotFoundError(
f"JSONL dataset config not found. Expected " "JSONL dataset config not found. Expected "
f"dataset_config.json alongside *.jsonl files, pass " "dataset_config.json alongside *.jsonl files, pass "
f"tokenizer_path= for the built-in messages config, or " "tokenizer_path= for the built-in messages config, or "
f"use processor= for lazy on-the-fly tokenisation." "use processor= for lazy on-the-fly tokenisation."
) )
store.load(load_path, transform=transform, **kwargs) store.load(load_path, transform=transform, **kwargs)
else: else:
@@ -492,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),
+1 -1
View File
@@ -217,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}"
-122
View File
@@ -1,122 +0,0 @@
"""Streaming IterableDataset for pre-training with shard-level shuffle.
Unlike the map-style datasets, the streaming dataset yields windows
sequentially through each data shard — no random access, no sampler.
Each DataLoader worker independently streams its assigned shard subset,
giving better OS page-cache locality for large-scale (TB+) datasets.
Key properties:
- Implements ``torch.utils.data.IterableDataset``.
- ``__len__`` returns total window count so ``compute_total_steps`` works.
- Shard-level shuffle with deterministic seed (reproducible across runs).
- Distributed: each rank gets a disjoint subset of shards.
- Multi-worker: each worker within a rank gets a disjoint subset.
"""
import random
from typing import Iterator, Optional
import torch
import torch.distributed as dist
from torch import Tensor
from torch.utils.data import IterableDataset
from astrai.dataset.storage import Store
def _resolve_rank_and_world_size() -> tuple[int, int]:
if dist.is_available() and dist.is_initialized():
return dist.get_rank(), dist.get_world_size()
return 0, 1
def _total_windows(token_count, window_size, stride):
if token_count <= window_size:
return 0
return (token_count - 1 - window_size) // stride + 1
class StreamingSeqDataset(IterableDataset):
"""Streaming next-token prediction dataset.
Yields ``{"input_ids": [L], "target_ids": [L]}`` dicts by sliding a
window sequentially through each data shard. Shards are shuffled
deterministically. Distributed and multi-worker DataLoader modes are
supported: each consumer gets a disjoint shard subset.
Args:
store: Already-loaded Store with a ``"sequence"`` key.
window_size: Context length per sample.
stride: Step between consecutive windows (default: window_size).
shuffle: Shuffle shard order.
seed: Base seed for deterministic shard shuffle.
"""
def __init__(
self,
store: Store,
window_size: int,
stride: Optional[int] = None,
shuffle: bool = True,
seed: int = 42,
rank: Optional[int] = None,
world_size: Optional[int] = None,
):
super().__init__()
if window_size <= 0:
raise ValueError("window_size must be positive")
self.store = store
self.window_size = window_size
self.stride = stride if stride is not None else window_size
self.shuffle = shuffle
self.seed = seed
self._rank, self._world_size = (
rank,
world_size if rank is not None else _resolve_rank_and_world_size(),
)
if "sequence" not in store.keys:
raise KeyError(
f"Store is missing required key 'sequence'; "
f"available keys: {sorted(store.keys)}"
)
@property
def num_samples(self) -> int:
return _total_windows(self.store.token_count, self.window_size, self.stride)
def __len__(self) -> int:
return self.num_samples
def __iter__(self) -> Iterator[dict[str, Tensor]]:
segments = self.store._data["sequence"]
n_shards = len(segments)
indices = list(range(n_shards))
if self.shuffle:
rng = random.Random(self.seed)
rng.shuffle(indices)
worker_info = torch.utils.data.get_worker_info()
if worker_info is None:
num_consumers = self._world_size
consumer_id = self._rank
else:
num_consumers = self._world_size * worker_info.num_workers
consumer_id = self._rank * worker_info.num_workers + worker_info.id
my_shards = [
i for idx, i in enumerate(indices) if idx % num_consumers == consumer_id
]
for shard_idx in my_shards:
segment = segments[shard_idx]
seq_len = segment.shape[0]
for begin in range(0, seq_len - self.window_size, self.stride):
end = begin + self.window_size
yield {
"input_ids": torch.as_tensor(segment[begin:end], dtype=torch.long),
"target_ids": torch.as_tensor(
segment[begin + 1 : end + 1], dtype=torch.long
),
}
+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",
+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",
]
@@ -21,9 +21,20 @@ Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
... ...
Thread-safe via ``contextvars`` each scheduler thread gets its own Thread-safe via ``contextvars`` each scheduler thread gets its own
active backend. ``get_backend()`` returns the active one, falling back active backend. Backend resolution follows a strict precedence:
to a process-wide default (cuda > flash > torch, overridable via
``ASTR_BACKEND``). 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]`` 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]``. (blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
@@ -32,29 +43,35 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
import contextvars import contextvars
import enum import enum
import functools import functools
import importlib import logging
import os import os
import threading import threading
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from typing import TYPE_CHECKING, Optional, Union from typing import TYPE_CHECKING, Dict, Optional, Tuple, Union
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.extension.attention_ops import ( from astrai.extension.loader import is_available
from astrai.extension.ops.attention import (
attn_paged_decode, attn_paged_decode,
attn_paged_prefill, attn_paged_prefill,
) )
from astrai.extension.loader import is_available
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
try:
import flash_attn as _flash_attn
except Exception:
_flash_attn = None
if TYPE_CHECKING: if TYPE_CHECKING:
from astrai.inference.cache import KVCache from astrai.inference.cache import KVCache
logger = logging.getLogger(__name__)
_default_backend: Optional["AttentionBackend"] = None
_default_backend_lock = threading.Lock() _default_backend_lock = threading.Lock()
_env_backend_name: Optional[str] = None _env_backend_name: Optional[str] = None
_env_backend: Optional["AttentionBackend"] = None _env_backend: Optional["AttentionBackend"] = None
@@ -62,12 +79,16 @@ _current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = (
contextvars.ContextVar("attn_backend", default=None) 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) @functools.lru_cache(maxsize=1)
def flash_attn_available() -> bool: def flash_attn_available() -> bool:
if not torch.cuda.is_available(): if not torch.cuda.is_available():
return False return False
fa = _get_flash_attn() fa = _flash_attn
if fa is None: if fa is None:
return False return False
@@ -90,14 +111,6 @@ def flash_attn_available() -> bool:
return False 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): class ATTN_BACKEND(enum.Enum):
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``.""" """Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
@@ -106,54 +119,40 @@ class ATTN_BACKEND(enum.Enum):
FLASH = "flash" FLASH = "flash"
def _priority_backends() -> list["AttentionBackend"]: def _instance(backend_cls: type) -> "AttentionBackend":
"""Available backends in priority order: cuda -> flash -> torch.""" """Return the canonical singleton instance for a backend class.
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
Backends hold no per-instance state, so a single cached instance is
def _backend_supports( safe and avoids per-call allocation on the attention hot path.
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): backend = _singletons.get(backend_cls)
return ( if backend is None:
kv_cache is not None backend = backend_cls()
and q.dtype == torch.bfloat16 _singletons[backend_cls] = backend
and q.size(-1) in (32, 64, 128, 256) return backend
)
if isinstance(backend, FlashAttnBackend):
if not flash_attn_available(): @functools.lru_cache(maxsize=1)
return False def _priority_backends() -> Tuple["AttentionBackend", ...]:
if q.dtype not in (torch.float16, torch.bfloat16): """Available backends in priority order: cuda -> flash -> torch.
return False
if q.size(1) == 1 and kv_cache is not None: Computed once (machine availability cannot change at runtime) and
return True cached forever; the tuple always ends with ``TorchNativeBackend``,
if attn_mask is None or is_causal: which is unconditionally available.
return True """
return attn_mask.dim() == 4 return tuple(
return True _instance(cls)
for cls in (CudaBackend, FlashAttnBackend, TorchNativeBackend)
if cls.available()
)
def _resolve_default_backend() -> "AttentionBackend": def _resolve_default_backend() -> "AttentionBackend":
"""Pick the highest-priority available backend (cuda -> flash -> torch). """Pick the highest-priority available backend (cuda -> flash -> torch).
Resolved lazily on first ``get_backend()`` and cached. Per-call Resolved lazily on first use and cached via ``_priority_backends``.
capability fallback happens in ``attention()``, so the default is Per-call capability fallback happens in ``attention()``, so the
safe for training and fp32 models. default is safe for training and fp32 models.
""" """
return _priority_backends()[0] return _priority_backends()[0]
@@ -168,9 +167,14 @@ def _environment_backend() -> Optional["AttentionBackend"]:
with _default_backend_lock: with _default_backend_lock:
if name != _env_backend_name: if name != _env_backend_name:
try: try:
_env_backend = AttentionBackendFactory.create(name) _env_backend = _resolve_backend(name)
except (ValueError, RuntimeError): except (ValueError, RuntimeError):
_env_backend = None _env_backend = None
logger.warning(
"ASTR_BACKEND=%r is not a registered attention backend; "
"falling back to default resolution",
name,
)
_env_backend_name = name _env_backend_name = name
return _env_backend return _env_backend
@@ -178,43 +182,46 @@ def _environment_backend() -> Optional["AttentionBackend"]:
def _resolve_backend( def _resolve_backend(
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None, backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
) -> "AttentionBackend": ) -> "AttentionBackend":
"""Resolve a backend configuration, defaulting to the process policy.""" """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 backend is not None:
if isinstance(backend, ATTN_BACKEND): if isinstance(backend, ATTN_BACKEND):
return AttentionBackendFactory.create(backend.value) return _instance(AttentionBackendFactory.get_component_class(backend.value))
if isinstance(backend, str): if isinstance(backend, str):
return AttentionBackendFactory.create(backend) return _instance(AttentionBackendFactory.get_component_class(backend))
if isinstance(backend, type) and issubclass(backend, AttentionBackend): if isinstance(backend, type) and issubclass(backend, AttentionBackend):
return backend() return _instance(backend)
if isinstance(backend, AttentionBackend): if isinstance(backend, AttentionBackend):
return backend return backend
raise TypeError( raise TypeError(
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, " f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
f"or instance, got {type(backend).__name__}" f"or instance, got {type(backend).__name__}"
) )
return _resolve_default_backend()
global _default_backend
if _default_backend is None:
with _default_backend_lock:
if _default_backend is None:
_default_backend = _resolve_default_backend()
return _default_backend
def get_backend( def get_backend(
use_default: bool = True, use_default: bool = True,
) -> Optional["AttentionBackend"]: ) -> Optional["AttentionBackend"]:
"""Return the context override, optionally falling back to the process default. """Resolve the active backend: explicit context > env > default.
``ASTR_BACKEND`` is a process-wide override and takes precedence over the An ``attn_backend(...)`` context is the caller's explicit choice and
context value. Pass ``use_default=False`` at request submission to retain always wins. ``ASTR_BACKEND`` is a process-wide override consulted
only an environment override or the caller's :func:`attn_backend` value. 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.
""" """
return ( context_backend = _current_backend.get()
_environment_backend() if context_backend is not None:
or _current_backend.get() return context_backend
or (_resolve_backend() if use_default else None) env_backend = _environment_backend()
) if env_backend is not None:
return env_backend
return _resolve_default_backend() if use_default else None
@contextmanager @contextmanager
@@ -243,38 +250,16 @@ def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
def repeat_kv(x: Tensor, n_rep: int) -> Tensor: def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
"""Expand KV heads to match Q heads for GQA.""" """Expand KV heads to match Q heads for GQA."""
bs, slen, n_heads, head_dim = x.shape
if n_rep == 1: if n_rep == 1:
return x return x
n_heads, head_dim = x.shape[-2:]
return ( return (
x[:, :, :, None, :] x.unsqueeze(-2)
.expand(bs, slen, n_heads, n_rep, head_dim) .expand(*x.shape[:-2], n_heads, n_rep, head_dim)
.reshape(bs, slen, n_heads * n_rep, head_dim) .reshape(*x.shape[:-2], 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( def attention(
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
@@ -283,13 +268,24 @@ def attention(
layer_id: int = 0, layer_id: int = 0,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
) -> Tensor: ) -> Tensor:
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``. """Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
Delegates to the active backend (set via ``with attn_backend(...)``). 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 Handles KV cache I/O, GQA head expansion, and causal masking so the
caller only needs to provide projected q/k/v. 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: Args:
q: [batch, q_len, n_heads, head_dim] (blhd) q: [batch, q_len, n_heads, head_dim] (blhd)
k: [batch, q_len, n_kv_heads, head_dim] (blhd) k: [batch, q_len, n_kv_heads, head_dim] (blhd)
@@ -298,28 +294,43 @@ def attention(
layer_id: transformer layer index for buffer access. layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask (SDPA-compatible). attn_mask: pre-built attention mask (SDPA-compatible).
is_causal: whether to apply causal masking. 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: Returns:
[batch, q_len, n_heads * head_dim] [batch, q_len, n_heads * head_dim]
""" """
backend = get_backend() if backend is not None:
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal): selected = _resolve_backend(backend)
explicit = get_backend(use_default=False) explicit = True
if explicit is not None: 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( raise RuntimeError(
f"Explicitly-set backend {type(backend).__name__} cannot " f"Explicitly-set backend {type(selected).__name__} cannot "
f"handle this attention call (shape={q.shape}, " f"handle this attention call (shape={q.shape}, "
f"dtype={q.dtype}, kv_cache={'none' if kv_cache is None else 'present'}, " 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"attn_mask={'none' if attn_mask is None else 'present'}). "
f"Remove the attn_backend() context or switch to a compatible backend." f"Remove the attn_backend() context or switch to a compatible backend."
) )
for candidate in _priority_backends(): selected = next(
if isinstance(candidate, type(backend)): (
continue candidate
if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal): for candidate in _priority_backends()
backend = candidate if candidate.supports_call(q, kv_cache, attn_mask, is_causal, fwd)
break ),
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal) _instance(TorchNativeBackend),
)
return selected.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
class AttentionBackend(ABC): class AttentionBackend(ABC):
@@ -329,6 +340,17 @@ class AttentionBackend(ABC):
``fwd_prefill`` (q_len > 1, with or without cache). The public ``fwd_prefill`` (q_len > 1, with or without cache). The public
``forward`` method dispatches based on q_len. ``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:: Three equivalent ways to activate a backend::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
@@ -346,6 +368,30 @@ class AttentionBackend(ABC):
def __exit__(self, *exc) -> None: def __exit__(self, *exc) -> None:
_current_backend.reset(self._token) _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( def forward(
self, self,
q: Tensor, q: Tensor,
@@ -355,6 +401,7 @@ class AttentionBackend(ABC):
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
fwd: Optional[str] = None,
) -> Tensor: ) -> Tensor:
"""Dispatch to decode or extend based on q_len. """Dispatch to decode or extend based on q_len.
@@ -370,9 +417,11 @@ class AttentionBackend(ABC):
Returns: Returns:
[batch, q_len, n_heads * head_dim] [batch, q_len, n_heads * head_dim]
""" """
if kv_cache is not None and q.size(1) == 1: if fwd == "decode":
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal) 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) 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 @abstractmethod
def fwd_decode( def fwd_decode(
@@ -428,8 +477,18 @@ class TorchNativeBackend(AttentionBackend):
runs SDPA directly on the projected q/k/v. runs SDPA directly on the projected q/k/v.
""" """
@staticmethod @classmethod
def supports(**kwargs) -> bool: 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 return True
def fwd_decode( def fwd_decode(
@@ -466,23 +525,52 @@ class TorchNativeBackend(AttentionBackend):
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
if kv_cache is not None: if q.ndim == 4:
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)
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()
)
n_rep = q.size(2) // k.size(2) if kv_cache is None or kv_cache.qo_indptr is None:
if n_rep > 1: raise ValueError("packed attention requires KV cache metadata")
k = repeat_kv(k, n_rep) kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
v = repeat_kv(v, n_rep) kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
outputs = []
out = F.scaled_dot_product_attention( n_rep = q.size(1) // k.size(1)
q.permute(0, 2, 1, 3), for i in range(kv_cache.req_pool_indices.numel()):
k.permute(0, 2, 1, 3), q_start = int(kv_cache.qo_indptr[i])
v.permute(0, 2, 1, 3), q_end = int(kv_cache.qo_indptr[i + 1])
attn_mask, indices = kv_cache.req_to_token[
is_causal=is_causal, kv_cache.req_pool_indices[i], : kv_cache.seq_lens[i]
) ]
out = out.permute(0, 2, 1, 3).contiguous().flatten(2) k_i = kv_cache.k_buffer[layer_id, indices]
return out 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) @AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
@@ -503,16 +591,37 @@ class CudaBackend(AttentionBackend):
Raises ``RuntimeError`` if the required kernel is not available. Raises ``RuntimeError`` if the required kernel is not available.
""" """
@staticmethod # Head dims supported by the CUDA kernels (single source of truth).
def supports(**kwargs) -> bool: HEAD_DIMS = (32, 64, 128, 256)
head_dim = kwargs.get("head_dim", -1)
@classmethod
def available(cls) -> bool:
return ( return (
torch.cuda.is_available() torch.cuda.is_available()
and head_dim in (32, 64, 128, 256)
and is_available("attn_paged_decode") and is_available("attn_paged_decode")
and is_available("attn_paged_prefill") 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 @staticmethod
def supports_graph() -> bool: def supports_graph() -> bool:
return True return True
@@ -530,27 +639,23 @@ class CudaBackend(AttentionBackend):
if kv_cache is None: if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=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 kv_indptr = kv_cache.kv_indptr
out = attn_paged_decode( out = attn_paged_decode(
q_3d, q,
kv_cache.k_buffer[layer_id], kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id], kv_cache.v_buffer[layer_id],
kv_cache.req_to_token, kv_cache.req_to_token,
kv_cache.req_pool_indices, kv_cache.req_pool_indices,
kv_indptr, kv_indptr,
new_k=k,
new_v=v,
is_causal=True, is_causal=True,
o_part_buf=kv_cache.decode_o_part, o_part_buf=kv_cache.decode_o_part,
ml_part_buf=kv_cache.decode_ml_part, ml_part_buf=kv_cache.decode_ml_part,
out_buf=kv_cache.decode_out, out_buf=kv_cache.decode_out,
) )
return out.unsqueeze(1).flatten(2) return out
def fwd_prefill( def fwd_prefill(
self, self,
@@ -565,52 +670,63 @@ class CudaBackend(AttentionBackend):
if kv_cache is None: if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)") raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
loc = kv_cache.out_cache_loc.reshape(-1) loc = kv_cache.out_cache_loc
kv_cache.k_buffer[layer_id].index_copy_( kv_cache.k_buffer[layer_id, loc] = k
0, loc, k.reshape(-1, k.size(2), k.size(3)) kv_cache.v_buffer[layer_id, loc] = v
)
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( out = attn_paged_prefill(
q_flat, q,
kv_cache.k_buffer[layer_id], kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id], kv_cache.v_buffer[layer_id],
kv_cache.req_to_token, kv_cache.req_to_token,
kv_cache.req_pool_indices, kv_cache.req_pool_indices,
kv_indptr, kv_cache.kv_indptr,
qo_indptr, kv_cache.qo_indptr,
kv_cache.q_tile_to_batch,
kv_cache.q_tile_to_index,
attn_mask, attn_mask,
is_causal=is_causal, is_causal=is_causal,
) )
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2) return out
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value) @AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
class FlashAttnBackend(AttentionBackend): class FlashAttnBackend(AttentionBackend):
"""FlashAttention backend via the optional ``flash-attn`` package. """FlashAttention backend via the optional ``flash-attn`` package.
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``, Decode (q_len=1, contiguous cache): writes K/V to the pool, gathers
which reads K/V directly from the flat pool via cache_batch_idx + flat K/V via the ``req_to_token`` page table, and calls
cache_seqlens no materialized KV gather. ``flash_attn_varlen_func`` over the ragged batch
(``qo_indptr``/``kv_indptr``).
Prefill / non-contiguous decode: falls back to KV gather + Prefill: packed 3-D calls share the ``flash_attn_varlen_func`` path;
``flash_attn_func``. dense 4-D calls go through ``flash_attn_func`` (mask-free only).
""" """
@staticmethod @classmethod
def supports(**kwargs) -> bool: def available(cls) -> bool:
return flash_attn_available() 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( def fwd_decode(
self, self,
q: Tensor, q: Tensor,
@@ -621,7 +737,7 @@ class FlashAttnBackend(AttentionBackend):
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal) return self._forward_packed(q, k, v, kv_cache, layer_id)
def fwd_prefill( def fwd_prefill(
self, self,
@@ -633,36 +749,29 @@ class FlashAttnBackend(AttentionBackend):
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal) 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( def _forward_dense(
self, self,
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> 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) n_rep = q.size(2) // k.size(2)
if n_rep > 1: if n_rep > 1:
k = repeat_kv(k, n_rep) k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep) v = repeat_kv(v, n_rep)
if attn_mask is not None and not is_causal and attn_mask.dim() != 4: if attn_mask is not None:
raise ValueError( raise ValueError(
"FlashAttnBackend does not support a custom attention mask; " "FlashAttnBackend cannot handle a custom attention mask; "
"use a causal mask or select TorchNativeBackend." "use a causal mask or select TorchNativeBackend."
) )
fa = _get_flash_attn() fa = _flash_attn
if fa is None: if fa is None:
raise RuntimeError( raise RuntimeError(
"FlashAttnBackend requires the optional 'flash-attn' package. " "FlashAttnBackend requires the optional 'flash-attn' package. "
@@ -672,11 +781,11 @@ class FlashAttnBackend(AttentionBackend):
q.contiguous(), q.contiguous(),
k.contiguous(), k.contiguous(),
v.contiguous(), v.contiguous(),
causal=is_causal or (attn_mask is not None and attn_mask.dim() == 4), causal=is_causal,
) )
return out.contiguous().flatten(2) return out.contiguous()
def _decode_with_kvcache( def _forward_packed(
self, self,
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
@@ -684,22 +793,27 @@ class FlashAttnBackend(AttentionBackend):
kv_cache: "KVCache", kv_cache: "KVCache",
layer_id: int, layer_id: int,
) -> Tensor: ) -> Tensor:
max_batch = kv_cache.req_to_token.size(0) fa = _flash_attn
max_seq = kv_cache.req_to_token.size(1) if fa is None or not hasattr(fa, "flash_attn_varlen_func"):
n_kv = k.size(2) raise RuntimeError("packed inference requires flash_attn_varlen_func")
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
k_cache = kv_cache.k_buffer[layer_id].view(max_batch, max_seq, n_kv, k.size(3)) kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
v_cache = kv_cache.v_buffer[layer_id].view(max_batch, max_seq, n_kv, v.size(3)) page_table = kv_cache.req_to_token[
kv_cache.req_pool_indices, : kv_cache.max_len
fa = _get_flash_attn() ]
out = fa.flash_attn_with_kvcache( positions = torch.arange(kv_cache.max_len, device=q.device)
q=q, indices = page_table[positions.unsqueeze(0) < kv_cache.seq_lens.unsqueeze(1)]
k_cache=k_cache, k_flat = kv_cache.k_buffer[layer_id, indices].contiguous()
v_cache=v_cache, v_flat = kv_cache.v_buffer[layer_id, indices].contiguous()
k=k, out = fa.flash_attn_varlen_func(
v=v, q.contiguous(),
cache_seqlens=(kv_cache.seq_lens - 1).to(torch.int32), k_flat,
cache_batch_idx=kv_cache.req_pool_indices.to(torch.int32), 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, causal=True,
) )
return out.flatten(2) 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)
+362 -196
View File
@@ -1,126 +1,191 @@
"""FP8 training: scaling state and aten::linear dispatch. """FP8 training: scaling recipes, per-tensor state, and aten::linear dispatch.
Layered (see also ``fp8_ops.py`` for the CUDA interface adapter): Layered (see ``ops/fp8.py`` for the CUDA interface adapter):
1. ``ops.fp8`` — the only module touching the pybind.
1. Kernel interface: "fp8_ops" — the only module touching the pybind. 2. This module (strategy layer): scaling *recipes* (TE-style delayed scaling
2. Training state (this module): per-tensor scales, amax history, delayed or dynamic current-amax scaling), per-tensor scales + amax history, and the
scaling, and the ``fp8_autocast`` context (TE-style, like ``fp8_autocast`` context manager (like ``torch.autocast``).
``torch.autocast``). 3. aten::linear integration: registers the CUDA + AutogradCUDA impls.
3. aten::linear integration (this module): registers the CUDA impl and the
M/N alignment guard.
Usage:: Usage::
from astrai.extension.fp8 import fp8_autocast from astrai.extension.fp8 import fp8_autocast
with fp8_autocast(enabled=True, fp8_format="hybrid"):
with fp8_autocast(enabled=True):
logits = model(input_ids) logits = model(input_ids)
loss.backward() loss.backward() # fp8 backward runs anywhere; fwd captured state on the node
Importing this module registers the aten::linear CUDA implementation. 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.
""" """
from contextlib import contextmanager import functools
from contextvars import ContextVar, Token
from dataclasses import dataclass
from enum import Enum
from typing import Dict, List, NamedTuple, Optional
import torch import torch
from torch.library import Library from torch.library import Library
from astrai.extension.fp8_ops import ( from astrai.extension.ops.fp8 import mm_fp8, quantize, quantize_dual
linear_backward_scaled,
linear_forward_scaled,
)
E4M3_MAX = 448.0 # Max representable value per FP8 format (E4M3: 448, E5M2: 57344).
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
# --------------------------------------------------------------------------- class FP8Format(str, Enum):
# Layer 2: training state (scales, amax history, delayed scaling, autocast) """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
class FP8TensorMeta: @dataclass
"""Scales + amax state for one weight tensor and its paired activations. class FP8Recipe:
"""Scale-from-amax policy: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``.
- weight: delayed scale from a 16-step amax history window (TE style) ``dynamic=False`` (default) is TE-style delayed scaling: max over the
- x/g: delayed one step, reuse the quantize kernel's free atomic amax 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.
""" """
__slots__ = ( history_len: int = 16
"scale", margin: int = 0
"scale_inv", dynamic: bool = False
"amax_history",
"idx",
"x_scale",
"x_scale_inv",
"g_scale",
"g_scale_inv",
)
def __init__(self, device: torch.device, update_interval: int): def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
self.scale = torch.ones(1, device=device, dtype=torch.float32) peak = amax.max()
self.scale_inv = torch.ones(1, device=device, dtype=torch.float32) return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
self.amax_history = torch.ones(
update_interval, device=device, dtype=torch.float32
) 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.idx = 0
self.x_scale = torch.ones(1, device=device, dtype=torch.float32) self.initialized = False
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
def record(self, amax: torch.Tensor) -> None: def advance(self) -> None:
"""Push the latest amax into the ring buffer (device-side copy, no sync).""" """Rotate to the next history slot after metadata update."""
self.amax_history[self.idx] = amax.reshape(()) self.idx = (self.idx + 1) % self.hist.numel()
self.idx = (self.idx + 1) % self.amax_history.numel()
def refresh(self) -> None: def seed(self, t: torch.Tensor, fmt: str) -> None:
"""Recompute scale from the amax history window (delayed scaling).""" amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
amax = self.amax_history.max() self.hist.fill_(amax)
if amax > 0: self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
self.scale.copy_(amax / E4M3_MAX) self.initialized = True
self.scale_inv.copy_(E4M3_MAX / amax)
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: class FP8State:
"""Global fp8 training state, TE-style.""" """Global fp8 training state: per-tensor metas + out-of-region defaults.
def __init__(self, update_interval: int = 16): The active ``(enabled, recipe, fp8_format)`` triple is a ``ContextVar``
self.enabled = False set by ``fp8_autocast`` (see ``_active``/``_current_config``); these plain
self.update_interval = update_interval attributes are the persistent defaults applied outside any region —
self.step_count = 0 ``fp8_linear_enable`` writes ``default_enabled``. The metas registry is
self._metas: dict[tuple, FP8TensorMeta] = {} shared across threads (GIL-protected); fp8 backward runs on autograd
self._last_device: torch.device | None = None engine threads and only touches metas captured on ``ctx`` at forward time.
"""
def _get_device(self, t: torch.Tensor) -> torch.device: def __init__(self):
if self._last_device is None: self.default_enabled = False
self._last_device = t.device self.default_recipe: FP8Recipe = FP8Recipe()
return t.device self.default_format: FP8Format = FP8Format.HYBRID
self._metas: Dict[tuple, FP8TensorMeta] = {}
def get_weight_meta(self, w: torch.Tensor) -> FP8TensorMeta: def get_weight_meta(self, w: torch.Tensor, recipe: FP8Recipe) -> FP8TensorMeta:
key = (w.data_ptr(), w.shape, w.dtype) key = (w.data_ptr(), w.shape, w.dtype)
meta = self._metas.get(key) meta = self._metas.get(key)
if meta is None: if meta is None:
meta = FP8TensorMeta(self._get_device(w), self.update_interval) meta = FP8TensorMeta(
_ScaleRing(w.device, recipe),
_ScaleRing(w.device, recipe),
_ScaleRing(w.device, recipe),
)
self._metas[key] = meta self._metas[key] = meta
return meta return meta
def step(self) -> None:
"""Advance the counter and refresh all weight scales every N steps."""
self.step_count += 1
if self.step_count % self.update_interval == 0:
for meta in self._metas.values():
meta.refresh()
def reset(self) -> None: def reset(self) -> None:
self.enabled = False """Restore construction defaults (switch, recipe, format) and drop all
self.step_count = 0 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() self._metas.clear()
self._last_device = None
# Global singleton: autograd backward runs on the engine worker threads, so # Process-wide singleton; per-thread/per-region state lives in _active_config.
# thread-local state would lose the fp8 flag during loss.backward(). The GIL
# protects Python-side mutation; the CUDA kernels take their own mutex.
_state = FP8State() _state = FP8State()
@@ -128,120 +193,245 @@ def fp8_state() -> FP8State:
return _state return _state
@contextmanager def _active() -> Optional[_ActiveConfig]:
def fp8_autocast(enabled: bool = True, update_interval: int = 16): """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. """Autocast-style context: fp8 linear dispatch on this thread.
Usage:: Mirrors ``torch.autocast`` — a class-based, reentrant, nestable context
over thread-local state::
with fp8_autocast(enabled=True): with fp8_autocast(enabled=True, fp8_format="hybrid"):
logits = model(input_ids) # aten::linear -> fp8 path logits = model(input_ids) # aten::linear -> fp8 path
loss.backward() loss.backward() # fp8 backward; state was captured at forward time
The scale-update counter advances once per ``enter`` (one training step), Nesting follows torch: each ``__enter__`` pushes the new active config, each
refreshing weight scales from their amax history every ``update_interval``. ``__exit__`` restores the previous one, and a nested ``enabled=False`` region
simply disables dispatch inside it. The instance doubles as a decorator.
""" """
state = fp8_state()
prev_enabled = state.enabled
prev_interval = state.update_interval
state.enabled = enabled
state.update_interval = update_interval
try:
if enabled:
state.step()
yield
finally:
state.enabled = prev_enabled
state.update_interval = prev_interval
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 _update_delayed_scale(scale, scale_inv, amax) -> None: def __enter__(self) -> "fp8_autocast":
"""scale = amax / 448 for the *next* call (device-side, no sync).""" self._tokens.append(_active_config.set(self._config))
amax_f = amax.reshape(()).to(torch.float32).clamp_min(1e-12) return self
scale.copy_(amax_f / E4M3_MAX)
scale_inv.copy_(E4M3_MAX / amax_f)
def __exit__(self, exc_type, exc_val, exc_tb) -> bool:
token = self._tokens.pop()
_active_config.reset(token)
return False
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None): def __call__(self, func):
"""TE-style scaled fp8 linear forward (called from the aten::linear impl). @functools.wraps(func)
def decorate(*args, **kwargs):
with self:
return func(*args, **kwargs)
x uses the delayed scale of its paired weight meta (amax from the previous return decorate
forward of this linear); the quantize kernel emits the current amax for the
next step. No extra abs/max reduce.
"""
if bias is None:
bias = torch.empty(0, device=x.device, dtype=x.dtype)
state = fp8_state()
meta = state.get_weight_meta(w)
amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
out = linear_forward_scaled(
x,
w,
bias,
meta.x_scale,
meta.scale,
meta.x_scale_inv,
meta.scale_inv,
amax_x,
amax_w,
)
meta.record(amax_w)
_update_delayed_scale(meta.x_scale, meta.x_scale_inv, amax_x)
return out
def fp8_linear_backward(g, x, w, masks):
"""TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
state = fp8_state()
meta = state.get_weight_meta(w)
amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
out = linear_backward_scaled(
g,
x,
w,
masks,
meta.g_scale,
meta.scale,
meta.x_scale,
meta.g_scale_inv,
meta.scale_inv,
meta.x_scale_inv,
amax_g,
)
_update_delayed_scale(meta.g_scale, meta.g_scale_inv, amax_g)
return out
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Layer 3: aten::linear integration # 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: def fp8_linear_enable(enabled: bool = True) -> None:
"""Toggle fp8 dispatch for aten::linear (global; backward runs on engine """Toggle fp8 dispatch for aten::linear globally (the out-of-region default;
worker threads, so a thread-local flag would be lost during backward).""" ``fp8_autocast`` regions override it thread-locally)."""
fp8_state().enabled = enabled fp8_state().default_enabled = enabled
def fp8_linear_enabled() -> bool: def fp8_linear_enabled() -> bool:
return fp8_state().enabled """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: def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
"""cuBLASLt fp8 requires M % 16 == 0 and N % 16 == 0 (K is padded).""" """Shape guard for the fp8 path. Unlike a strict 16-alignment requirement,
m = x.numel() // x.size(-1) the kernels handle unaligned M/N via boundary checks (slower but correct) —
return m % 16 == 0 and w.size(0) % 16 == 0 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): def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
if ( if (
fp8_linear_enabled() _active() is not None
and x.dtype == torch.bfloat16 and x.dtype is torch.bfloat16
and w.dtype == torch.bfloat16 and w.dtype is torch.bfloat16
and _fp8_supported(x, w) and _fp8_supported(x, w)
): ):
return fp8_linear_forward(x, w, bias) return _LinearFp8.apply(x, w, bias)
return torch.ops.aten.linear.default.redispatch( return torch.ops.aten.linear.default.redispatch(
torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd), torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd),
x, x,
@@ -250,35 +440,11 @@ def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
) )
def _linear_backward_cuda_impl(input_tensor, grad_output, weight, output_mask):
if (
fp8_linear_enabled()
and weight.dtype == torch.bfloat16
and _fp8_supported(grad_output, weight)
):
return fp8_linear_backward(grad_output, input_tensor, weight, list(output_mask))
compute_dtype = weight.dtype
grad = grad_output.to(compute_dtype)
grad_2d = grad.reshape(-1, weight.size(0))
input_2d = input_tensor.reshape(-1, input_tensor.size(-1)).to(compute_dtype)
grad_input = (
torch.mm(grad_2d, weight)
if output_mask[0]
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
)
grad_weight = (
torch.mm(grad_2d.t(), input_2d)
if output_mask[1]
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
)
grad_bias = (
grad.sum(dim=0)
if output_mask[2]
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
)
return grad_input.reshape_as(input_tensor), grad_weight, grad_bias
_lib = Library("aten", "IMPL", "CUDA") _lib = Library("aten", "IMPL", "CUDA")
_lib.impl("linear", _linear_cuda_impl) _lib.impl("linear", _linear_cuda_impl)
_lib.impl("linear_backward", _linear_backward_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)
-73
View File
@@ -1,73 +0,0 @@
"""FP8 CUDA kernel interface adapter (the only module touching the pybind.
Isolates the ``fp8_mm`` CUDA extension behind stable Python functions:
- availability / dtype checks and clear errors
- torch.library ``custom::fp8_mm`` registration (meta + CPU fallback)
- quantize-in-GEMM primitives used by ``fp8.py`` training state
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
this module is stateless.
"""
import torch
from torch.library import custom_op
from astrai.extension.loader import get_module, is_available
def _mod():
if not is_available("fp8_mm"):
raise RuntimeError(
"CUDA kernel 'fp8_mm' is not available. Build with CSRC_KERNELS=true."
)
return get_module("fp8_mm")
@custom_op("custom::fp8_mm", mutates_args=())
def fp8_mm(
a: torch.Tensor, b: torch.Tensor, sx: torch.Tensor, sw: torch.Tensor
) -> torch.Tensor:
"""FP8 e4m3 GEMM: a[M,K] x b[N,K] -> bf16[M,N] (pre-scaled inputs)."""
@fp8_mm.register_fake
def _fp8_mm_fake(a, b, sx, sw):
return torch.empty((a.size(0), b.size(1)), device=a.device, dtype=torch.bfloat16)
@fp8_mm.register_kernel("cuda")
def _fp8_mm_cuda(a, b, sx, sw):
return _mod().fp8_mm(a, b)
@fp8_mm.register_kernel("cpu")
def _fp8_mm_cpu(a, b, sx, sw):
return torch.mm(a.float(), b.float().t()).to(torch.bfloat16)
def linear_forward_scaled(x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w):
"""Quantize x/w with per-tensor scales + cuBLASLt GEMM + bias -> bf16.
x/w: [..., K] / [N, K] bf16; sx/sw: f32 scale tensors (device scalars);
sx_inv/sw_inv: 1/scale; amax_x/amax_w: f32 buffers receiving max-abs.
"""
if not (x.dtype == torch.bfloat16 and w.dtype == torch.bfloat16):
raise TypeError(f"fp8 forward requires bf16 inputs, got {x.dtype}/{w.dtype}")
return _mod().fp8_linear_forward_scaled(
x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w
)
def linear_backward_scaled(g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g):
"""dX = g @ W, dW = g^T @ X, dB = sum(g) with per-tensor scales."""
if not (
g.dtype == torch.bfloat16
and x.dtype == torch.bfloat16
and w.dtype == torch.bfloat16
):
raise TypeError(
f"fp8 backward requires bf16 inputs, got {g.dtype}/{x.dtype}/{w.dtype}"
)
return _mod().fp8_linear_backward_scaled(
g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g
)
+63 -23
View File
@@ -1,43 +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] = []
"fp8_mm", 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,6 +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,
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,
@@ -113,9 +107,11 @@ 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
new_k: current-token K to append, [batch, n_kv_heads, head_dim]
new_v: current-token V to append, same shape as new_k
mask: 2D [batch, max_context_len] (bool, True=keep) or None 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)
@@ -125,15 +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,
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,
@@ -150,6 +148,8 @@ 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,
is_causal: bool = False, is_causal: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
@@ -163,19 +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
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,
@@ -183,6 +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,
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)
+2 -65
View File
@@ -12,45 +12,10 @@ Modules:
- engine.py: Facade (InferenceEngine) - engine.py: Facade (InferenceEngine)
""" """
from astrai.inference.cache import (
Allocator,
KVCache,
KVStorage,
PagePool,
RadixCache,
ReqToTokenPool,
TaskCacheManager,
page_hash,
)
from astrai.inference.engine import InferenceEngine from astrai.inference.engine import InferenceEngine
from astrai.inference.network import ( from astrai.inference.network import get_app, run_server
AnthropicMessage,
BaseToolParser,
ChatCompletionRequest,
ChatMessage,
FunctionDef,
GenContext,
MessagesRequest,
ProtocolHandler,
SimpleJsonToolParser,
StopChecker,
ToolDef,
ToolParserFactory,
get_app,
run_server,
)
from astrai.inference.network.anthropic import AnthropicResponseBuilder
from astrai.inference.network.openai import OpenAIResponseBuilder
from astrai.inference.runtime.executor import Executor from astrai.inference.runtime.executor import Executor
from astrai.inference.runtime.sample import ( from astrai.inference.runtime.sample import sample
BaseSamplingStrategy,
FrequencyPenaltyStrategy,
SamplingPipeline,
TemperatureStrategy,
TopKStrategy,
TopPStrategy,
sample,
)
from astrai.inference.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
@@ -62,35 +27,7 @@ __all__ = [
"Task", "Task",
"TaskManager", "TaskManager",
"TaskStatus", "TaskStatus",
"Allocator",
"KVCache",
"KVStorage",
"PagePool",
"RadixCache",
"ReqToTokenPool",
"TaskCacheManager",
"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",
] ]
+3 -11
View File
@@ -27,7 +27,7 @@ class ReqToTokenPool:
self.size = size self.size = size
self.max_context_len = max_context_len self.max_context_len = max_context_len
self.req_to_token = torch.zeros( self.req_to_token = torch.zeros(
(size, max_context_len), dtype=torch.long, device=device (size, max_context_len), dtype=torch.int32, device=device
) )
self.free_slots = list(range(size)) self.free_slots = list(range(size))
self._lock = threading.Lock() self._lock = threading.Lock()
@@ -72,16 +72,6 @@ class KVStorage:
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype (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 @dataclass
class KVCache: class KVCache:
@@ -99,6 +89,8 @@ class KVCache:
max_len: int = 0 max_len: int = 0
kv_indptr: Optional[Tensor] = None kv_indptr: Optional[Tensor] = None
qo_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_o_part: Optional[Tensor] = None
decode_ml_part: Optional[Tensor] = None decode_ml_part: Optional[Tensor] = None
decode_out: Optional[Tensor] = None decode_out: Optional[Tensor] = None
+44 -13
View File
@@ -25,7 +25,7 @@ from astrai.inference.cache.strategy import (
RadixCache, RadixCache,
TaskCacheState, TaskCacheState,
) )
from astrai.inference.workspace import InferenceWorkspace from astrai.inference.workspace import Q_TILE_ROWS, InferenceWorkspace
# Re-export everything so existing ``from astrai.inference.cache import ...`` # Re-export everything so existing ``from astrai.inference.cache import ...``
# continues to work unchanged after the file split. # continues to work unchanged after the file split.
@@ -115,6 +115,8 @@ class PagePool:
self.contiguous = n_tokens is None self.contiguous = n_tokens is None
self.n_tokens = max_batch_size * max_seq_len if self.contiguous else n_tokens 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._storage = KVStorage(
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
@@ -124,7 +126,10 @@ class PagePool:
if self.contiguous: if self.contiguous:
for i in range(max_batch_size): for i in range(max_batch_size):
self._req_pool.req_to_token[i] = torch.arange( self._req_pool.req_to_token[i] = torch.arange(
i * max_seq_len, (i + 1) * max_seq_len, device=device i * max_seq_len,
(i + 1) * max_seq_len,
dtype=torch.int32,
device=device,
) )
self._strategy: AllocationStrategy = ContiguousStrategy() self._strategy: AllocationStrategy = ContiguousStrategy()
else: else:
@@ -184,7 +189,7 @@ class PagePool:
kvp_buf[: b + 1] += inc_buf[: b + 1] kvp_buf[: b + 1] += inc_buf[: b + 1]
else: else:
rpi_buf[:b].copy_( rpi_buf[:b].copy_(
torch.tensor(req_indices, dtype=torch.long, device=device) torch.tensor(req_indices, dtype=torch.int32, device=device)
) )
sl_buf[:b].copy_(torch.tensor(seq_lens, 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[: b + 1].zero_()
@@ -195,24 +200,48 @@ class PagePool:
kv_indptr = kvp_buf[: b + 1] kv_indptr = kvp_buf[: b + 1]
if start_pos is not None: if start_pos is not None:
# ---- prefill: out_cache_loc covers prefix range [start_pos:seq_len] ---- # Packed prefill concatenates each request's query tokens.
seq_len = seq_lens[0] q_lens = [seq_len - start_pos for seq_len in seq_lens]
out_cache_loc = self._req_pool.req_to_token[ if any(q_len <= 0 for q_len in q_lens):
req_pool_indices, start_pos:seq_len raise ValueError("prefill sequence lengths must exceed start_pos")
] out_cache_loc = torch.cat(
q_len = seq_len - start_pos [
workspace.qo_indptr[: b + 1].copy_( self._req_pool.req_to_token[
torch.arange(b + 1, dtype=torch.int32, device=device) * q_len 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] 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 decode_o_part = decode_ml_part = decode_out = None
else: else:
# ---- decode: out_cache_loc is a single column (last position) ---- # ---- decode: out_cache_loc is a single column (last position) ----
write_pos = seq_lens_t - 1 write_pos = seq_lens_t - 1
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1) loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
ocl_buf[:b].copy_(loc) ocl_buf[:b].copy_(loc)
out_cache_loc = ocl_buf[:b] out_cache_loc = ocl_buf[:b].reshape(-1)
qo_indptr = None 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_o_part = getattr(workspace, "decode_o_part", None)
decode_ml_part = getattr(workspace, "decode_ml_part", None) decode_ml_part = getattr(workspace, "decode_ml_part", None)
decode_out = getattr(workspace, "decode_out", None) decode_out = getattr(workspace, "decode_out", None)
@@ -227,6 +256,8 @@ class PagePool:
max_len=max(seq_lens), max_len=max(seq_lens),
kv_indptr=kv_indptr, kv_indptr=kv_indptr,
qo_indptr=qo_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_o_part=decode_o_part,
decode_ml_part=decode_ml_part, decode_ml_part=decode_ml_part,
decode_out=decode_out, decode_out=decode_out,
-2
View File
@@ -16,8 +16,6 @@ from abc import ABC, abstractmethod
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Callable, Dict, List, Optional, OrderedDict from typing import Callable, Dict, List, Optional, OrderedDict
import torch
from astrai.inference.cache.buffer import ReqToTokenPool from astrai.inference.cache.buffer import ReqToTokenPool
# ---- data contract: per-task slot state ---- # ---- data contract: per-task slot state ----
+2 -3
View File
@@ -156,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
-22
View File
@@ -148,17 +148,6 @@ class MetricsCollector:
self._completed.append(timing) self._completed.append(timing)
self._accumulate(timing) self._accumulate(timing)
def clear(self):
"""Reset all state (e.g. on engine shutdown)."""
self._timings.clear()
self._completed.clear()
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
# timing scopes # timing scopes
@contextmanager @contextmanager
@@ -180,17 +169,6 @@ class MetricsCollector:
t._decode_steps += 1 t._decode_steps += 1
t._decode_total_s += dt t._decode_total_s += dt
# access
def get_timing(self, task_id: str) -> Optional[TaskTiming]:
"""Return the timing record for *task_id* (active or completed)."""
if task_id in self._timings:
return self._timings[task_id]
for t in self._completed:
if t.task_id == task_id:
return t
return None
# aggregate stats # aggregate stats
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
+62 -49
View File
@@ -7,7 +7,7 @@ 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 (
CudaBackend, CudaBackend,
get_backend, get_backend,
) )
@@ -67,11 +67,15 @@ class DecodeSteadyState:
When the same ordered task set decodes one token per step, sampling When the same ordered task set decodes one token per step, sampling
params and task signature are reused; only positions advance by 1. 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 task_sig: tuple
positions: list[int] positions: list[int]
sampling_info: SamplingBatchInfo 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:
@@ -118,13 +122,13 @@ def _warmup_cuda_graphs(
timed("warmup prefill", logger), timed("warmup prefill", logger),
): ):
kv = task_cache.bind([tid], ws, start_pos=0) kv = task_cache.bind([tid], ws, start_pos=0)
ids_in = torch.arange(warmup_len, device=dev).unsqueeze(0) ids_in = torch.arange(warmup_len, device=dev)
pos_in = ids_in pos_in = ids_in
model( model(
ids_in, ids_in,
input_mask=pos_in.unsqueeze(-1) >= torch.arange(warmup_len, device=dev),
kv_cache=kv, kv_cache=kv,
position_ids=pos_in, position_ids=pos_in,
fwd="prefill",
) )
task_cache.task_free(tid) task_cache.task_free(tid)
@@ -159,15 +163,14 @@ def _warmup_cuda_graphs(
for tid in task_ids: for tid in task_ids:
task_cache.task_extend(tid, seq_pos) task_cache.task_extend(tid, seq_pos)
kv = task_cache.bind(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:
@@ -206,8 +209,8 @@ class Executor:
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
backend = get_backend() backend = get_backend()
self._graph_supported = backend.supports_graph() and CudaBackend.supports( self._graph_supported = backend.supports_graph() and (
head_dim=head_dim 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,
@@ -251,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 = [
@@ -285,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,
@@ -308,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(),
@@ -329,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.task_cache.bind( 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
@@ -363,24 +372,30 @@ 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] cached = self._decode_cache
sig_match = cached is not None and cached.task_sig == task_sig
if sig_match and cached.last_tokens is not None:
with torch.inference_mode():
input_ids = ws.fill_input_ids_from_device(cached.last_tokens)
else:
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) kv_cache = self.task_cache.bind(task_ids, ws)
task_sig = tuple(task_ids) reuse_decode_state = self.task_cache.bind_was_steady and sig_match
reuse_decode_state = (
self.task_cache.bind_was_steady
and self._decode_cache is not None
and self._decode_cache.task_sig == task_sig
)
if reuse_decode_state: if reuse_decode_state:
info = self._decode_cache.sampling_info info = self._decode_cache.sampling_info
ws.position_ids[:b] += 1 ws.position_ids[:b] += 1
@@ -391,9 +406,6 @@ class Executor:
) )
self._decode_cache = DecodeSteadyState(task_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)
# ---- forward (graph replay or live run + capture) ---- # ---- forward (graph replay or live run + capture) ----
use_graph = ( use_graph = (
@@ -402,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),
@@ -413,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
+22 -30
View File
@@ -79,29 +79,21 @@ class InferenceScheduler:
if backend is None: if backend is None:
self._backend = None self._backend = None
default_backend = get_backend() active_backend = get_backend()
self._backend_name = type(default_backend).__name__
with attn_backend(default_backend):
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,
)
else: else:
with attn_backend(backend): active_backend = backend
with attn_backend(active_backend):
if backend is not None:
self._backend = get_backend() self._backend = get_backend()
self._backend_name = type(self._backend).__name__ self._backend_name = type(get_backend()).__name__
self._executor = Executor( self._executor = Executor(
model=model, model=model,
kv_cache=self._cache, kv_cache=self._cache,
task_cache=self._task_cache, task_cache=self._task_cache,
device=self.device, device=self.device,
dtype=self.dtype, dtype=self.dtype,
enable_cuda_graph=enable_cuda_graph, 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
@@ -283,12 +275,7 @@ class InferenceScheduler:
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)
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)
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():
@@ -304,15 +291,20 @@ 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
self._abort_and_clear(free_waiting=True)
if torch.cuda.is_available():
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(): for task in self._task_mgr.get_active_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP) self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_cache.task_free(task.task_id) self._task_cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks(): for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP) self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_cache.task_free(task.task_id) if free_waiting:
self._task_cache.task_free(task.task_id)
self._task_mgr.clear_queues() self._task_mgr.clear_queues()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def run_batch( def run_batch(
self, self,
+22 -9
View File
@@ -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.
+11 -1
View File
@@ -1,6 +1,15 @@
import logging import logging
import os 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"): def setup_logging(level: str = "INFO"):
"""Attach a StreamHandler to the ``astrai`` logger (idempotent). """Attach a StreamHandler to the ``astrai`` logger (idempotent).
@@ -18,9 +27,10 @@ def setup_logging(level: str = "INFO"):
level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper() level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper()
logger.setLevel(getattr(logging, level_name, logging.INFO)) logger.setLevel(getattr(logging, level_name, logging.INFO))
handler = logging.StreamHandler() handler = logging.StreamHandler()
handler.addFilter(_DistributedContextFilter())
handler.setFormatter( handler.setFormatter(
logging.Formatter( logging.Formatter(
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s", "%(asctime)s | %(levelname)-8s | rank=%(rank)2s/%(world_size)-2s | %(name)-32s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S", datefmt="%Y-%m-%d %H:%M:%S",
) )
) )
+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
+12 -11
View File
@@ -5,8 +5,7 @@ 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.cache import KVCache from astrai.inference.cache import KVCache
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
@@ -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))
+2
View File
@@ -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()
+14 -1
View File
@@ -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
+2 -18
View File
@@ -94,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")
+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):
+23 -6
View File
@@ -8,6 +8,7 @@ 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.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
@@ -15,7 +16,13 @@ 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
@@ -126,9 +133,18 @@ class TrainContextBuilder:
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():
state.model_config = load_json(config_path) state.model_config = adapt_config(load_json(config_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:
if 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.state_dict = checkpoint.state_dict
state.model_config = checkpoint.config or state.model_config state.model_config = checkpoint.config or state.model_config
if self._resume: if self._resume:
@@ -140,8 +156,10 @@ class TrainContextBuilder:
checkpoint.consumed_samples // per_step * per_step checkpoint.consumed_samples // per_step * per_step
) )
state.checkpoint = checkpoint state.checkpoint = checkpoint
if not state.model_config and hasattr(cfg.model_fn(), "config"): if not state.model_config:
state.model_config = cfg.model_fn().config.to_dict() model = cfg.model_fn()
if hasattr(model, "config"):
state.model_config = model.config.to_dict()
return state return state
def _create_context( def _create_context(
@@ -204,7 +222,6 @@ class TrainContextBuilder:
def _create_dataloaders( def _create_dataloaders(
self, context: TrainContext, train_dataset, val_dataset self, context: TrainContext, train_dataset, val_dataset
) -> None: ) -> None:
cfg = self.config
sampler_offset = context.consumed_samples // context.world_size sampler_offset = context.consumed_samples // context.world_size
if self._resume and sampler_offset > 0: if self._resume and sampler_offset > 0:
samples_per_replica = ( samples_per_replica = (
@@ -261,7 +278,7 @@ class TrainContextBuilder:
def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict: def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict:
cfg = self.config cfg = self.config
kwargs = dict(cfg.extra_kwargs) kwargs = dict(cfg.strategy_kwargs)
kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef) kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"): if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"):
kwargs["ref_model"] = create_ref_model( kwargs["ref_model"] = create_ref_model(
+37 -6
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 fp8_mm) # 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})
@@ -61,9 +95,6 @@ foreach(name ${KERNELS})
"${PYTHON_INCLUDE_DIR}") "${PYTHON_INCLUDE_DIR}")
target_link_libraries(${name} PRIVATE ${TORCH_LIBS}) target_link_libraries(${name} PRIVATE ${TORCH_LIBS})
if(${name} STREQUAL "fp8_mm")
target_link_libraries(${name} PRIVATE CUDA::cublasLt)
endif()
target_link_options(${name} PRIVATE "-Wl,-rpath,${TORCH_LIB_DIR}") target_link_options(${name} PRIVATE "-Wl,-rpath,${TORCH_LIB_DIR}")
target_compile_options(${name} PRIVATE target_compile_options(${name} PRIVATE
+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,
@@ -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
@@ -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;
} }
@@ -141,3 +146,6 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride; int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
p.o_ptr[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
@@ -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 ----
@@ -123,9 +128,9 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
for (int it = 0; it < ntiles; it++) { for (int it = 0; it < ntiles; it++) {
if (it + 1 == ntiles) if (it + 1 == ntiles)
cp_async_wait_group<0>(); astrai::cp_async_wait_group<0>();
else else
cp_async_wait_group<STAGES - 1>(); astrai::cp_async_wait_group<STAGES - 1>();
__syncwarp(); __syncwarp();
process_tile(it, it & (STAGES - 1)); process_tile(it, it & (STAGES - 1));
__syncwarp(); __syncwarp();
@@ -136,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);
@@ -178,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
constexpr int ROWS = Traits::BR * WARPS; // share each block's K/V stream (~HB× less global K/V traffic).
dim3 grid(KV::host_q_blocks(p, ROWS), p.q_head, // Each head gets WPH = WARPS/HB 16-row chunks per block, so per-head
KV::kPaged ? 1 : 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 = (HEAD_DIM == 32) ? 4 : 8, ROWS = 32, P_BC = 32; constexpr int G = (HEAD_DIM == 32) ? 4 : 8, ROWS = 64, P_BC = 32;
dim3 grid(KV::host_q_blocks(p, ROWS), p.q_head, dim3 grid(QSchedule::host_q_blocks(p, ROWS), p.q_head,
KV::kPaged ? 1 : 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
} }
@@ -206,3 +242,6 @@ static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t
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
@@ -130,6 +132,8 @@ inline void attn_pack_params(
p.q_ptr = (const T*)q.data_ptr(); p.q_ptr = (const T*)q.data_ptr();
p.k_ptr = (const T*)k.data_ptr(); p.k_ptr = (const T*)k.data_ptr();
p.v_ptr = (const T*)v.data_ptr(); p.v_ptr = (const T*)v.data_ptr();
p.new_k_ptr = nullptr;
p.new_v_ptr = nullptr;
p.o_ptr = nullptr; p.o_ptr = nullptr;
p.o_part = nullptr; p.o_part = nullptr;
p.ml_part = nullptr; p.ml_part = nullptr;
@@ -148,6 +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,
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,
@@ -160,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]");
@@ -184,12 +191,39 @@ inline void attn_pack_paged_decode_params(
p.k_ptr = (const T*)k_cache.data_ptr(); p.k_ptr = (const T*)k_cache.data_ptr();
p.v_ptr = (const T*)v_cache.data_ptr(); p.v_ptr = (const T*)v_cache.data_ptr();
p.q_ptr = (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);
TORCH_CHECK(new_k.has_value() == new_v.has_value(),
"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;
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);
@@ -226,6 +260,8 @@ 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 causal_offset, int64_t causal_offset,
double scale, double scale,
@@ -236,13 +272,19 @@ 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]");
@@ -259,6 +301,10 @@ inline void attn_pack_paged_prefill_params(
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_l_stride = (int)q.stride(0); p.q_l_stride = (int)q.stride(0);
p.q_h_stride = (int)q.stride(1); p.q_h_stride = (int)q.stride(1);
@@ -266,11 +312,16 @@ inline void attn_pack_paged_prefill_params(
p.k_ptr = (const T*)k_cache.data_ptr(); p.k_ptr = (const T*)k_cache.data_ptr();
p.v_ptr = (const T*)v_cache.data_ptr(); p.v_ptr = (const T*)v_cache.data_ptr();
p.new_k_ptr = nullptr;
p.new_v_ptr = nullptr;
p.q_ptr = (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 = 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.causal_offset = (int)causal_offset; p.causal_offset = (int)causal_offset;
@@ -307,3 +358,6 @@ inline void attn_pack_paged_prefill_params(
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,62 +70,18 @@ __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
@@ -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]);
} }
} }
} }
@@ -290,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,6 +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,
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,
@@ -21,6 +25,7 @@ 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,
new_k, new_v,
mask, causal_offset, scale, p); mask, causal_offset, scale, p);
torch::Tensor O; torch::Tensor O;
@@ -71,6 +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("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,6 +11,8 @@ 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 causal_offset, int64_t causal_offset,
double scale double scale
@@ -18,8 +22,9 @@ 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,
q_tile_to_batch, q_tile_to_index, mask,
causal_offset, scale, p); 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());
@@ -39,6 +44,8 @@ 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("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_prefill( torch::Tensor attn_prefill(
torch::Tensor q, torch::Tensor q,
@@ -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,13 +29,13 @@ __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 batch, q_tile; int batch, q_tile;
if (!map_q_block<ROWS, KV>(p, batch, q_tile)) QSchedule::map_block(p, batch, q_tile);
return;
int q_head = blockIdx.y; int q_head = blockIdx.y;
int gpos = threadIdx.x; // 0..G-1 (which d-chunk) int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
@@ -47,8 +44,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
// 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);
@@ -56,7 +53,7 @@ __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_l_stride + gpos * DPT * p.q_d_stride; int q_off = q_base + q_row * p.q_l_stride + gpos * DPT * p.q_d_stride;
@@ -90,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;
} }
@@ -154,3 +152,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
p.o_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl); p.o_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl);
} }
} }
} // namespace attention
} // namespace astrai
@@ -1,39 +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;
int batch, q_tile; const int HB = min(G, Traits::WARPS); // q heads packed per block
if (!map_q_block<Traits::BR * Traits::WARPS, KV>(p, batch, q_tile)) const int WPH = Traits::WARPS / HB; // 16-row chunks per head
return; const int BPG = (G + HB - 1) / HB; // blocks per GQA group
const int kv_head = q_head / (p.q_head / p.kv_head); const int chunk = warp % WPH;
const int qrow0 = (q_tile * Traits::WARPS + warp) * Traits::BR;
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
@@ -42,7 +63,7 @@ __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;
@@ -60,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;
q_tile * 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) {
@@ -83,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 ----
@@ -98,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);
@@ -137,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_ptr[o_base + qr0 * p.q_l_stride + d * p.q_d_stride]) = 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_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v; &p.o_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v;
} }
} }
} }
} // namespace attention
} // namespace astrai
-69
View File
@@ -1,69 +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 {
// 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;
int use_mask;
// pointers
const T* __restrict__ q_ptr;
const T* __restrict__ k_ptr;
const T* __restrict__ v_ptr;
T* __restrict__ o_ptr;
const bool* __restrict__ mask;
// 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 mask_b_stride;
int mask_h_stride;
int mask_l_stride;
// Paged K/V addressing
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 for decode
int max_context_len; // req_to_token stride (dim 1)
// Decode split-KV workspace
int num_splits;
AT* __restrict__ o_part;
AT* __restrict__ ml_part;
};
-212
View File
@@ -1,212 +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_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.
// ============================================================================
// 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_b_stride + kv_head*kv_h_stride
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_blocks(const AttentionParams<bf16>& p, int rows) {
return (p.q_len + rows - 1) / rows;
}
template <int ROWS>
HOST_DEV_FORCEINLINE bool map_q_tile(const AttentionParams<bf16>&,
int flat_tile, int grid_batch,
int& batch, int& q_tile) {
batch = grid_batch;
q_tile = flat_tile;
return true;
}
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_l_stride)
HOST_DEV_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;
}
// 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_b_stride + q_head * p.q_h_stride;
}
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_b_stride + kv_head * p.kv_h_stride;
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_l_stride + d * p.kv_d_stride;
return {&p.k_ptr[g_off], &p.v_ptr[g_off], valid};
}
};
// ---- Paged (SGLang-style flat pool) K/V ----
struct PagedKV {
static constexpr bool kPaged = true;
HOST_DEV_FORCEINLINE int host_q_blocks(const AttentionParams<bf16>& p, int rows) {
// sum(ceil(q_len[b] / rows)) <= ceil(total_q / rows) + batch - 1.
return (p.q_len + rows - 1) / rows + p.batch - 1;
}
template <int ROWS>
HOST_DEV_FORCEINLINE bool map_q_tile(const AttentionParams<bf16>& p,
int flat_tile, int,
int& batch, int& q_tile) {
int tile_base = 0;
for (int b = 0; b < p.batch; ++b) {
int len = p.qo_indptr[b + 1] - p.qo_indptr[b];
int tiles = (len + ROWS - 1) / ROWS;
if (flat_tile < tile_base + tiles) {
batch = b;
q_tile = flat_tile - tile_base;
return true;
}
tile_base += tiles;
}
return false;
}
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_l_stride + q_head * p.q_h_stride;
}
// 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_l_stride + q_head * p.q_h_stride;
}
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_ptr[gmem_off], &p.v_ptr[gmem_off], ok};
}
};
// ---- Q-block mapping ----
// Contiguous grids map directly to (batch, q_tile). Paged grids flatten the
// ragged Q tiles, so one thread resolves the request and broadcasts it.
template <int ROWS, typename KV>
__device__ __forceinline__ bool map_q_block(
const AttentionParams<bf16>& p, int& batch, int& q_tile) {
if constexpr (!KV::kPaged) {
batch = blockIdx.z;
q_tile = blockIdx.x;
return true;
} else {
__shared__ int mapped_batch;
__shared__ int mapped_q_tile;
if ((threadIdx.x | threadIdx.y) == 0) {
mapped_batch = -1;
KV::template map_q_tile<ROWS>(
p, blockIdx.x, blockIdx.z, mapped_batch, mapped_q_tile);
}
__syncthreads();
batch = mapped_batch;
q_tile = mapped_q_tile;
return batch >= 0;
}
}
-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
-441
View File
@@ -1,441 +0,0 @@
// FP8 e4m3 matrix multiply via cuBLASLt (sm89 TN layout).
//
// cuBLASLt exposes fp8 kernels only for op(A)=T, op(B)=N on Ada; we exploit
// the identity: row-major a[M,K] == A^T as col-major [K,M] (zero copy), and
// row-major wT[N,K] == B as col-major [K,N] (zero copy). The col-major
// result D[M,N] is C^T in row-major terms, so we transpose the output once.
//
// Inputs arrive pre-scaled fp8 e4m3 tensors; output is unscaled fp32.
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublasLt.h>
#include <cuda_fp8.h>
#include <cstdint>
#include <mutex>
#include <unordered_map>
static std::recursive_mutex g_mutex;
static cublasLtHandle_t g_handle = nullptr;
static cublasLtMatmulDesc_t g_desc = nullptr;
static cublasLtMatrixLayout_t g_layout_a = nullptr;
static cublasLtMatrixLayout_t g_layout_b = nullptr;
static cublasLtMatrixLayout_t g_layout_c = nullptr;
static cublasLtMatmulPreference_t g_pref = nullptr;
static void* g_workspace = nullptr;
static size_t g_ws_size = 0;
struct ShapeKey {
int64_t m;
int64_t k;
int64_t n;
bool operator==(const ShapeKey& other) const {
return m == other.m && k == other.k && n == other.n;
}
};
struct ShapeKeyHash {
size_t operator()(const ShapeKey& s) const {
size_t h = std::hash<int64_t>()(s.m);
h ^= std::hash<int64_t>()(s.k) + 0x9e3779b9 + (h << 6) + (h >> 2);
h ^= std::hash<int64_t>()(s.n) + 0x9e3779b9 + (h << 6) + (h >> 2);
return h;
}
};
using AlgoCache = std::unordered_map<ShapeKey, cublasLtMatmulAlgo_t, ShapeKeyHash>;
static void create_matmul_config(cublasLtMatmulDesc_t* desc,
cublasLtMatrixLayout_t* layout_a,
cublasLtMatrixLayout_t* layout_b,
cublasLtMatrixLayout_t* layout_c) {
cublasOperation_t ta = CUBLAS_OP_T, tb = CUBLAS_OP_N;
TORCH_CHECK(cublasLtMatmulDescCreate(desc, CUBLAS_COMPUTE_32F, CUDA_R_32F) ==
CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatmulDescSetAttribute(
*desc, CUBLASLT_MATMUL_DESC_TRANSA, &ta, sizeof(ta)) ==
CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatmulDescSetAttribute(
*desc, CUBLASLT_MATMUL_DESC_TRANSB, &tb, sizeof(tb)) ==
CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatrixLayoutCreate(layout_a, CUDA_R_8F_E4M3, 1, 1, 1) ==
CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatrixLayoutCreate(layout_b, CUDA_R_8F_E4M3, 1, 1, 1) ==
CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatrixLayoutCreate(layout_c, CUDA_R_16BF, 1, 1, 1) ==
CUBLAS_STATUS_SUCCESS);
}
static void ensure_cublas_lt() {
std::lock_guard<std::recursive_mutex> lock(g_mutex);
if (g_handle) {
return;
}
TORCH_CHECK(cublasLtCreate(&g_handle) == CUBLAS_STATUS_SUCCESS);
create_matmul_config(&g_desc, &g_layout_a, &g_layout_b, &g_layout_c);
TORCH_CHECK(cublasLtMatmulPreferenceCreate(&g_pref) == CUBLAS_STATUS_SUCCESS);
size_t ws = 16 * 1024 * 1024;
TORCH_CHECK(cublasLtMatmulPreferenceSetAttribute(
g_pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &ws, sizeof(ws)) ==
CUBLAS_STATUS_SUCCESS);
}
static cublasStatus_t get_algo_cached(int64_t m, int64_t k, int64_t n,
AlgoCache* cache,
cublasLtMatmulAlgo_t* algo);
static void fp8_gemm_into(torch::Tensor lhs, torch::Tensor rhs, torch::Tensor out,
int64_t m, int64_t k, int64_t n,
const float* a_scale, const float* b_scale,
cudaStream_t stream);
static const float k_scale_one = 1.0f;
static void set_layout(cublasLtMatrixLayout_t layout, int64_t rows, int64_t cols,
int64_t ld) {
TORCH_CHECK(cublasLtMatrixLayoutSetAttribute(layout, CUBLASLT_MATRIX_LAYOUT_ROWS,
&rows, sizeof(rows)) ==
CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatrixLayoutSetAttribute(layout, CUBLASLT_MATRIX_LAYOUT_COLS,
&cols, sizeof(cols)) ==
CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatrixLayoutSetAttribute(layout, CUBLASLT_MATRIX_LAYOUT_LD,
&ld, sizeof(ld)) ==
CUBLAS_STATUS_SUCCESS);
}
torch::Tensor fp8_mm(torch::Tensor a, torch::Tensor b) {
TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required");
TORCH_CHECK(a.scalar_type() == torch::kFloat8_e4m3fn, "a must be float8_e4m3fn");
TORCH_CHECK(b.scalar_type() == torch::kFloat8_e4m3fn, "b must be float8_e4m3fn");
TORCH_CHECK(a.dim() == 2 && b.dim() == 2, "2D tensors required");
const at::cuda::OptionalCUDAGuard guard(a.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto a_c = a.contiguous();
auto b_c = b.contiguous();
int64_t m = a_c.size(0), k = a_c.size(1), n = b_c.size(0);
TORCH_CHECK(b_c.size(1) == k, "inner dim mismatch");
auto buf = torch::empty({m, n}, a_c.options().dtype(torch::kBFloat16));
ensure_cublas_lt();
fp8_gemm_into(a_c, b_c, buf, m, k, n, &k_scale_one, &k_scale_one,
stream.stream());
return buf;
}
// ---------------------------------------------------------------------------
// Quantize: bf16 * scale_inv -> fp8, one atomicMax amax per kernel call.
// amax_ptr must be zeroed before launch; float-bits atomicMax works because
// |v| >= 0 has a monotonic IEEE bit pattern.
// ---------------------------------------------------------------------------
template <typename T8>
__device__ __forceinline__ T8 cast_fp8(float v);
template <>
__device__ __forceinline__ __nv_fp8_e4m3 cast_fp8<__nv_fp8_e4m3>(float v) {
return __nv_fp8_e4m3(v);
}
template <>
__device__ __forceinline__ __nv_fp8_e5m2 cast_fp8<__nv_fp8_e5m2>(float v) {
return __nv_fp8_e5m2(v);
}
template <typename T8>
__global__ void quantize_kernel(const __nv_bfloat16* __restrict__ src,
const float* __restrict__ scale_inv,
T8* __restrict__ dst,
float* __restrict__ amax_ptr, int64_t n) {
int64_t i = blockIdx.x * (int64_t)blockDim.x + threadIdx.x;
float amax = 0.f;
if (i < n) {
float v = __bfloat162float(src[i]) * *scale_inv;
dst[i] = cast_fp8<T8>(v);
amax = fabsf(v);
}
for (int off = 16; off; off >>= 1)
amax = fmaxf(amax, __shfl_xor_sync(0xffffffffu, amax, off));
__shared__ float sm[8];
if ((threadIdx.x & 31) == 0) sm[threadIdx.x >> 5] = amax;
__syncthreads();
if (threadIdx.x == 0) {
float m = 0.f;
for (int w = 0; w < blockDim.x / 32; ++w) m = fmaxf(m, sm[w]);
atomicMax(reinterpret_cast<unsigned*>(amax_ptr), __float_as_uint(m));
}
}
// Same but with a transpose (rows x cols bf16 row-major -> fp8 [cols, rows]).
template <typename T8>
__global__ void transpose_quantize_kernel(
const __nv_bfloat16* __restrict__ src, const float* __restrict__ scale_inv,
T8* __restrict__ dst, float* __restrict__ amax_ptr, int64_t rows,
int64_t cols) {
__shared__ T8 tile[32][33];
int64_t x = blockIdx.x * 32 + threadIdx.x;
int64_t y = blockIdx.y * 32 + threadIdx.y;
float amax = 0.f;
for (int j = 0; j < 32; j += 8) {
if (x < cols && y + j < rows) {
float v = __bfloat162float(src[(y + j) * cols + x]) * *scale_inv;
tile[threadIdx.y + j][threadIdx.x] = cast_fp8<T8>(v);
amax = fmaxf(amax, fabsf(v));
}
}
__syncthreads();
x = blockIdx.y * 32 + threadIdx.x;
y = blockIdx.x * 32 + threadIdx.y;
for (int j = 0; j < 32; j += 8) {
if (x < rows && y + j < cols) {
dst[(y + j) * rows + x] = tile[threadIdx.x][threadIdx.y + j];
}
}
for (int off = 16; off; off >>= 1)
amax = fmaxf(amax, __shfl_xor_sync(0xffffffffu, amax, off));
__shared__ float sm[8];
if ((threadIdx.x & 31) == 0) sm[threadIdx.x >> 5] = amax;
__syncthreads();
if (threadIdx.x == 0) {
float m = 0.f;
for (int w = 0; w < blockDim.x / 32; ++w) m = fmaxf(m, sm[w]);
atomicMax(reinterpret_cast<unsigned*>(amax_ptr), __float_as_uint(m));
}
}
__global__ void bias_add_bf16_kernel(
__nv_bfloat16* __restrict__ dst, const __nv_bfloat16* __restrict__ bias,
int64_t total, int64_t n) {
// GEMM and output use the same row-major [M,N] layout.
int64_t idx = blockIdx.x * (int64_t)blockDim.x + threadIdx.x;
if (idx >= total) return;
float v = __bfloat162float(dst[idx]);
dst[idx] = __float2bfloat16(v + __bfloat162float(bias[idx % n]));
}
static cublasStatus_t get_algo_cached(int64_t m, int64_t k, int64_t n,
AlgoCache* cache,
cublasLtMatmulAlgo_t* algo) {
std::lock_guard<std::recursive_mutex> lock(g_mutex);
ShapeKey key{m, k, n};
auto it = cache->find(key);
if (it != cache->end()) {
*algo = it->second;
return CUBLAS_STATUS_SUCCESS;
}
cublasLtMatmulHeuristicResult_t heur;
int returned = 0;
cublasStatus_t st = cublasLtMatmulAlgoGetHeuristic(
g_handle, g_desc, g_layout_a, g_layout_b, g_layout_c, g_layout_c, g_pref, 1,
&heur, &returned);
if (st != CUBLAS_STATUS_SUCCESS || returned == 0)
return CUBLAS_STATUS_NOT_SUPPORTED;
if (heur.workspaceSize > g_ws_size) {
if (g_workspace) cudaFree(g_workspace);
TORCH_CHECK(cudaMalloc(&g_workspace, heur.workspaceSize) == cudaSuccess);
g_ws_size = heur.workspaceSize;
}
cache->emplace(key, heur.algo);
*algo = heur.algo;
return CUBLAS_STATUS_SUCCESS;
}
static void fp8_gemm_into(torch::Tensor lhs, torch::Tensor rhs, torch::Tensor out,
int64_t m, int64_t k, int64_t n,
const float* a_scale, const float* b_scale,
cudaStream_t stream) {
std::lock_guard<std::recursive_mutex> lock(g_mutex);
set_layout(g_layout_a, k, n, k); // param A = rhs (op=T -> [N,K])
set_layout(g_layout_b, k, m, k); // param B = lhs (op=N -> [K,M])
set_layout(g_layout_c, n, m, n); // col-major [N,M] == row-major [M,N]
// Per-tensor FP32 scales applied inside the GEMM:
// D = alpha * A_SCALE * B_SCALE * A * B (alpha = 1).
TORCH_CHECK(cublasLtMatmulDescSetAttribute(
g_desc, CUBLASLT_MATMUL_DESC_A_SCALE_POINTER, &a_scale,
sizeof(a_scale)) == CUBLAS_STATUS_SUCCESS);
TORCH_CHECK(cublasLtMatmulDescSetAttribute(
g_desc, CUBLASLT_MATMUL_DESC_B_SCALE_POINTER, &b_scale,
sizeof(b_scale)) == CUBLAS_STATUS_SUCCESS);
float alpha = 1.0f, beta = 0.0f;
static AlgoCache cache;
cublasLtMatmulAlgo_t algo;
cublasStatus_t st = get_algo_cached(m, k, n, &cache, &algo);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS,
"cublasLtMatmulAlgoGetHeuristic failed: ", cublasLtGetStatusName(st));
st = cublasLtMatmul(g_handle, g_desc, &alpha, rhs.data_ptr(), g_layout_a,
lhs.data_ptr(), g_layout_b, &beta, out.data_ptr(), g_layout_c,
out.data_ptr(), g_layout_c, &algo, g_workspace, g_ws_size,
stream);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS,
"cublasLtMatmul failed: ", cublasLtGetStatusName(st));
}
// ---------------------------------------------------------------------------
// Scaled FP8 linear forward: quantize x/w with per-tensor scales -> cublasLt
// GEMM (scales applied inside) -> bias in-place -> bf16 [..., N].
// sx/sw: f32 scale tensors (device scalars); sx_inv/sw_inv: 1/scale.
// amax_x/amax_w: f32 buffers receiving max-abs of the quantized tensors.
// ---------------------------------------------------------------------------
torch::Tensor fp8_linear_forward_scaled(torch::Tensor x, torch::Tensor w,
torch::Tensor bias, torch::Tensor sx,
torch::Tensor sw, torch::Tensor sx_inv,
torch::Tensor sw_inv,
torch::Tensor amax_x,
torch::Tensor amax_w) {
TORCH_CHECK(x.is_cuda() && w.is_cuda(), "CUDA tensors required");
TORCH_CHECK(x.dtype() == torch::kBFloat16 && w.dtype() == torch::kBFloat16,
"x and w must be bf16");
const at::cuda::OptionalCUDAGuard guard(x.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto x_c = x.reshape({-1, w.size(1)}).contiguous();
auto w_c = w.contiguous();
int64_t m = x_c.size(0), k = x_c.size(1), n = w_c.size(0);
TORCH_CHECK(w_c.size(1) == k, "inner dim mismatch");
ensure_cublas_lt();
const float* sx_ptr = sx.data_ptr<float>();
const float* sw_ptr = sw.data_ptr<float>();
const float* sxi_ptr = sx_inv.data_ptr<float>();
const float* swi_ptr = sw_inv.data_ptr<float>();
float* amax_x_ptr = amax_x.data_ptr<float>();
float* amax_w_ptr = amax_w.data_ptr<float>();
C10_CUDA_CHECK(cudaMemsetAsync(amax_x_ptr, 0, sizeof(float), stream.stream()));
C10_CUDA_CHECK(cudaMemsetAsync(amax_w_ptr, 0, sizeof(float), stream.stream()));
auto x8 = torch::empty({m, k}, x_c.options().dtype(torch::kFloat8_e4m3fn));
auto w8 = torch::empty({n, k}, w_c.options().dtype(torch::kFloat8_e4m3fn));
int64_t block = 256;
quantize_kernel<__nv_fp8_e4m3>
<<<(unsigned)((m * k + block - 1) / block), block, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(x_c.data_ptr()), sxi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(x8.data_ptr()), amax_x_ptr, m * k);
quantize_kernel<__nv_fp8_e4m3>
<<<(unsigned)((n * k + block - 1) / block), block, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(w_c.data_ptr()), swi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(w8.data_ptr()), amax_w_ptr, n * k);
C10_CUDA_CHECK(cudaGetLastError());
auto out = torch::empty({m, n}, x_c.options());
fp8_gemm_into(x8, w8, out, m, k, n, sw_ptr, sx_ptr, stream.stream());
if (bias.defined() && bias.numel() > 0) {
TORCH_CHECK(bias.scalar_type() == torch::kBFloat16 && bias.numel() == n,
"bias must be bf16 with shape [N]");
bias_add_bf16_kernel<<<(unsigned)((m * n + block - 1) / block), block, 0, stream>>>(
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(bias.data_ptr()), m * n, n);
C10_CUDA_CHECK(cudaGetLastError());
}
std::vector<int64_t> shape(x.sizes().begin(), x.sizes().end() - 1);
shape.push_back(n);
return out.reshape(shape);
}
// ---------------------------------------------------------------------------
// Scaled FP8 linear backward: dX = g @ W, dW = g^T @ X, dB = sum(g).
// Scales: g uses sg (immediate), w/x reuse the forward scales.
// ---------------------------------------------------------------------------
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scaled(
torch::Tensor g, torch::Tensor x, torch::Tensor w,
std::vector<int64_t> masks, torch::Tensor sg, torch::Tensor sw,
torch::Tensor sx, torch::Tensor sg_inv, torch::Tensor sw_inv,
torch::Tensor sx_inv, torch::Tensor amax_g) {
const at::cuda::OptionalCUDAGuard guard(g.device());
TORCH_CHECK(g.dtype() == torch::kBFloat16 && x.dtype() == torch::kBFloat16 &&
w.dtype() == torch::kBFloat16,
"g, x, and w must be bf16");
auto stream = at::cuda::getCurrentCUDAStream();
auto g_c = g.reshape({-1, w.size(0)}).contiguous();
auto x_c = x.reshape({-1, x.size(-1)}).contiguous();
auto w_c = w.contiguous();
int64_t m = g_c.size(0);
int64_t n = w.size(0);
int64_t k = w.size(1);
TORCH_CHECK(x_c.size(0) == m && x_c.size(1) == k && g_c.size(1) == n,
"backward shape mismatch");
auto grad_input = torch::empty_like(x);
auto grad_weight = torch::empty_like(w);
auto grad_bias = torch::empty({0}, g_c.options().dtype(g.dtype()));
ensure_cublas_lt();
const float* sg_ptr = sg.data_ptr<float>();
const float* sw_ptr = sw.data_ptr<float>();
const float* sx_ptr = sx.data_ptr<float>();
const float* sgi_ptr = sg_inv.data_ptr<float>();
const float* swi_ptr = sw_inv.data_ptr<float>();
const float* sxi_ptr = sx_inv.data_ptr<float>();
float* amax_g_ptr = amax_g.data_ptr<float>();
C10_CUDA_CHECK(cudaMemsetAsync(amax_g_ptr, 0, sizeof(float), stream.stream()));
auto fp8_options = g_c.options().dtype(torch::kFloat8_e4m3fn);
auto g8 = torch::empty({m, n}, fp8_options);
auto gt8 = masks[1] ? torch::empty({n, m}, fp8_options) : torch::Tensor();
auto wt8 = masks[0] ? torch::empty({k, n}, fp8_options) : torch::Tensor();
auto xt8 = masks[1] ? torch::empty({k, m}, fp8_options) : torch::Tensor();
int64_t block = 256;
quantize_kernel<__nv_fp8_e4m3>
<<<(unsigned)((m * n + block - 1) / block), block, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(g_c.data_ptr()), sgi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(g8.data_ptr()), amax_g_ptr, m * n);
dim3 threads(32, 8);
if (masks[0]) {
dim3 blocks((k + 31) / 32, (n + 31) / 32);
transpose_quantize_kernel<__nv_fp8_e4m3>
<<<blocks, threads, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(w_c.data_ptr()), swi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(wt8.data_ptr()), amax_g_ptr,
n, k);
fp8_gemm_into(g8, wt8, grad_input.reshape({m, k}), m, n, k, sg_ptr,
sw_ptr, stream.stream());
}
if (masks[1]) {
dim3 g_blocks((n + 31) / 32, (m + 31) / 32);
dim3 x_blocks((k + 31) / 32, (m + 31) / 32);
transpose_quantize_kernel<__nv_fp8_e4m3>
<<<g_blocks, threads, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(g_c.data_ptr()), sgi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(gt8.data_ptr()), amax_g_ptr,
m, n);
transpose_quantize_kernel<__nv_fp8_e4m3>
<<<x_blocks, threads, 0, stream.stream()>>>(
reinterpret_cast<const __nv_bfloat16*>(x_c.data_ptr()), sxi_ptr,
reinterpret_cast<__nv_fp8_e4m3*>(xt8.data_ptr()), amax_g_ptr,
m, k);
fp8_gemm_into(gt8, xt8, grad_weight, n, m, k, sg_ptr, sx_ptr,
stream.stream());
}
C10_CUDA_CHECK(cudaGetLastError());
if (masks[2]) {
grad_bias = g_c.sum(0).to(g.dtype());
}
return std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>(
grad_input, grad_weight, grad_bias);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("fp8_mm", &fp8_mm, py::arg("a"), py::arg("b"),
"FP8 e4m3 GEMM: a[M,K] x b[N,K] -> bf16[M,N] (pre-scaled inputs)");
m.def("fp8_linear_forward_scaled", &fp8_linear_forward_scaled,
py::arg("x"), py::arg("w"), py::arg("bias"), py::arg("sx"),
py::arg("sw"), py::arg("sx_inv"), py::arg("sw_inv"),
py::arg("amax_x"), py::arg("amax_w"),
"Scaled FP8 linear forward: quantize with per-tensor scales + "
"cublasLt GEMM (scales applied inside) + bias -> bf16");
m.def("fp8_linear_backward_scaled", &fp8_linear_backward_scaled,
py::arg("g"), py::arg("x"), py::arg("w"), py::arg("masks"),
py::arg("sg"), py::arg("sw"), py::arg("sx"), py::arg("sg_inv"),
py::arg("sw_inv"), py::arg("sx_inv"), py::arg("amax_g"),
"Scaled FP8 linear backward: dX = g*sw @ W, dW = (g*sx)^T @ X, "
"dB = sum(g)");
}
@@ -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"
); );
} }
+88 -46
View File
@@ -7,18 +7,40 @@
#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 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); } }; 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]
// req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B] // req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B]
// 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)
@@ -27,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;
@@ -35,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] *
@@ -66,7 +88,7 @@ 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_l_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,
@@ -78,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++) {
@@ -89,7 +111,7 @@ static void cpu_paged_prefill_ref(
for (int kj = 0; kj < lim; kj++) { for (int kj = 0; kj < lim; kj++) {
if (mask && !mask[b * mask_l_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] *
@@ -149,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);
@@ -181,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++) {
@@ -191,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);
@@ -216,7 +238,7 @@ 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.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -278,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;
@@ -312,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++) {
@@ -321,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);
@@ -351,7 +373,7 @@ 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.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -417,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);
@@ -446,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++) {
@@ -455,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);
@@ -483,8 +505,11 @@ 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.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -497,6 +522,8 @@ static int run_prefill_test(int B, int Hq, int Hkv,
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = 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.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; p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
@@ -523,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;
} }
@@ -546,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);
@@ -577,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++) {
@@ -586,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);
@@ -619,7 +647,11 @@ 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.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -632,6 +664,8 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = 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.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; p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
@@ -659,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;
} }
@@ -667,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);
@@ -696,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);
@@ -709,7 +744,7 @@ 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.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -749,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);
@@ -769,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);
@@ -786,7 +821,11 @@ 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.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -798,6 +837,8 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = 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.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; p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
auto launch = [&]() { auto launch = [&]() {
@@ -827,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() {
@@ -933,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();
+17 -6
View File
@@ -7,7 +7,9 @@ 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 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); } }; struct PrefillDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_prefill<H>(p, 0); } };
@@ -56,7 +58,7 @@ 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);
@@ -118,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;
@@ -136,7 +139,7 @@ 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);
@@ -183,7 +186,7 @@ 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);
@@ -229,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},
@@ -257,7 +266,7 @@ 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);
@@ -324,7 +333,9 @@ int main() {
{ {
const int configs[][7] = { const int configs[][7] = {
{1,2,1,64,128,32,0}, // scalar fallback D=32 {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;
}
+11 -7
View File
@@ -1,5 +1,6 @@
services: services:
server: server:
image: astrai:latest
build: build:
context: . context: .
dockerfile: Dockerfile dockerfile: Dockerfile
@@ -9,16 +10,18 @@ services:
USER_GID: ${ASTRAI_GID:-1000} USER_GID: ${ASTRAI_GID:-1000}
user: "${ASTRAI_UID:-1000}:${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"]
@@ -29,6 +32,7 @@ services:
restart: unless-stopped restart: unless-stopped
server-cpu: server-cpu:
image: astrai:latest
profiles: [cpu] profiles: [cpu]
build: build:
context: . context: .
@@ -39,9 +43,9 @@ services:
USER_GID: ${ASTRAI_GID:-1000} USER_GID: ${ASTRAI_GID:-1000}
user: "${ASTRAI_UID:-1000}:${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"]
@@ -52,6 +56,7 @@ services:
restart: unless-stopped restart: unless-stopped
trainer: trainer:
image: astrai:latest
profiles: [train] profiles: [train]
build: build:
context: . context: .
@@ -72,9 +77,8 @@ services:
- BASE_MODEL=${BASE_MODEL:-/models/base} - BASE_MODEL=${BASE_MODEL:-/models/base}
- CHECKPOINT_ROOT=/checkpoints - CHECKPOINT_ROOT=/checkpoints
- TRAIN_GPU_COUNT=${TRAIN_GPU_COUNT:-all} - TRAIN_GPU_COUNT=${TRAIN_GPU_COUNT:-all}
- TRAIN_PARALLEL_MODE=${TRAIN_PARALLEL_MODE:-auto}
- CUDA_VISIBLE_DEVICES - CUDA_VISIBLE_DEVICES
- NCCL_P2P_DISABLE
- NCCL_NET_GDR_LEVEL
entrypoint: ["bash", "/app/scripts/docker/train-entrypoint.sh"] entrypoint: ["bash", "/app/scripts/docker/train-entrypoint.sh"]
ipc: ${TRAIN_IPC_MODE:-host} ipc: ${TRAIN_IPC_MODE:-host}
stop_grace_period: ${TRAIN_STOP_GRACE_PERIOD:-10m} stop_grace_period: ${TRAIN_STOP_GRACE_PERIOD:-10m}
+6 -3
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
``` ```
@@ -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` |
### 贡献 ### 贡献
+46 -13
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
@@ -816,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
@@ -845,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
@@ -888,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
@@ -926,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 {
@@ -1316,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
@@ -1419,7 +1450,7 @@ classDiagram
Task --> TaskStatus Task --> TaskStatus
InferenceEngine --> AutoModel InferenceEngine --> AutoModel
Executor --> AutoModel Executor --> AutoModel
Executor --> AutoTokenizer Executor --> TaskCacheManager
TaskManager --> AutoTokenizer TaskManager --> AutoTokenizer
``` ```
@@ -1436,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, 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 |
@@ -1456,11 +1488,12 @@ classDiagram
| **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`, `CudaBackend`, `FlashAttnBackend`, `TorchNativeBackend` | 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
@@ -1468,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 (cuda > flash > torch priority; `ASTR_BACKEND` env var overrides default; `TorchNativeBackend` fallback). Rotary embedding auto-dispatches to CUDA kernel when available, else torch complex multiply. 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`
@@ -1476,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
+338 -43
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,124 @@ 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
- **`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`). Default on GPU.
- **`FlashAttnBackend`**: Optional flash-attn dispatch with `flash_attn_with_kvcache` fast path. - **`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) - **`TorchNativeBackend`**: SDPA with indirect KV cache gather (always-available fallback)
Default priority: cuda > flash > torch. Set ``ASTR_BACKEND=cuda|torch_native|flash`` Default priority: cuda > flash > torch. Set ``ASTR_BACKEND=cuda|torch_native|flash``
@@ -102,11 +320,20 @@ with attn_backend(ATTN_BACKEND.CUDA):
engine.generate("hello") engine.generate("hello")
``` ```
`CudaBackend` falls back to `FlashAttnBackend` (when flash-attn is installed and supports the input dtype) or `TorchNativeBackend` otherwise. 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
@@ -115,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):
``` ```
@@ -127,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:
@@ -140,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
@@ -162,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
+111
View File
@@ -0,0 +1,111 @@
# Containerized Serving Deployment
AstrAI uses one serving YAML as the declaration for both host-side container
runtime settings and in-container server settings. `scripts/serve.sh` wraps the
Compose commands so preflight validation and container lifecycle stay
consistent with the trainer.
## Architecture
```text
serve.yaml
├── runtime parsed on the host before Docker starts
└── server parsed by server.py inside the container
scripts/serve.sh preflight, Compose wrapper, lifecycle
└── docker-compose.yml GPU passthrough, mounts, image, port mapping
└── server.py --config /run/astrai/serve.yaml
```
`scripts/docker/serve_runtime.py` reads `runtime:` plus the two container-side
values Compose needs (`server.port` for the port mapping, `server.device` for
the preflight GPU check). `scripts/tools/server.py --config` reads `server:`.
Explicit CLI arguments to `server.py` override `server:` YAML values.
## Runtime Schema
```yaml
runtime:
job_name: serve
port: 8000
paths:
param: ./params
gpu:
enabled: true # false → cpu profile (server-cpu service)
devices: all # all | [0]
container:
cuda_tag: cu128
# environment:
# TOKENIZERS_PARALLELISM: "false"
server:
host: 0.0.0.0
port: 8000
device: cuda # cuda | cpu
dtype: bfloat16 # bfloat16 | float16 | float32
max_batch_size: 16
max_seq_len: null # falls back to model config
```
- Relative paths resolve from the YAML file's directory, not the current shell.
- `runtime.port` is the host publish port; `server.port` is the port the
container listens on. The Compose mapping is
`${SERVE_PORT}:${SERVE_CONTAINER_PORT}`.
- `runtime.gpu.enabled: true` (default) selects the `server` service with an
NVIDIA device reservation; `false` selects `server-cpu` (no GPU passthrough).
When disabled, `server.device` must be `cpu`.
- `runtime.gpu.devices` is `all` (default) or a single-device list such as `[0]`;
the list becomes `CUDA_VISIBLE_DEVICES`. Compose passes `count: all`; the
env var performs the only filtering.
- `environment` values are explicitly passed to the serving container. Keep
host-specific settings here; they are not universal defaults.
- `server.device` must agree with `runtime.gpu.enabled`; `preflight` enforces it.
## Fixed Container Paths
| Runtime path | Container path | Access |
|---|---|---|
| `runtime.paths.param` | `/app/params` | read-only |
| the selected YAML | `/run/astrai/serve.yaml` | read-only |
`server.param_path` is optional: the server default is
`project_root/params`, which is exactly `/app/params` inside the container
(the working directory is `/app`). Set it explicitly only when serving from a
different location; in Docker it must be a container path.
## Operations
The config argument defaults to `./serve.yaml`:
```bash
bash scripts/serve.sh init [CONFIG]
bash scripts/serve.sh preflight [CONFIG]
bash scripts/serve.sh up [CONFIG]
bash scripts/serve.sh run [CONFIG]
bash scripts/serve.sh down [CONFIG]
bash scripts/serve.sh restart [CONFIG]
bash scripts/serve.sh logs [CONFIG]
bash scripts/serve.sh status [CONFIG]
```
`preflight` validates Docker, the model directory
(`config.json` + `model.safetensors`), GPU/device consistency, and the
rendered Compose configuration. `up` starts the container detached; `run`
keeps it in the foreground. Both reuse the existing image; run
`bash scripts/serve.sh build [CONFIG]` after code changes. The wrapper manages a fixed container name
(`astrai-server` or `astrai-server-<job_name>`); the plain
`docker compose up -d` / `docker compose --profile cpu up -d` path keeps
working with defaults (port 8000, `./params`).
## Hard Rules
1. Keep Docker settings in `runtime` and server settings in `server`.
2. Filter GPUs once: Compose passes `count: all`; a `devices`
list becomes `CUDA_VISIBLE_DEVICES`.
3. `runtime.gpu.enabled: false` requires `server.device: cpu`.
4. In Docker, `server.port` must match the published container port (default
`8000`); change `runtime.port` to publish on a different host port.
5. The image user is built with the host UID/GID so the mounted model
directory stays readable.
> Document Update Time: 2026-08-22
+185 -38
View File
@@ -1,57 +1,204 @@
# Containerized Training Deployment # Containerized Training Deployment
Rules for running AstrAI distributed training in containers, distilled from real deployment failures. Read before touching `Dockerfile`, `docker-compose.yml`, `scripts/train.sh`, `train-entrypoint.sh`. AGENTS.md mirrors this locally; this file is the committed version. AstrAI uses one training YAML as the declaration for both host-side container
runtime settings and in-container training settings. Do not invoke the trainer
with raw `docker compose up`; use `scripts/train.sh` so preflight validation,
checkpoint recovery, and graceful shutdown remain active.
## Architecture ## Architecture
``` ```text
scripts/train.sh host-side CLI: env loading, preflight, compose wrapper, lifecycle train.yaml
── docker-compose.yml GPU passthrough, mounts, in-container env vars, entrypoint ── runtime parsed on the host before Docker starts
└── train-entrypoint.sh GPU-count resolution, parallel-mode selection, auto-resume └── model/data/... parsed by train.py inside the container
scripts/train.sh preflight, Compose wrapper, lifecycle, timer
└── docker-compose.yml GPU passthrough, mounts, image, container limits
└── scripts/docker/train-entrypoint.sh process count, parallel mode, auto-resume
└── train.py --config /run/astrai/train.yaml └── train.py --config /run/astrai/train.yaml
``` ```
| Layer | Responsible for | NOT responsible for | The two parsers deliberately own different sections. `scripts/docker/train_runtime.py`
|-------|-----------------|---------------------| reads only `runtime`; `scripts/tools/train.py` reads only
| `train.sh` | host paths, `.env.train`, preflight, lifecycle | training args, GPU selection, parallel mode | `model/data/parallel/training/ckpt/log`. Explicit trainer arguments after `--`
| compose | GPU passthrough, mounts, in-container env (NCCL) | training args (beyond `TRAIN_*` forwarding) | override training YAML values.
| entrypoint | `--ckpt_dir/--nprocs/--parallel_mode/--param_path`, resume | hyperparameters (YAML/CLI) |
| `train.yaml` | hyperparameters (`_merge_yaml_into_kwargs`, CLI wins) | container paths, process count |
## Path Conventions ## Runtime Schema
| Host var | Container | Perm | Purpose | ```yaml
|---|---|---|---| runtime:
| `TRAIN_DATA_DIR` | `/data` | ro | dataset (`data_root_path` must be `/data`) | job_name: astrai-train
| `TRAIN_MODEL_DIR` | `/models/base` | ro | base model (`config.json` + `model.safetensors`) | paths:
| `TRAIN_CHECKPOINT_DIR` | `/checkpoints` | rw | checkpoint root, per-`TRAIN_JOB_NAME` subdirs | data: ./data
| `TRAIN_CONFIG_FILE` | `/run/astrai/train.yaml` | ro | training YAML (mounted only on `start`) | model: ./params
| code | `/app` | image | **not a mount**; rebuild image for code changes | checkpoints: ./checkpoints
gpu:
devices: all
parallel_mode: auto # one GPU: none; multiple GPUs: ddp
container:
cuda_tag: cu128
ipc: host
stop_grace_period: 10m
stop_timeout_seconds: 600
checkpoint_keep_last: 5
# max_duration_hours: 12
# Optional; entries are passed verbatim into the trainer container
# (see "Per-Job Environment"):
# environment:
# ASTR_LOG_LEVEL: DEBUG
# ASTR_BACKEND: torch_native
```
## Hard Rules - Relative paths resolve from the YAML file's directory, not the current shell.
- `devices` is either `all` or a non-empty physical GPU index list. Compose
passes all GPUs once; `CUDA_VISIBLE_DEVICES` performs the only filtering.
- The process count is derived from `devices`. With `all`, the entrypoint uses
`torch.cuda.device_count()` after Docker starts.
- `parallel_mode: auto` selects `none` for one GPU and `ddp` for multiple GPUs.
Use `fsdp` explicitly when model sharding is required.
- To select specific physical GPUs, replace `all` with a list such as
`devices: [0, 1]`.
- `environment` entries apply only to the job defined by this YAML file, not to
the host or to other jobs. Keep the section omitted unless this job's GPU
selection needs it; see [Per-Job Environment](#per-job-environment).
- `max_duration_hours` starts a detached host timer that calls the same graceful
`stop` command. A manual stop cancels the timer.
1. **Filter GPUs once**: compose passes the full physical set (`count: all`); `CUDA_VISIBLE_DEVICES` filters inside by physical index. Never `count: N` + physical indices (double filter leaves 1 card → `device_id out of range`). ## Per-Job Environment
2. **In-container UID = host UID**: Dockerfile builds the user via `USER_UID/USER_GID` args; `train.sh` injects `ASTRAI_UID/GID` (bash `UID` is readonly). compose `user:` alone does not create the /etc/passwd entry — torch's `getpass.getuser()` then dies with `uid not found`.
3. **In-container env vars are explicit**: `.env.train` (`--env-file`) is only compose's interpolation dictionary — never reaches the container. A var arrives only via a value-less `environment` entry (`- VAR`, read from the calling process env). `runtime.environment` is scoped to one job. `start` passes only the entries of
4. **NCCL hang workaround** (this host): `NCCL_P2P_DISABLE=1` + `NCCL_NET_GDR_LEVEL=0` must be in-container. the config file it was given, so a variable reaches exactly the GPUs declared
5. **Checkpoint complete =** `meta.json + config.json + model.safetensors + optimizer.pt + scheduler.pt`; `start` auto-resumes the latest complete one. in that file's `runtime.gpu.devices` and nothing else. Two jobs on the same
6. **tqdm is silent without a TTY**: add `disable=False` in `astrai/trainer/train_callback.py`; `metric.jsonl` (per step) works as progress evidence regardless. machine can therefore differ: a job whose GPUs have working peer-to-peer keeps
the section omitted, a job whose GPUs cross broken PCIe/NVLink paths declares
the NCCL workarounds, and a job on an NVSwitch fabric can pin the NVLink fast
path on.
Because of that scoping, the effective pattern is one YAML per GPU group
rather than one shared YAML that gets edited whenever the device list changes:
```yaml
# train-local.yaml: GPUs with working peer-to-peer; nothing to declare
runtime:
gpu:
devices: [0, 1]
# train-cross-pcie.yaml: this GPU set crosses broken paths, so only this job
# declares the workarounds (confirm first; see docs/guides/distributed.md)
runtime:
gpu:
devices: [4, 5, 6, 7]
environment:
NCCL_P2P_DISABLE: "1"
NCCL_NET_GDR_LEVEL: "0"
```
The same mechanism carries positive tuning, not just workarounds. On an
NVSwitch node (Hopper-class GPUs with fabric manager running), NVLink SHARP
multicast (NVLS) is the fast allreduce path and NCCL enables it automatically
where supported. A job may pin it on explicitly and raise channel parallelism
when benchmarks show the NVLink bandwidth is underused:
```yaml
# train-nvlink.yaml: NVSwitch node; keep the disables OUT and pin the fast
# path on instead (verify support with NCCL_DEBUG=INFO first)
runtime:
gpu:
devices: [0, 1, 2, 3]
environment:
NCCL_NVLS_ENABLE: "1"
NCCL_MIN_NCHANNELS: "8"
# NCCL_ALGO: NVLS # force one algorithm; unsupported values fail loudly
```
NVLS requires NVSwitch multicast support; on plain NVLink bridges or PCIe-only
sets, keep the section omitted and let NCCL pick Ring/Tree with P2P. Newer
drivers list the actual interconnect and NVLS support directly in
`nvidia-smi topo -m`, so check that before assuming.
Confirm a variable is needed before adding it, and only in the YAML of the job
that hits the problem:
```bash
nvidia-smi topo -m # check P2P support between exactly the selected GPUs
NCCL_DEBUG=INFO # confirm NCCL transport errors before disabling them
```
See `docs/guides/distributed.md` for what each troubleshooting variable
disables. The two directions are mutually exclusive: `NCCL_P2P_DISABLE` and
`NCCL_NET_GDR_LEVEL` remove bandwidth and must never appear in the same
environment as the NVLink entries above.
Semantics:
- Values must be scalars and are rendered with `str()`, so quote them
explicitly (`"1"`, `"0"`) instead of relying on YAML booleans or numbers.
- A `null` value exports the name with an empty value.
- This section is the only path for extra host variables into the trainer
container; variables exported in the host shell do not pass through Compose.
## Fixed Container Paths
| Runtime path | Container path | Access |
|---|---|---|
| `runtime.paths.data` | `/data` | read-only |
| `runtime.paths.model` | `/models/base` | read-only |
| `runtime.paths.checkpoints` | `/checkpoints` | read-write |
| the selected YAML | `/run/astrai/train.yaml` | read-only |
Training configuration must therefore use `data_root_path: /data`. The source
code is baked into `/app`; `start` reuses the existing image, so run
`bash scripts/train.sh build [CONFIG]` after code changes.
## Operations ## Operations
The config argument defaults to `./train.yaml`:
```bash ```bash
bash scripts/train.sh init # first run: dirs + .env.train (edit per machine) bash scripts/train.sh init [CONFIG]
bash scripts/train.sh preflight # validate Docker/paths/GPU/model/YAML/compose bash scripts/train.sh preflight [CONFIG]
bash scripts/train.sh start # build + start in background (auto-resume) bash scripts/train.sh start [CONFIG]
bash scripts/train.sh start --foreground -- --dry-run # print plan only bash scripts/train.sh start [CONFIG] --foreground -- --dry-run
bash scripts/train.sh logs | status | stop | restart bash scripts/train.sh logs [CONFIG]
bash scripts/train.sh clean --keep 5 # prune old checkpoints (--force to delete) bash scripts/train.sh status [CONFIG]
bash scripts/train.sh stop [CONFIG]
bash scripts/train.sh restart [CONFIG]
bash scripts/train.sh clean [CONFIG] --keep 5
bash scripts/train.sh clean [CONFIG] --keep 5 --force
``` ```
## Files `init` creates the declared runtime directories but does not generate or mutate
the YAML. `preflight` validates Docker, paths, base model files, checkpoint
writability, GPU configuration, and rendered Compose configuration.
- `docker-compose.yml` — trainer service: `count: all`, `ASTRAI_UID/GID` build args + `user:`, env whitelist, mounts ## Checkpoint Recovery
- `Dockerfile` — production stage builds user from `USER_UID/USER_GID`; `ENV HOME=/home/astrai`; `USER astrai`
- `scripts/train.sh``load_env` filters `UID=` lines (readonly var); `compose()` injects `ASTRAI_UID/GID` Checkpoints are stored below
- `scripts/docker/train-entrypoint.sh` — GPU-count resolution, parallel mode, resume `runtime.paths.checkpoints/<job_name>/epoch_<N>_step_<N>`. A checkpoint is
- `.env.train`, `train.yaml` — host-specific; templates from `scripts/train.sh init`; scientific-notation floats (`2e-5`) parse correctly since train.py uses the YAML 1.2 float schema complete only when it contains:
```text
meta.json
config.json
model.safetensors
optimizer.pt
scheduler.pt
```
`start` resumes the latest complete checkpoint and ignores partial writes. If no
complete checkpoint exists, `/models/base/config.json` and
`/models/base/model.safetensors` are required. `stop` sends `SIGTERM`; the
trainer finishes at a batch boundary and saves an emergency checkpoint before
the Docker timeout expires.
## Hard Rules
1. Keep Docker settings in `runtime` and trainer settings in the remaining YAML sections.
2. Filter GPUs once: Compose passes `count: all`; `devices` becomes `CUDA_VISIBLE_DEVICES`.
3. Do not force DDP for a model that requires FSDP; declare the mode explicitly.
4. Do not use `kill -9` for routine shutdown; use `scripts/train.sh stop CONFIG`.
5. The image user is built with the host UID/GID so mounted checkpoints retain usable ownership.
6. Scope `runtime.environment` to the job YAML that needs it; do not copy NCCL
workarounds into every config.
> Document Update Time: 2026-08-29
+15 -6
View File
@@ -176,14 +176,21 @@ Three-layer separation (SGLang-inspired):
### Attention Backend ### Attention Backend
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`): The extension package separates mechanism from policy:
- **`CudaBackend`** (default): decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool). Falls back to `FlashAttnBackend` when dtype unsupported. - `astrai/extension/ops/` contains stateless wrappers that invoke one exact compiled kernel and fail when it is unavailable.
- **`FlashAttnBackend`**: optional flash-attn dispatch with `flash_attn_with_kvcache` fast path for contiguous cache; falls back to KV gather + `flash_attn_func`. - `astrai/extension/backend/` owns capability checks, implementation selection, fallback, and KV cache I/O.
- Model and inference code use the stable `astrai.extension` API instead of selecting ops directly.
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/backend/attention.py`):
- **`CudaBackend`** (default when supported): decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool).
- **`FlashAttnBackend`**: optional flash-attn dispatch; inference paths gather flat K/V from the pool via `req_to_token` and call `flash_attn_varlen_func` over the ragged batch (fp16/bf16 only); dense mask-free training calls use `flash_attn_func`.
- **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`. - **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
- Default priority: cuda > flash > torch. Set `ASTR_BACKEND=cuda|torch_native|flash` to override. - The `attention(...)` entry point uses cuda > flash > torch priority and chooses another compatible backend when an automatically selected backend cannot handle a call.
- Resolution precedence is: explicit `attn_backend(...)` context > `ASTR_BACKEND` env > default. An explicit `attn_backend(...)` selection is strict (incompatible calls raise); `ASTR_BACKEND` is a default-level override that falls back to a 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 embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch. Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/backend/rotary.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.
Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`: Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`:
@@ -196,6 +203,8 @@ with attn_backend(ATTN_BACKEND.CUDA):
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)`.
Direct imports from `astrai.extension.ops` are reserved for low-level kernel tests and code that intentionally requires a specific compiled implementation. They do not provide fallback.
## Mask Algorithm Internals ## Mask Algorithm Internals
### Template mode (`template: true`) ### Template mode (`template: true`)
@@ -253,4 +262,4 @@ total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset. This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
> Document Update Time: 2026-08-02 > Document Update Time: 2026-08-16
+24 -6
View File
@@ -26,17 +26,24 @@ This guide walks you through installing AstrAI, downloading a model, running inf
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
# Basic install (pure PyTorch, no custom CUDA kernels) # Kernels auto-build when nvcc + CUDA are detected; skip with CSRC_KERNELS=false
pip install -e . pip install -e .
# With CUDA kernels (optional, for fused attention and rotary embedding) # Force the CUDA kernel build (fused attention, rotary embedding, FP8 GEMM)
# CSRC_KERNELS=true pip install -e . --no-build-isolation # CSRC_KERNELS=true pip install -e . --no-build-isolation
# With dev dependencies (pytest, ruff) # With dev dependencies (pytest, ruff)
# pip install -e ".[dev]" # pip install -e ".[dev]"
``` ```
> **CUDA kernels** are opt-in at build time (`CSRC_KERNELS=true`). Once built, `CudaBackend` is the default attention backend on GPU (cuda > flash > torch priority). Override via `ASTR_BACKEND` env var or `attn_backend()` context manager. Fused rotary embedding kernel is auto-dispatched when available. Skip for CPU-only usage. > **CUDA kernels** build automatically when `nvcc` is on `PATH` and
> `torch.cuda.is_available()` returns `True`; set `CSRC_KERNELS=false` to skip
> them, or `CSRC_KERNELS=true` to force them (required when building in an
> isolated environment with `--no-build-isolation`). Once built, `CudaBackend`
> is the default attention backend on GPU (cuda > flash > torch priority).
> Override via `ASTR_BACKEND` env var or `attn_backend()` context manager.
> Fused rotary embedding kernel is auto-dispatched when available. Skip for
> CPU-only usage.
## 2. Download Model Weights ## 2. Download Model Weights
@@ -58,6 +65,14 @@ The model directory contains:
- `model.safetensors` — model weights - `model.safetensors` — model weights
- `tokenizer.json` + `tokenizer_config.json` — tokenizer files (including chat template) - `tokenizer.json` + `tokenizer_config.json` — tokenizer files (including chat template)
External HuggingFace checkpoints of the LLaMA layout (e.g. `meta-llama/...`,
`mistralai/...`, `Qwen/Qwen2-...`) can be loaded directly: `AutoModel.from_pretrained`
auto-detects HF `model_type` / key names (`input_layernorm`, `gate_proj`, MoE
`experts.<j>` ...) and converts config and weights in place. Dense and MoE
(Mixtral / DeepSeek-V3 layout) FFNs are supported; MLA attention
(DeepSeek-V2/V3) and biased projections (`attention_bias`) are not. Pass
`weights_format="astrai"` to skip conversion, or `"hf"` to force it.
## 3. Run Inference ## 3. Run Inference
### Interactive Chat (Simplest) ### Interactive Chat (Simplest)
@@ -175,8 +190,9 @@ python scripts/tools/train.py \
```bash ```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3 export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1 # Only if this host's NCCL transport is broken; see docs/guides/distributed.md:
export NCCL_NET_GDR_LEVEL=0 # export NCCL_P2P_DISABLE=1
# export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \ python scripts/tools/train.py \
--train_type=seq \ --train_type=seq \
@@ -250,5 +266,7 @@ docker compose up -d
| Multi-GPU DDP / FSDP | [Distributed Guide](guides/distributed.md) | | Multi-GPU DDP / FSDP | [Distributed Guide](guides/distributed.md) |
| System architecture | [Architecture](developer/architecture.md) | | System architecture | [Architecture](developer/architecture.md) |
| Data pipeline internals | [Data Flow](developer/dataflow.md) | | Data pipeline internals | [Data Flow](developer/dataflow.md) |
| YAML-driven containerized serving | [Docker Serving](developer/docker-serving.md) |
| YAML-driven containerized training | [Docker Training](developer/docker-training.md) |
> Document Update Time: 2026-07-31 > Document Update Time: 2026-08-22
+63 -10
View File
@@ -54,13 +54,26 @@ KVCache
├── out_cache_loc [batch, seq_len] — write indices for this forward ├── out_cache_loc [batch, seq_len] — write indices for this forward
├── max_len int — max(seq_lens), avoids GPU sync in decode ├── max_len int — max(seq_lens), avoids GPU sync in decode
├── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step ├── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
── qo_indptr [batch + 1] int32 — prefix sum of per-request q_lens (prefill), precomputed once per step ── qo_indptr [batch + 1] int32 — prefix sum of per-request q_lens (prefill), precomputed once per step
├── q_tile_to_batch [num_q_tiles] int32 — prefill: Q tile → request (precomputed once per step)
├── q_tile_to_index [num_q_tiles] int32 — prefill: Q tile → request-local tile index
├── decode_o_part [batch, n_heads, head_dim] — decode split-K partial output buffer
├── decode_ml_part [batch, n_heads] — decode split-K partial max/logsum buffer
└── decode_out [batch, n_heads, head_dim] — decode output accumulator
``` ```
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather. Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
## Attention Backend ## Attention Backend
Inference code calls the policy API exported by `astrai.extension`. The
extension implementation is split into two layers:
- `astrai.extension.backend` owns capability checks, backend selection,
fallback, and KV cache I/O.
- `astrai.extension.ops` contains direct wrappers around compiled CUDA kernels;
these wrappers raise if a kernel is unavailable and do not fall back.
Attention computation (cache I/O + SDPA/kernel dispatch) is decoupled from the model via `AttentionBackend` ABC: Attention computation (cache I/O + SDPA/kernel dispatch) is decoupled from the model via `AttentionBackend` ABC:
``` ```
@@ -70,8 +83,11 @@ AttentionBackend (ABC)
└── TorchNativeBackend SDPA + indirect KV cache gather (always-available fallback) └── TorchNativeBackend SDPA + indirect KV cache gather (always-available fallback)
``` ```
Default priority: cuda > flash > torch. Set ``ASTR_BACKEND=cuda|torch_native|flash`` Default priority is cuda > flash > torch. Automatic selection may choose a
to override. compatible fallback for a particular call. Set
`ASTR_BACKEND=cuda|torch_native|flash` to override the default process-wide;
an explicit `attn_backend(...)` context still takes precedence over the env
override.
Select via context manager (mirrors `torch.nn.attention.sdpa_kernel`): Select via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
@@ -82,15 +98,23 @@ with attn_backend(ATTN_BACKEND.CUDA):
engine.generate("hello") engine.generate("hello")
``` ```
`CudaBackend` decode path: writes K/V to cache, then calls `attn_paged_decode` with `page_size=1` — the `req_to_token` table serves directly as the page table, each token slot is a single-token "page". No explicit K/V gather needed. Environment and context selections are strict: if the selected backend cannot
handle the call, inference raises an error rather than silently switching.
`CudaBackend` decode path: writes K/V via `new_k`/`new_v` while calling `attn_paged_decode` — the `req_to_token` table serves directly as the page table (conceptually a single-token "page" per slot, i.e. `page_size=1`; the op itself takes no `page_size` argument). No explicit K/V gather needed.
`CudaBackend` prefill path: writes K/V, then calls `attn_paged_prefill` — a ragged-batch (paged) prefill kernel that reads K/V directly from the flat pool via `req_to_token`, addressing each request's `q_len`/`kv_len` through `qo_indptr` and `kv_indptr`. No explicit K/V gather needed. `CudaBackend` prefill path: writes K/V, then calls `attn_paged_prefill` — a ragged-batch (paged) prefill kernel that reads K/V directly from the flat pool via `req_to_token`, addressing each request's `q_len`/`kv_len` through `qo_indptr` and `kv_indptr`. No explicit K/V gather needed.
Fallback: when `CudaBackend` cannot handle an input (wrong dtype or head_dim), `FlashAttnBackend` is tried next (if installed), then `TorchNativeBackend`. Fallback: when `CudaBackend` cannot handle an input (wrong dtype or head_dim), `FlashAttnBackend` is tried next (if installed), then `TorchNativeBackend`.
This fallback is performed by the public `attention(...)` policy entry point
only when no backend was explicitly selected. Import from
`astrai.extension.ops` only for direct kernel tests or when failure on a missing
kernel is the intended behavior.
### Rotary Embedding Backend ### Rotary Embedding Backend
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches: Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/backend/rotary.py`, which auto-dispatches:
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, the input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode) - **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, the input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
- **Torch fallback**: complex multiply path (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available - **Torch fallback**: complex multiply path (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
@@ -154,6 +178,31 @@ InferenceEngine
`GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`. `GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`.
## Launching the Server
`scripts/tools/server.py` accepts every option as a CLI flag or from a YAML
config file (`--config serve.yaml`); explicit CLI flags override YAML values.
The YAML `server:` section mirrors the flags:
```yaml
server:
host: 0.0.0.0
port: 8000
device: cuda
dtype: bfloat16
max_batch_size: 16
max_seq_len: null
```
```bash
python scripts/tools/server.py --config serve.yaml
python scripts/tools/server.py --config serve.yaml --port 9000 # CLI wins
```
In Docker, `scripts/serve.sh` drives the same YAML (a `runtime:` section
controls ports/GPU/mounts); see
[Docker Serving](../developer/docker-serving.md).
## HTTP API ## HTTP API
``` ```
@@ -211,14 +260,18 @@ The HTTP protocols and direct engine API have distinct request models and defaul
| `max_tokens` | Optional[int] | 2048 | Max generation length | | `max_tokens` | Optional[int] | 2048 | Max generation length |
| `stream` | Optional[bool] | False | Stream output | | `stream` | Optional[bool] | False | Stream output |
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences | | `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
| `n` | Optional[int] | 1 | Number of choices requested | | `n` | Optional[int] | 1 | Accepted for API compatibility, **ignored** (always returns a single choice) |
| `presence_penalty` | Optional[float] | 0.0 | Presence penalty (-2.0 to 2.0) | | `presence_penalty` | Optional[float] | 0.0 | Accepted for API compatibility, **ignored** |
| `frequency_penalty` | Optional[float] | 0.0 | Frequency penalty (-2.0 to 2.0) | | `frequency_penalty` | Optional[float] | 0.0 | Frequency penalty (-2.0 to 2.0) |
| `logit_bias` | Optional[Dict[int, float]] | None | Per-token logit bias | | `logit_bias` | Optional[Dict[int, float]] | None | Accepted for API compatibility, **ignored** |
| `user` | Optional[str] | None | End-user identifier | | `user` | Optional[str] | None | Accepted for API compatibility, **ignored** |
| `tools` | Optional[List[ToolDef]] | None | Tool definitions for function calling | | `tools` | Optional[List[ToolDef]] | None | Tool definitions for function calling |
| `tool_choice` | Optional[Union[str, Dict[str, Any]]] | `"auto"` | Tool selection mode or explicit tool choice | | `tool_choice` | Optional[Union[str, Dict[str, Any]]] | `"auto"` | Tool selection mode or explicit tool choice |
> `n`, `presence_penalty`, `logit_bias`, and `user` are validated by the request
> model but ignored by the server (a warning is logged when a non-default value
> is supplied).
**Anthropic** (`MessagesRequest`): **Anthropic** (`MessagesRequest`):
| Param | Type | Default | Description | | Param | Type | Default | Description |
@@ -329,4 +382,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
print(token) print(token)
``` ```
> Document Update Time: 2026-07-31 > Document Update Time: 2026-08-22
+19 -2
View File
@@ -28,7 +28,7 @@
|-----------|-------------|---------| |-----------|-------------|---------|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 | | `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 | | `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm for clipping; the current CLI requires a positive number | 1.0 | | `--max_grad_norm` | Maximum gradient norm for clipping; `TrainConfig` validates it as positive (or `None`) | 1.0 |
### Optimizer ### Optimizer
@@ -203,6 +203,7 @@ nohup python scripts/tools/train.py \
| Parameter | Type | Default | Description | | Parameter | Type | Default | Description |
|-----------|------|---------|-------------| |-----------|------|---------|-------------|
| `--config`, `-c` | path | `None` | Serving YAML config. CLI flags override YAML values |
| `--host` | str | `0.0.0.0` | Host address | | `--host` | str | `0.0.0.0` | Host address |
| `--port` | int | `8000` | Port number | | `--port` | int | `8000` | Port number |
| `--param_path` | path | `project_root/params` | Path to model parameters | | `--param_path` | path | `project_root/params` | Path to model parameters |
@@ -217,6 +218,22 @@ Usage:
python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloat16 python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloat16
``` ```
YAML config (a `server:` section; explicit CLI flags override YAML values):
```bash
python scripts/tools/server.py --config serve.yaml
```
```yaml
server:
host: 0.0.0.0
port: 8000
device: cuda
dtype: bfloat16
max_batch_size: 16
max_seq_len: null
```
`serve.yaml` also carries a `runtime:` section for the Docker wrapper; see
[Docker Serving](../developer/docker-serving.md).
See [Inference Guide](inference.md) for HTTP API documentation. See [Inference Guide](inference.md) for HTTP API documentation.
## Generate (`generate.py`) ## Generate (`generate.py`)
@@ -264,4 +281,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
--- ---
> Document Update Time: 2026-07-20 > Document Update Time: 2026-08-22
+4 -1
View File
@@ -31,12 +31,15 @@ classifiers = [
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" } urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
[project.optional-dependencies] [project.optional-dependencies]
dev = ["pytest==9.0.2", "ruff", "httpx2"] dev = ["pytest==9.0.2", "ruff", "httpx"]
flash = ["flash-attn>=2.6"] flash = ["flash-attn>=2.6"]
[tool.setuptools.packages.find] [tool.setuptools.packages.find]
where = ["."] where = ["."]
[tool.setuptools.package-data]
"astrai.extension.lib" = ["*.so"]
[tool.setuptools.dynamic] [tool.setuptools.dynamic]
version = { attr = "astrai.__version__" } version = { attr = "astrai.__version__" }
+155
View File
@@ -0,0 +1,155 @@
"""Parse the host-side runtime section of a serving configuration.
The Compose wrapper needs a few container-side values on the host as well:
``server.port`` (the port the container listens on) and ``server.device``
(used by the preflight GPU consistency check). Everything else under
``server:`` is owned by ``scripts/tools/server.py --config`` inside the
container.
"""
import argparse
import re
import shlex
from pathlib import Path
import yaml
ENV_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
def _mapping(value, name: str) -> dict:
if value is None:
return {}
if not isinstance(value, dict):
raise ValueError(f"{name} must be a mapping")
return value
def _path(value, name: str, config_dir: Path) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"runtime.paths.{name} is required")
path = Path(value).expanduser()
if not path.is_absolute():
path = config_dir / path
return str(path.resolve())
def _port(value, name: str) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(f"{name} must be an integer")
if not 1 <= value <= 65535:
raise ValueError(f"{name} must be between 1 and 65535")
return value
def load_runtime(config_path: str) -> dict[str, str]:
path = Path(config_path).resolve()
with path.open(encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise ValueError("serving configuration must be a mapping")
runtime = _mapping(config.get("runtime"), "runtime")
if not runtime:
raise ValueError("top-level runtime section is required")
paths = _mapping(runtime.get("paths"), "paths")
gpu = _mapping(runtime.get("gpu"), "gpu")
container = _mapping(runtime.get("container"), "container")
environment = _mapping(runtime.get("environment"), "environment")
server = _mapping(config.get("server"), "server")
job_name = runtime.get("job_name", "")
if job_name and not isinstance(job_name, str):
raise ValueError("runtime.job_name must be a string")
if job_name and not re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9._-]*", job_name):
raise ValueError(
"runtime.job_name must use letters, numbers, dot, underscore, or dash"
)
port = _port(runtime.get("port", 8000), "runtime.port")
container_port = _port(server.get("port", 8000), "server.port")
device = server.get("device", "cuda")
if not isinstance(device, str) or not device.strip():
raise ValueError("server.device must be a string")
gpu_enabled = gpu.get("enabled", True)
if not isinstance(gpu_enabled, bool):
raise ValueError("runtime.gpu.enabled must be a boolean")
devices = gpu.get("devices", "all")
visible_devices = None
if gpu_enabled:
if devices == "all":
pass
elif isinstance(devices, list) and len(devices) == 1:
text = str(devices[0])
if not text.isdigit():
raise ValueError(
"runtime.gpu.devices entries must be non-negative integers"
)
visible_devices = text
else:
raise ValueError(
"runtime.gpu.devices must be 'all' or a single-device list such as [0]"
)
else:
if devices != "all":
raise ValueError(
"runtime.gpu.devices is ignored when runtime.gpu.enabled is false"
)
if device != "cpu":
raise ValueError(
"server.device must be 'cpu' when runtime.gpu.enabled is false"
)
values = {
"SERVE_JOB_NAME": job_name,
"SERVE_PORT": str(port),
"SERVE_CONTAINER_PORT": str(container_port),
"SERVE_PARAM_DIR": _path(paths.get("param", "./params"), "param", path.parent),
"SERVE_GPU_ENABLED": "true" if gpu_enabled else "false",
"SERVE_DEVICE": device,
"CUDA_TAG": str(container.get("cuda_tag", "cu128")),
}
if visible_devices is not None:
values["CUDA_VISIBLE_DEVICES"] = visible_devices
for name, value in environment.items():
if not isinstance(name, str) or not ENV_NAME.fullmatch(name):
raise ValueError(f"invalid runtime.environment name: {name!r}")
if value is not None and not isinstance(value, (str, int, float, bool)):
raise ValueError(f"runtime.environment.{name} must be a scalar")
values["environment"] = environment
return values
def shell_exports(runtime: dict[str, str]) -> str:
return "\n".join(
f"export {name}={shlex.quote(value)}"
for name, value in runtime.items()
if name != "environment"
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("command", choices=("exports", "environment"))
parser.add_argument("config")
args = parser.parse_args()
try:
runtime = load_runtime(args.config)
except (OSError, ValueError, yaml.YAMLError) as exc:
parser.error(str(exc))
if args.command == "exports":
print(shell_exports(runtime))
return
for name, value in runtime["environment"].items():
rendered = "" if value is None else str(value)
print(f"{name}={rendered}", end="\0")
if __name__ == "__main__":
main()
+17 -6
View File
@@ -10,12 +10,27 @@ CHECKPOINT_DIR="${CHECKPOINT_ROOT}/${TRAIN_JOB_NAME}"
BASE_MODEL="${BASE_MODEL:-/models/base}" BASE_MODEL="${BASE_MODEL:-/models/base}"
TRAIN_CONFIG="${TRAIN_CONFIG:-}" TRAIN_CONFIG="${TRAIN_CONFIG:-}"
TRAIN_GPU_COUNT="${TRAIN_GPU_COUNT:-all}" TRAIN_GPU_COUNT="${TRAIN_GPU_COUNT:-all}"
TRAIN_PARALLEL_MODE="${TRAIN_PARALLEL_MODE:-auto}"
validate_job_name "${TRAIN_JOB_NAME}" validate_job_name "${TRAIN_JOB_NAME}"
if [[ "${TRAIN_GPU_COUNT}" == "all" ]]; then if [[ "${TRAIN_GPU_COUNT}" == "all" ]]; then
TRAIN_GPU_COUNT="$(python -c 'import torch; print(torch.cuda.device_count())')" TRAIN_GPU_COUNT="$(python -c 'import torch; print(torch.cuda.device_count())')"
fi fi
[[ "${TRAIN_GPU_COUNT}" =~ ^[1-9][0-9]*$ ]] || die "No visible GPU found" [[ "${TRAIN_GPU_COUNT}" =~ ^[1-9][0-9]*$ ]] || die "No visible GPU found"
if [[ "${TRAIN_PARALLEL_MODE}" == "auto" ]]; then
if (( TRAIN_GPU_COUNT > 1 )); then
TRAIN_PARALLEL_MODE=ddp
else
TRAIN_PARALLEL_MODE=none
fi
fi
[[ "${TRAIN_PARALLEL_MODE}" =~ ^(none|ddp|fsdp)$ ]] || die "Invalid parallel mode: ${TRAIN_PARALLEL_MODE}"
if [[ "${TRAIN_PARALLEL_MODE}" == "none" ]] && (( TRAIN_GPU_COUNT != 1 )); then
die "Parallel mode none requires exactly one GPU"
fi
if [[ "${TRAIN_PARALLEL_MODE}" != "none" ]] && (( TRAIN_GPU_COUNT < 2 )); then
die "Parallel mode ${TRAIN_PARALLEL_MODE} requires at least two GPUs"
fi
if [[ -n "${TRAIN_CONFIG}" ]]; then if [[ -n "${TRAIN_CONFIG}" ]]; then
[[ -f "${TRAIN_CONFIG}" ]] || die "Training config not found: ${TRAIN_CONFIG}" [[ -f "${TRAIN_CONFIG}" ]] || die "Training config not found: ${TRAIN_CONFIG}"
fi fi
@@ -36,11 +51,7 @@ if [[ -n "${TRAIN_CONFIG}" ]]; then
train_args+=(--config "${TRAIN_CONFIG}") train_args+=(--config "${TRAIN_CONFIG}")
fi fi
if (( TRAIN_GPU_COUNT > 1 )); then train_args+=(--parallel_mode "${TRAIN_PARALLEL_MODE}")
train_args+=(--parallel_mode ddp)
else
train_args+=(--parallel_mode none)
fi
if [[ -n "${latest_checkpoint}" ]]; then if [[ -n "${latest_checkpoint}" ]]; then
log_info "Resuming ${TRAIN_JOB_NAME} from ${latest_checkpoint}" log_info "Resuming ${TRAIN_JOB_NAME} from ${latest_checkpoint}"
@@ -52,7 +63,7 @@ else
train_args+=(--param_path "${BASE_MODEL}") train_args+=(--param_path "${BASE_MODEL}")
fi fi
log_info "GPUs=${TRAIN_GPU_COUNT}, checkpoints=${CHECKPOINT_DIR}" log_info "GPUs=${TRAIN_GPU_COUNT}, parallel=${TRAIN_PARALLEL_MODE}, checkpoints=${CHECKPOINT_DIR}"
# Replace the shell so the container init forwards SIGTERM to the trainer. # Replace the shell so the container init forwards SIGTERM to the trainer.
exec "${train_args[@]}" "$@" exec "${train_args[@]}" "$@"
+153
View File
@@ -0,0 +1,153 @@
"""Parse the host-side runtime section of a training configuration."""
import argparse
import math
import re
import shlex
from pathlib import Path
import yaml
ENV_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
PARALLEL_MODES = {"auto", "none", "ddp", "fsdp"}
def _mapping(value, name: str) -> dict:
if value is None:
return {}
if not isinstance(value, dict):
raise ValueError(f"runtime.{name} must be a mapping")
return value
def _path(value, name: str, config_dir: Path) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"runtime.paths.{name} is required")
path = Path(value).expanduser()
if not path.is_absolute():
path = config_dir / path
return str(path.resolve())
def load_runtime(config_path: str) -> dict[str, str]:
path = Path(config_path).resolve()
with path.open(encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise ValueError("training configuration must be a mapping")
runtime = _mapping(config.get("runtime"), "runtime")
if not runtime:
raise ValueError("top-level runtime section is required")
paths = _mapping(runtime.get("paths"), "paths")
gpu = _mapping(runtime.get("gpu"), "gpu")
container = _mapping(runtime.get("container"), "container")
environment = _mapping(runtime.get("environment"), "environment")
job_name = runtime.get("job_name")
if not isinstance(job_name, str) or not re.fullmatch(
r"[A-Za-z0-9][A-Za-z0-9._-]*", job_name
):
raise ValueError(
"runtime.job_name must use letters, numbers, dot, underscore, or dash"
)
devices = gpu.get("devices", "all")
visible_devices = None
if devices == "all":
gpu_count = "all"
elif isinstance(devices, list) and devices:
normalized = []
for device in devices:
text = str(device)
if not text.isdigit():
raise ValueError(
"runtime.gpu.devices entries must be non-negative integers"
)
normalized.append(text)
if len(set(normalized)) != len(normalized):
raise ValueError("runtime.gpu.devices must not contain duplicates")
gpu_count = str(len(normalized))
visible_devices = ",".join(normalized)
else:
raise ValueError("runtime.gpu.devices must be 'all' or a non-empty list")
parallel_mode = str(gpu.get("parallel_mode", "auto"))
if parallel_mode not in PARALLEL_MODES:
raise ValueError("runtime.gpu.parallel_mode must be auto, none, ddp, or fsdp")
if gpu_count != "all":
count = int(gpu_count)
if parallel_mode == "none" and count != 1:
raise ValueError("parallel_mode none requires exactly one GPU")
if parallel_mode in {"ddp", "fsdp"} and count < 2:
raise ValueError(
f"parallel_mode {parallel_mode} requires at least two GPUs"
)
max_hours = container.get("max_duration_hours", 0)
try:
max_seconds = math.ceil(float(max_hours) * 3600) if max_hours else 0
except (TypeError, ValueError) as exc:
raise ValueError(
"runtime.container.max_duration_hours must be a number"
) from exc
if max_seconds < 0:
raise ValueError("runtime.container.max_duration_hours must not be negative")
values = {
"TRAIN_JOB_NAME": job_name,
"TRAIN_DATA_DIR": _path(paths.get("data"), "data", path.parent),
"TRAIN_MODEL_DIR": _path(paths.get("model"), "model", path.parent),
"TRAIN_CHECKPOINT_DIR": _path(
paths.get("checkpoints"), "checkpoints", path.parent
),
"TRAIN_GPU_COUNT": gpu_count,
"TRAIN_PARALLEL_MODE": parallel_mode,
"CUDA_TAG": str(container.get("cuda_tag", "cu128")),
"TRAIN_IPC_MODE": str(container.get("ipc", "host")),
"TRAIN_STOP_GRACE_PERIOD": str(container.get("stop_grace_period", "10m")),
"TRAIN_STOP_TIMEOUT": str(container.get("stop_timeout_seconds", 600)),
"CHECKPOINT_KEEP_LAST": str(container.get("checkpoint_keep_last", 5)),
"TRAIN_MAX_DURATION_SECONDS": str(max_seconds),
}
if visible_devices is not None:
values["CUDA_VISIBLE_DEVICES"] = visible_devices
for name, value in environment.items():
if not isinstance(name, str) or not ENV_NAME.fullmatch(name):
raise ValueError(f"invalid runtime.environment name: {name!r}")
if value is not None and not isinstance(value, (str, int, float, bool)):
raise ValueError(f"runtime.environment.{name} must be a scalar")
values["environment"] = environment
return values
def shell_exports(runtime: dict[str, str]) -> str:
return "\n".join(
f"export {name}={shlex.quote(value)}"
for name, value in runtime.items()
if name != "environment"
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("command", choices=("exports", "environment"))
parser.add_argument("config")
args = parser.parse_args()
try:
runtime = load_runtime(args.config)
except (OSError, ValueError, yaml.YAMLError) as exc:
parser.error(str(exc))
if args.command == "exports":
print(shell_exports(runtime))
return
for name, value in runtime["environment"].items():
rendered = "" if value is None else str(value)
print(f"{name}={rendered}", end="\0")
if __name__ == "__main__":
main()
+194
View File
@@ -0,0 +1,194 @@
#!/usr/bin/env bash
set -euo pipefail
ROOT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)"
source "${ROOT_DIR}/scripts/docker/lib/train-common.sh"
COMPOSE_BASE=(
docker compose
--project-directory "${ROOT_DIR}"
--file "${ROOT_DIR}/docker-compose.yml"
)
usage() {
cat <<'EOF'
Usage: scripts/serve.sh <command> [CONFIG] [options]
CONFIG defaults to ./serve.yaml. The same file declares host runtime settings
under `runtime:` and server settings under `server:`.
Commands:
init [CONFIG] Create the model directory
preflight [CONFIG] Validate Docker, paths, GPU, and Compose
build [CONFIG] Build the serving image
up [CONFIG] Start the server container (detached)
run [CONFIG] Start the server container (foreground)
down [CONFIG] Stop and remove the server container
restart [CONFIG] Down, then up
logs [CONFIG] Follow server logs
status [CONFIG] Show container status
EOF
}
resolve_path() {
if [[ "$1" = /* ]]; then
printf '%s\n' "$1"
else
printf '%s/%s\n' "${ROOT_DIR}" "${1#./}"
fi
}
load_config() {
CONFIG_FILE="$(resolve_path "$1")"
[[ -f "${CONFIG_FILE}" ]] || die "Serving config not found: ${CONFIG_FILE}"
require_command python3
python3 -c 'import yaml' >/dev/null 2>&1 ||
die "PyYAML is required on the host (install python3-yaml)"
local exports
exports="$(python3 "${ROOT_DIR}/scripts/docker/serve_runtime.py" exports "${CONFIG_FILE}")" ||
die "Failed to load runtime configuration"
eval "${exports}"
if [[ -n "${SERVE_JOB_NAME}" ]]; then
validate_job_name "${SERVE_JOB_NAME}"
fi
}
compose() {
if [[ -n "${CUDA_VISIBLE_DEVICES:-}" ]]; then
ASTRAI_UID="$(id -u)" ASTRAI_GID="$(id -g)" "${COMPOSE_BASE[@]}" "$@"
else
ASTRAI_UID="$(id -u)" ASTRAI_GID="$(id -g)" \
env -u CUDA_VISIBLE_DEVICES "${COMPOSE_BASE[@]}" "$@"
fi
}
container_name() {
if [[ -n "${SERVE_JOB_NAME}" ]]; then
printf 'astrai-server-%s\n' "${SERVE_JOB_NAME}"
else
printf 'astrai-server\n'
fi
}
service_name() {
if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then
printf 'server-cpu\n'
else
printf 'server\n'
fi
}
set_profile_args() {
PROFILE_ARGS=()
if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then
PROFILE_ARGS=(--profile cpu)
fi
}
init_environment() {
mkdir -p "${SERVE_PARAM_DIR}"
log_info "Model: ${SERVE_PARAM_DIR}"
}
preflight() {
require_command docker
docker info >/dev/null 2>&1 || die "Docker daemon is unavailable"
[[ -d "${SERVE_PARAM_DIR}" ]] || die "Model directory not found: ${SERVE_PARAM_DIR}"
[[ -s "${SERVE_PARAM_DIR}/config.json" ]] ||
die "Model config not found: ${SERVE_PARAM_DIR}/config.json"
[[ -s "${SERVE_PARAM_DIR}/model.safetensors" ]] ||
die "Model weights not found: ${SERVE_PARAM_DIR}/model.safetensors"
if [[ "${SERVE_GPU_ENABLED}" == "false" ]] && [[ "${SERVE_DEVICE}" != "cpu" ]]; then
die "runtime.gpu.enabled is false but server.device is '${SERVE_DEVICE}'; use server.device: cpu"
fi
compose config --quiet
log_info "Preflight passed (service: $(service_name), device: ${SERVE_DEVICE})"
}
runtime_environment_args() {
RUNTIME_ENV_ARGS=()
local pair
while IFS= read -r -d '' pair; do
RUNTIME_ENV_ARGS+=(--env "${pair}")
done < <(python3 "${ROOT_DIR}/scripts/docker/serve_runtime.py" environment "${CONFIG_FILE}")
}
start_server() {
local foreground="$1"
shift
local container running
local -a run_options
preflight
runtime_environment_args
set_profile_args
container="$(container_name)"
running="$(docker inspect --format '{{.State.Running}}' "${container}" 2>/dev/null || true)"
[[ "${running}" != "true" ]] || die "Server is already running: ${container}"
docker rm "${container}" >/dev/null 2>&1 || true
run_options=(
--volume "${CONFIG_FILE}:/run/astrai/serve.yaml:ro"
"${RUNTIME_ENV_ARGS[@]}"
)
if [[ "${foreground}" == "true" ]]; then
compose "${PROFILE_ARGS[@]}" run --rm --service-ports \
"${run_options[@]}" "$(service_name)" \
python -m scripts.tools.server --config /run/astrai/serve.yaml "$@"
else
compose "${PROFILE_ARGS[@]}" run -d --service-ports \
--name "${container}" "${run_options[@]}" "$(service_name)" \
python -m scripts.tools.server --config /run/astrai/serve.yaml "$@"
log_info "Server started; run scripts/serve.sh logs ${CONFIG_FILE} to follow it"
fi
}
stop_server() {
local container
container="$(container_name)"
docker stop --timeout 30 "${container}" >/dev/null 2>&1 ||
log_warn "Server container is not running"
docker rm "${container}" >/dev/null 2>&1 || true
}
show_status() {
docker ps -a --filter "name=^/$(container_name)$"
}
main() {
local command="${1:-}" config="${SERVE_CONFIG_FILE:-${ROOT_DIR}/serve.yaml}"
[[ -n "${command}" ]] || { usage; exit 1; }
shift || true
if [[ "${command}" =~ ^(help|-h|--help)$ ]]; then
usage
return
fi
if [[ $# -gt 0 && "$1" != --* ]]; then
config="$1"
shift
fi
load_config "${config}"
case "${command}" in
init) init_environment ;;
preflight) preflight ;;
build)
set_profile_args
preflight
compose "${PROFILE_ARGS[@]}" build "$(service_name)"
;;
up) start_server false "$@" ;;
run) start_server true "$@" ;;
down) stop_server ;;
restart) stop_server; start_server false ;;
logs) docker logs -f --tail "${SERVE_LOG_TAIL:-200}" "$(container_name)" ;;
status) show_status ;;
*) die "Unknown command: ${command}" ;;
esac
}
main "$@"
+25 -41
View File
@@ -13,6 +13,7 @@ from astrai.inference.engine import InferenceEngine
from astrai.inference.runtime.graph import CudaGraphContext from astrai.inference.runtime.graph import CudaGraphContext
from astrai.inference.workspace import InferenceWorkspace from astrai.inference.workspace import InferenceWorkspace
from astrai.model import AutoModel, AutoRegressiveLM from astrai.model import AutoModel, AutoRegressiveLM
from astrai.serialization import adapt_config
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
_DTYPES = ["bfloat16", "float16", "float32"] _DTYPES = ["bfloat16", "float16", "float32"]
@@ -118,16 +119,11 @@ class GenerationBenchmark:
workspace: InferenceWorkspace, workspace: InferenceWorkspace,
) -> list: ) -> list:
input_ids = torch.randint( input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, prompt_len), device=self.device 0, self.config.vocab_size, (batch_size * prompt_len,), device=self.device
)
position_ids = (
torch.arange(0, prompt_len, dtype=torch.long, device=self.device)
.unsqueeze(0)
.expand(batch_size, -1)
)
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_len, device=self.device
) )
position_ids = torch.arange(
prompt_len, dtype=torch.long, device=self.device
).repeat(batch_size)
task_ids = [f"bench_{i}" for i in range(batch_size)] task_ids = [f"bench_{i}" for i in range(batch_size)]
for tid in task_ids: for tid in task_ids:
@@ -137,9 +133,9 @@ class GenerationBenchmark:
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="prefill",
) )
torch.cuda.synchronize() torch.cuda.synchronize()
return task_ids return task_ids
@@ -154,24 +150,20 @@ class GenerationBenchmark:
): ):
batch_size = len(task_ids) batch_size = len(task_ids)
input_ids = torch.randint( input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, 1), device=self.device 0, self.config.vocab_size, (batch_size,), device=self.device
) )
position_ids = torch.tensor( position_ids = torch.tensor(
[[seq_len] for _ in range(batch_size)], dtype=torch.long, device=self.device [seq_len] * batch_size, dtype=torch.long, device=self.device
) )
total_len = seq_len + 1
for tid in task_ids: for tid in task_ids:
task_cache.task_extend(tid, seq_len) task_cache.task_extend(tid, seq_len)
input_mask = position_ids[:, :, None] >= torch.arange(
total_len, device=self.device
)
kv_cache = task_cache.bind(task_ids, workspace, self.device) kv_cache = task_cache.bind(task_ids, workspace, self.device)
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="decode",
) )
def run_prefill_benchmark( def run_prefill_benchmark(
@@ -188,25 +180,23 @@ class GenerationBenchmark:
task_cache.task_alloc(tid, list(range(prompt_length))) task_cache.task_alloc(tid, list(range(prompt_length)))
input_ids = torch.randint( input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, prompt_length), device=self.device 0,
) self.config.vocab_size,
position_ids = ( (batch_size * prompt_length,),
torch.arange(0, prompt_length, dtype=torch.long, device=self.device) device=self.device,
.unsqueeze(0)
.expand(batch_size, -1)
)
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_length, device=self.device
) )
position_ids = torch.arange(
prompt_length, dtype=torch.long, device=self.device
).repeat(batch_size)
kv_cache = task_cache.bind(task_ids, workspace, self.device, start_pos=0) kv_cache = task_cache.bind(task_ids, workspace, self.device, start_pos=0)
for _ in range(3): for _ in range(3):
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="prefill",
) )
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -215,9 +205,9 @@ class GenerationBenchmark:
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids, position_ids=position_ids,
fwd="prefill",
) )
torch.cuda.synchronize() torch.cuda.synchronize()
elapsed = time.perf_counter() - t0 elapsed = time.perf_counter() - t0
@@ -311,37 +301,29 @@ class GenerationBenchmark:
) )
b = batch_size b = batch_size
input_ids_buf = torch.zeros(b, 1, dtype=torch.long, device=self.device) input_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device)
position_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device) position_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device)
arange = torch.arange(max_seq_len, device=self.device)
gctx = CudaGraphContext(enabled=True) gctx = CudaGraphContext(enabled=True)
graph_key = (b,) graph_key = (b,)
def _decode_graph_step(seq_len): def _decode_graph_step(seq_len):
input_ids_buf.copy_( input_ids_buf.copy_(
torch.randint(0, self.config.vocab_size, (b, 1), device=self.device) torch.randint(0, self.config.vocab_size, (b,), device=self.device)
) )
position_ids_buf[:] = seq_len position_ids_buf[:] = seq_len
for tid in task_ids: for tid in task_ids:
task_cache.task_extend(tid, seq_len) task_cache.task_extend(tid, seq_len)
kv_cache = task_cache.bind(task_ids, workspace, self.device) kv_cache = task_cache.bind(task_ids, workspace, self.device)
input_mask = torch.ge(
position_ids_buf[:, None],
arange,
out=workspace.input_mask[:b, 0, :max_seq_len],
)
input_mask = input_mask.unsqueeze(1)
with torch.inference_mode(), attn_backend(self.backend): with torch.inference_mode(), attn_backend(self.backend):
return gctx.forward( return gctx.forward(
self.model, self.model,
key=graph_key, key=graph_key,
input_ids=input_ids_buf, input_ids=input_ids_buf,
input_mask=input_mask,
kv_cache=kv_cache, kv_cache=kv_cache,
position_ids=position_ids_buf.unsqueeze(1), position_ids=position_ids_buf,
fwd="decode",
) )
for i in range(5): for i in range(5):
@@ -497,7 +479,9 @@ def benchmark_command(
if ckpt is not None: if ckpt is not None:
click.echo(f"Loading model from {ckpt} ...") click.echo(f"Loading model from {ckpt} ...")
config = ConfigFactory.load( config = ConfigFactory.load(
json.loads((Path(ckpt) / "config.json").read_text(encoding="utf-8-sig")) adapt_config(
json.loads((Path(ckpt) / "config.json").read_text(encoding="utf-8-sig"))
)
) )
model = AutoModel.from_pretrained(ckpt) model = AutoModel.from_pretrained(ckpt)
else: else:
+128 -2
View File
@@ -2,16 +2,105 @@ from pathlib import Path
import click import click
import torch import torch
import yaml
from click.core import ParameterSource
from astrai.inference import run_server from astrai.inference import run_server
_DTYPES = ["bfloat16", "float16", "float32"] _DTYPES = ["bfloat16", "float16", "float32"]
_SERVER_KEYS = (
"host",
"port",
"reload",
"param_path",
"device",
"dtype",
"max_batch_size",
"max_seq_len",
)
def _merge_yaml_into_kwargs(
config_path: str,
passed_kwargs: dict,
explicit_keys: set[str] | None = None,
) -> dict:
"""Merge Click defaults, YAML server values, then explicit CLI values."""
with open(config_path, encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise click.UsageError(f"Serving config must be a mapping: {config_path}")
server = config.get("server") or {}
if not isinstance(server, dict):
raise click.UsageError("top-level server section must be a mapping")
unknown = sorted(set(server) - set(_SERVER_KEYS))
if unknown:
click.echo(
f"Warning: ignoring unknown server config keys: {', '.join(unknown)}",
err=True,
)
merged = dict(passed_kwargs)
merged.update({key: server[key] for key in _SERVER_KEYS if key in server})
if explicit_keys is None:
explicit_keys = set(passed_kwargs)
for key in explicit_keys:
if key in passed_kwargs:
merged[key] = passed_kwargs[key]
return merged
def _as_int(value, name: str) -> int | None:
if value is None:
return None
if isinstance(value, bool):
raise click.UsageError(f"{name} must be an integer")
try:
return int(value)
except (TypeError, ValueError):
raise click.UsageError(f"{name} must be an integer, got {value!r}") from None
def _resolve_server_config(
config_path: str,
passed_kwargs: dict,
explicit_keys: set[str] | None = None,
) -> dict:
"""Merge YAML values, then coerce and validate the resolved settings.
``explicit_keys`` are CLI flags that win over YAML; when None, YAML values
win over Click defaults.
"""
merged = _merge_yaml_into_kwargs(config_path, passed_kwargs, explicit_keys or set())
resolved = dict(merged)
resolved["port"] = _as_int(resolved["port"], "server.port") or 8000
resolved["max_batch_size"] = (
_as_int(resolved["max_batch_size"], "server.max_batch_size") or 16
)
resolved["max_seq_len"] = _as_int(resolved["max_seq_len"], "server.max_seq_len")
resolved["reload"] = bool(resolved["reload"])
if resolved["dtype"] not in _DTYPES:
raise click.UsageError(
f"server.dtype must be one of {', '.join(_DTYPES)}, got {resolved['dtype']!r}"
)
return resolved
@click.command(name="serve", help="Launch inference server (OpenAI-compatible API).") @click.command(name="serve", help="Launch inference server (OpenAI-compatible API).")
@click.option(
"--config",
"-c",
"config_path",
type=click.Path(exists=True, dir_okay=False),
default=None,
help="Serving YAML config. CLI flags override YAML values.",
)
@click.option("--host", default="0.0.0.0", help="Host address.") @click.option("--host", default="0.0.0.0", help="Host address.")
@click.option("--port", type=int, default=8000, help="Port number.") @click.option("--port", type=int, default=8000, help="Port number.")
@click.option("--reload", is_flag=True, help="Enable auto-reload for development.") @click.option(
"--reload", is_flag=True, default=False, help="Enable auto-reload for development."
)
@click.option( @click.option(
"--param_path", "--param_path",
type=click.Path(exists=True), type=click.Path(exists=True),
@@ -37,10 +126,47 @@ _DTYPES = ["bfloat16", "float16", "float32"]
default=None, default=None,
help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.", help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.",
) )
@click.pass_context
def server_command( def server_command(
host, port, reload, param_path, device, dtype, max_batch_size, max_seq_len ctx,
config_path,
host,
port,
reload,
param_path,
device,
dtype,
max_batch_size,
max_seq_len,
): ):
"""Launch inference server (OpenAI-compatible API).""" """Launch inference server (OpenAI-compatible API)."""
if config_path:
passed_kwargs = {
"host": host,
"port": port,
"reload": reload,
"param_path": param_path,
"device": device,
"dtype": dtype,
"max_batch_size": max_batch_size,
"max_seq_len": max_seq_len,
}
explicit_keys = {
key
for key in passed_kwargs
if ctx.get_parameter_source(key) is ParameterSource.COMMANDLINE
}
resolved = _resolve_server_config(config_path, passed_kwargs, explicit_keys)
host = resolved["host"]
port = resolved["port"]
reload = resolved["reload"]
param_path = resolved["param_path"]
device = resolved["device"]
dtype = resolved["dtype"]
max_batch_size = resolved["max_batch_size"]
max_seq_len = resolved["max_seq_len"]
click.echo(f"Config: {config_path}")
dtype_map = { dtype_map = {
"bfloat16": torch.bfloat16, "bfloat16": torch.bfloat16,
"float16": torch.float16, "float16": torch.float16,

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