124 Commits
Author SHA1 Message Date
ViperEkura 6f09b1d2ee docs : clarify radix cache architecture
- document exact page-aligned radix prefix matching
- explain partial-page ownership and materialized KV boundaries
- remove bilingual wording from project overview
2026-08-06 11:50:45 +08:00
ViperEkura b2230fefd8 feat : add radix prefix cache
- replace hash-only lookup with page-granular radix matching
- keep partial pages private and cache only materialized KV prefixes
- integrate completed-request caching and add radix behavior tests
2026-08-06 11:45:52 +08:00
ViperEkura 654e6eb0d1 fix : correct prefill sampling and record alignment
- sample the first token from prefill logits without duplicating the prompt tail
- reject incomplete multi-output records before preprocessing alignment
- cover cached generation and partial DPO records with regression tests
2026-08-05 22:20:29 +08:00
ViperEkura a317a4756b refactor: stateless MoE routing with grouped dispatch
- replace per-expert mask scan with sort+bincount grouped dispatch
- carry router stats in forward output instead of module state
- keep MoE diagnostics working under DDP/FSDP wrappers
- remove unused _load_balancing_loss helper
2026-08-05 18:42:12 +08:00
ViperEkura 9b7e6c205f feat: add moe auxloss and metrics 2026-08-05 18:12:28 +08:00
ViperEkura 602b5ce216 docs : add project capability overview
- summarize the end-to-end model lifecycle
- add matching capability tables in both READMEs
2026-08-05 15:47:42 +08:00
ViperEkura 8152760b5f refactor : use factory for attention backends
- register built-in backends through BaseFactory
- derive benchmark choices from registered backends
- cover string selection and invalid backend names
2026-08-05 15:37:22 +08:00
ViperEkura 8c052c99ee feat: add optional FlashAttention (FA2/FA3) backend
- add FlashAttnBackend (ATTN_BACKEND.FLASH) using flash_attn_func with KV-cache gather + GQA, mirroring TorchNativeBackend
- add flash_attn_available() probe gated on compute capability plus a real-kernel smoke test, cached at first use
- lazy-import flash-attn via importlib so it stays an optional dependency, raising clear errors when unusable
- add 'flash' optional extra (flash-attn>=2.6) and export the new backend
2026-08-05 15:27:26 +08:00
ViperEkura 2667b8116d refactor: unify paged and contiguous attention kernels via KVSource policy
- merge AttentionParams and PagedAttentionParams into one struct
- add attn_kv_source.cuh with ContigKV/PagedKV addressing policies
- template prefill/decode kernels (MMA + scalar) on the KV policy, deleting the four duplicated attn_paged_*.cuh variants
- template dispatcher launchers on KV; single combine kernel
- verify: all correctness tests pass and SASS matches baseline (no perf regression)
2026-08-05 14:06:13 +08:00
ViperEkura 6dffb0305a fix: satisfy ruff format and import lint in setup.py
- Merge nested if for CUDA version mismatch check
- Convert try-except-pass to return None (S110)
- Apply ruff format
2026-08-04 21:32:33 +08:00
ViperEkura 49a9c6b3d2 build: migrate CUDA kernel build to CMake
Replace torch CUDAExtension/ParallelBuildExtension with a CMake-based build. Each kernel compiles as an independent pybind11 module in parallel via cmake --build -j, outputting to astrai/extension/lib.

- Add csrc/CMakeLists.txt (5 kernel targets, torch/pybind11 linking)
- setup.py: _CMakeBuildExt invokes cmake; auto-detect CUDA arch via torch
- Remove csrc/build.py (REGISTRY/build flags now in CMakeLists)
- Fix rel-err eps in attn_test.cu (1e-8 -> 1e-4, bf16 scale)
- Update docs/developer/cuda_kernels.md build section
- .gitignore: allow csrc/CMakeLists.txt
2026-08-04 21:27:22 +08:00
ViperEkura cdf9145ecf docs: align CUDA kernel and RoPE docs with code
- Fix rotary docs to describe cos/sin freqs_cis table, not complex buffer
- Replace attn_prefill with attn_paged_prefill for the CudaBackend path
- Register attn_paged_prefill in kernel overview, layout, and module list
- Add qo_indptr and InferenceWorkspace to architecture class diagram
- Add FrequencyPenaltyStrategy to sampling design patterns
2026-08-03 20:54:40 +08:00
ViperEkura 85f0461b3b docs: update license refs from GPL-3.0 to Apache-2.0 2026-08-03 20:21:36 +08:00
ViperEkura 9f0e9195f7 Update LICENSE 2026-08-03 20:18:27 +08:00
ViperEkura 88751d0b08 refactor: share prefill+decode step between scheduler paths
- Extract _step() as the single prefill-group + task_extend + decode primitive
- _run_generation_loop and run_batch now both call it, so the two cannot drift
- run_batch now records prefix hashes (paged mode) and uses input order for
  decode, matching the loop thread
2026-08-03 13:45:27 +08:00
ViperEkura d0e5d910de perf: reduce remaining per-step allocations
- hoist prefill qo_indptr into the workspace so CudaBackend.fwd_prefill does not rebuild it per layer
- cache has_freq in SamplingBatchInfo to drop the per-step GPU any() sync
- drop pin_memory host staging for input_ids; sync copy suffices for a small batch
2026-08-03 01:10:06 +08:00
ViperEkura a03504a280 perf: preallocate inference decode buffers
- add InferenceWorkspace with fixed-shape per-step buffers (input_ids, decode mask, KV bind metadata) for CUDA-graph capture
- bind_tasks derives seq_lens from the pool's own _task_len tracking, dropping the seq_lens parameter
- update decode metadata in-place (position_ids, seq_lens, kv_indptr) instead of re-allocating per step
- task_extend advances _task_len in contiguous mode so the pool tracks current length
- skip log_softmax when logprobs are not requested
2026-08-03 00:55:26 +08:00
ViperEkura d033b2ef0f perf: cache per-step decode tensor construction
- SamplingBatchInfo: sample params built once per task set (top_k int32, pinned async H2D)
- position_ids advances by +1 on steady-state decode instead of re-building
- DecodeBindCache: bind_tasks increments seq_lens/kv_indptr, reuses req_pool_indices
- saves ~240us of python/launch overhead per decode step
2026-08-02 20:32:53 +08:00
ViperEkura 8447f88f61 fix: size KV pool from prompt/gen args in benchmark
- Drop hardcoded CACHE_MAX_SEQ=2048 which overflowed at long prompts
- Size prefill pool to prompt_length and decode pool to prompt+5+gen*num_trials
- Unblocks decode/prefill benchmark at prompt 4096+ (was KV cache index OOB)
2026-08-02 16:25:27 +08:00
ViperEkura b1b65a657e perf: target 512 grid blocks for decode split-K
- compute_num_splits used 2*sm/base, undersplitting at large batch
- single-warp decode blocks host ~11/SM, not 1/2-SM, so B=16 got 3 splits when 8 was optimal
- Grid search on L20: bandwidth saturates near 256-512 total blocks; target 512
- Pass num_passes into base_blocks for the non-paged decode to match the paged path
- B=16 kv=2048: 0.0230->0.0157ms (-32%); paged B=16: 0.0527->0.0243ms (-54%); B=32: 0.0406->0.0241ms (-41%)
2026-08-02 16:10:40 +08:00
ViperEkura 3439e3104e perf: launch CUDA kernels on torch's current stream
- Thread a cudaStream_t through attn dispatchers onto torch's current stream
- Scope the device guard to the entry function so kernels run on tensor device
- DISPATCH_HEAD_DIM now forwards varargs so stream reaches each dispatch
- Parallelize CPU reference kernels with OpenMP (paged test 31s -> 7s)
- Merge decode/prefill standalone tests into attn_test.cu with correctness tables
- Drop bench error column (CPU ref too slow at large sizes)
- Update cuda_kernels.md for the merged test layout
2026-08-02 13:20:14 +08:00
ViperEkura 288ba20db1 docs: audit non-CUDA documentation
- Aligns CLI and strategy metric contracts
- Refreshes architecture, dataflow, preprocessing, distributed, and eval guides
- Corrects links, TOCs, defaults, and repository paths
2026-08-02 07:39:24 +08:00
ViperEkura 020e2eff4e refactor: emit strategy metrics as floats
- Converts detached strategy metrics before returning loss output
- Removes redundant item conversion from the trainer loop
- Updates the documented contract and regression tests
2026-08-02 06:38:28 +08:00
ViperEkura 1c7369f293 feat: add MoE auxiliary loss metrics
- Propagates MoE load-balancing loss through model outputs
- Logs task, auxiliary, and weighted losses across strategies
- Computes only explicitly requested callback metrics
- Preserves tensor compute_loss API and adds regression tests
2026-08-02 06:30:43 +08:00
ViperEkura 0fc1b1bd46 feat: extend DeepSeek MoE configuration 2026-08-02 05:30:40 +08:00
ViperEkura d7db37a70f fix: preserve MoE routing defaults 2026-08-02 05:30:26 +08:00
Gaolingx 6d98bb4f9f 20260801-moe model impl
need to add aux loss for load balancing
2026-08-01 22:48:58 +08:00
ViperEkura 925cbedc93 feat: scalar paged prefill fallback and decode causal fix
- Add scalar paged prefill kernel mirroring split-Q MMA indexing for sm<80
- Wire scalar path into dispatch_paged_prefill under ASTRAI_NO_MMA
- Fix paged decode scalar causal mask dropping all kv>0 for decode
2026-08-01 16:52:01 +08:00
ViperEkura fda82ee232 perf: drop redundant smem zero-init in paged decode kernel
- Removes per-step STAGES*BC*LD smem clear loop (2 buffers x 24 layers)
- cp.async predicated load + softmax mask already exclude padding slots,
  matching the paged prefill kernel which never zero-inits
- Standalone and extension tests pass; decode step time unchanged
2026-08-01 16:17:34 +08:00
ViperEkura 4b25664c79 perf: precompute kv_indptr once per decode step
- bind_tasks builds kv_indptr (prefix sum of seq_lens) a single time
- fwd_decode/fwd_prefill reuse it instead of rebuilding per layer
- Removes 24 cumsum launches per decode step (was ~1ms/step at B=4)
- Decode B=4: 9.60 -> 7.82 ms/step (-18.5%), +22.8% tok/s
2026-08-01 16:09:26 +08:00
ViperEkura a27c8a819d test: prune low-value and duplicate tests
- Remove tautological test_trainer assertions that never trained
- Drop grpo isfinite-only smokes and merge frozen-model checks via parametrize
- Merge duplicate tool_parser cases (find/streaming/factory) with parametrize
- Collapse duplicate dataset store/detect_format tests
- Remove misleading scheduler/task tests that asserted the opposite of their names
- Merge signal-handler SIGTERM/SIGINT into one parametrized case
- Drop cross-file grpo strategy duplication kept in online_strategy
2026-08-01 16:01:20 +08:00
ViperEkura 91acaf4b0b refactor: unify attention mask to single attn_mask tensor
- CudaBackend.fwd_decode passes attn_mask directly instead of kv_cache.decode_mask
- TorchNativeBackend derives pos_mask from attn_mask[:,0,0] on decode
- Drop decode_mask and page_table fields from KVCache and bind_tasks
2026-08-01 15:49:26 +08:00
ViperEkura 41dcf0feb9 feat: SGLang-style paged attention kernels replace page-table path
- PagedAttentionParams uses flat KV pool + req_to_token + kv_indptr/qo_indptr instead of page_table
- MMA split-KV decode and split-Q prefill kernels with indirect ragged-batch addressing
- Prefill kernel accepts 4D mask (causal-aware); decode kernel supports 2D mask
- CudaBackend is inference-only: kv_cache=None raises, no torch fallback
- benchmark.py: required --ckpt, --backend/--compare options
- Parallel build isolates build-temp/build-lib per subprocess
- Standalone test covers decode/prefill with mask, 27 cases pass
2026-08-01 15:41:25 +08:00
ViperEkura 9960f79920 feat: parallel kernel build via BUILD_PARALLEL env var
- Add ParallelBuildExtension that dispatches each extension to a subprocess
- 4 extensions compile concurrently (3m34s → 1m1s on L20, ~3.5x faster)
- Default 8 workers, override with BUILD_PARALLEL=N
2026-08-01 12:34:48 +08:00
ViperEkura 7feeb0b93e refactor: replace magic layout ints with TensorLayout enum
- Add TensorLayout enum (C++ + Python) to replace magic layout ints
- Add C10_CUDA_CHECK post-launch error checking to all kernel entries
- Add CUDAGuard + freqs_cis shape validation to rotary_emb.cu
- Cache SM count to eliminate per-call cudaDeviceGetAttribute
- Add DISPATCH_CAUSAL_MASK macro to deduplicate dispatcher if/else
- Convert mask type hints from X|None to Optional[X]
2026-08-01 11:05:52 +08:00
ViperEkura 3639b50b4a chore: bump version to 1.3.12 2026-08-01 09:22:16 +08:00
ViperEkura d855c09cf3 fix: use torch.optim.AdamW in ManoAdamW instead of NAdamW
- ManoAdamW now uses torch.optim.AdamW(fused=True, betas=(0.9, 0.95)) matching MuonAdamW, eliminating a confounding variable in optimizer comparison experiments
- only NoraNAdamW retains NAdamW, which is correct per the Nora paper design
2026-08-01 09:20:54 +08:00
ViperEkura d6bfb09863 feat: add grad_snr metric with EMA-based gradient SNR tracking
- add GradSNRTracker to metric_util.py computing SNR = E[g]^2 / Var(g) via per-parameter EMA moments
- add grad_snr_tracker field to TrainContext (instantiated by default)
- register grad_snr in MetricCallback, update tracker on each optimizer step before metrics are recorded
- add grad_snr to default --metrics in train.py CLI
2026-08-01 08:54:44 +08:00
ViperEkura 6db276f37a feat: add Mano manifold optimizer (mano_adamw)
- implement Mano (v2) with axis-rotating tangent projection and manifold normalization, replacing Newton-Schulz iteration
- composite ManoAdamW reuses partition_optimizer_parameters and composite helpers
- register mano_adamw in OptimizerFactory, export Mano and ManoAdamW
- add --mano_momentum and --mano_nesterov CLI options in Optimizer group
- add mano_adamw hyperparameters branch in train.py
- document mano_adamw in params.md
- add tests for single-step projection, axis alternation, factory registration, closure, and resume
2026-08-01 08:51:08 +08:00
ViperEkura 6c76c16480 feat: group train CLI options in --help output
- add GroupedOption/GroupedCommand (no third-party dep) that tags each option with a group label and renders help in labeled sections
- add opt() shorthand wrapping click.option with cls=GroupedOption
- tag all ~55 options into 10 groups aligned with params.md chapters
2026-08-01 08:40:22 +08:00
ViperEkura 11073bd1d2 refactor: extract composite optimizer helpers and unify naming
- add astrai/optim/composite.py with shared step/zero_grad/state_dict/param_groups helpers and OptimizerFactory
- rename MuonMix to MuonAdamW (matches registered name muon_adamw) and file to muon_adamw.py
- use @OptimizerFactory.register decorator in each optimizer module instead of post-import registration in __init__
- fix closure being invoked once per sub-optimizer in MuonAdamW.step (now exactly once via composite_step)
- NoraNAdamW.step now forwards closure correctly
2026-08-01 08:07:45 +08:00
ViperEkura 25c9e81b2b refactor: keep muon_adamw as default optimizer and drop nora docs
- revert CLI/create_optimizer/display defaults to muon_adamw
- revert README, README-zh-CN, params.md to pre-merge state
2026-08-01 07:51:51 +08:00
ViperEkura ffbd9b57c9 Merge branch 'codex/nora-nadamw-default' into experiment
feat: add Nora+NAdamW optimizer with factory-based optimizer selection
2026-08-01 07:49:30 +08:00
QueenAmish 04899a2b15 Make Nora+NAdamW the default optimizer 2026-07-31 23:16:39 +08:00
ViperEkura 530d280e33 perf: remove split partials memset and overlap decode tile loads
- alloc_split_partials now uses torch::empty: the split kernel writes every slot it owns, so the per-call zeros/full memset was pure overhead (2 kernels per layer per step)
- decode split-KV MMA kernels now run a true multi-stage cp.async pipeline (wait_group<STAGES-1> instead of wait_group<0>), keeping STAGES-1 tile loads in flight; the old wait_group<0> serialized load and compute so deeper STAGES made no difference
- add a fallback path when ntiles < STAGES to avoid a race on the last tile
2026-07-31 22:37:44 +08:00
ViperEkura 21ddead238 fix: stabilize paged decode attention kernels
- zero-fill split partials so combine skips unwritten splits deterministically
- skip loading masked KV in paged decode kernels to avoid 0*NaN output poisoning
- zero-fill shared memory tile buffers to prevent stale NaN leaking into softmax
2026-07-31 21:01:12 +08:00
ViperEkura 7aa5ed09d9 refactor: unify rotary embedding interface and update docs
- Merge cos/sin into single freqs_cis tensor [batch, seq, dim/2, 2] throughout the pipeline: RotaryEmbedding buffer, forward return type, apply_rotary_emb signature, CUDA kernel interface
- CUDA kernel now takes freqs_cis directly and reads cos/sin via stride offset internally, eliminating Python-side slice/copy overhead
- Kernel interface: rotary_emb(x, freqs_cis) replaces rotary_emb(x, cos, sin)
- All call sites pass rotary_emb as Tensor (was tuple), type annotations consistent
- Update build threads from 8 to 16
- Fix all docs: get-started, inference, training, cuda_kernels, architecture, internals — reflect new rotary interface, KVCache fields, rotary backend dispatch, .so path, kernel registry count, file layout
2026-07-31 16:52:25 +08:00
ViperEkura 75411ce0cc fix: skip CUDA rotary kernel when grad is enabled
- apply_rotary_emb now checks torch.is_grad_enabled() before dispatching to CUDA kernel
- Training (grad enabled) uses torch complex multiply path which supports autograd backward
- Inference (inference_mode/no_grad) uses CUDA kernel as before
- Without this fix, training backward would crash — the CUDA kernel has no autograd backward()
2026-07-31 15:43:15 +08:00
ViperEkura 9f83d982ec refactor: move compiled kernel .so files into extension/lib
- CUDAExtension module names changed from astrai.extension.<name> to astrai.extension.lib.<name>
- Compiled .so files now land in astrai/extension/lib/ instead of alongside Python source
- loader.py imports from .lib.<name> subpackage
- Add astrai/extension/lib/__init__.py to make lib a proper package
- Separates compiled artifacts from Python source for cleaner directory structure
2026-07-31 15:36:32 +08:00
ViperEkura 3e67b4f88d perf: add fused CUDA rotary embedding kernel
- Single-kernel rotary embedding (cos/sin lookup + rotation) replaces PyTorch complex-multiply path (3 kernel launches + f32 upcast per call)
- RotaryEmbedding now stores cos_table/sin_table and returns (cos, sin) f32 tuple instead of a complex tensor
- apply_rotary_emb in rotary_backend.py auto-dispatches: CUDA kernel if available, else torch complex-multiply fallback; backend-agnostic (both attention backends benefit)
- Kernel: 256-thread blocks, grid-stride loop, vectorized __nv_bfloat162 load/store, f32 compute, bf16 out
- Standalone kernel 6-9x faster than torch across decode/prefill shapes, max diff 0 (decode) to 3e-2 (large prefill, bf16)
- Benchmark (L20, bf16, CUDA backend): B=1 9.48->7.25ms (+31%), B=4 10.73->7.67ms (+40%), B=8 10.77->7.81ms (+38%), B=16 10.79->7.83ms (+38%)
2026-07-31 15:27:31 +08:00
ViperEkura 50cfd0d555 perf: reduce decode overhead in scheduler and executor
- Precompute page_table and decode_mask on KVCache once per step in PagePool.bind_tasks, instead of per-layer in CudaBackend/TorchNativeBackend
- Skip frequency penalty history tensor construction when all penalties are 0 in Executor.execute_decode
- Omit FrequencyPenaltyStrategy from sampling pipeline when penalty is 0
- Deduplicate get_active_tasks calls in scheduler loop (3 to 1), remove redundant sorted() on decode tasks
- Benchmark (L20, bf16, CUDA backend): B=1 9.48->9.40ms (+1%), B=4 10.73->9.89ms (+8.6%), B=8 10.77->10.13ms (+6.4%)
2026-07-31 14:50:16 +08:00
ViperEkura 5756054d38 build: parametrize CUDA version for wheels and docker
- Add cu128/cu130 build matrix to release workflow
- Parametrize Dockerfile and docker-compose with CUDA_TAG build arg
- Allow csrc/ and setup.py in docker context via .dockerignore
- Add nvcc/torch CUDA version mismatch preflight warning in setup.py
- Add cuda_toolkit_version() helper in csrc/build.py
- Use at::IntArrayRef explicitly to fix ATen overload ambiguity
- Guard kernels with CUDART_VERSION >= 11020 check
- Remove invalid [tool.pip] section from pyproject.toml
2026-07-31 14:10:55 +08:00
ViperEkura 738cb8f128 fix: broadcast ref/old model state_dict for FSDP
- Add broadcast_state_dict to sync state_dict from rank-0 to all ranks
- Fix create_ref_model returning None on non-rank-0 under FSDP
- Fix sync_old_model only updating old_model on rank-0 under FSDP
- Split skip_no_cuda/skip_no_kernel markers and hoist to top-level conftest
- Add distributed tests for broadcast_state_dict and create_ref_model
2026-07-31 08:32:22 +08:00
ViperEkura 28d1bd07cf style: unify decode expf to __expf
- attn_decode_split_kv.cuh: 4 expf -> __expf
- attn_paged_decode_split_kv.cuh: 4 expf -> __expf
- --use_fast_math makes expf emit __expf anyway, so no behavior change
- aligns decode with prefill/mma kernels that already use __expf
2026-07-31 00:19:18 +08:00
ViperEkura 02625739fe perf: increase eval batch sizes and add max_seq_len
- humaneval/ifeval: default batch_size 64, add --max_seq_len=4096
- mmlu: batch 4 questions x 4 choices per forward, add --batch_size
- ppl: default batch_size 64
2026-07-30 23:55:37 +08:00
ViperEkura f688cd9c5a fix: update benchmark to use checkpoint loading and CudaBackend 2026-07-30 22:54:45 +08:00
ViperEkura 8055027df7 perf: enable paged MMA kernel for page_size=1
- Replace per-tile page lookup with per-element lookup in load_tile
- Remove page_ok gate and scalar fallback in launch_paged_decode_mma
- Unified path works for any page_size (L1-cached when page_size >= BC)
- HBM BW: 12% → 73%, decode throughput: 2,250 → 2,606 tok/s (B=32)
- Scales to 5,232 tok/s at B=128 (2.54x vs torch native)
2026-07-30 22:06:41 +08:00
ViperEkura 3067a8e1a6 feat: unify attention backend with multi-dim mask support
- Add attention() functional entry delegating to active backend
- GQA/MLA forward calls attention() instead of inline cache/SDPA
- CUDA kernels support 2D/3D/4D mask via mask_h_stride field
- CudaBackend.fwd_decode builds 2D padding mask for mixed seq_lens
- KVCache.max_len precomputed in bind_tasks to avoid GPU sync
- batch==1 decode short-circuits mask=None
- Split tests into conftest, test_backend, test_backend_equivalence, test_kernel_mask
- 440 tests pass, L20 decode 1.44-1.60x speedup vs torch native
2026-07-30 20:38:34 +08:00
ViperEkura 97114b95a4 docs: update for attention backend and extension API
- Remove stale 'not yet wired' references
- Add AttentionBackend/CudaBackend sections to cuda_kernels.md, internals.md, inference.md
- Add astrai.extension to architecture.md module table and design patterns
- Update get-started.md: CUDA kernels activatable via attn_backend()
2026-07-30 18:50:16 +08:00
ViperEkura 32fd03a025 feat: add CudaBackend and rename to fwd_decode/fwd_prefill
- CudaBackend: paged decode via attn_paged_decode, prefill via attn_prefill
- Decode uses req_to_token as page_table with page_size=1
- Falls back to TorchNativeBackend when kernel unavailable
- Rename forward_decode/forward_extend to fwd_decode/fwd_prefill
- Register ATTN_BACKEND.CUDA in _BACKEND_REGISTRY
2026-07-30 18:45:33 +08:00
ViperEkura 21bf37dd83 refactor: unify extension API to blhd layout and is_causal
- Rename ops.py to attention_ops.py
- Remove layout/scale params: fixed blhd, auto scale
- Replace causal_offset with is_causal bool
- Move SDPA fallback to backend, ops only calls CUDA kernels
- Update __init__.py exports
2026-07-30 18:39:20 +08:00
ViperEkura 5b67d5865a feat: add AttentionBackend ABC with context manager
- AttentionBackend ABC with forward_decode/forward_extend dispatch
- TorchNativeBackend: SDPA with indirect KV cache gather
- attn_backend() context manager + ATTN_BACKEND enum (mirrors sdpa_kernel)
- ContextVar-based thread-safe backend switching
- get_backend() falls back to default TorchNativeBackend singleton
2026-07-30 18:20:27 +08:00
ViperEkura df979b4469 refactor: use single-index access and update docs for cache architecture
- Replace all buffer[layer_id][loc] double indexing with buffer[layer_id, loc] single advanced indexing in cache.py and attention.py
- Revert KVStorage buffers back to 4D [n_layers, size, n_kv_heads, head_dim], remove leftover 3D reshape/view in MLA path
- Update docs/guides/inference.md, docs/developer/internals.md, docs/developer/architecture.md to reflect new PagePool/KVStorage/ReqToTokenPool/KVCache classes
2026-07-30 17:47:04 +08:00
ViperEkura deb2d7e127 refactor: rebuild KV cache with three-layer separation architecture
- Replace CacheView/ContiguousCache/PageCache with SGLang-inspired design: KVStorage (flat token-level NHD buffers [n_layers, size, H, D]), ReqToTokenPool (index table [req_idx, pos] -> token_slot), Allocator + PrefixCache (slot allocation with LRU and prefix sharing)
- Add KVCache as pure dataclass passed to model: k_buffer, v_buffer, req_to_token, req_pool_indices, seq_lens, out_cache_loc
- PagePool orchestrates all three layers, supports contiguous mode (pre-allocated per-request blocks, default) and paged mode (page_size=1 or >1 with dynamic allocation and prefix caching)
- Attention layers now do raw buffer indexing instead of opaque write/gather method calls on CacheView objects
- Update executor.bind_tasks signature: seq_lens list + start_pos
- Rename paged_cache -> kv_cache throughout model/ and inference/
2026-07-30 17:19:06 +08:00
ViperEkura fc47319240 refactor: simplify BaseFactory and separate ModelFactory from AutoModel
- Extract _resolve_base_type and _validate_component as module-level helpers
- Replace ForwardRef._evaluate private API with eval in module namespace
- Remove broad except Exception in __init_subclass__, _component_base always set
- Replace direct _entries mutation in strategy.py with register() call form
- Remove dead TOKENIZER_CLASSES registry from AutoTokenizer
- Extract ModelFactory(BaseFactory[nn.Module]) as pure factory
- AutoModel now inherits only nn.Module, no factory state
- Move @AutoModel.register to @ModelFactory.register in transformer.py and encoder.py
2026-07-30 09:38:20 +08:00
ViperEkura 22cf798d81 feat: add field and model validators to config classes
- TrainConfig: enum validators (strategy, parallel_mode, backend, start_method, compile_mode), positive/non-negative/range validators, model_validator requiring reward_model_fn for online RL strategies
- AutoRegressiveLMConfig/EncoderConfig: attn_type, ffn_type enum validators
- ProcessingConfig: packing_strategy, truncation_mode enums, positive int validators
- OutputConfig: storage_format, position_ids_mode enum validators
2026-07-30 08:41:14 +08:00
ViperEkura 164be9708b refactor: migrate config system to Pydantic dataclasses
- Replace hand-rolled BaseConfig (from_dict/to_dict/_coerce/_unwrap_optional) with pydantic.dataclasses
- from_dict now uses cls(**d), to_dict uses dataclasses.asdict + json.dumps filter
- TrainConfig: required fields are now truly required (no default=None), delete manual validate()/__post_init__
- Remove dead required() helper and metadata={'help': ...} annotations
- Fix gradient_checkpointing_modules type from List[str] to List[type]
- Add pydantic>=2.0 as direct dependency in pyproject.toml
- Add numpy-style Parameters docstrings to all config classes
- Enable use_attribute_docstrings in BaseConfig for schema generation
- LoRAConfig also migrated to pydantic dataclass
2026-07-30 08:25:32 +08:00
ViperEkura 6a97524db4 refactor: inline parallel utils into executor module
- Move create_ref_model from astrai/parallel/utils.py into executor.py
- Remove unused ColumnParallelLinear/RowParallelLinear (module.py)
- Update imports in strategy.py and train_context.py
- Drop unused astrai.parallel.utils and astrai.parallel.module
2026-07-30 07:54:54 +08:00
ViperEkura c8b1e40f71 docs: restructure to docs/, add guides and developer docs
- Rename assets/ to docs/, split into guides/ and developer/
- Add get-started.md: installation + 5-step quickstart
- Add guides/evaluation.md: 7 eval scripts with CLI args
- Add guides/distributed.md: DDP/FSDP, gradient accumulation, NCCL
- Add developer/internals.md: loss formulas, RoPE, KV cache math
- Add developer/cuda_kernels.md: build system, benchmarks, file layout
- Fix storage_format doc in preprocessing.md
- Update cross-references in README.md, README-zh-CN.md, Dockerfile
2026-07-30 00:49:04 +08:00
ViperEkura bcaa2d1ae0 fix: FSDP unwrap_model collective op and None guard
- unshard() and full_tensor() are collective ops, all ranks must participate
- Old code returned None on non-rank-0 before calling unshard, causing deadlock
- Fix: all ranks unshard/full_tensor, only rank-0 keeps the result
- Move create_ref_model to parallel/utils.py, accept executor+model directly
- Guard create_ref_model and sync_old_model against None on non-rank-0
2026-07-29 23:41:10 +08:00
ViperEkura 8206afefd9 fix: FSDP clip_grad_norm and default reshard_after_forward=False
- FSDP params are DTensors sharded across ranks
- torch.nn.utils.clip_grad_norm_ computes LOCAL norm only
- Each rank would clip by a different factor, causing gradient divergence
- Fix: compute local norm, all-reduce squared sum, sqrt for global norm
- Default reshard_after_forward=False (forward then backward makes reshard redundant)
- Reduces per-step time by ~19% (1033ms to 839ms on 2xL20)
2026-07-29 23:27:10 +08:00
ViperEkura 646b1b0f46 refactor: replace FSDP with FSDP2 as default parallel backend
- Remove FSDPExecutor (FullyShardedDataParallel wrapper)
- Rename FSDP2Executor to FSDPExecutor, register as 'fsdp'
- Remove 'fsdp2' from CLI choices, make 'fsdp' the default parallel_mode
- Pass after_wrap to executor.prepare for compile-after-wrap ordering
- Update architecture.md, params.md, AGENTS.md references
- FSDP2 uses per-module fully_shard: no FlatParameter, better compile compat
2026-07-29 23:09:37 +08:00
ViperEkura 8150ab6c32 feat: add torch.compile CLI option for training
- Add --compile flag (default/reduce-overhead/max-autotune)
- Apply torch.compile in _before_wrap before DDP/FSDP wrapping
- Profiling shows MFU 85.5% -> 88.5% (+3%), time -3.2%, memory -7.9%
2026-07-29 22:06:51 +08:00
ViperEkura 0b0693a0a2 fix: make ChatTemplate picklable for spawn multiprocessing
- Add __getstate__/__setstate__ to drop cached _compiled Jinja2 template
- Jinja2 Template.root_render_func is a dynamic closure unpicklable by reference
- cached_property rebuilds the template lazily on first render after unpickle
2026-07-29 13:24:13 +08:00
ViperEkura 115192c67c refactor: remove H5 storage backend in favor of mmap bin
- Remove H5Store, H5Writer, save_h5/load_h5 and h5py dependency
- MmapStore (bin) is the sole pre-tokenized storage backend
- Move setup_logging after imports to fix E402 in __init__.py
- Clean up unused imports across test files
- Move inline test imports to file top
2026-07-29 12:50:27 +08:00
ViperEkura c2b04d8458 refactor: align generate.py params with engine API
- Remove --max_tokens, let scheduler use max_seq_len - prompt_len
- Rename --cache_len to --max_seq_len to match engine naming
- Unify sampling defaults to 0.8/50/0.95
2026-07-29 09:47:53 +08:00
ViperEkura db487ab48b feat: append EOS to response in IFD evaluation
- Add EOS token at end of response in both conditional and unconditional passes so model also predicts when response should end
- New --append_eos/--no-append_eos CLI flag (default: enabled) with graceful fallback when tokenizer has no EOS
2026-07-28 22:22:59 +08:00
ViperEkura a95794d3db perf: use Rust-native DecodeStream for O(n) streaming decode
- Replace hand-rolled StreamDecoder (O(n^2) full-history re-decode per token) with tokenizers.decoders.DecodeStream
- Keep O(1) bounded token buffer internally via prefix drain instead of accumulating all token IDs
- Simplify flush_remaining to no-op since stream always emits completed text per step
- Benchmark on 8000 tokens: 2305ms -> 3.9ms (~592x speedup)
2026-07-28 14:32:10 +08:00
ViperEkura 39f84f3b4c refactor: move signal_handler from parallel/ to top-level for broader reuse 2026-07-28 10:36:17 +08:00
ViperEkura 9f7cf50c56 fix: keep metric logs cumulative instead of segmental in each checkpoint 2026-07-28 09:18:48 +08:00
ViperEkura d9a0c72149 feat: store metric logs inside each checkpoint dir, remove log_dir config 2026-07-28 00:22:29 +08:00
ViperEkura 5ab18bec48 fix: correct epoch computation on resume to avoid redoing whole epoch 2026-07-28 00:01:29 +08:00
ViperEkura 2e29ed45d3 perf: shrink decode tile to BC=16 for higher occupancy
- BC=32→16 halves smem (32KB→16KB for D=128), doubling blocks/SM (3→6)
- D=256 now fits STAGES=2 double-buffer in 32KB, eliminating 176-byte spill
- min_tiles_per_split=2 avoids excessive split overhead on small kv
- paged decode: require page_size multiple of BC so tiles stay page-aligned

Benchmark (L20 sm_89, D=128):
- B=1 kv=4096: 0.0134→0.0122ms (+9% BW)
- B=16 kv=2048: 0.0434→0.0352ms (+23% BW)
- B=32 kv=1024: 0.0343→0.0282ms (+22% BW)
2026-07-27 22:44:02 +08:00
ViperEkura 5ba21f4eb3 refactor: eliminate test duplication via shared helpers
- Add tests/helpers.py with shared config, dataset, tokenizer, executor, and assertion helpers
- Replace 15 copies of device one-liner with session-scoped fixture
- Collapse 5 near-identical Dataset subclasses into RandomTokenDataset
- Remove duplicate _make_config/_make_model/_make_frozen and FakeTokenizer/FakeExecutor definitions
- Make test_callbacks and test_early_stopping use existing train_config_factory
- Replace 6 duplicate meta.json read blocks with load_shard_meta
- Fix mkdtemp leaks in test_lora.py with TemporaryDirectory
2026-07-27 22:34:53 +08:00
ViperEkura c26a47b0df docs: sync docs with current code after refactor
- architecture: remove TaskManager.max_prompt_len (merged into max_seq_len in 53c804e)
- dataflow: fix DatasetFactory.load param name max_position_embeddings -> max_len
- params: add fsdp2 to parallel_mode, add --max_seq_len to server, add 4 missing generate options
- preprocessing: add missing batch_size config field
2026-07-27 21:43:29 +08:00
ViperEkura b1a87b22bb feat: add --device flag for GPU-accelerated SVD, default to cuda 2026-07-27 08:53:40 +08:00
ViperEkura 07625057f2 feat : add setup_logging with hierarchical astrai logger
- setup_logging(): attach handler only to astrai logger, not root
- all astrai.* sub-module loggers inherit automatically
- controlled by ASTR_LOG_LEVEL env var (default INFO)
- called in if __name__ == '__main__' of each CLI script
2026-07-27 08:13:48 +08:00
ViperEkura 53c804e233 refactor : merge max_prompt_len into max_seq_len, replace assert with raise
- Engine/Scheduler/TaskManager: merge max_prompt_len into max_seq_len
- train.py: replace bare assert with ValueError/FileNotFoundError
- server.py: add --max_seq_len CLI option
- engine.py: remove dead page_size param
2026-07-27 08:05:11 +08:00
ViperEkura 05c7432964 chore: remove AGENTS.md 2026-07-27 07:21:40 +08:00
ViperEkura 4de42d83c2 refactor: migrate scripts from argparse to click, add YAML config support
- Replace argparse with click in all scripts (train, server, generate,
  preprocess, benchmark)
- Add --config YAML support to train.py with CLI flag override
- Add --dry-run mode to validate config before training
- Add type annotations throughout benchmark.py
- Unify docstring format across all commands
- Remove redundant deps httpx, requests, pyyaml, rich from pyproject.toml
- Net -346 lines while adding YAML config support
2026-07-27 06:55:46 +08:00
ViperEkura b99485f462 chore: bump version to 1.3.11 2026-07-27 01:23:13 +08:00
ViperEkura 20041d7aa9 perf: extend MMA decode to arbitrary GQA ratio, add launch bounds, vectorize combine
- Multi-pass MMA: encode pass in grid blockIdx.x, compute q_head0/G in-kernel
- Fixes crash for G>32 (previously block(32,G) exceeded 1024 threads)
- Fixes alloc_split_partials using uninitialized num_splits (MAX_SPLITS=32)
- __launch_bounds__ on all MMA and prefill kernels for better register allocation
- 4x vectorized combine kernel (4 head_dim per thread)
- uint4 vectorized K loads in scalar decode kernels
- cp.async .L2::128B cache hint for K/V tile streaming
- Extract warp_reduce_sum, bf16, MAX_SPLITS to attn_warp_utils.cuh
2026-07-27 00:35:34 +08:00
ViperEkura 59248032dc chore: fix ruff lint warnings and signal handling edge cases
- Fix pre-existing ruff lint warnings (F401, F541, F841, E741)
- Exclude .md/.json/.yml from ruff format check
- Unblock SIGTERM/SIGINT via pthread_sigmask in early signal handler
- Do not restore SIG_DFL on unregister to prevent pending signal kills
2026-07-25 21:08:30 +08:00
ViperEkura ceadc34ea9 feat: auto-checkpoint on SIGTERM/SIGINT with DDP support
- Register SIGTERM/SIGINT handlers in training loop, set stop flag on signal
- Check stop_requested at each epoch/batch boundary, break and call on_error to save checkpoint
- LocalStrategy parent forwards signal to child processes via terminate(), waits up to 600s for graceful exit
- TrainContext gains threading.Event-based stop_requested/request_stop
- Tests verify SIGTERM/SIGINT trigger checkpoint save with exit code 0, works on both CPU and GPU
2026-07-25 20:40:54 +08:00
ViperEkura 8ab5631446 fix: correct online rollout lifecycle 2026-07-23 19:01:37 +08:00
ViperEkura 99b5d2b2da perf: batch tokenizer preprocessing 2026-07-23 18:42:19 +08:00
ViperEkura 021e6f3788 style: apply ruff formatting to FSDP2 changes 2026-07-23 16:30:10 +08:00
ViperEkura 4e38183e86 fix: make FSDP2 executor work with ABC+Generic model hierarchy
- Wrap each child module individually, skip root (CPython layout
  conflict between ABC+Generic and FSDP2 __class__ assignment)
- Remove manual unshard in clip_grad_norm (DTensor compatible)
- Fix _no_sync to iterate modules() instead of checking root
- Add reshard after unwrap_model
- Guard __init_subclass__ type resolution against dynamic subclasses
- Add fsdp2 to --parallel_mode CLI choices
2026-07-23 16:11:02 +08:00
ViperEkura 4eeb23e2b3 fix: use copy-on-write mmap mode to silence non-writable tensor warning 2026-07-22 17:37:41 +08:00
ViperEkura ef8783b7e3 fix: separate attn_mask and loss_mask in get_logprobs, compose causal masking in strategy
- add loss_mask parameter to get_logprobs to decouple attention from loss masking
- DPO/GRPO strategies compose key-padding + causal mask before model forward
- prevents prompt tokens from being masked out of attention and missing causal masking
2026-07-21 23:47:28 +08:00
ViperEkura 60d7ee614a fix: improve attention kernel numerical stability and test precision checks
- use fmaf() for V-accumulation in scalar decode paths to reduce rounding
- delay scale multiplication to after dot-product in scalar prefill
- unify __expf/expf across MMA and scalar paths for consistent numerics
- harmonize divide-by-zero guards to 1e-20f
- add both absolute and relative error checks in standalone tests (atol=0.01, rtol=0.01)
2026-07-21 23:05:17 +08:00
ViperEkura f7a16efc9d refactor: extract shared dispatcher header, unify MMA/scalar dispatch format
- Merge 3 duplicated dispatch blocks into single attn_dispatchers.cuh
- Merge compute_num_splits from attn_utils.cuh into dispatcher header
- All dim3 grid/block declarations and <<<>>> launches are single-line
- Production .cu files (35-42 loc) only handle torch wrapping + pybind11
- Test files include dispatcher header directly, removing all #ifndef ASTRAI_NO_MMA duplication
2026-07-21 22:21:39 +08:00
ViperEkura a01e8bbe98 refactor: adopt FA2-style KernelTraits + compile-time causal/mask dispatch
- Introduce KernelTraits<HEAD_DIM, BC, WARPS, STAGES> compile-time config bundle, replacing scattered <KD, NC8, KT2, ...> template params
- Template all MMA and scalar kernels on IsCausal/HasMask bools to eliminate inner-loop runtime branches
- Dispatch to 4-path IsCausal/HasMask kernel variants at entry points based on p.causal_offset and p.use_mask
- Update standalone test files with new kernel signatures, add causal test cases
- Fix duplicate using bf16 in MMA kernels that include attn_mma_utils.cuh
2026-07-21 21:52:46 +08:00
ViperEkura ccf728a1b7 perf: eliminate GPU syncs in contiguous cache write/gather hot paths
- Replace .tolist() calls with _total_len in gather(); move _slot_len updates from per-layer write to once-per-step bind_tasks

- Use torch.as_tensor instead of torch.tensor in decode penalty history construction
2026-07-21 16:40:56 +08:00
ViperEkura f1b4b05d08 feat: add attention dimension dispatch 2026-07-21 12:52:45 +08:00
ViperEkura 0c86c89af4 refactor : align config field names with Hugging Face
- dim -> hidden_size, n_layers -> num_hidden_layers
- dim_ffn -> intermediate_size, n_heads -> num_attention_heads
- n_kv_heads -> num_key_value_heads, max_len -> max_position_embeddings
- norm_eps -> rms_norm_eps, tie_weight -> tie_word_embeddings
- update model, inference, training, scripts, tests, docs
2026-07-20 22:05:31 +08:00
ViperEkura d7ac66fb73 refactor: simplify attention mask handling 2026-07-20 20:36:16 +08:00
ViperEkura a6e920fdb0 Merge pull request #20 from ccx1324/lora-device-fix
fix: LoRA device mismatch and checkpoint resume
2026-07-20 19:38:10 +08:00
ViperEkura 958df58f9d refactor: unify tokenizer encode and apply_chat_template for batch support
- encode(str) single-thread, encode(List[str]) Rust parallel encode_batch
- apply_chat_template accepts single Messages or List[Messages] for batch
- add Message/Messages type aliases at module level
2026-07-20 17:49:26 +08:00
ViperEkura e0f102c4d9 feat: support SFT directly from JSONL without dataset_config.json
- JsonlStore falls back to built-in messages config when no config file found and tokenizer_path is provided
- DatasetFactory.load forwards tokenizer_path to store for SFT/SEQ+jsonl
- assistant turns train, other roles masked, position_ids doc_reset
2026-07-20 17:25:09 +08:00
ccx1324andccx 5a942527b2 fix: inject LoRA before loading checkpoint state_dict
move inject_lora() before load_state_dict in _before_wrap so that
  LoRA adapter weights from a checkpoint are properly restored on
  training resume. Previously, inject happened after load, causing
  lora_A/lora_B keys to be silently ignored (strict=False).

  Co-Authored-By: ccx1324 <2424441089@qq.com>
2026-07-20 17:00:24 +08:00
ViperEkura 37a3036934 refactor: split LoRA param init into local vars 2026-07-20 16:23:47 +08:00
ViperEkura 121a7bf8b4 Merge pull request #19 from ccx1324/lora-device-fix
fix: create LoRA parameters on base weight device instead of CPU
2026-07-20 16:14:30 +08:00
ccx1324andClaude Opus 4.7 a5678c9185 fix: create LoRA parameters on base weight device instead of CPU
When `inject_lora()` replaces Linear layers with LoRALinear after the model
has been moved to CUDA, the new lora_A and lora_B parameters were always
created on CPU, causing a device mismatch error during the forward pass.

Now lora_A and lora_B are created on the same device and dtype as the
parent weight, matching the model's current device.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-20 16:11:48 +08:00
ViperEkura 2c50b3cf37 ci: preserve both release wheel artifacts 2026-07-20 15:33:43 +08:00
ViperEkura eee7f54789 docs: sync training and architecture guides 2026-07-20 15:23:30 +08:00
ViperEkura 06eeeead79 refactor: map instruction/input/output to chat roles
- RolloutGenerator._instruction_to_messages builds system/user/assistant list (instruction->system, input->user, output->assistant), replacing single-user-turn concatenation
- Remove _iter_samples helper; _prepare_prompts zips parallel list-of-strings fields directly per the collate_fn contract
- Tests adopt a system-aware chat template and pin the three-field role mapping
- Drop unused imports caught by ruff F401 (torch.Tensor in scheduler.py, iter_raw_records in pipeline.py, Tuple in evaluate_rouge.py)
2026-07-20 13:55:25 +08:00
ViperEkura e8ff7f5321 fix: use batch_per_device for rollout scheduler batch sizing
- train_context.py referenced non-existent cfg.batch_size, replaced with cfg.batch_per_device
- default group_size lowered from 8 to 1: without a group concept (DPO), scheduler batch equals batch_per_device; rollout-based DPO can opt in via extra_kwargs['group_size']>=2
- inline expressions (rollout_batch_size, max_seq_len) extracted for readability
- add tests/trainer/test_online_e2e.py: end-to-end online_dpo via Trainer.train, exercising KV-cache-backed rollout path
2026-07-20 13:32:04 +08:00
ViperEkura a6e1f26cd4 refactor: simplify sample return_logprobs path
- SamplingPipeline.sample gains return_logprobs; both greedy and multinomial paths now share a single log_softmax+gather instead of duplicating the sampling logic
- module-level sample() becomes a thin forwarder instead of re-implementing the three-branch logic
- eliminates ~10 lines of duplicated softmax/gather code; no caller-facing API change
2026-07-20 13:16:18 +08:00
ViperEkura 95c43368ae refactor: unify rollout onto inference engine KV-cache path
- RolloutGenerator now delegates prefill/decode to InferenceScheduler.run_batch (sync API, no background thread), sharing one KV-cache code path with the inference server and eliminating O(n^2) recompute in rollout
- Add sample(return_logprobs=) and Executor.execute_decode(return_logprobs=) to expose behaviour-policy log-probs through the engine; Task gains output_logprobs
- RolloutResult now subclasses RawRollout (adds rewards only), removing duplicated fields
- RolloutRunner.__call__ returns (result, is_fresh) instead of relying on object identity, removing the fragile refresh-detection contract
- Remove O(n^2) generate_responses helper and dead code (_tokenize_prompts, unused old_model arg)
- train_context.py wires InferenceScheduler directly instead of hand-rolling SamplingPipeline
- Tests: +11 covering return_logprobs, run_batch, and KV-cache-backed rollout semantics; 404 pass
2026-07-20 12:52:20 +08:00
ViperEkura 754624acf0 feat: add online rollout framework for RL strategies
- RolloutRunner: generate + score responses with cached re-rollout trigger
- BaseStrategy.__call__ switches online/offline via runner injection
- GRPO/DPO implement prepare_from_rollout; aliases online_grpo/online_dpo
- TrainConfig + train.py add rollout params and CLI flags
- Tests cover generate_responses, RolloutRunner cache, shared __call__
2026-07-20 03:49:56 +08:00
ViperEkura 0b6a17330f feat: add FSDP2Executor using torch.distributed.fsdp.fully_shard API
- New FSDP2Executor registers as 'fsdp2' in ExecutorFactory, using per-module fully_shard() instead of FSDP1 FlatParameter wrapper
- FSDP2 preserves original Parameter objects as DTensors, eliminating use_orig_params=True hack
- FSDP2Executor implements _no_sync via set_requires_gradient_sync, clip_grad_norm via unshard, unwrap_model via DTensor.full_tensor
- Drop **_extra/**_ddp_only_kwargs fallbacks in BaseExecutor/FSDPExecutor/FSDP2Executor, replaced by parallel_mode-aware executor_kwargs dispatch in train.py (ddp-only kwargs only passed when parallel_mode=ddp)
- Export FSDP2Executor in astrai.parallel.__init__
2026-07-20 01:46:25 +08:00
ViperEkura 74b9308883 refactor: pass model_fn/optimizer_fn to executor.prepare
- BaseExecutor.prepare now takes factories and instantiates model via model_fn(), runs before_wrap hook, wraps DDP/FSDP, then builds optimizer/scheduler on the wrapped model
- optimizer/scheduler creation moved into executor.prepare, eliminating the old 'create-then-wrap' hack reliance on use_orig_params=True
- FSDPExecutor/BaseExecutor accept **_extra kwargs to tolerate DDP-only keys (broadcast_buffers, gradient_as_bucket_view) being forwarded via executor_kwargs
- dataloader builds stay external; executor only handles model/optimizer/scheduler
- train_context.py rewritten to load checkpoint state_dict before prepare via a before_wrap closure
2026-07-20 01:32:05 +08:00
ViperEkura e5f9b1a3a9 fix: default max_grad_norm to 1.0 and drop None branch 2026-07-20 01:08:13 +08:00
159 changed files with 15669 additions and 6908 deletions
+3 -1
View File
@@ -4,6 +4,8 @@
# Allow necessary files # Allow necessary files
!astrai/ !astrai/
!scripts/ !scripts/
!assets/ !docs/
!csrc/
!setup.py
!pyproject.toml !pyproject.toml
!README.md !README.md
+38 -9
View File
@@ -23,24 +23,33 @@ jobs:
with: with:
name: pure-wheel name: pure-wheel
path: dist/*.whl path: dist/*.whl
if-no-files-found: error
build-cuda-linux: build-cuda-linux:
name: Build CUDA wheel (Linux) name: Build CUDA wheel (Linux, ${{ matrix.cuda_tag }})
runs-on: ubuntu-latest runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
include:
- cuda_tag: "cu128"
cuda_ver: "12.8.0"
- cuda_tag: "cu130"
cuda_ver: "13.0.0"
steps: steps:
- uses: actions/checkout@v4 - uses: actions/checkout@v4
- uses: actions/setup-python@v5 - uses: actions/setup-python@v5
with: with:
python-version: "3.12" python-version: "3.12"
- name: Install torch (CUDA 12.8) - name: Install torch (${{ matrix.cuda_tag }})
run: | run: |
pip install torch --index-url https://download.pytorch.org/whl/cu128 pip install torch --index-url https://download.pytorch.org/whl/${{ matrix.cuda_tag }}
- name: Setup CUDA - name: Setup CUDA (${{ matrix.cuda_ver }})
uses: Jimver/cuda-toolkit@v0.2.35 uses: Jimver/cuda-toolkit@v0.2.35
with: with:
cuda: "12.8.0" cuda: "${{ matrix.cuda_ver }}"
- name: Build wheel (with CUDA kernels) - name: Build wheel (with CUDA kernels)
run: | run: |
@@ -48,8 +57,9 @@ jobs:
- uses: actions/upload-artifact@v4 - uses: actions/upload-artifact@v4
with: with:
name: cuda-wheel-linux name: cuda-wheel-linux-${{ matrix.cuda_tag }}
path: dist/*.whl path: dist/*.whl
if-no-files-found: error
release: release:
name: Attach wheels to release name: Attach wheels to release
@@ -58,14 +68,33 @@ jobs:
permissions: permissions:
contents: write contents: write
steps: steps:
- uses: actions/download-artifact@v4 - name: Download pure-Python wheel
uses: actions/download-artifact@v4
with: with:
pattern: "*-wheel" name: pure-wheel
path: release-assets/pure
- name: Download CUDA wheels (all variants)
uses: actions/download-artifact@v4
with:
pattern: cuda-wheel-linux-*
merge-multiple: true merge-multiple: true
path: release-assets/cuda
- name: Verify release assets
shell: bash
run: |
set -euo pipefail
pure_wheels=(release-assets/pure/*.whl)
cuda_wheels=(release-assets/cuda/*.whl)
test "${#pure_wheels[@]}" -eq 1
test "${#cuda_wheels[@]}" -ge 1
- name: Create release & upload assets - name: Create release & upload assets
uses: softprops/action-gh-release@v2 uses: softprops/action-gh-release@v2
with: with:
files: ./*.whl files: |
release-assets/pure/*.whl
release-assets/cuda/*.whl
tag_name: ${{ github.ref_name }} tag_name: ${{ github.ref_name }}
generate_release_notes: true generate_release_notes: true
+2 -1
View File
@@ -9,6 +9,7 @@
!scripts/**/*.py !scripts/**/*.py
!tests/**/*.py !tests/**/*.py
!csrc/**/*.py !csrc/**/*.py
!csrc/CMakeLists.txt
!csrc/**/*.cu !csrc/**/*.cu
!csrc/**/*.h !csrc/**/*.h
@@ -24,7 +25,7 @@
!/.dockerignore !/.dockerignore
!/Dockerfile !/Dockerfile
!/docker-compose.yml !/docker-compose.yml
!/assets/** !/docs/**
!/CONTRIBUTING.md !/CONTRIBUTING.md
!/LICENSE !/LICENSE
!/pyproject.toml !/pyproject.toml
+10 -8
View File
@@ -20,9 +20,6 @@ Run the following checks **in order** — CI will reject if any fail.
ruff format . ruff format .
``` ```
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
> Always review the diff after formatting.
### 2. Import sorting ### 2. Import sorting
```bash ```bash
@@ -44,7 +41,7 @@ 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 `%TEMP%`. Clean them manually if needed.
### 4. (Optional) Full pre-commit check ### 4. (Optional) Full pre-commit check script
If you have Git Bash available: If you have Git Bash available:
@@ -52,12 +49,17 @@ If you have Git Bash available:
bash scripts/pre_commit.sh bash scripts/pre_commit.sh
``` ```
This runs format check, import sort check, and tests in one go. The script installs development dependencies by default, then runs the format
check, import sort check, and tests. If dependencies are already installed, use:
```bash
bash scripts/pre_commit.sh --skip-deps
```
## Commit Style ## Commit Style
``` ```
fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description (~50 chars) type: short description (~50 chars)
- bullet point body (each ~60 chars) - bullet point body (each ~60 chars)
``` ```
@@ -73,7 +75,7 @@ fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description
|---------|-------|-----| |---------|-------|-----|
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` | | `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` |
| `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging | | `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging |
| Pre-commit hook rejects | Tests or lint failed | Fix individually, do not `--no-verify` | | Pre-commit check script fails | Dependency install, tests, or lint failed | Fix the failing step; use `--skip-deps` only when dependencies are already installed |
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually | | Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
## Submitting Changes ## Submitting Changes
@@ -93,7 +95,7 @@ fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description
## License ## License
By contributing, you agree that your contributions will be licensed under the [GPL-3.0 License](LICENSE). By contributing, you agree that your contributions will be licensed under the [Apache-2.0 License](LICENSE).
--- ---
+12 -2
View File
@@ -1,8 +1,16 @@
# AstrAI Dockerfile - Multi-stage Build (Optimized) # AstrAI Dockerfile - Multi-stage Build (Optimized)
#
# CUDA version selection:
# docker build -t astrai .
# docker build -t astrai --build-arg CUDA_TAG=cu128 .
# docker build -t astrai --build-arg CUDA_TAG=cu130 .
# Default: cu128
# Build stage - use base image with minimal build tools # Build stage - use base image with minimal build tools
FROM ubuntu:24.04 AS builder FROM ubuntu:24.04 AS builder
ARG CUDA_TAG=cu128
WORKDIR /app WORKDIR /app
# Install Python 3.12 and minimal build dependencies # Install Python 3.12 and minimal build dependencies
@@ -20,10 +28,12 @@ ENV PATH="/opt/venv/bin:$PATH"
# Copy source code and install (deps read from pyproject.toml) # Copy source code and install (deps read from pyproject.toml)
COPY astrai/ ./astrai/ COPY astrai/ ./astrai/
COPY csrc/ ./csrc/
COPY setup.py .
COPY pyproject.toml . COPY pyproject.toml .
RUN pip install --no-cache-dir --upgrade pip \ RUN pip install --no-cache-dir --upgrade pip \
&& pip install --no-cache-dir . \ && pip install --no-cache-dir . \
--extra-index-url https://download.pytorch.org/whl/cu128 --extra-index-url "https://download.pytorch.org/whl/${CUDA_TAG}"
# Production stage # Production stage
FROM ubuntu:24.04 AS production FROM ubuntu:24.04 AS production
@@ -43,7 +53,7 @@ ENV PATH="/opt/venv/bin:$PATH"
# Copy application code # Copy application code
COPY astrai/ ./astrai/ COPY astrai/ ./astrai/
COPY scripts/ ./scripts/ COPY scripts/ ./scripts/
COPY assets/ ./assets/ COPY docs/ ./docs/
COPY pyproject.toml . COPY pyproject.toml .
COPY README.md . COPY README.md .
+193 -666
View File
@@ -1,674 +1,201 @@
GNU GENERAL PUBLIC LICENSE Apache License
Version 3, 29 June 2007 Version 2.0, January 2004
http://www.apache.org/licenses/
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
Everyone is permitted to copy and distribute verbatim copies TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
of this license document, but changing it is not allowed.
1. Definitions.
Preamble
"License" shall mean the terms and conditions for use, reproduction,
The GNU General Public License is a free, copyleft license for and distribution as defined by Sections 1 through 9 of this document.
software and other kinds of works.
"Licensor" shall mean the copyright owner or entity authorized by
The licenses for most software and other practical works are designed the copyright owner that is granting the License.
to take away your freedom to share and change the works. By contrast,
the GNU General Public License is intended to guarantee your freedom to "Legal Entity" shall mean the union of the acting entity and all
share and change all versions of a program--to make sure it remains free other entities that control, are controlled by, or are under common
software for all its users. We, the Free Software Foundation, use the control with that entity. For the purposes of this definition,
GNU General Public License for most of our software; it applies also to "control" means (i) the power, direct or indirect, to cause the
any other work released this way by its authors. You can apply it to direction or management of such entity, whether by contract or
your programs, too. otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
When we speak of free software, we are referring to freedom, not
price. Our General Public Licenses are designed to make sure that you "You" (or "Your") shall mean an individual or Legal Entity
have the freedom to distribute copies of free software (and charge for exercising permissions granted by this License.
them if you wish), that you receive source code or can get it if you
want it, that you can change the software or use pieces of it in new "Source" form shall mean the preferred form for making modifications,
free programs, and that you know you can do these things. including but not limited to software source code, documentation
source, and configuration files.
To protect your rights, we need to prevent others from denying you
these rights or asking you to surrender the rights. Therefore, you have "Object" form shall mean any form resulting from mechanical
certain responsibilities if you distribute copies of the software, or if transformation or translation of a Source form, including but
you modify it: responsibilities to respect the freedom of others. not limited to compiled object code, generated documentation,
and conversions to other media types.
For example, if you distribute copies of such a program, whether
gratis or for a fee, you must pass on to the recipients the same "Work" shall mean the work of authorship, whether in Source or
freedoms that you received. You must make sure that they, too, receive Object form, made available under the License, as indicated by a
or can get the source code. And you must show them these terms so they copyright notice that is included in or attached to the work
know their rights. (an example is provided in the Appendix below).
Developers that use the GNU GPL protect your rights with two steps: "Derivative Works" shall mean any work, whether in Source or Object
(1) assert copyright on the software, and (2) offer you this License form, that is based on (or derived from) the Work and for which the
giving you legal permission to copy, distribute and/or modify it. editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
For the developers' and authors' protection, the GPL clearly explains of this License, Derivative Works shall not include works that remain
that there is no warranty for this free software. For both users' and separable from, or merely link (or bind by name) to the interfaces of,
authors' sake, the GPL requires that modified versions be marked as the Work and Derivative Works thereof.
changed, so that their problems will not be attributed erroneously to
authors of previous versions. "Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
Some devices are designed to deny users access to install or run to that Work or Derivative Works thereof, that is intentionally
modified versions of the software inside them, although the manufacturer submitted to Licensor for inclusion in the Work by the copyright owner
can do so. This is fundamentally incompatible with the aim of or by an individual or Legal Entity authorized to submit on behalf of
protecting users' freedom to change the software. The systematic the copyright owner. For the purposes of this definition, "submitted"
pattern of such abuse occurs in the area of products for individuals to means any form of electronic, verbal, or written communication sent
use, which is precisely where it is most unacceptable. Therefore, we to the Licensor or its representatives, including but not limited to
have designed this version of the GPL to prohibit the practice for those communication on electronic mailing lists, source code control systems,
products. If such problems arise substantially in other domains, we and issue tracking systems that are managed by, or on behalf of, the
stand ready to extend this provision to those domains in future versions Licensor for the purpose of discussing and improving the Work, but
of the GPL, as needed to protect the freedom of users. excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
Finally, every program is threatened constantly by software patents.
States should not allow patents to restrict development and use of "Contributor" shall mean Licensor and any individual or Legal Entity
software on general-purpose computers, but in those that do, we wish to on behalf of whom a Contribution has been received by Licensor and
avoid the special danger that patents applied to a free program could subsequently incorporated within the Work.
make it effectively proprietary. To prevent this, the GPL assures that
patents cannot be used to render the program non-free. 2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
The precise terms and conditions for copying, distribution and worldwide, non-exclusive, no-charge, royalty-free, irrevocable
modification follow. copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
TERMS AND CONDITIONS Work and such Derivative Works in Source or Object form.
0. Definitions. 3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
"This License" refers to version 3 of the GNU General Public License. worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
"Copyright" also means copyright-like laws that apply to other kinds of use, offer to sell, sell, import, and otherwise transfer the Work,
works, such as semiconductor masks. where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
"The Program" refers to any copyrightable work licensed under this Contribution(s) alone or by combination of their Contribution(s)
License. Each licensee is addressed as "you". "Licensees" and with the Work to which such Contribution(s) was submitted. If You
"recipients" may be individuals or organizations. institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
To "modify" a work means to copy from or adapt all or part of the work or a Contribution incorporated within the Work constitutes direct
in a fashion requiring copyright permission, other than the making of an or contributory patent infringement, then any patent licenses
exact copy. The resulting work is called a "modified version" of the granted to You under this License for that Work shall terminate
earlier work or a work "based on" the earlier work. as of the date such litigation is filed.
A "covered work" means either the unmodified Program or a work based 4. Redistribution. You may reproduce and distribute copies of the
on the Program. Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
To "propagate" a work means to do anything with it that, without meet the following conditions:
permission, would make you directly or secondarily liable for
infringement under applicable copyright law, except executing it on a (a) You must give any other recipients of the Work or
computer or modifying a private copy. Propagation includes copying, Derivative Works a copy of this License; and
distribution (with or without modification), making available to the
public, and in some countries other activities as well. (b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
To "convey" a work means any kind of propagation that enables other
parties to make or receive copies. Mere interaction with a user through (c) You must retain, in the Source form of any Derivative Works
a computer network, with no transfer of a copy, is not conveying. that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
An interactive user interface displays "Appropriate Legal Notices" excluding those notices that do not pertain to any part of
to the extent that it includes a convenient and prominently visible the Derivative Works; and
feature that (1) displays an appropriate copyright notice, and (2)
tells the user that there is no warranty for the work (except to the (d) If the Work includes a "NOTICE" text file as part of its
extent that warranties are provided), that licensees may convey the distribution, then any Derivative Works that You distribute must
work under this License, and how to view a copy of this License. If include a readable copy of the attribution notices contained
the interface presents a list of user commands or options, such as a within such NOTICE file, excluding those notices that do not
menu, a prominent item in the list meets this criterion. pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
1. Source Code. as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
The "source code" for a work means the preferred form of the work within a display generated by the Derivative Works, if and
for making modifications to it. "Object code" means any non-source wherever such third-party notices normally appear. The contents
form of a work. of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
A "Standard Interface" means an interface that either is an official notices within Derivative Works that You distribute, alongside
standard defined by a recognized standards body, or, in the case of or as an addendum to the NOTICE text from the Work, provided
interfaces specified for a particular programming language, one that that such additional attribution notices cannot be construed
is widely used among developers working in that language. as modifying the License.
The "System Libraries" of an executable work include anything, other You may add Your own copyright statement to Your modifications and
than the work as a whole, that (a) is included in the normal form of may provide additional or different license terms and conditions
packaging a Major Component, but which is not part of that Major for use, reproduction, or distribution of Your modifications, or
Component, and (b) serves only to enable use of the work with that for any such Derivative Works as a whole, provided Your use,
Major Component, or to implement a Standard Interface for which an reproduction, and distribution of the Work otherwise complies with
implementation is available to the public in source code form. A the conditions stated in this License.
"Major Component", in this context, means a major essential component
(kernel, window system, and so on) of the specific operating system 5. Submission of Contributions. Unless You explicitly state otherwise,
(if any) on which the executable work runs, or a compiler used to any Contribution intentionally submitted for inclusion in the Work
produce the work, or an object code interpreter used to run it. by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
The "Corresponding Source" for a work in object code form means all Notwithstanding the above, nothing herein shall supersede or modify
the source code needed to generate, install, and (for an executable the terms of any separate license agreement you may have executed
work) run the object code and to modify the work, including scripts to with Licensor regarding such Contributions.
control those activities. However, it does not include the work's
System Libraries, or general-purpose tools or generally available free 6. Trademarks. This License does not grant permission to use the trade
programs which are used unmodified in performing those activities but names, trademarks, service marks, or product names of the Licensor,
which are not part of the work. For example, Corresponding Source except as required for reasonable and customary use in describing the
includes interface definition files associated with source files for origin of the Work and reproducing the content of the NOTICE file.
the work, and the source code for shared libraries and dynamically
linked subprograms that the work is specifically designed to require, 7. Disclaimer of Warranty. Unless required by applicable law or
such as by intimate data communication or control flow between those agreed to in writing, Licensor provides the Work (and each
subprograms and other parts of the work. Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
The Corresponding Source need not include anything that users implied, including, without limitation, any warranties or conditions
can regenerate automatically from other parts of the Corresponding of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
Source. PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
The Corresponding Source for a work in source code form is that risks associated with Your exercise of permissions under this License.
same work.
8. Limitation of Liability. In no event and under no legal theory,
2. Basic Permissions. whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
All rights granted under this License are granted for the term of negligent acts) or agreed to in writing, shall any Contributor be
copyright on the Program, and are irrevocable provided the stated liable to You for damages, including any direct, indirect, special,
conditions are met. This License explicitly affirms your unlimited incidental, or consequential damages of any character arising as a
permission to run the unmodified Program. The output from running a result of this License or out of the use or inability to use the
covered work is covered by this License only if the output, given its Work (including but not limited to damages for loss of goodwill,
content, constitutes a covered work. This License acknowledges your work stoppage, computer failure or malfunction, or any and all
rights of fair use or other equivalent, as provided by copyright law. other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
You may make, run and propagate covered works that you do not
convey, without conditions so long as your license otherwise remains 9. Accepting Warranty or Additional Liability. While redistributing
in force. You may convey covered works to others for the sole purpose the Work or Derivative Works thereof, You may choose to offer,
of having them make modifications exclusively for you, or provide you and charge a fee for, acceptance of support, warranty, indemnity,
with facilities for running those works, provided that you comply with or other liability obligations and/or rights consistent with this
the terms of this License in conveying all material for which you do License. However, in accepting such obligations, You may act only
not control copyright. Those thus making or running the covered works on Your own behalf and on Your sole responsibility, not on behalf
for you must do so exclusively on your behalf, under your direction of any other Contributor, and only if You agree to indemnify,
and control, on terms that prohibit them from making any copies of defend, and hold each Contributor harmless for any liability
your copyrighted material outside their relationship with you. incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
Conveying under any other circumstances is permitted solely under
the conditions stated below. Sublicensing is not allowed; section 10
makes it unnecessary.
3. Protecting Users' Legal Rights From Anti-Circumvention Law.
No covered work shall be deemed part of an effective technological
measure under any applicable law fulfilling obligations under article
11 of the WIPO copyright treaty adopted on 20 December 1996, or
similar laws prohibiting or restricting circumvention of such
measures.
When you convey a covered work, you waive any legal power to forbid
circumvention of technological measures to the extent such circumvention
is effected by exercising rights under this License with respect to
the covered work, and you disclaim any intention to limit operation or
modification of the work as a means of enforcing, against the work's
users, your or third parties' legal rights to forbid circumvention of
technological measures.
4. Conveying Verbatim Copies.
You may convey verbatim copies of the Program's source code as you
receive it, in any medium, provided that you conspicuously and
appropriately publish on each copy an appropriate copyright notice;
keep intact all notices stating that this License and any
non-permissive terms added in accord with section 7 apply to the code;
keep intact all notices of the absence of any warranty; and give all
recipients a copy of this License along with the Program.
You may charge any price or no price for each copy that you convey,
and you may offer support or warranty protection for a fee.
5. Conveying Modified Source Versions.
You may convey a work based on the Program, or the modifications to
produce it from the Program, in the form of source code under the
terms of section 4, provided that you also meet all of these conditions:
a) The work must carry prominent notices stating that you modified
it, and giving a relevant date.
b) The work must carry prominent notices stating that it is
released under this License and any conditions added under section
7. This requirement modifies the requirement in section 4 to
"keep intact all notices".
c) You must license the entire work, as a whole, under this
License to anyone who comes into possession of a copy. This
License will therefore apply, along with any applicable section 7
additional terms, to the whole of the work, and all its parts,
regardless of how they are packaged. This License gives no
permission to license the work in any other way, but it does not
invalidate such permission if you have separately received it.
d) If the work has interactive user interfaces, each must display
Appropriate Legal Notices; however, if the Program has interactive
interfaces that do not display Appropriate Legal Notices, your
work need not make them do so.
A compilation of a covered work with other separate and independent
works, which are not by their nature extensions of the covered work,
and which are not combined with it such as to form a larger program,
in or on a volume of a storage or distribution medium, is called an
"aggregate" if the compilation and its resulting copyright are not
used to limit the access or legal rights of the compilation's users
beyond what the individual works permit. Inclusion of a covered work
in an aggregate does not cause this License to apply to the other
parts of the aggregate.
6. Conveying Non-Source Forms.
You may convey a covered work in object code form under the terms
of sections 4 and 5, provided that you also convey the
machine-readable Corresponding Source under the terms of this License,
in one of these ways:
a) Convey the object code in, or embodied in, a physical product
(including a physical distribution medium), accompanied by the
Corresponding Source fixed on a durable physical medium
customarily used for software interchange.
b) Convey the object code in, or embodied in, a physical product
(including a physical distribution medium), accompanied by a
written offer, valid for at least three years and valid for as
long as you offer spare parts or customer support for that product
model, to give anyone who possesses the object code either (1) a
copy of the Corresponding Source for all the software in the
product that is covered by this License, on a durable physical
medium customarily used for software interchange, for a price no
more than your reasonable cost of physically performing this
conveying of source, or (2) access to copy the
Corresponding Source from a network server at no charge.
c) Convey individual copies of the object code with a copy of the
written offer to provide the Corresponding Source. This
alternative is allowed only occasionally and noncommercially, and
only if you received the object code with such an offer, in accord
with subsection 6b.
d) Convey the object code by offering access from a designated
place (gratis or for a charge), and offer equivalent access to the
Corresponding Source in the same way through the same place at no
further charge. You need not require recipients to copy the
Corresponding Source along with the object code. If the place to
copy the object code is a network server, the Corresponding Source
may be on a different server (operated by you or a third party)
that supports equivalent copying facilities, provided you maintain
clear directions next to the object code saying where to find the
Corresponding Source. Regardless of what server hosts the
Corresponding Source, you remain obligated to ensure that it is
available for as long as needed to satisfy these requirements.
e) Convey the object code using peer-to-peer transmission, provided
you inform other peers where the object code and Corresponding
Source of the work are being offered to the general public at no
charge under subsection 6d.
A separable portion of the object code, whose source code is excluded
from the Corresponding Source as a System Library, need not be
included in conveying the object code work.
A "User Product" is either (1) a "consumer product", which means any
tangible personal property which is normally used for personal, family,
or household purposes, or (2) anything designed or sold for incorporation
into a dwelling. In determining whether a product is a consumer product,
doubtful cases shall be resolved in favor of coverage. For a particular
product received by a particular user, "normally used" refers to a
typical or common use of that class of product, regardless of the status
of the particular user or of the way in which the particular user
actually uses, or expects or is expected to use, the product. A product
is a consumer product regardless of whether the product has substantial
commercial, industrial or non-consumer uses, unless such uses represent
the only significant mode of use of the product.
"Installation Information" for a User Product means any methods,
procedures, authorization keys, or other information required to install
and execute modified versions of a covered work in that User Product from
a modified version of its Corresponding Source. The information must
suffice to ensure that the continued functioning of the modified object
code is in no case prevented or interfered with solely because
modification has been made.
If you convey an object code work under this section in, or with, or
specifically for use in, a User Product, and the conveying occurs as
part of a transaction in which the right of possession and use of the
User Product is transferred to the recipient in perpetuity or for a
fixed term (regardless of how the transaction is characterized), the
Corresponding Source conveyed under this section must be accompanied
by the Installation Information. But this requirement does not apply
if neither you nor any third party retains the ability to install
modified object code on the User Product (for example, the work has
been installed in ROM).
The requirement to provide Installation Information does not include a
requirement to continue to provide support service, warranty, or updates
for a work that has been modified or installed by the recipient, or for
the User Product in which it has been modified or installed. Access to a
network may be denied when the modification itself materially and
adversely affects the operation of the network or violates the rules and
protocols for communication across the network.
Corresponding Source conveyed, and Installation Information provided,
in accord with this section must be in a format that is publicly
documented (and with an implementation available to the public in
source code form), and must require no special password or key for
unpacking, reading or copying.
7. Additional Terms.
"Additional permissions" are terms that supplement the terms of this
License by making exceptions from one or more of its conditions.
Additional permissions that are applicable to the entire Program shall
be treated as though they were included in this License, to the extent
that they are valid under applicable law. If additional permissions
apply only to part of the Program, that part may be used separately
under those permissions, but the entire Program remains governed by
this License without regard to the additional permissions.
When you convey a copy of a covered work, you may at your option
remove any additional permissions from that copy, or from any part of
it. (Additional permissions may be written to require their own
removal in certain cases when you modify the work.) You may place
additional permissions on material, added by you to a covered work,
for which you have or can give appropriate copyright permission.
Notwithstanding any other provision of this License, for material you
add to a covered work, you may (if authorized by the copyright holders of
that material) supplement the terms of this License with terms:
a) Disclaiming warranty or limiting liability differently from the
terms of sections 15 and 16 of this License; or
b) Requiring preservation of specified reasonable legal notices or
author attributions in that material or in the Appropriate Legal
Notices displayed by works containing it; or
c) Prohibiting misrepresentation of the origin of that material, or
requiring that modified versions of such material be marked in
reasonable ways as different from the original version; or
d) Limiting the use for publicity purposes of names of licensors or
authors of the material; or
e) Declining to grant rights under trademark law for use of some
trade names, trademarks, or service marks; or
f) Requiring indemnification of licensors and authors of that
material by anyone who conveys the material (or modified versions of
it) with contractual assumptions of liability to the recipient, for
any liability that these contractual assumptions directly impose on
those licensors and authors.
All other non-permissive additional terms are considered "further
restrictions" within the meaning of section 10. If the Program as you
received it, or any part of it, contains a notice stating that it is
governed by this License along with a term that is a further
restriction, you may remove that term. If a license document contains
a further restriction but permits relicensing or conveying under this
License, you may add to a covered work material governed by the terms
of that license document, provided that the further restriction does
not survive such relicensing or conveying.
If you add terms to a covered work in accord with this section, you
must place, in the relevant source files, a statement of the
additional terms that apply to those files, or a notice indicating
where to find the applicable terms.
Additional terms, permissive or non-permissive, may be stated in the
form of a separately written license, or stated as exceptions;
the above requirements apply either way.
8. Termination.
You may not propagate or modify a covered work except as expressly
provided under this License. Any attempt otherwise to propagate or
modify it is void, and will automatically terminate your rights under
this License (including any patent licenses granted under the third
paragraph of section 11).
However, if you cease all violation of this License, then your
license from a particular copyright holder is reinstated (a)
provisionally, unless and until the copyright holder explicitly and
finally terminates your license, and (b) permanently, if the copyright
holder fails to notify you of the violation by some reasonable means
prior to 60 days after the cessation.
Moreover, your license from a particular copyright holder is
reinstated permanently if the copyright holder notifies you of the
violation by some reasonable means, this is the first time you have
received notice of violation of this License (for any work) from that
copyright holder, and you cure the violation prior to 30 days after
your receipt of the notice.
Termination of your rights under this section does not terminate the
licenses of parties who have received copies or rights from you under
this License. If your rights have been terminated and not permanently
reinstated, you do not qualify to receive new licenses for the same
material under section 10.
9. Acceptance Not Required for Having Copies.
You are not required to accept this License in order to receive or
run a copy of the Program. Ancillary propagation of a covered work
occurring solely as a consequence of using peer-to-peer transmission
to receive a copy likewise does not require acceptance. However,
nothing other than this License grants you permission to propagate or
modify any covered work. These actions infringe copyright if you do
not accept this License. Therefore, by modifying or propagating a
covered work, you indicate your acceptance of this License to do so.
10. Automatic Licensing of Downstream Recipients.
Each time you convey a covered work, the recipient automatically
receives a license from the original licensors, to run, modify and
propagate that work, subject to this License. You are not responsible
for enforcing compliance by third parties with this License.
An "entity transaction" is a transaction transferring control of an
organization, or substantially all assets of one, or subdividing an
organization, or merging organizations. If propagation of a covered
work results from an entity transaction, each party to that
transaction who receives a copy of the work also receives whatever
licenses to the work the party's predecessor in interest had or could
give under the previous paragraph, plus a right to possession of the
Corresponding Source of the work from the predecessor in interest, if
the predecessor has it or can get it with reasonable efforts.
You may not impose any further restrictions on the exercise of the
rights granted or affirmed under this License. For example, you may
not impose a license fee, royalty, or other charge for exercise of
rights granted under this License, and you may not initiate litigation
(including a cross-claim or counterclaim in a lawsuit) alleging that
any patent claim is infringed by making, using, selling, offering for
sale, or importing the Program or any portion of it.
11. Patents.
A "contributor" is a copyright holder who authorizes use under this
License of the Program or a work on which the Program is based. The
work thus licensed is called the contributor's "contributor version".
A contributor's "essential patent claims" are all patent claims
owned or controlled by the contributor, whether already acquired or
hereafter acquired, that would be infringed by some manner, permitted
by this License, of making, using, or selling its contributor version,
but do not include claims that would be infringed only as a
consequence of further modification of the contributor version. For
purposes of this definition, "control" includes the right to grant
patent sublicenses in a manner consistent with the requirements of
this License.
Each contributor grants you a non-exclusive, worldwide, royalty-free
patent license under the contributor's essential patent claims, to
make, use, sell, offer for sale, import and otherwise run, modify and
propagate the contents of its contributor version.
In the following three paragraphs, a "patent license" is any express
agreement or commitment, however denominated, not to enforce a patent
(such as an express permission to practice a patent or covenant not to
sue for patent infringement). To "grant" such a patent license to a
party means to make such an agreement or commitment not to enforce a
patent against the party.
If you convey a covered work, knowingly relying on a patent license,
and the Corresponding Source of the work is not available for anyone
to copy, free of charge and under the terms of this License, through a
publicly available network server or other readily accessible means,
then you must either (1) cause the Corresponding Source to be so
available, or (2) arrange to deprive yourself of the benefit of the
patent license for this particular work, or (3) arrange, in a manner
consistent with the requirements of this License, to extend the patent
license to downstream recipients. "Knowingly relying" means you have
actual knowledge that, but for the patent license, your conveying the
covered work in a country, or your recipient's use of the covered work
in a country, would infringe one or more identifiable patents in that
country that you have reason to believe are valid.
If, pursuant to or in connection with a single transaction or
arrangement, you convey, or propagate by procuring conveyance of, a
covered work, and grant a patent license to some of the parties
receiving the covered work authorizing them to use, propagate, modify
or convey a specific copy of the covered work, then the patent license
you grant is automatically extended to all recipients of the covered
work and works based on it.
A patent license is "discriminatory" if it does not include within
the scope of its coverage, prohibits the exercise of, or is
conditioned on the non-exercise of one or more of the rights that are
specifically granted under this License. You may not convey a covered
work if you are a party to an arrangement with a third party that is
in the business of distributing software, under which you make payment
to the third party based on the extent of your activity of conveying
the work, and under which the third party grants, to any of the
parties who would receive the covered work from you, a discriminatory
patent license (a) in connection with copies of the covered work
conveyed by you (or copies made from those copies), or (b) primarily
for and in connection with specific products or compilations that
contain the covered work, unless you entered into that arrangement,
or that patent license was granted, prior to 28 March 2007.
Nothing in this License shall be construed as excluding or limiting
any implied license or other defenses to infringement that may
otherwise be available to you under applicable patent law.
12. No Surrender of Others' Freedom.
If conditions are imposed on you (whether by court order, agreement or
otherwise) that contradict the conditions of this License, they do not
excuse you from the conditions of this License. If you cannot convey a
covered work so as to satisfy simultaneously your obligations under this
License and any other pertinent obligations, then as a consequence you may
not convey it at all. For example, if you agree to terms that obligate you
to collect a royalty for further conveying from those to whom you convey
the Program, the only way you could satisfy both those terms and this
License would be to refrain entirely from conveying the Program.
13. Use with the GNU Affero General Public License.
Notwithstanding any other provision of this License, you have
permission to link or combine any covered work with a work licensed
under version 3 of the GNU Affero General Public License into a single
combined work, and to convey the resulting work. The terms of this
License will continue to apply to the part which is the covered work,
but the special requirements of the GNU Affero General Public License,
section 13, concerning interaction through a network will apply to the
combination as such.
14. Revised Versions of this License.
The Free Software Foundation may publish revised and/or new versions of
the GNU General Public License from time to time. Such new versions will
be similar in spirit to the present version, but may differ in detail to
address new problems or concerns.
Each version is given a distinguishing version number. If the
Program specifies that a certain numbered version of the GNU General
Public License "or any later version" applies to it, you have the
option of following the terms and conditions either of that numbered
version or of any later version published by the Free Software
Foundation. If the Program does not specify a version number of the
GNU General Public License, you may choose any version ever published
by the Free Software Foundation.
If the Program specifies that a proxy can decide which future
versions of the GNU General Public License can be used, that proxy's
public statement of acceptance of a version permanently authorizes you
to choose that version for the Program.
Later license versions may give you additional or different
permissions. However, no additional obligations are imposed on any
author or copyright holder as a result of your choosing to follow a
later version.
15. Disclaimer of Warranty.
THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY
APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT
HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY
OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO,
THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR
PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM
IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF
ALL NECESSARY SERVICING, REPAIR OR CORRECTION.
16. Limitation of Liability.
IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING
WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS
THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY
GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE
USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF
DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD
PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS),
EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF
SUCH DAMAGES.
17. Interpretation of Sections 15 and 16.
If the disclaimer of warranty and limitation of liability provided
above cannot be given local legal effect according to their terms,
reviewing courts shall apply local law that most closely approximates
an absolute waiver of all civil liability in connection with the
Program, unless a warranty or assumption of liability accompanies a
copy of the Program in return for a fee.
END OF TERMS AND CONDITIONS END OF TERMS AND CONDITIONS
How to Apply These Terms to Your New Programs APPENDIX: How to apply the Apache License to your work.
If you develop a new program, and you want it to be of the greatest To apply the Apache License to your work, attach the following
possible use to the public, the best way to achieve this is to make it boilerplate notice, with the fields enclosed by brackets "[]"
free software which everyone can redistribute and change under these terms. replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
To do so, attach the following notices to the program. It is safest Copyright [yyyy] [name of copyright owner]
to attach them to the start of each source file to most effectively
state the exclusion of warranty; and each file should have at least
the "copyright" line and a pointer to where the full notice is found.
<one line to give the program's name and a brief idea of what it does.> Licensed under the Apache License, Version 2.0 (the "License");
Copyright (C) <year> <name of author> you may not use this file except in compliance with the License.
You may obtain a copy of the License at
This program is free software: you can redistribute it and/or modify http://www.apache.org/licenses/LICENSE-2.0
it under the terms of the GNU General Public License as published by
the Free Software Foundation, either version 3 of the License, or
(at your option) any later version.
This program is distributed in the hope that it will be useful, Unless required by applicable law or agreed to in writing, software
but WITHOUT ANY WARRANTY; without even the implied warranty of distributed under the License is distributed on an "AS IS" BASIS,
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
GNU General Public License for more details. See the License for the specific language governing permissions and
limitations under the License.
You should have received a copy of the GNU General Public License
along with this program. If not, see <https://www.gnu.org/licenses/>.
Also add information on how to contact you by electronic and paper mail.
If the program does terminal interaction, make it output a short
notice like this when it starts in an interactive mode:
<program> Copyright (C) <year> <name of author>
This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'.
This is free software, and you are welcome to redistribute it
under certain conditions; type `show c' for details.
The hypothetical commands `show w' and `show c' should show the appropriate
parts of the General Public License. Of course, your program's commands
might be different; for a GUI interface, you would use an "about box".
You should also get your employer (if you work as a programmer) or school,
if any, to sign a "copyright disclaimer" for the program, if necessary.
For more information on this, and how to apply and follow the GNU GPL, see
<https://www.gnu.org/licenses/>.
The GNU General Public License does not permit incorporating your program
into proprietary programs. If your program is a subroutine library, you
may consider it more useful to permit linking proprietary applications with
the library. If this is what you want to do, use the GNU Lesser General
Public License instead of this License. But first, please read
<https://www.gnu.org/licenses/why-not-lgpl.html>.
+33 -22
View File
@@ -1,6 +1,6 @@
<div align="center"> <div align="center">
<img src="assets/images/logo.png" width="auto" alt="Logo"> <img src="docs/images/logo.png" width="auto" alt="Logo">
<p> <p>
<strong>A lightweight Transformer training & inference framework</strong> <strong>A lightweight Transformer training & inference framework</strong>
</p> </p>
@@ -8,7 +8,7 @@
<div align="center"> <div align="center">
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python"> <img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license"> <img src="https://img.shields.io/badge/license-Apache--2.0-blue.svg" alt="license">
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release"> <img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars"> <img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks"> <img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
@@ -17,7 +17,7 @@
<div align="center"> <div align="center">
<a href="#english">English</a> • <a href="#english">English</a> •
<a href="assets/docs/README-zh-CN.md">中文</a> • <a href="docs/README-zh-CN.md">中文</a> •
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> • <a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> • <a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
<a href="https://huggingface.co/ViperEkura">HuggingFace</a> <a href="https://huggingface.co/ViperEkura">HuggingFace</a>
@@ -27,7 +27,7 @@
## 📖 Table of Contents ## 📖 Table of Contents
- [Features](#features) - [Overview](#overview)
- [Getting Started](#getting-started) - [Getting Started](#getting-started)
- [Demo](#demo) - [Demo](#demo)
- [Documentation](#documentation) - [Documentation](#documentation)
@@ -40,15 +40,19 @@
<a id="english"></a> <a id="english"></a>
## English ## English
### Features ### Overview
- 🚀 **High Performance**: Optimized for both training and inference with efficient parallelization. AstrAI is an end-to-end Transformer framework for building, training, evaluating, and serving models. It provides a compact PyTorch codebase for the complete model lifecycle, from declarative data preprocessing and distributed training to continuous-batching inference and OpenAI/Anthropic-compatible APIs.
- 🔧 **Flexible**: Support for seq/sft/dpo/grpo training, customizable model architectures.
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos. | Area | Capabilities |
- 📦 **Lightweight**: Minimal dependencies, easy to deploy. |---|---|
- 🔬 **ResearchFriendly**: Modular design, easy to experiment with new ideas. | **Models** | Autoregressive language models and embedding models with GQA, MLA, MoE, RoPE, and extensible attention/FFN components |
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading. | **Training** | Pre-training (`seq`), supervised fine-tuning (`sft`), DPO, and GRPO with gradient accumulation, checkpointing, DDP, and FSDP |
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box. | **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 |
| **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 |
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
### Getting Started ### Getting Started
@@ -56,6 +60,8 @@ End-to-end walkthrough in 5 steps:
**1. Install** **1. Install**
AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `scripts/tools/generate.py`, generation evaluations, and the generation demos require CUDA; CPU support is limited to components with an explicit CPU device path, such as the HTTP server and direct-scoring evaluations.
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
@@ -132,7 +138,7 @@ Check out the demos in the `scripts/demo/` folder:
# Download model weights (required before running demos) # Download model weights (required before running demos)
python scripts/demo/download.py # model → params/ python scripts/demo/download.py # model → params/
# Interactive streaming chat (multi-turn, maintains history) # Single-turn interactive streaming prompt loop (no conversation history)
python scripts/demo/stream_chat.py python scripts/demo/stream_chat.py
# Type your message after >>, type !exit to quit # Type your message after >>, type !exit to quit
@@ -183,7 +189,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker Compose (GPU, default) # Docker Compose (GPU, default)
docker compose up -d docker compose up -d
# Docker Compose (CPU only) # Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
docker compose --profile cpu up -d docker compose --profile cpu up -d
``` ```
@@ -213,18 +219,23 @@ curl -X POST http://localhost:8000/v1/messages \
curl http://localhost:8000/health curl http://localhost:8000/health
``` ```
See [Inference Guide](assets/docs/inference.md) for SSE streaming format, error codes, and stats endpoint. See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error codes, and stats endpoint.
### Documentation ### Documentation
| Document | Description | | Document | Description |
|----------|-------------| |----------|-------------|
| [CLI Reference](./assets/docs/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) | | [Get Started](./docs/get-started.md) | Installation and quickstart |
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns | | [CLI Reference](./docs/guides/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas | | [Preprocessing](./docs/guides/preprocessing.md) | Declarative JSON-driven data preprocessing |
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API | | [Training](./docs/guides/training.md) | Training loop, strategies & formulas |
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture | | [Inference](./docs/guides/inference.md) | KVCache, continuous batching, sampling & HTTP API |
| [Preprocessing](./assets/docs/preprocessing.md) | Declarative JSON-driven data preprocessing | | [Evaluation](./docs/guides/evaluation.md) | HumanEval, MMLU, PPL, ROUGE, IFD, IFEval |
| [Distributed](./docs/guides/distributed.md) | Multi-GPU DDP / FSDP training |
| [Architecture](./docs/developer/architecture.md) | System architecture, class diagram & design patterns |
| [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 |
| [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks |
### Contributing ### Contributing
@@ -245,7 +256,7 @@ For major changes, please open an issue first to discuss what you would like to
### License ### License
This project is licensed under the [GPL-3.0 License](LICENSE). This project is licensed under the [Apache-2.0 License](LICENSE).
--- ---
-130
View File
@@ -1,130 +0,0 @@
# Data Flow
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](preprocessing.md).
## Contents
- [Overview](#overview)
- [Data Preparation](#data-preparation) — tokenization, format detection, backends
- [Data Keys by Training Type](#data-keys-by-training-type)
- [Dataset Architecture](#dataset-architecture)
- [Sampler](#sampler)
- [DataLoader](#dataloader)
## Overview
```
JSONL Lines → Pipeline (mask builder) → Tokenized Tensors
.h5 or .bin storage
Store.load()
Store.fetch(begin, end, keys)
BaseDataset.__getitem__(idx)
Sampler → DataLoader → Training / Inference
```
## Data Preparation
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or binary (`.bin` + `meta.json`) files with keyed tensor groups.
### Tokenization
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](preprocessing.md)), and produces flat token sequences:
```python
# Per JSONL line: messages → chat template → token IDs + loss mask
tokens = tokenizer.encode(rendered_text) # List[int]
loss_mask = [0, 0, 0, 1, 1, 1, 1, 1, 1] # 0=masked, 1=train
# Stored as flat tensors, packed with other lines by packing strategy
```
The output `meta.json` records the storage format, key names, dtype, total token count, and tensor shapes for each shard.
### Format Detection
`detect_format(load_path)` inspects the path:
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5``"h5"`, `.jsonl``"jsonl"`, unknown suffix raises `ValueError`
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json``"bin"`, or `*.jsonl` + `dataset_config.json``"jsonl"`
### Store Backends
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
```
StoreFactory.create("h5") → H5Store
StoreFactory.create("bin") → MmapStore
StoreFactory.create("jsonl") → JsonlStore
```
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
## Data Keys by Training Type
| Type | Storage Keys | Access Mode |
|------|-------------|-------------|
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
## Dataset Architecture
```
DatasetFactory.load(train_type, load_path, window_size, stride=None,
storage_type=None, tokenizer_path=None,
max_len=2048, store=None)
→ BaseDataset.load(load_path, storage_type=None)
→ detect_format(load_path)
→ StoreFactory.create(storage_type)
→ Store.load(load_path)
→ _normalize(raw) # base Store, shared by both backends
→ Store._data[Dict[str, List[Tensor]]]
+ _cum[Dict[str, List[int]]] (stream mode)
+ _offsets[Dict[str, List[int]]] (record mode)
Stream datasets (SEQ/SFT):
BaseDataset.__getitem__(idx)
→ get_index(idx) → [begin, end)
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
Record datasets (DPO/GRPO via RecordDataset):
RecordDataset.__getitem__(idx)
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
```
Class hierarchy: `BaseDataset``SEQDataset` / `SFTDataset` (stream); `BaseDataset``RecordDataset``DPODataset` / `GRPODataset` (record).
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
## Sampler
`ResumableDistributedSampler` supports checkpoint-aware distributed sampling:
- Tracks `start_epoch` / `start_iter` for resume
- Shuffle via `torch.Generator(seed + epoch)`
- Per-replica index slicing for DDP
## DataLoader
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
> Document Update Time: 2026-07-19
-252
View File
@@ -1,252 +0,0 @@
# Inference
## Contents
- [KV Cache](#kv-cache)
- [KVCache System](#kvcache-system)
- [Continuous Batching](#continuous-batching)
- [Sampling](#sampling-strategy-pattern)
- [Protocol Handlers](#protocol-handlers-strategy-pattern)
- [Engine & GenerateResult](#engine--generateresult)
- [HTTP API](#http-api) — endpoints, SSE, errors, stats
- [Engine API](#engine-api)
## KV Cache
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
$$
o_n = \sum_j \text{softmax}\left(\frac{q_n k_j}{\sqrt{d_k}}\right) v_j
$$
RoPE is applied **before** KV cache write, not after — otherwise position encoding drift occurs.
## KVCache System
Seven classes working together, with two concrete cache implementations:
### ContiguousCache (default)
```
ContiguousCache (simple contiguous per-slot cache)
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
```
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, n_kv_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
### PageCache (paged with prefix sharing)
```
PageCache (paged KV cache with prefix sharing, alternative)
├── PagePool orchestrates page allocation + prefix matching
│ ├── Allocator bitmask-based page allocator + ref-count + LRU
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
├── TaskTable maps task_id → page_table + cached token count
├── Storage k_cache / v_cache tensors (n_layers × n_pages × page_size × n_kv_heads × head_dim)
└── PageCacheView bundles Storage + page_table + total_len for attention layers
```
`isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
## Continuous Batching
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
```
1. Cleanup → Remove finished tasks, free KV cache slots/pages
2. Refill → Pop from waiting_queue, task_alloc resources, activate
3. Prefill → Group by (prompt_len, start_pos), run full forward
4. Decode → Run single-token forward for each same-position group
```
## Sampling (Strategy Pattern)
```
BaseSamplingStrategy (ABC)
├── TemperatureStrategy
├── TopKStrategy
├── TopPStrategy
└── SamplingPipeline
```
`SamplingPipeline` composes them: Temperature → Top-K → Top-P → softmax → multinomial.
`sample()` is a convenience shortcut for one-shot usage.
## Protocol Handlers (Strategy Pattern)
```python
class ProtocolHandler: # concrete orchestrator
def __init__(self, request, engine, builder): ...
async def handle(self):
prompt, ctx, stops = builder.prepare(request, engine)
agen = engine.generate_async(prompt, ...)
if stream: self._handle_stream(agen, ctx, stops)
else: return await self._handle_non_stream(agen, ctx, stops)
```
`ResponseBuilder` (ABC): `prepare()`, `format_stream_start()`, `format_chunk()`, `format_stream_end()`, `format_response()`.
`OpenAIResponseBuilder``/v1/chat/completions`, `AnthropicResponseBuilder``/v1/messages`.
Adding a protocol = one builder file, no handler subclassing needed.
## Engine & GenerateResult
```
InferenceEngine
├── generate(prompt, stream, ...) → str | List[str] | Generator
├── generate_with_request(req) → same
├── generate_async(prompt, ...) → AsyncGenerator
├── get_stats() → Dict
└── shutdown()
```
`GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`.
## HTTP API
```
POST /v1/chat/completions OpenAI
POST /v1/messages Anthropic
GET /health {"status":"ok","model_loaded":true}
GET /stats scheduler statistics
```
### OpenAI
```bash
curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
```
Response:
```json
{
"id": "chatcmpl-abc123",
"object": "chat.completion",
"created": 1717000000,
"model": "astrai",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Hello!"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}
}
```
Streaming SSE: `object: "chat.completion.chunk"` — starts with role delta, then token chunks, ends with finish chunk + usage stats, then `data: [DONE]`.
### Anthropic
```bash
curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json" \
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
```
Supports `stop_sequences` and streaming via `event: content_block_delta`.
### GenerationRequest Parameters
| Param | Type | Default | Description |
|-------|------|---------|-------------|
| `messages` | List[dict] | required | Chat messages (role, content) |
| `top_k` | int | 50 | Top-k count |
| `top_p` | float | 1.0 | Nucleus threshold |
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
| `max_tokens` | Optional[int] | None | Max generation length |
| `stream` | bool | False | Stream output |
### SSE Streaming Format
**OpenAI** (`/v1/chat/completions`, `stream=true`):
```
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
"choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":0,"model":"astrai",
"choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}
data: {"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}
data: [DONE]
```
**Anthropic** (`/v1/messages`, `stream=true`):
```
event: message_start
data: {"type":"message_start","message":{"id":"msg_...","model":"astrai","role":"assistant",
"content":[],"usage":{"input_tokens":0}}}
event: content_block_start
data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}
event: content_block_delta
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}
event: content_block_stop
data: {"type":"content_block_stop","index":0}
event: message_delta
data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{...}}
event: message_stop
data: {"type":"message_stop"}
```
### Error Responses
The server returns standard HTTP status codes. Pydantic validation errors (e.g. missing required fields)
are handled automatically by FastAPI with 422 status. The only application-level error is engine initialization:
| Status | Meaning |
|--------|---------|
| 200 | Success |
| 422 | Unprocessable entity (Pydantic validation) |
| 503 | Service unavailable (model not loaded, engine not ready) |
Error response body (503):
```json
{
"detail": "Engine not initialized"
}
```
### Stats Endpoint
```
GET /stats
```
Response:
```json
{
"total_tasks": 128,
"total_tokens": 10240,
"active_tasks": 3,
"waiting_queue": 2
}
```
## Engine API
```python
# Non-streaming
engine.generate("Hello", stream=False) # -> str
engine.generate(["A", "B"], stream=False) # -> List[str]
# Streaming
engine.generate("Hello", stream=True) # -> Generator[str]
engine.generate(["A", "B"], stream=True) # -> Generator[Tuple[int, str]]
# Async
async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[str]
print(token)
```
> Document Update Time: 2026-07-09
+29 -1
View File
@@ -1,6 +1,9 @@
__version__ = "1.3.10" __version__ = "1.3.12"
__author__ = "ViperEkura" __author__ = "ViperEkura"
import logging
import os
from astrai.config import ( from astrai.config import (
AutoRegressiveLMConfig, AutoRegressiveLMConfig,
BaseModelConfig, BaseModelConfig,
@@ -53,6 +56,30 @@ from astrai.trainer import (
Trainer, Trainer,
) )
def setup_logging(level: str = "INFO"):
"""Attach a handler to the ``astrai`` logger (only, not root).
Call once per process, e.g. at the top of CLI scripts.
Set ``ASTR_LOG_LEVEL`` to override the default ``INFO``.
"""
_logger = logging.getLogger("astrai")
if _logger.handlers:
return
_level = getattr(
logging, os.environ.get("ASTR_LOG_LEVEL", level).upper(), logging.INFO
)
_logger.setLevel(_level)
_handler = logging.StreamHandler()
_handler.setFormatter(
logging.Formatter(
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
)
_logger.addHandler(_handler)
__all__ = [ __all__ = [
"AutoRegressiveLM", "AutoRegressiveLM",
"AutoRegressiveLMConfig", "AutoRegressiveLMConfig",
@@ -94,5 +121,6 @@ __all__ = [
"only_on_rank", "only_on_rank",
"run_server", "run_server",
"sample", "sample",
"setup_logging",
"spawn_parallel_fn", "spawn_parallel_fn",
] ]
+17 -77
View File
@@ -1,92 +1,32 @@
import json import json
from dataclasses import MISSING, dataclass, fields from dataclasses import asdict
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Optional, Self, Union, get_type_hints from typing import Any, Dict, Self, Union
from pydantic import ConfigDict
from pydantic.dataclasses import dataclass
@dataclass @dataclass(config=ConfigDict(use_attribute_docstrings=True))
class BaseConfig: class BaseConfig:
def to_dict(self) -> Dict[str, Any]: def to_dict(self) -> Dict[str, Any]:
d = {} result = {}
for fld in fields(self): for k, v in asdict(self).items():
v = getattr(self, fld.name) if isinstance(v, tuple):
if isinstance(v, (str, int, float, bool)): v = list(v)
d[fld.name] = v
elif v is None:
d[fld.name] = None
elif isinstance(v, (dict, list, tuple)):
try: try:
val = list(v) if isinstance(v, tuple) else v json.dumps(v)
json.dumps(val) result[k] = v
d[fld.name] = val
except (TypeError, ValueError): except (TypeError, ValueError):
# Skip non-serializable runtime objects (e.g. model_fn, dataset).
# TrainConfig mixes hyperparams with callables/datasets; only the
# JSON-serializable subset is written to checkpoint meta.
pass pass
elif isinstance(v, BaseConfig): return result
d[fld.name] = v.to_dict()
elif hasattr(v, "__dataclass_fields__"):
sub = {}
for f in fields(v):
a = getattr(v, f.name)
sub[f.name] = list(a) if isinstance(a, tuple) else a
d[fld.name] = sub
return d
@classmethod @classmethod
def from_dict(cls, d: Dict[str, Any]) -> Self: def from_dict(cls, d: Dict[str, Any]) -> Self:
hints = get_type_hints(cls) return cls(**d)
inst = cls.__new__(cls)
for fld in fields(cls):
if fld.name in d:
v = d[fld.name]
target = cls._unwrap_optional(hints.get(fld.name))
if target is not None:
try:
v = cls._coerce(v, target)
except (TypeError, ValueError):
pass
object.__setattr__(inst, fld.name, v)
elif fld.default is not MISSING:
object.__setattr__(inst, fld.name, fld.default)
elif fld.default_factory is not MISSING:
object.__setattr__(inst, fld.name, fld.default_factory())
else:
object.__setattr__(inst, fld.name, None)
return inst
@staticmethod
def _unwrap_optional(tp) -> Optional[type]:
if tp is None:
return None
origin = getattr(tp, "__origin__", None)
if origin is not None:
args = getattr(tp, "__args__", ())
non_none = [a for a in args if a is not type(None)]
return non_none[0] if non_none else None
return tp
@staticmethod
def _coerce(value: Any, target_type: type) -> Any:
if target_type is bool and isinstance(value, bool):
return value
if (
target_type is int
and isinstance(value, (int, float))
and not isinstance(value, bool)
):
return int(value)
if (
target_type is float
and isinstance(value, (int, float))
and not isinstance(value, bool)
):
return float(value)
if target_type is str and isinstance(value, str):
return value
if isinstance(value, target_type):
return value
if isinstance(value, dict) and issubclass(target_type, BaseConfig):
return target_type.from_dict(value)
raise TypeError
@classmethod @classmethod
def from_file(cls, path: Union[str, Path]) -> Self: def from_file(cls, path: Union[str, Path]) -> Self:
+122 -26
View File
@@ -1,9 +1,14 @@
from dataclasses import dataclass
from typing import Any, Dict, Optional from typing import Any, Dict, Optional
from pydantic import field_validator
from pydantic.dataclasses import dataclass
from astrai.config.base import BaseConfig from astrai.config.base import BaseConfig
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
_ATTN_TYPES = frozenset({"gqa", "mla"})
_FFN_TYPES = frozenset({"mlp", "moe"})
class ConfigFactory(BaseFactory[BaseConfig]): class ConfigFactory(BaseFactory[BaseConfig]):
"""Factory that dispatches config classes by ``model_type``.""" """Factory that dispatches config classes by ``model_type``."""
@@ -17,7 +22,12 @@ class ConfigFactory(BaseFactory[BaseConfig]):
@dataclass @dataclass
class BaseModelConfig(BaseConfig): class BaseModelConfig(BaseConfig):
"""Base config with ``model_type`` dispatch and file I/O.""" """Base config with ``model_type`` dispatch and file I/O.
Args:
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
"""
model_type: Optional[str] = None model_type: Optional[str] = None
neftune_alpha: float = 0.0 neftune_alpha: float = 0.0
@@ -26,57 +36,143 @@ class BaseModelConfig(BaseConfig):
@dataclass @dataclass
@ConfigFactory.register("autoregressive_lm") @ConfigFactory.register("autoregressive_lm")
class AutoRegressiveLMConfig(BaseModelConfig): class AutoRegressiveLMConfig(BaseModelConfig):
"""Configuration for autoregressive language model.""" """Configuration for autoregressive language model.
Args:
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
tie_word_embeddings (Optional[bool]): Whether to tie embedding and lm_head weights. Defaults to None.
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
kv_lora_rank (Optional[int]): KV compression rank, MLA only. Defaults to None.
qk_nope_head_dim (Optional[int]): Non-RoPE head dimension, MLA only. Defaults to None.
qk_rope_head_dim (Optional[int]): RoPE head dimension, MLA only. Defaults to None.
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
n_routed_experts (Optional[int]): Number of routed experts, MoE only. Defaults to None.
n_shared_experts (Optional[int]): Number of shared experts, MoE only. Defaults to None.
n_activated_experts (Optional[int]): Number of activated experts per token, MoE only. Defaults to None.
topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None.
moe_intermediate_size (Optional[int]): Expert hidden dim, defaults to intermediate_size if None. MoE only.
shared_expert_intermediate_size (Optional[int]): Shared expert hidden dim, defaults to intermediate_size if None. MoE only.
norm_topk_prob (bool): Normalize top-k routing probabilities. Defaults to True.
decoder_sparse_step (int): Frequency of MoE layers, 1=every layer. Defaults to 1.
mlp_only_layers (Optional[list[int]]): Layer indices using dense MLP instead of MoE. Defaults to None.
"""
vocab_size: Optional[int] = None vocab_size: Optional[int] = None
dim: Optional[int] = None hidden_size: Optional[int] = None
n_layers: Optional[int] = None num_hidden_layers: Optional[int] = None
norm_eps: Optional[float] = None rms_norm_eps: Optional[float] = None
dim_ffn: Optional[int] = None intermediate_size: Optional[int] = None
tie_weight: Optional[bool] = None tie_word_embeddings: Optional[bool] = None
max_position_embeddings: Optional[int] = None
max_len: Optional[int] = None
rope_theta: Optional[float] = None rope_theta: Optional[float] = None
rope_scaling: Optional[dict] = None rope_scaling: Optional[dict] = None
attn_type: str = "gqa" attn_type: str = "gqa"
n_heads: Optional[int] = None num_attention_heads: Optional[int] = None
n_kv_heads: Optional[int] = None num_key_value_heads: Optional[int] = None
use_qk_norm: Optional[bool] = None use_qk_norm: Optional[bool] = None
use_gated_attention: Optional[bool] = None use_gated_attention: Optional[bool] = None
kv_lora_rank: Optional[int] = None kv_lora_rank: Optional[int] = None
qk_nope_head_dim: Optional[int] = None qk_nope_head_dim: Optional[int] = None
qk_rope_head_dim: Optional[int] = None qk_rope_head_dim: Optional[int] = None
ffn_type: str = "mlp" ffn_type: str = "mlp"
n_routed_experts: Optional[int] = None n_routed_experts: Optional[int] = None
n_shared_experts: Optional[int] = None n_shared_experts: Optional[int] = None
n_activated_experts: Optional[int] = None n_activated_experts: Optional[int] = None
topk_method: Optional[str] = None topk_method: Optional[str] = None
moe_intermediate_size: Optional[int] = None
shared_expert_intermediate_size: Optional[int] = None
norm_topk_prob: bool = True
decoder_sparse_step: int = 1
mlp_only_layers: Optional[list[int]] = None
moe_aux_loss_coef: float = 0.01
@field_validator("attn_type")
def _validate_attn_type(cls, v: str) -> str:
if v not in _ATTN_TYPES:
raise ValueError(
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
)
return v
@field_validator("ffn_type")
def _validate_ffn_type(cls, v: str) -> str:
if v not in _FFN_TYPES:
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
return v
@field_validator("decoder_sparse_step")
def _validate_decoder_sparse_step(cls, v: int) -> int:
if v < 1:
raise ValueError(f"decoder_sparse_step must be at least 1, got {v}")
return v
@dataclass @dataclass
@ConfigFactory.register("embedding") @ConfigFactory.register("embedding")
class EncoderConfig(BaseModelConfig): class EncoderConfig(BaseModelConfig):
"""Configuration for embedding encoder model.""" """Configuration for embedding encoder model.
Args:
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
pooling_type (Optional[str]): Pooling strategy for embedding, e.g. 'mean', 'cls'. Defaults to None.
normalize_embeddings (Optional[bool]): Whether to L2-normalize output embeddings. Defaults to None.
"""
vocab_size: Optional[int] = None vocab_size: Optional[int] = None
dim: Optional[int] = None hidden_size: Optional[int] = None
n_layers: Optional[int] = None num_hidden_layers: Optional[int] = None
norm_eps: Optional[float] = None rms_norm_eps: Optional[float] = None
dim_ffn: Optional[int] = None intermediate_size: Optional[int] = None
max_position_embeddings: Optional[int] = None
max_len: Optional[int] = None
rope_theta: Optional[float] = None rope_theta: Optional[float] = None
rope_scaling: Optional[dict] = None rope_scaling: Optional[dict] = None
attn_type: str = "gqa" attn_type: str = "gqa"
n_heads: Optional[int] = None num_attention_heads: Optional[int] = None
n_kv_heads: Optional[int] = None num_key_value_heads: Optional[int] = None
use_qk_norm: Optional[bool] = None use_qk_norm: Optional[bool] = None
use_gated_attention: Optional[bool] = None use_gated_attention: Optional[bool] = None
ffn_type: str = "mlp" ffn_type: str = "mlp"
pooling_type: Optional[str] = None pooling_type: Optional[str] = None
normalize_embeddings: Optional[bool] = None normalize_embeddings: Optional[bool] = None
@field_validator("attn_type")
def _validate_attn_type(cls, v: str) -> str:
if v not in _ATTN_TYPES:
raise ValueError(
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
)
return v
@field_validator("ffn_type")
def _validate_ffn_type(cls, v: str) -> str:
if v not in _FFN_TYPES:
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
return v
+86 -43
View File
@@ -5,11 +5,19 @@ modes, both driven declaratively through ``input.sections`` or
``input.sources``. ``input.sources``.
""" """
from dataclasses import dataclass, field from dataclasses import field
from typing import Dict, List, Optional from typing import Dict, List, Optional
from pydantic import field_validator
from pydantic.dataclasses import dataclass
from astrai.config.base import BaseConfig from astrai.config.base import BaseConfig
_PACKING_STRATEGIES = frozenset({"simple", "bfd", "bfd_split"})
_TRUNCATION_MODES = frozenset({"keep_start", "keep_end"})
_STORAGE_FORMATS = frozenset({"bin", "jsonl"})
_POSITION_IDS_MODES = frozenset({"none", "doc_reset", "continuous"})
@dataclass @dataclass
class InputConfig(BaseConfig): class InputConfig(BaseConfig):
@@ -25,6 +33,10 @@ class InputConfig(BaseConfig):
"chosen": {"sections": [{"field": "chosen", ...}]}, "chosen": {"sections": [{"field": "chosen", ...}]},
"rejected": {"sections": [{"field": "rejected", ...}]}, "rejected": {"sections": [{"field": "rejected", ...}]},
}}} }}}
Args:
sections (Optional[List[Dict]]): Section list for single-output mode. Defaults to None.
sources (Optional[Dict[str, Dict]]): Source map for multi-output mode, DPO/GRPO. Defaults to None.
""" """
sections: Optional[List[Dict]] = None sections: Optional[List[Dict]] = None
@@ -33,63 +45,67 @@ class InputConfig(BaseConfig):
@dataclass @dataclass
class ProcessingConfig(BaseConfig): class ProcessingConfig(BaseConfig):
"""Processing configuration. """Processing configuration for tokenization and packing.
Parameters Args:
---------- max_seq_len (int): Maximum sequence length. Defaults to 2048.
max_seq_len : int min_chars (int): Minimum number of characters to keep. Defaults to 50.
Maximum sequence length (default: 2048). max_chars (int): Maximum number of characters to keep. Defaults to 2_000_000.
min_chars : int max_items (Optional[int]): Maximum number of items to process, None=unlimited. Defaults to None.
Minimum number of characters to keep (default: 50). batch_size (int): Number of records tokenized together. Defaults to 256.
max_chars : int packing_strategy (str): How to pack sequences: 'simple', 'bfd', or 'bfd_split'. Defaults to "simple".
Maximum number of characters to keep (default: 2_000_000). max_packed_len (int): Maximum length of a packed bin. Defaults to 8192.
max_items : Optional[int] truncation_mode (str): How to truncate over-length sequences: 'keep_start' or 'keep_end'. Defaults to "keep_start".
Maximum number of items to process (default: None, unlimited).
packing_strategy : str
How to pack sequences into a contiguous stream.
- ``"simple"``: sequential concatenation (default, backward compatible).
- ``"bfd"``: best-fit decreasing bin packing, minimises wasted tokens.
- ``"bfd_split"``: BFD with over-length sequences split into chunks.
max_packed_len : int
Maximum length of a packed bin. Sequences longer than this are
truncated or split depending on ``packing_strategy`` (default: 8192).
truncation_mode : str
How to truncate sequences longer than ``max_packed_len``.
- ``"keep_start"``: keep the first ``max_packed_len`` tokens (default).
- ``"keep_end"``: keep the last ``max_packed_len`` tokens.
""" """
max_seq_len: int = 2048 max_seq_len: int = 2048
min_chars: int = 50 min_chars: int = 50
max_chars: int = 2_000_000 max_chars: int = 2_000_000
max_items: Optional[int] = None max_items: Optional[int] = None
batch_size: int = 256
packing_strategy: str = "simple" packing_strategy: str = "simple"
max_packed_len: int = 8192 max_packed_len: int = 8192
truncation_mode: str = "keep_start" truncation_mode: str = "keep_start"
@field_validator("packing_strategy")
def _validate_packing_strategy(cls, v: str) -> str:
if v not in _PACKING_STRATEGIES:
raise ValueError(
f"packing_strategy must be one of {sorted(_PACKING_STRATEGIES)}, got {v!r}"
)
return v
@field_validator("truncation_mode")
def _validate_truncation_mode(cls, v: str) -> str:
if v not in _TRUNCATION_MODES:
raise ValueError(
f"truncation_mode must be one of {sorted(_TRUNCATION_MODES)}, got {v!r}"
)
return v
@field_validator("max_seq_len", "batch_size", "max_packed_len")
def _validate_positive_int(cls, v: int) -> int:
if v <= 0:
raise ValueError(f"must be positive, got {v}")
return v
@field_validator("min_chars")
def _validate_non_negative(cls, v: int) -> int:
if v < 0:
raise ValueError(f"min_chars must be non-negative, got {v}")
return v
@dataclass @dataclass
class OutputConfig(BaseConfig): class OutputConfig(BaseConfig):
"""Output configuration. """Output configuration for storage.
Parameters Args:
---------- domain_key (Optional[str]): Domain key for the output store. Defaults to None.
domain_key : Optional[str] storage_format (str): Storage format: 'bin' or 'jsonl'. Defaults to "bin".
Domain key for the output store (default: None). max_tokens_per_shard (int): Maximum tokens per shard before splitting. Defaults to 100_000_000.
storage_format : str dtype (Dict[str, str]): Per-key dtype overrides, e.g. {"input_ids": "int32"}. Defaults to {}.
Storage format, one of ``"bin"``, ``"jsonl"`` (default: ``"bin"``). position_ids_mode (str): Position ids mode: 'none', 'doc_reset', or 'continuous'. Defaults to "doc_reset".
max_tokens_per_shard : int
Maximum tokens per shard before splitting (default: 100_000_000).
dtype : Dict[str, str]
Per-key dtype overrides, e.g. ``{"input_ids": "int32"}`` (default: {}).
position_ids_mode : Optional[str]
How to compute position_ids in packed sequences.
- ``"none"``: do not generate (default).
- ``"doc_reset"``: reset to 0 at each document boundary.
- ``"continuous"``: sequential 0, 1, 2, ... (pretrain, single doc).
""" """
domain_key: Optional[str] = None domain_key: Optional[str] = None
@@ -98,9 +114,36 @@ class OutputConfig(BaseConfig):
dtype: Dict[str, str] = field(default_factory=dict) dtype: Dict[str, str] = field(default_factory=dict)
position_ids_mode: str = "doc_reset" position_ids_mode: str = "doc_reset"
@field_validator("storage_format")
def _validate_storage_format(cls, v: str) -> str:
if v not in _STORAGE_FORMATS:
raise ValueError(
f"storage_format must be one of {sorted(_STORAGE_FORMATS)}, got {v!r}"
)
return v
@field_validator("position_ids_mode")
def _validate_position_ids_mode(cls, v: str) -> str:
if v not in _POSITION_IDS_MODES:
raise ValueError(
f"position_ids_mode must be one of {sorted(_POSITION_IDS_MODES)}, got {v!r}"
)
return v
@dataclass @dataclass
class PipelineConfig(BaseConfig): class PipelineConfig(BaseConfig):
"""Top-level preprocessing pipeline config.
Args:
version (int): Config schema version. Defaults to 1.
input (InputConfig): Input mapping config.
mask (Dict[str, str]): Per-field mask labels, e.g. {"system": "mask", "assistant": "train"}. Defaults to {}.
mask_default (str): Default mask label for unlisted fields. Defaults to "mask".
preprocessing (ProcessingConfig): Processing config.
output (OutputConfig): Output config.
"""
version: int = 1 version: int = 1
input: InputConfig = field(default_factory=InputConfig) input: InputConfig = field(default_factory=InputConfig)
mask: Dict[str, str] = field(default_factory=dict) mask: Dict[str, str] = field(default_factory=dict)
+196 -133
View File
@@ -1,7 +1,9 @@
from dataclasses import dataclass, field, fields from dataclasses import field
from typing import Any, Callable, Dict, List, Optional from typing import Any, Callable, Dict, List, Optional
import torch.nn as nn import torch.nn as nn
from pydantic import ConfigDict, field_validator, model_validator
from pydantic.dataclasses import dataclass
from torch.optim import Optimizer from torch.optim import Optimizer
from torch.optim.lr_scheduler import LRScheduler from torch.optim.lr_scheduler import LRScheduler
from torch.utils.data import Dataset from torch.utils.data import Dataset
@@ -9,147 +11,208 @@ 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"})
def required(**kw): _PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
return {"required": True, **kw} _BACKENDS = frozenset({"nccl", "gloo"})
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
@dataclass @dataclass(config=ConfigDict(arbitrary_types_allowed=True))
class TrainConfig(BaseConfig): class TrainConfig(BaseConfig):
# basic setting """Training configuration.
model_fn: Callable[[], nn.Module] = field(
default=None, metadata=required(help="Model factory for training.")
)
strategy: str = field(default=None, metadata=required(help="Training strategy."))
dataset: Dataset = field(
default=None, metadata=required(help="Dataset for training.")
)
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
default=None, metadata=required(help="Optimizer factory for training.")
)
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
default=None, metadata=required(help="Scheduler factory for training.")
)
n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."})
batch_per_device: int = field(
default=4, metadata={"help": "Batch size per device."}
)
grad_accum_steps: int = field(
default=1, metadata={"help": "Number of iterations between steps."}
)
max_grad_norm: Optional[float] = field(
default=None,
metadata={"help": "Maximum gradient norm. None disables clipping."},
)
gradient_checkpointing_modules: List[str] = field(
default_factory=list,
metadata={"help": "Module types to enable activation checkpointing for."},
)
# checkpoint setting Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."}) Only JSON-serializable fields are written to checkpoint meta via to_dict().
start_samples: int = field(
default=0,
metadata={
"help": "Start samples count (per rank). Superseded by checkpoint consumed_samples."
},
)
ckpt_dir: str = field(
default="./checkpoint", metadata={"help": "Checkpoint directory."}
)
ckpt_interval: int = field(
default=5000,
metadata={"help": "Number of optimizer steps between checkpoints."},
)
# lora setting Args:
lora: Optional[LoRAConfig] = field( model_fn (Callable[[], nn.Module]): Model factory for training.
default=None, strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
metadata={"help": "LoRA config. None means full fine-tuning."}, dataset (Dataset): Dataset for training.
) optimizer_fn (Callable[[nn.Module], Optimizer]): Optimizer factory for training.
optimizer_name (Optional[str]): Serializable built-in optimizer identifier. Defaults to None.
optimizer_hyperparameters (Dict[str, Any]): Serializable optimizer settings. Defaults to {}.
scheduler_fn (Callable[[Optimizer], LRScheduler]): Scheduler factory for training.
n_epoch (int): Number of epochs for training. Defaults to 1.
batch_per_device (int): Batch size per device. Defaults to 4.
grad_accum_steps (int): Number of iterations between optimizer steps. Defaults to 1.
max_grad_norm (Optional[float]): Maximum gradient norm. None disables clipping. Defaults to 1.0.
gradient_checkpointing_modules (List[type]): Module types to enable activation checkpointing for. Defaults to [].
compile_mode (Optional[str]): torch.compile mode: 'default', 'reduce-overhead', 'max-autotune', or None. Defaults to None.
start_epoch (int): Start epoch for training. Defaults to 0.
start_samples (int): Start samples count (per rank). Superseded by checkpoint consumed_samples. Defaults to 0.
ckpt_dir (str): Checkpoint directory. Defaults to "./checkpoint".
ckpt_interval (int): Number of optimizer steps between checkpoints. Defaults to 5000.
lora (Optional[LoRAConfig]): LoRA config. None means full fine-tuning. Defaults to None.
metrics (List[str]): Metrics to record during training. Defaults to ["loss", "lr", "grad_norm"].
random_seed (int): Random seed. Defaults to 3407.
num_workers (int): Number of workers for dataloader. Defaults to 0.
prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None.
pin_memory (bool): Pin memory for dataloader. Defaults to False.
collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None.
nprocs (int): Number of processes for distributed training. Defaults to 1.
backend (str): Distributed training backend. Defaults to "nccl".
master_addr (str): Master address for distributed training. Defaults to "localhost".
master_port (str): Master port for distributed training. Defaults to "29500".
parallel_mode (str): Parallel strategy: none, ddp, fsdp. Defaults to "none".
start_method (str): Multiprocessing start method: spawn/fork/forkserver. Defaults to "spawn".
device_type (str): Device type for distributed training. Defaults to "cuda".
val_dataset (Optional[Dataset]): Dataset for validation. Defaults to None.
val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. Defaults to None.
val_step (int): Number of optimizer steps between validation runs. Defaults to 1000.
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
moe_aux_loss_coef (float): Weight applied to the MoE load-balancing loss. Defaults to 0.01.
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7.
rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
rollout_top_p (float): Top-p (nucleus) filtering for online rollout. Defaults to 0.9.
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.
executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}.
extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}.
"""
# metric setting model_fn: Callable[[], nn.Module]
log_dir: str = field( strategy: str
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."} dataset: Dataset
) optimizer_fn: Callable[[nn.Module], Optimizer]
metrics: List[str] = field( scheduler_fn: Callable[[Optimizer], LRScheduler]
default_factory=lambda: ["loss", "lr", "grad_norm"], optimizer_name: Optional[str] = None
metadata={"help": "Metrics to record during training."}, optimizer_hyperparameters: Dict[str, Any] = field(default_factory=dict)
) n_epoch: int = 1
batch_per_device: int = 4
grad_accum_steps: int = 1
max_grad_norm: Optional[float] = 1.0
gradient_checkpointing_modules: List[type] = field(default_factory=list)
compile_mode: Optional[str] = None
# dataloader setting start_epoch: int = 0
random_seed: int = field(default=3407, metadata={"help": "Random seed."}) start_samples: int = 0
num_workers: int = field( ckpt_dir: str = "./checkpoint"
default=0, metadata={"help": "Number of workers for dataloader."} ckpt_interval: int = 5000
)
prefetch_factor: Optional[int] = field(
default=None, metadata={"help": "Prefetch factor for dataloader."}
)
pin_memory: bool = field(
default=False, metadata={"help": "Pin memory for dataloader."}
)
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
default=None,
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
)
# distributed training lora: Optional[LoRAConfig] = None
nprocs: int = field(
default=1, metadata={"help": "Number of processes for distributed training."}
)
backend: str = field(
default="nccl", metadata={"help": "Distributed training backend."}
)
master_addr: str = field(
default="localhost",
metadata={"help": "Master address for distributed training."},
)
master_port: str = field(
default="29500", metadata={"help": "Master port for distributed training."}
)
parallel_mode: str = field(
default="none",
metadata={"help": "Parallel strategy: none, ddp, fsdp."},
)
start_method: str = field(
default="spawn",
metadata={"help": "Multiprocessing start method (spawn/fork/forkserver)."},
)
# others metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"])
device_type: str = field(
default="cuda", metadata={"help": "Device type for distributed training."}
)
val_dataset: Optional[Dataset] = field(
default=None, metadata={"help": "Dataset for validation."}
)
val_split: Optional[float] = field(
default=None,
metadata={
"help": "Ratio to split from training dataset for validation (e.g. 0.05). Ignored if val_dataset is set."
},
)
val_step: int = field(
default=1000,
metadata={"help": "Number of optimizer steps between validation runs."},
)
neftune_alpha: float = field(
default=0.0,
metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."},
)
executor_kwargs: Dict[str, Any] = field( random_seed: int = 3407
default_factory=dict, num_workers: int = 0
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."}, prefetch_factor: Optional[int] = None
) pin_memory: bool = False
extra_kwargs: Dict[str, Any] = field( collate_fn: Optional[Callable[[List[Any]], Any]] = None
default_factory=dict, metadata={"help": "Other arguments."}
)
def __post_init__(self): nprocs: int = 1
self.validate() backend: str = "nccl"
master_addr: str = "localhost"
master_port: str = "29500"
parallel_mode: str = "none"
start_method: str = "spawn"
def validate(self): device_type: str = "cuda"
for fld in fields(self): val_dataset: Optional[Dataset] = None
if fld.metadata.get("required") and getattr(self, fld.name) is None: val_split: Optional[float] = None
raise ValueError(f"TrainConfig.{fld.name} is required but got None.") val_step: int = 1000
neftune_alpha: float = 0.0
moe_aux_loss_coef: float = 0.01
rollout_interval: int = 512
rollout_temperature: float = 0.7
rollout_top_k: int = 0
rollout_top_p: float = 0.9
rollout_max_tokens: int = 1024
reward_model_fn: Optional[Callable] = None
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
extra_kwargs: Dict[str, Any] = field(default_factory=dict)
@field_validator("strategy")
def _validate_strategy(cls, v: str) -> str:
if v not in _TRAIN_TYPES:
raise ValueError(
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}"
)
return v
@field_validator("parallel_mode")
def _validate_parallel_mode(cls, v: str) -> str:
if v not in _PARALLEL_MODES:
raise ValueError(
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}"
)
return v
@field_validator("backend")
def _validate_backend(cls, v: str) -> str:
if v not in _BACKENDS:
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}")
return v
@field_validator("start_method")
def _validate_start_method(cls, v: str) -> str:
if v not in _START_METHODS:
raise ValueError(
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}"
)
return v
@field_validator("compile_mode")
def _validate_compile_mode(cls, v: Optional[str]) -> Optional[str]:
if v is not None and v not in _COMPILE_MODES:
raise ValueError(
f"compile_mode must be one of {sorted(_COMPILE_MODES)} or None, got {v!r}"
)
return v
@field_validator(
"n_epoch",
"batch_per_device",
"grad_accum_steps",
"ckpt_interval",
"val_step",
"rollout_interval",
"rollout_max_tokens",
)
def _validate_positive_int(cls, v: int) -> int:
if v <= 0:
raise ValueError(f"must be positive, got {v}")
return v
@field_validator("rollout_temperature")
def _validate_positive_float(cls, v: float) -> float:
if v <= 0:
raise ValueError(f"must be positive, got {v}")
return v
@field_validator("rollout_top_p")
def _validate_top_p(cls, v: float) -> float:
if not 0 < v <= 1:
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
return v
@field_validator(
"rollout_top_k", "num_workers", "neftune_alpha", "moe_aux_loss_coef"
)
def _validate_non_negative(cls, v):
if v < 0:
raise ValueError(f"must be non-negative, got {v}")
return v
@field_validator("max_grad_norm")
def _validate_max_grad_norm(cls, v: Optional[float]) -> Optional[float]:
if v is not None and v <= 0:
raise ValueError(f"max_grad_norm must be positive or None, got {v}")
return v
@field_validator("val_split")
def _validate_val_split(cls, v: Optional[float]) -> Optional[float]:
if v is not None and not 0 < v < 1:
raise ValueError(f"val_split must be in (0, 1) or None, got {v}")
return v
@model_validator(mode="after")
def _validate_online_strategy(self) -> "TrainConfig":
if self.strategy.startswith("online_") and self.reward_model_fn is None:
raise ValueError(
f"reward_model_fn is required for online RL strategy {self.strategy!r}"
)
return self
-6
View File
@@ -6,7 +6,6 @@ from astrai.dataset.dataset import (
) )
from astrai.dataset.sampler import RDSampler from astrai.dataset.sampler import RDSampler
from astrai.dataset.storage import ( from astrai.dataset.storage import (
H5Store,
JsonlStore, JsonlStore,
MmapStore, MmapStore,
Recordable, Recordable,
@@ -17,9 +16,7 @@ from astrai.dataset.storage import (
) )
from astrai.serialization import ( from astrai.serialization import (
load_bin, load_bin,
load_h5,
save_bin, save_bin,
save_h5,
) )
__all__ = [ __all__ = [
@@ -31,12 +28,9 @@ __all__ = [
"Streamable", "Streamable",
"Recordable", "Recordable",
"StoreFactory", "StoreFactory",
"H5Store",
"MmapStore", "MmapStore",
"JsonlStore", "JsonlStore",
"detect_format", "detect_format",
"save_h5",
"load_h5",
"save_bin", "save_bin",
"load_bin", "load_bin",
"RDSampler", "RDSampler",
+18 -6
View File
@@ -190,7 +190,8 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
- rewards: [G] - rewards: [G]
Output: Output:
- prompts: [B, P_max] - prompts: [B, P_max], left-padded
- prompt_mask: [B, P_max]
- responses: [B, G, R_max] - responses: [B, G, R_max]
- masks: [B, G, R_max] - masks: [B, G, R_max]
- rewards: [B, G] - rewards: [B, G]
@@ -201,13 +202,15 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
R_max = max(r.size(0) for b in batch for r in b["responses"]) R_max = max(r.size(0) for b in batch for r in b["responses"])
prompts = torch.zeros(B, P_max, dtype=torch.long) prompts = torch.zeros(B, P_max, dtype=torch.long)
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
responses = torch.zeros(B, G, R_max, dtype=torch.long) responses = torch.zeros(B, G, R_max, dtype=torch.long)
masks = torch.zeros(B, G, R_max, dtype=torch.bool) masks = torch.zeros(B, G, R_max, dtype=torch.bool)
rewards = torch.zeros(B, G, dtype=torch.float32) rewards = torch.zeros(B, G, dtype=torch.float32)
for i, b in enumerate(batch): for i, b in enumerate(batch):
p_len = b["prompts"].size(0) p_len = b["prompts"].size(0)
prompts[i, :p_len] = b["prompts"] prompts[i, -p_len:] = b["prompts"]
prompt_mask[i, -p_len:] = True
rewards[i, : b["rewards"].size(0)] = b["rewards"] rewards[i, : b["rewards"].size(0)] = b["rewards"]
for g in range(min(G, len(b["responses"]))): for g in range(min(G, len(b["responses"]))):
r_len = b["responses"][g].size(0) r_len = b["responses"][g].size(0)
@@ -217,6 +220,7 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
return { return {
"prompts": prompts, "prompts": prompts,
"prompt_mask": prompt_mask,
"responses": responses, "responses": responses,
"masks": masks, "masks": masks,
"rewards": rewards, "rewards": rewards,
@@ -310,7 +314,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
stream datasets (SEQ/SFT). Record datasets ignore it. stream datasets (SEQ/SFT). Record datasets ignore it.
stride: Stride between consecutive stream samples stride: Stride between consecutive stream samples
(default: same as *window_size*). (default: same as *window_size*).
storage_type: Storage backend ("h5", "bin", "jsonl") or storage_type: Storage backend ("bin", "jsonl") or
None for auto-detection. None for auto-detection.
tokenizer_path: Path to tokenizer for lazy JSONL tokenizer_path: Path to tokenizer for lazy JSONL
tokenisation (record datasets only). tokenisation (record datasets only).
@@ -346,7 +350,15 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
if processor is not None: if processor is not None:
store.load(load_path, processor=processor, **kwargs) store.load(load_path, processor=processor, **kwargs)
else: else:
store.load(load_path, **kwargs) load_kwargs = dict(kwargs)
if (
tokenizer_path is not None
and storage_type == "jsonl"
and train_type in ("seq", "sft")
and "tokenizer_path" not in load_kwargs
):
load_kwargs["tokenizer_path"] = tokenizer_path
store.load(load_path, **load_kwargs)
return cls.create(train_type, store=store) return cls.create(train_type, store=store)
@@ -372,7 +384,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
"""Build an on-the-fly tokenisation processor if applicable. """Build an on-the-fly tokenisation processor if applicable.
Only raw JSONL + record datasets (DPO/GRPO) need a processor; Only raw JSONL + record datasets (DPO/GRPO) need a processor;
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT) pre-tokenised backends (bin) and stream datasets (SEQ/SFT)
return ``None`` so no tokenizer is loaded. return ``None`` so no tokenizer is loaded.
""" """
if tokenizer_path is None or storage_type != "jsonl": if tokenizer_path is None or storage_type != "jsonl":
@@ -439,7 +451,7 @@ class DPODataset(BaseDataset):
Two loading paths (handled by :class:`DatasetFactory`): Two loading paths (handled by :class:`DatasetFactory`):
- **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record - **Pre-tokenized** (bin): ``store.load(path)`` reads per-record
tensors; ``__getitem__`` returns them directly. tensors; ``__getitem__`` returns them directly.
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor - **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
via :func:`dpo_processor` that tokenises on the fly — no packing, via :func:`dpo_processor` that tokenises on the fly — no packing,
+45 -55
View File
@@ -10,7 +10,6 @@ Architecture (composition over inheritance):
Streamable (mixin) — raw token slice fetch(begin, end, keys) Streamable (mixin) — raw token slice fetch(begin, end, keys)
Recordable (mixin) — raw record slice fetch_record(idx, keys) Recordable (mixin) — raw record slice fetch_record(idx, keys)
H5Store(Store, Streamable, Recordable)
MmapStore(Store, Streamable, Recordable) MmapStore(Store, Streamable, Recordable)
JsonlStore(Store, Streamable, Recordable) JsonlStore(Store, Streamable, Recordable)
@@ -36,9 +35,9 @@ control. ``store.token_count`` is the total stream token count (what
``len(store)`` used to mean in the legacy stream-only API). ``len(store)`` used to mean in the legacy stream-only API).
``segments_are_records`` (class attribute on each Store subclass) ``segments_are_records`` (class attribute on each Store subclass)
tells ``_normalize`` whether segments are inherently per-record (H5/ tells ``_normalize`` whether segments are inherently per-record (JSONL)
JSONL) or opaque shards (bin). Record access for bin relies on or opaque shards (bin). Record access for bin relies on ``_offsets``
``_offsets`` instead. instead.
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps :class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
raw records and defers tokenisation to ``fetch_record`` — used by DPO raw records and defers tokenisation to ``fetch_record`` — used by DPO
@@ -56,12 +55,12 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.config.preprocess_config import PipelineConfig
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.preprocessing.transform import TokenizeTransform from astrai.preprocessing.transform import TokenizeTransform
from astrai.serialization import ( from astrai.serialization import (
load_bin, load_bin,
load_bin_offsets, load_bin_offsets,
load_h5,
) )
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -82,19 +81,10 @@ def detect_format(load_path: str) -> str:
root = Path(load_path) root = Path(load_path)
if root.is_file(): if root.is_file():
suffix = root.suffix.lower() suffix = root.suffix.lower()
if suffix in (".h5", ".hdf5"):
return "h5"
if suffix == ".jsonl": if suffix == ".jsonl":
return "jsonl" return "jsonl"
raise ValueError(f"Unsupported file format: {suffix}") raise ValueError(f"Unsupported file format: {suffix}")
h5_files = [
Path(p)
for pattern in ("*.h5", "*.hdf5")
for p in glob.glob(str(root / "**" / pattern), recursive=True)
]
if h5_files:
return "h5"
bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)] bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)]
if bin_files: if bin_files:
has_meta = (root / "meta.json").exists() or len( has_meta = (root / "meta.json").exists() or len(
@@ -184,7 +174,7 @@ class Store(ABC):
"""Number of records available via :meth:`fetch_record`. """Number of records available via :meth:`fetch_record`.
Non-zero only when the backing layout provides per-record Non-zero only when the backing layout provides per-record
indexing (H5/JSONL segments or bin ``_offsets``). indexing (JSONL segments or bin ``_offsets``).
""" """
return self._num_records return self._num_records
@@ -268,7 +258,7 @@ class Store(ABC):
Record mode: if *offsets* is provided (bin layout), Record mode: if *offsets* is provided (bin layout),
``_offsets[key]`` stores cumulative per-record offsets into the ``_offsets[key]`` stores cumulative per-record offsets into the
single concatenated segment. Otherwise, when single concatenated segment. Otherwise, when
``segments_are_records`` is True (H5/JSONL), ``_data[key]`` is ``segments_are_records`` is True (JSONL), ``_data[key]`` is
a per-record list and ``fetch_record`` indexes it directly. a per-record list and ``fetch_record`` indexes it directly.
Nested keys (GRPO ``responses``/``masks`` as Nested keys (GRPO ``responses``/``masks`` as
@@ -304,7 +294,7 @@ class Store(ABC):
logger.warning( logger.warning(
"Key '%s' has %d segments with offsets — record mode " "Key '%s' has %d segments with offsets — record mode "
"disabled for this key (multi-shard bin+offsets not " "disabled for this key (multi-shard bin+offsets not "
"supported). Merge shards or use H5/JSONL.", "supported). Merge shards or use JSONL.",
key, key,
len(segs), len(segs),
) )
@@ -329,7 +319,7 @@ class Streamable:
Stateless trait relying on ``self._data``, ``self._cum``, Stateless trait relying on ``self._data``, ``self._cum``,
``self._length`` maintained by :class:`Store`. Stream mode is ``self._length`` maintained by :class:`Store`. Stream mode is
active when the owning store has ``window_size > 0``; for stores active when the owning store has ``window_size > 0``; for stores
that can also serve record access (H5/JSONL/bin+offsets), the that can also serve record access (JSONL/bin+offsets), the
``fetch_record`` API from :class:`Recordable` is used instead. ``fetch_record`` API from :class:`Recordable` is used instead.
""" """
@@ -414,33 +404,6 @@ class StoreFactory(BaseFactory["Store"]):
"""Factory for creating Store instances by type name.""" """Factory for creating Store instances by type name."""
@StoreFactory.register("h5")
class H5Store(Store, Streamable, Recordable):
"""HDF5-based storage backend (pre-tokenized data).
Each key is stored as a group of per-record datasets (``data_0``,
``data_1``, …). Supports both access modes:
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
across concatenated records via ``_cum`` — used by SEQ/SFT.
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
``window_size == 0``) index ``_data[key]`` directly — used by
DPO/GRPO.
"""
segments_are_records = True
def __init__(
self,
window_size: int = 0,
stride: Optional[int] = None,
):
super().__init__(window_size=window_size, stride=stride)
def load(self, path: str, **kwargs):
self._normalize(load_h5(path))
@StoreFactory.register("bin") @StoreFactory.register("bin")
class MmapStore(Store, Streamable, Recordable): class MmapStore(Store, Streamable, Recordable):
"""Memory-mapped binary storage backend. """Memory-mapped binary storage backend.
@@ -545,18 +508,29 @@ class JsonlSource:
@StoreFactory.register("jsonl") @StoreFactory.register("jsonl")
class JsonlStore(Store, Streamable, Recordable): class JsonlStore(Store, Streamable, Recordable):
"""JSONL reader with two tokenisation modes. """JSONL reader with eager/lazy tokenisation modes.
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl`` A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
files plus (optionally) a ``dataset_config.json`` describing the files plus (optionally) a ``dataset_config.json`` describing the
tokenization pipeline. tokenization pipeline.
Two modes, selected at :meth:`load` time: Three ways to supply an eager transform (first match wins):
- **Eager** (default): applies a :class:`TokenizeTransform` to every - **Explicit** (``transform=``): caller-built
record at load time and registers per-key tensors via :class:`TokenizeTransform` applied eagerly.
``_normalize``. Both ``fetch`` (stream) and ``fetch_record`` - **Config file**: ``dataset_config.json`` alongside the ``*.jsonl``
(record) work. files — loaded via :meth:`TokenizeTransform.from_config_file`.
- **Default messages** (``tokenizer_path=`` given, no config file):
a built-in chatml config that tokenises the ``messages`` field,
masking every role except ``assistant`` (loss on assistant only).
Lets SFT/SEQ train straight from a chat-style JSONL directory
without a hand-written config.
Two tokenisation modes, selected at :meth:`load` time:
- **Eager** (default): applies the transform to every record at load
time and registers per-key tensors via ``_normalize``. Both
``fetch`` (stream) and ``fetch_record`` (record) work.
- **Lazy** (``processor=fn`` passed): keeps raw records and defers - **Lazy** (``processor=fn`` passed): keeps raw records and defers
tokenisation to ``fetch_record``. Only record access works — tokenisation to ``fetch_record``. Only record access works —
``len(store)`` returns ``num_records``; stream primitives raise. ``len(store)`` returns ``num_records``; stream primitives raise.
@@ -565,6 +539,16 @@ class JsonlStore(Store, Streamable, Recordable):
CONFIG_NAME = "dataset_config.json" CONFIG_NAME = "dataset_config.json"
segments_are_records = True segments_are_records = True
_DEFAULT_MESSAGES_CONFIG = {
"version": 1,
"input": {
"sections": [{"field": "messages", "action": "$role", "template": True}]
},
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
"mask_default": "mask",
"output": {"position_ids_mode": "doc_reset"},
}
def __init__( def __init__(
self, self,
window_size: int = 0, window_size: int = 0,
@@ -587,14 +571,20 @@ class JsonlStore(Store, Streamable, Recordable):
if transform is None: if transform is None:
root = Path(path) root = Path(path)
config_path = root / self.CONFIG_NAME if root.is_dir() else None config_path = root / self.CONFIG_NAME if root.is_dir() else None
if config_path is None or not config_path.exists(): if config_path is not None and config_path.exists():
transform = TokenizeTransform.from_config_file(str(config_path))
else:
tokenizer_path = kwargs.get("tokenizer_path")
if not tokenizer_path:
raise FileNotFoundError( raise FileNotFoundError(
f"JSONL dataset config not found. Expected " f"JSONL dataset config not found. Expected "
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an " f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
f"explicit transform, or pass processor= for lazy " f"explicit transform, pass processor= for lazy "
f"on-the-fly tokenisation." f"on-the-fly tokenisation, or pass tokenizer_path= to "
f"use the built-in messages config."
) )
transform = TokenizeTransform.from_config_file(str(config_path)) config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
transform = TokenizeTransform(config, tokenizer_path)
transformed = transform.apply(records) transformed = transform.apply(records)
self._normalize(transformed) self._normalize(transformed)
+36 -10
View File
@@ -4,26 +4,52 @@ Public API:
- ``attn_decode`` — single-query decode attention - ``attn_decode`` — single-query decode attention
- ``attn_prefill`` — multi-query prefill attention - ``attn_prefill`` — multi-query prefill attention
- ``attn_paged_decode`` — paged decode attention (direct page-table access) - ``attn_paged_decode`` — paged decode attention (direct page-table access)
- ``AttentionBackend`` — ABC for attention computation strategies
- ``TorchNativeBackend`` — default SDPA backend with KV cache I/O
- ``CudaBackend`` — CUDA kernel backend with paged decode + prefill
Interface (shared by all wrappers): Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token (blhd). Scale is always ``1/sqrt(head_dim)``.
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
layout: "bhld" (default) or "blhd"
Causal and mask can coexist — both are applied simultaneously. Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
SDPA is handled by the attention backend, not the wrapper functions.
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
""" """
from astrai.extension.attention_backend import (
ATTN_BACKEND,
AttentionBackend,
AttentionBackendFactory,
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
attention,
attn_backend,
get_backend,
)
from astrai.extension.attention_ops import (
TensorLayout,
attn_decode,
attn_paged_decode,
attn_prefill,
)
from astrai.extension.loader import KERNEL_NAMES, is_available from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.ops import attn_decode, attn_paged_decode, attn_prefill from astrai.extension.rotary_backend import apply_rotary_emb
__all__ = [ __all__ = [
"ATTN_BACKEND",
"AttentionBackend",
"AttentionBackendFactory",
"CudaBackend",
"TorchNativeBackend",
"FlashAttnBackend",
"TensorLayout",
"attention",
"attn_backend",
"get_backend",
"attn_decode", "attn_decode",
"attn_paged_decode", "attn_paged_decode",
"attn_prefill", "attn_prefill",
"is_available", "is_available",
"KERNEL_NAMES", "KERNEL_NAMES",
"apply_rotary_emb",
] ]
+570
View File
@@ -0,0 +1,570 @@
"""Attention backend abstraction with context-manager switching.
The backend encapsulates KV cache I/O and attention computation. The
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
and output projection; the backend handles everything from "write K/V
to cache" through "SDPA output".
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
from astrai.extension import attn_backend, ATTN_BACKEND
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
engine.generate("hello")
# or with an instance:
with attn_backend(TorchNativeBackend()):
...
# or the shorthand (instance is itself a context manager):
with TorchNativeBackend():
...
Thread-safe via ``contextvars`` — each scheduler thread gets its own
active backend. ``get_backend()`` returns the active one, falling back
to a process-wide ``TorchNativeBackend`` singleton.
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
"""
import contextvars
import enum
import importlib
import threading
from abc import ABC, abstractmethod
from contextlib import contextmanager
from typing import Optional, Union
import torch
import torch.nn.functional as F
from torch import Tensor
from astrai.extension.attention_ops import (
attn_paged_decode,
attn_paged_prefill,
)
from astrai.factory import BaseFactory
from astrai.inference.core.cache import KVCache
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
"attn_backend"
)
_lock = threading.Lock()
_flash_available: Optional[bool] = None
def flash_attn_available() -> bool:
"""Return ``True`` if the optional ``flash-attn`` package is usable.
``flash-attn`` is not a hard dependency (declared only as an optional
extra and imported lazily), so this is checked at first use and cached.
The check is stronger than "import works": it also gates on the GPU
compute capability for the installed major version and smoke-tests a
real tiny kernel call, because wheels that import fine can still fail
at the first actual invocation (wrong arch build, torch mismatch, or a
missing ``flash_attn_func`` entry point). It never raises.
"""
global _flash_available
if _flash_available is None:
with _lock:
if _flash_available is None:
_flash_available = _flash_attn_check()
return _flash_available
_flash_attn_module = None
_flash_attn_import_tried = False
def _get_flash_attn():
"""Lazily import and cache the optional ``flash_attn`` module.
Uses ``importlib.import_module`` so no static import binds the name when
the package is absent. Returns the module object, or ``None`` if the
package is not installed or cannot be imported. Never raises.
"""
global _flash_attn_module, _flash_attn_import_tried
if not _flash_attn_import_tried:
_flash_attn_import_tried = True
try:
_flash_attn_module = importlib.import_module("flash_attn")
except Exception:
_flash_attn_module = None
return _flash_attn_module
def _flash_attn_check() -> bool:
if not torch.cuda.is_available():
return False
fa = _get_flash_attn()
if fa is None:
return False
# version + compute-capability gate:
# FlashAttention-2 kernels need sm_70+; FlashAttention-3 (tcgen05,
# sm_90/sm_100) needs sm_90+.
try:
major = int(fa.__version__.split(".")[0])
cc = torch.cuda.get_device_capability()
cc_num = cc[0] * 10 + cc[1]
except Exception:
major, cc_num = 0, 0
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
return False
# smoke-test the real kernel: a wheel that imports but was built for a
# different arch/torch fails here instead of at the first real forward.
try:
if not hasattr(fa, "flash_attn_func"):
return False
x = torch.zeros(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16)
out = fa.flash_attn_func(x, x, x, causal=True)
return bool(torch.isfinite(out).all().item())
except Exception:
return False
class ATTN_BACKEND(enum.Enum):
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
TORCH_NATIVE = "torch_native"
CUDA = "cuda"
FLASH = "flash"
def get_backend() -> "AttentionBackend":
"""Return the active backend for the current thread/context.
Falls back to a ``TorchNativeBackend`` singleton when no backend
has been activated via ``with``.
"""
try:
return _current_backend.get()
except LookupError:
return _default_backend
@contextmanager
def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
"""Context manager to select an attention backend.
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
registered name, ``ATTN_BACKEND`` enum value, backend class, or instance.
Examples::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
...
with attn_backend(TorchNativeBackend):
...
with attn_backend(TorchNativeBackend()):
...
"""
if isinstance(backend, ATTN_BACKEND):
instance = AttentionBackendFactory.create(backend.value)
elif isinstance(backend, str):
instance = AttentionBackendFactory.create(backend)
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
instance = backend()
elif isinstance(backend, AttentionBackend):
instance = backend
else:
raise TypeError(
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
f"or instance, "
f"got {type(backend).__name__}"
)
token = _current_backend.set(instance)
try:
yield instance
finally:
_current_backend.reset(token)
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
"""Expand KV heads to match Q heads for GQA."""
bs, slen, n_heads, head_dim = x.shape
if n_rep == 1:
return x
return (
x[:, :, :, None, :]
.expand(bs, slen, n_heads, n_rep, head_dim)
.reshape(bs, slen, n_heads * n_rep, head_dim)
)
def attention(
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache] = None,
layer_id: int = 0,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
Delegates to the active backend (set via ``with attn_backend(...)``).
Handles KV cache I/O, GQA head expansion, and causal masking so the
caller only needs to provide projected q/k/v.
Args:
q: [batch, q_len, n_heads, head_dim] (blhd)
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
kv_cache: cache dataclass, or None for training (no cache).
layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask (SDPA-compatible).
is_causal: whether to apply causal masking.
Returns:
[batch, q_len, n_heads * head_dim]
"""
backend = get_backend()
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
class AttentionBackend(ABC):
"""Abstract base for attention computation strategies.
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
``fwd_prefill`` (q_len > 1, with or without cache). The public
``forward`` method dispatches based on q_len.
Three equivalent ways to activate a backend::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
...
with attn_backend(TorchNativeBackend): # class
...
with TorchNativeBackend(): # instance
...
"""
def __enter__(self) -> "AttentionBackend":
self._token = _current_backend.set(self)
return self
def __exit__(self, *exc) -> None:
_current_backend.reset(self._token)
def forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Dispatch to decode or extend based on q_len.
Args:
q: [batch, q_len, n_heads, head_dim]
k: [batch, q_len, n_kv_heads, head_dim]
v: [batch, q_len, n_kv_heads, head_dim]
kv_cache: cache dataclass, or None for training (no cache).
layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask compatible with SDPA.
is_causal: whether to apply causal masking.
Returns:
[batch, q_len, n_heads * head_dim]
"""
if kv_cache is not None and q.size(1) == 1:
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
@abstractmethod
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Single-token decode with KV cache."""
@abstractmethod
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Multi-token prefill or training forward."""
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
"""Factory for registered attention backends."""
@AttentionBackendFactory.register(ATTN_BACKEND.TORCH_NATIVE.value)
class TorchNativeBackend(AttentionBackend):
"""Reference backend using torch SDPA with indirect KV cache indexing.
Writes new K/V into the cache buffers, gathers the full sequence K/V
via ``req_to_token`` indirect indexing, then calls
``F.scaled_dot_product_attention``.
For training (``kv_cache is None``), skips cache I/O entirely and
runs SDPA directly on the projected q/k/v.
"""
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def _forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is not None:
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]
# Zero out padding positions so gather never touches invalid slots.
# Decode: attn_mask[:,0,0] is exactly the per-position validity
# mask ([B, max_len], True=keep). Prefill: fall back to seq_lens.
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))
k = kv_cache.k_buffer[layer_id, indices]
v = kv_cache.v_buffer[layer_id, indices]
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
out = F.scaled_dot_product_attention(
q.permute(0, 2, 1, 3),
k.permute(0, 2, 1, 3),
v.permute(0, 2, 1, 3),
attn_mask,
is_causal=is_causal,
)
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
return out
_default_backend = TorchNativeBackend()
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
class CudaBackend(AttentionBackend):
"""CUDA kernel backend with direct KV cache access.
Decode path: writes K/V to the flat pool, then calls
``attn_paged_decode`` with req_to_token + kv_indptr.
Prefill path: writes K/V to the flat pool, then calls
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
kv_indptr.
``kv_cache is None`` (training) is not handled — use
``TorchNativeBackend`` for training.
Raises ``RuntimeError`` if the required kernel is not available.
"""
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
b = q.size(0)
q_3d = q.squeeze(1)
kv_indptr = kv_cache.kv_indptr
out = attn_paged_decode(
q_3d,
kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id],
kv_cache.req_to_token,
kv_cache.req_pool_indices,
kv_indptr,
kv_cache.max_len,
mask=attn_mask,
is_causal=is_causal,
)
return out.unsqueeze(1).flatten(2)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
b = q.size(0)
q_len = q.size(1)
kv_indptr = kv_cache.kv_indptr
qo_indptr = kv_cache.qo_indptr
q_flat = q.reshape(b * q_len, q.size(2), q.size(3))
out = attn_paged_prefill(
q_flat,
kv_cache.k_buffer[layer_id],
kv_cache.v_buffer[layer_id],
kv_cache.req_to_token,
kv_cache.req_pool_indices,
kv_indptr,
qo_indptr,
attn_mask,
q_len,
is_causal=is_causal,
)
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
class FlashAttnBackend(AttentionBackend):
"""FlashAttention (FA2/FA3) backend via the optional ``flash-attn`` package.
Uses the general ``flash_attn_func`` entry point for both prefill and
single-token decode, mirroring ``TorchNativeBackend``'s KV-cache gather.
This backend only does flash attention — inputs ``flash-attn`` cannot
express (missing package, custom attention mask, fp32, unsupported
head_dim) raise a clear error instead of silently falling back to torch.
For a torch fallback, select ``TorchNativeBackend`` instead.
"""
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def fwd_prefill(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def _forward(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
if kv_cache is not None:
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))
k = kv_cache.k_buffer[layer_id, indices]
v = kv_cache.v_buffer[layer_id, indices]
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
if attn_mask is not None and not is_causal:
raise ValueError(
"FlashAttnBackend does not support a custom attention mask; "
"use a causal mask or select TorchNativeBackend."
)
fa = _get_flash_attn()
if fa is None:
raise RuntimeError(
"FlashAttnBackend requires the optional 'flash-attn' package. "
"Install with `pip install flash-attn`."
)
out = fa.flash_attn_func(
q.contiguous(), k.contiguous(), v.contiguous(), causal=is_causal
)
return out.contiguous().flatten(2)
+185
View File
@@ -0,0 +1,185 @@
"""Attention kernel wrapper functions — one entry point per compiled kernel.
Each wrapper calls its CUDA kernel directly. If the kernel is not
available, raises ``RuntimeError``. Fallback to torch SDPA is the
responsibility of the attention backend, not this module.
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). Scale is always ``1/sqrt(head_dim)``.
Interface (all functions):
is_causal: True = causal mask; False = non-causal
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
"""
import enum
from typing import Optional
import torch
from astrai.extension.loader import _available, _modules
class TensorLayout(enum.IntEnum):
"""Q/K/V tensor layout, mirrors the C++ ``TensorLayout`` enum in ``attn_common.h``.
Kernels internally operate on BHLD; BLHD inputs are transposed at entry.
"""
BHLD = 0 # [batch, n_heads, seq_len, 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(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
) -> torch.Tensor:
"""GQA decode attention (q_len == 1).
Args:
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
is_causal: apply causal mask
Returns:
[batch, 1, n_heads, head_dim] (blhd, bf16)
"""
_check_available("attn_decode")
causal_offset = (k.size(1) - 1) if is_causal else -1
return _modules["attn_decode"].attn_decode(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
)
def attn_prefill(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
) -> torch.Tensor:
"""GQA prefill attention (q_len > 1).
Args:
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
is_causal: apply causal mask
Returns:
[batch, q_len, n_heads, head_dim] (blhd, bf16)
"""
_check_available("attn_prefill")
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
return _modules["attn_prefill"].attn_prefill(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
)
def attn_paged_decode(
q: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
kv_indptr: torch.Tensor,
max_seq_len: int,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
) -> torch.Tensor:
"""SGLang-style paged decode (q_len == 1, flat KV pool).
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
req_to_token indirect indexing. Each request has its own seq_len
(from kv_indptr), eliminating padding waste.
Args:
q: [batch, n_heads, head_dim] (bf16, 3D — no seq dim)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int64) — token -> slot
req_pool_indices: [batch] (int64) — rows into req_to_token
kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens
max_seq_len: max per-request seq_len (Python int, for split computation)
mask: 2D [batch, max_seq_len] (bool, True=keep) or None
is_causal: apply causal mask
Returns:
[batch, n_heads, head_dim] (bf16, 3D)
"""
_check_available("attn_paged_decode")
causal_offset = 0 if is_causal else -1
return _modules["attn_paged_decode"].attn_paged_decode(
q,
k_cache,
v_cache,
req_to_token,
req_pool_indices,
kv_indptr,
max_seq_len,
mask=mask,
causal_offset=causal_offset,
)
def attn_paged_prefill(
q: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
kv_indptr: torch.Tensor,
qo_indptr: torch.Tensor,
mask: Optional[torch.Tensor] = None,
max_q_len: int = 0,
is_causal: bool = False,
) -> torch.Tensor:
"""SGLang-style paged prefill (ragged batch, flat KV pool).
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
req_to_token. Supports ragged batches: each request has its own
q_len and kv_len, addressed via qo_indptr and kv_indptr.
Args:
q: [total_q, n_heads, head_dim] (bf16, 3D — flattened across requests)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int64)
req_pool_indices: [batch] (int64)
kv_indptr: [batch+1] (int32) — prefix sum of per-request kv_lens
qo_indptr: [batch+1] (int32) — prefix sum of per-request q_lens
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
max_q_len: max per-request q_len (Python int, for grid computation)
is_causal: apply causal mask
Returns:
[total_q, n_heads, head_dim] (bf16, 3D)
"""
_check_available("attn_paged_prefill")
causal_offset = 0 if is_causal else -1
return _modules["attn_paged_prefill"].attn_paged_prefill(
q,
k_cache,
v_cache,
req_to_token,
req_pool_indices,
kv_indptr,
qo_indptr,
mask,
max_q_len,
causal_offset=causal_offset,
)
+1
View File
@@ -0,0 +1 @@
"""Compiled CUDA kernel modules (``*.so``) live here, kept separate from Python source."""
+8 -2
View File
@@ -11,14 +11,20 @@ import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode"] KERNEL_NAMES = [
"attn_decode",
"attn_prefill",
"attn_paged_decode",
"attn_paged_prefill",
"rotary_emb",
]
_available: dict[str, bool] = {} _available: dict[str, bool] = {}
_modules: dict[str, object] = {} _modules: dict[str, object] = {}
for _name in KERNEL_NAMES: for _name in KERNEL_NAMES:
try: try:
_mod = importlib.import_module(f".{_name}", package=__package__) _mod = importlib.import_module(f".lib.{_name}", package=__package__)
_available[_name] = True _available[_name] = True
_modules[_name] = _mod _modules[_name] = _mod
except ImportError: except ImportError:
-246
View File
@@ -1,246 +0,0 @@
"""GQA attention wrapper functions — one entry point per compiled kernel.
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
available, otherwise falls back to ``torch`` SDPA.
Interface (all functions):
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
layout: "bhld" (default) or "blhd"
Add new kernel wrappers here; split into per-variant files only if this file
grows large.
"""
import math
import torch
import torch.nn.functional as F
from astrai.extension.loader import _available, _modules
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
def _parse_layout(layout: str | int) -> int:
if isinstance(layout, int):
return layout
code = _LAYOUT_CODES.get(layout.lower())
if code is None:
raise ValueError(
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
)
return code
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
if layout == 1:
return t.transpose(1, 2)
return t
def _expand_kv_heads(
k: torch.Tensor, v: torch.Tensor, q_head: int
) -> tuple[torch.Tensor, torch.Tensor]:
"""Expand K/V heads to match Q heads for GQA fallback."""
kv_head = k.size(1)
if kv_head == q_head:
return k, v
group = q_head // kv_head
k = k.repeat_interleave(group, dim=1)
v = v.repeat_interleave(group, dim=1)
return k, v
def _build_attn_mask(
q: torch.Tensor,
k: torch.Tensor,
mask: torch.Tensor | None,
causal_offset: int,
scale: float,
) -> tuple[torch.Tensor | None, float]:
"""Build SDPA-compatible attn_mask + resolved scale.
q and k must already be in b h l d layout.
Causal and mask can coexist: causal sets -inf above the diagonal, mask
sets -inf for padded positions. Both are OR'd into a single bool mask.
"""
q_len = q.size(2)
kv_len = k.size(2)
head_dim = q.size(3)
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
attn_mask = None
if mask is not None:
if mask.dim() == 2:
# [batch, kv_len] → [batch, 1, 1, kv_len]
attn_mask = mask[:, None, None, :]
elif mask.dim() == 3:
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
attn_mask = mask[:, None, :, :]
else:
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
if causal_offset >= 0:
batch = q.size(0)
# q row i attends to kv cols 0..(causal_offset + i)
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
causal_mask = causal_bool.unsqueeze(0).expand(
batch, -1, -1
) # [batch, q_len, kv_len]
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
if attn_mask is not None:
attn_mask = attn_mask | causal_mask
else:
attn_mask = causal_mask
return attn_mask, resolved_scale
def _torch_fallback(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None,
causal_offset: int,
scale: float,
q_layout: int,
kv_layout: int | None = None,
) -> torch.Tensor:
"""Reference attention via ``scaled_dot_product_attention``.
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
"""
if kv_layout is None:
kv_layout = q_layout
q = _to_bhld(q, q_layout)
k = _to_bhld(k, kv_layout)
v = _to_bhld(v, kv_layout)
k, v = _expand_kv_heads(k, v, q.size(1))
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
)
# Restore Q's original layout
if q_layout == 1:
out = out.transpose(1, 2)
return out
def _gather_kv_from_pages(
page_table: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
page_size: int,
kv_len: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
Shapes:
page_table : [batch, max_pages] (int64)
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
v_cache : same as k_cache
Returns:
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
"""
batch, max_pages = page_table.shape
_, ps, n_kv_heads, head_dim = k_cache.shape
if ps != page_size:
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
# Vectorized gather: build physical page + offset indices, then advanced-index
positions = torch.arange(kv_len, device=page_table.device)
logical_pages = positions // page_size # [kv_len]
page_offsets = positions % page_size # [kv_len]
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
k = k_cache[phys_pages, page_offsets]
v = v_cache[phys_pages, page_offsets]
return k, v
def attn_decode(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_decode"]:
return _modules["attn_decode"].attn_decode(
q,
k,
v,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_prefill(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_prefill"]:
return _modules["attn_prefill"].attn_prefill(
q,
k,
v,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_paged_decode(
q: torch.Tensor,
page_table: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
page_size: int,
kv_len: int,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_paged_decode"]:
return _modules["attn_paged_decode"].attn_paged_decode(
q,
page_table,
k_cache,
v_cache,
page_size,
kv_len,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
# Gathered K/V are always b l h d
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
return _torch_fallback(
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
)
+54
View File
@@ -0,0 +1,54 @@
"""Rotary embedding with auto-dispatch to CUDA kernel.
Single entry point ``apply_rotary_emb(x, freqs_cis)`` — uses the fused
CUDA kernel when available, falls back to torch complex multiply otherwise.
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16).
freqs_cis is [batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
"""
import torch
from torch import Tensor
from astrai.extension.loader import is_available
_cache = {"available": None}
def _cuda_available() -> bool:
if _cache["available"] is None:
_cache["available"] = is_available("rotary_emb")
return _cache["available"]
def _torch_apply(x: Tensor, freqs_cis: Tensor) -> Tensor:
cos, sin = freqs_cis[..., 0], freqs_cis[..., 1]
dtype = x.dtype
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
x_complex = torch.view_as_complex(x_)
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(2)
x_rotated = x_complex * freqs_cis_complex
x_out = torch.view_as_real(x_rotated).flatten(-2)
return x_out.to(dtype)
def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
"""Apply rotary embedding to x.
Args:
x: [batch, seq_len, n_heads, head_dim] (bf16)
freqs_cis: [batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs
Returns:
[batch, seq_len, n_heads, head_dim] (bf16)
"""
if (
_cuda_available()
and not torch.is_grad_enabled()
and x.is_cuda
and x.dtype == torch.bfloat16
):
from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary
return _cuda_rotary(x, freqs_cis)
return _torch_apply(x, freqs_cis)
+39
View File
@@ -0,0 +1,39 @@
"""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)
+41 -32
View File
@@ -13,41 +13,63 @@ from typing import (
Type, Type,
TypeVar, TypeVar,
Union, Union,
get_args,
get_origin,
) )
from typing import get_args as _get_args
from typing import get_origin as _get_origin
T = TypeVar("T") T = TypeVar("T")
def _resolve_type( def _resolve_base_type(
arg: Union[Type, str, ForwardRef], factory_cls: type arg: Union[Type, str, ForwardRef], factory_cls: type
) -> Optional[Type]: ) -> Optional[Type]:
"""Resolve a generic type-arg (str forward-ref, ForwardRef, or class).""" """Resolve the generic type-arg T to a concrete class.
if not isinstance(arg, (str, ForwardRef)):
- Concrete class (``BaseFactory[MyBase]``): returned directly.
- Forward reference (``BaseFactory["MyBase"]``): ``Base["X"]``
produces a ``ForwardRef("X")`` at class-creation time. We
extract the name and evaluate it in the factory module's
global namespace — the same mechanism ``typing.get_type_hints``
uses internally.
"""
if isinstance(arg, type):
return arg return arg
name = arg if isinstance(arg, str) else arg.__forward_arg__ if isinstance(arg, str):
if name == factory_cls.__name__: name = arg
return factory_cls elif isinstance(arg, ForwardRef):
name = arg.__forward_arg__
else:
return None
mod = sys.modules.get(factory_cls.__module__) mod = sys.modules.get(factory_cls.__module__)
if mod is None: if mod is None:
return None return None
ns = vars(mod) try:
return eval(name, vars(mod)) # noqa: S307
except NameError:
return None
if isinstance(arg, ForwardRef):
return arg._evaluate(ns, None, recursive_guard=frozenset())
return ns.get(name) def _validate_component(component_cls: Type, base: Optional[Type]) -> None:
"""Validate that *component_cls* inherits from *base*.
No-op when *base* is ``None`` (e.g. forward-ref resolution failed).
"""
if base is not None and not issubclass(component_cls, base):
raise TypeError(f"{component_cls.__name__} must inherit from {base.__name__}")
class BaseFactory(ABC, Generic[T]): class BaseFactory(ABC, Generic[T]):
"""Generic factory with decorator-based component registration. """Generic factory with decorator-based registration.
Create a factory by subclassing with the desired base type::
class MyFactory(BaseFactory[MyBase]): class MyFactory(BaseFactory[MyBase]):
pass pass
Register components with the ``register`` decorator::
@MyFactory.register("custom") @MyFactory.register("custom")
class CustomComponent(MyBase): class CustomComponent(MyBase):
... ...
@@ -64,10 +86,10 @@ class BaseFactory(ABC, Generic[T]):
def __init_subclass__(cls, **kwargs): def __init_subclass__(cls, **kwargs):
super().__init_subclass__(**kwargs) super().__init_subclass__(**kwargs)
for orig_base in getattr(cls, "__orig_bases__", ()): for orig_base in getattr(cls, "__orig_bases__", ()):
if _get_origin(orig_base) is BaseFactory: if get_origin(orig_base) is BaseFactory:
(arg,) = _get_args(orig_base) (arg,) = get_args(orig_base)
cls._entries = {} cls._entries = {}
cls._component_base = _resolve_type(arg, cls) cls._component_base = _resolve_base_type(arg, cls)
return return
@classmethod @classmethod
@@ -79,7 +101,7 @@ class BaseFactory(ABC, Generic[T]):
""" """
def decorator(component_cls: Type[T]) -> Type[T]: def decorator(component_cls: Type[T]) -> Type[T]:
cls._validate_component(component_cls) _validate_component(component_cls, cls._component_base)
if name in cls._entries: if name in cls._entries:
raise ValueError(f"Component '{name}' is already registered") raise ValueError(f"Component '{name}' is already registered")
cls._entries[name] = component_cls cls._entries[name] = component_cls
@@ -92,12 +114,11 @@ class BaseFactory(ABC, Generic[T]):
"""Create a component instance by name, filtering kwargs to match """Create a component instance by name, filtering kwargs to match
the component's ``__init__`` signature. the component's ``__init__`` signature.
""" """
entry = cls._entries.get(name) component_cls = cls._entries.get(name)
if entry is None: if component_cls is None:
raise ValueError( raise ValueError(
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}" f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
) )
component_cls = entry
sig = inspect.signature(component_cls.__init__) sig = inspect.signature(component_cls.__init__)
has_var_kwargs = any( has_var_kwargs = any(
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values() p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
@@ -111,18 +132,6 @@ class BaseFactory(ABC, Generic[T]):
kwargs = {k: v for k, v in kwargs.items() if k in valid} kwargs = {k: v for k, v in kwargs.items() if k in valid}
return component_cls(*args, **kwargs) return component_cls(*args, **kwargs)
@classmethod
def _validate_component(cls, component_cls: Type[T]):
"""Validate the decorated class inherits from the factory's base type.
Override for custom validation beyond ``issubclass``.
"""
base = cls._component_base
if base is not None and not issubclass(component_cls, base):
raise TypeError(
f"{component_cls.__name__} must inherit from {base.__name__}"
)
@classmethod @classmethod
def get_component_class(cls, name: str) -> Type[T]: def get_component_class(cls, name: str) -> Type[T]:
"""Get the registered component class without instantiating it.""" """Get the registered component class without instantiating it."""
+6 -16
View File
@@ -30,21 +30,16 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
from astrai.inference.core import ( from astrai.inference.core import (
STOP, STOP,
Allocator, Allocator,
CacheView,
ContiguousCache,
ContiguousCacheView,
Executor, Executor,
InferenceScheduler, InferenceScheduler,
KVCache, KVCache,
PageCache, KVStorage,
PageCacheView,
PagePool, PagePool,
PrefixCache, RadixCache,
Storage, ReqToTokenPool,
Task, Task,
TaskManager, TaskManager,
TaskStatus, TaskStatus,
TaskTable,
page_hash, page_hash,
) )
from astrai.inference.engine import GenerationRequest, InferenceEngine from astrai.inference.engine import GenerationRequest, InferenceEngine
@@ -68,16 +63,11 @@ __all__ = [
"TaskManager", "TaskManager",
"TaskStatus", "TaskStatus",
"Allocator", "Allocator",
"CacheView",
"KVCache", "KVCache",
"ContiguousCache", "KVStorage",
"ContiguousCacheView",
"PageCache",
"PageCacheView",
"PagePool", "PagePool",
"PrefixCache", "RadixCache",
"Storage", "ReqToTokenPool",
"TaskTable",
"page_hash", "page_hash",
"sample", "sample",
"BaseSamplingStrategy", "BaseSamplingStrategy",
+4
View File
@@ -110,6 +110,7 @@ def _create_engine(
device: str = "cuda", device: str = "cuda",
dtype: torch.dtype = torch.bfloat16, dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16, max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
) -> InferenceEngine: ) -> InferenceEngine:
if not param_path.exists(): if not param_path.exists():
raise FileNotFoundError(f"Parameter directory not found: {param_path}") raise FileNotFoundError(f"Parameter directory not found: {param_path}")
@@ -123,6 +124,7 @@ def _create_engine(
model=model, model=model,
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
max_seq_len=max_seq_len,
) )
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}") logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
return engine return engine
@@ -186,6 +188,7 @@ def run_server(
device: str = "cuda", device: str = "cuda",
dtype: torch.dtype = torch.bfloat16, dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16, max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
): ):
app = get_app() app = get_app()
app.state.server_config = { app.state.server_config = {
@@ -193,6 +196,7 @@ def run_server(
"dtype": dtype, "dtype": dtype,
"param_path": param_path, "param_path": param_path,
"max_batch_size": max_batch_size, "max_batch_size": max_batch_size,
"max_seq_len": max_seq_len,
} }
uvicorn.run( uvicorn.run(
app, app,
+10 -15
View File
@@ -22,13 +22,10 @@ class BaseToolParser(ABC):
Maintains streaming state internally so that each call to :meth:`feed` Maintains streaming state internally so that each call to :meth:`feed`
can diff against previously emitted content. can diff against previously emitted content.
Parameters Args:
---------- tools (list of dict, optional): Tool definitions from the request.
tools : list of dict, optional tool_choice (str): ``"auto"`` / ``"required"`` / ``"none"`` or a named
Tool definitions from the request. tool choice dict.
tool_choice : str
``"auto"`` / ``"required"`` / ``"none"`` or a named tool choice
dict.
""" """
def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"): def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"):
@@ -51,14 +48,12 @@ class BaseToolParser(ABC):
Returns an empty list when nothing new should be emitted. Returns an empty list when nothing new should be emitted.
Parameters Args:
---------- body (str): The complete accumulated generated text so far.
body : str current_token_ids (list of int, optional): All token IDs decoded
The complete accumulated generated text so far. into *body* (cumulative).
current_token_ids : list of int, optional delta_token_ids (list of int, optional): Only the token IDs for
All token IDs decoded into *body* (cumulative). this chunk.
delta_token_ids : list of int, optional
Only the token IDs for this chunk.
""" """
@abstractmethod @abstractmethod
+6 -16
View File
@@ -2,16 +2,11 @@
from astrai.inference.core.cache import ( from astrai.inference.core.cache import (
Allocator, Allocator,
CacheView,
ContiguousCache,
ContiguousCacheView,
KVCache, KVCache,
PageCache, KVStorage,
PageCacheView,
PagePool, PagePool,
PrefixCache, RadixCache,
Storage, ReqToTokenPool,
TaskTable,
page_hash, page_hash,
) )
from astrai.inference.core.executor import Executor from astrai.inference.core.executor import Executor
@@ -20,16 +15,11 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
__all__ = [ __all__ = [
"Allocator", "Allocator",
"CacheView",
"KVCache", "KVCache",
"ContiguousCache", "KVStorage",
"ContiguousCacheView",
"PageCache",
"PageCacheView",
"PagePool", "PagePool",
"PrefixCache", "RadixCache",
"Storage", "ReqToTokenPool",
"TaskTable",
"page_hash", "page_hash",
"Executor", "Executor",
"InferenceScheduler", "InferenceScheduler",
+502 -412
View File
@@ -1,16 +1,34 @@
"""KV cache architecture: three-layer separation (SGLang-inspired).
Layer 1 — KVStorage: flat token-level K/V buffers [n_layers, size, H, D]
Layer 2 — ReqToTokenPool: index table [req_idx, pos] → physical token slot
Layer 3 — Allocator: slot/page allocation with ref-counting and LRU
PagePool orchestrates all three plus RadixCache (prefix addressing).
KVCache is a pure dataclass passed to the model for direct buffer access.
Two modes:
- contiguous (default): pre-allocated per-request blocks, no dynamic alloc
- paged: shared pool with on-demand allocation, prefix caching support
"""
import threading import threading
from abc import ABC, abstractmethod
from collections import OrderedDict from collections import OrderedDict
from typing import Callable, Dict, List, Optional, Tuple from dataclasses import dataclass
from typing import Callable, Dict, List, Optional
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.inference.core.workspace import InferenceWorkspace
def page_hash(token_ids: List[int], page_idx: int, page_size: int) -> int:
def page_hash(
token_ids: List[int], page_idx: int, page_size: int, parent_hash: int = 0
) -> int:
start = page_idx * page_size start = page_idx * page_size
end = min(start + page_size, len(token_ids)) end = min(start + page_size, len(token_ids))
h = 0 h = parent_hash
for i in range(start, end): for i in range(start, end):
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
return h return h
@@ -67,467 +85,539 @@ class Allocator:
self._lru.move_to_end(idx) self._lru.move_to_end(idx)
class PrefixCache: class RadixNode:
"""Hash-based prefix matching: maps page hashes to physical page indices.""" """A page-aligned edge in the CPU-side prefix radix."""
__slots__ = ("parent", "children", "page_idx", "tokens", "lock_ref")
def __init__(self, parent=None, tokens=(), page_idx=None):
self.parent = parent
self.children: Dict[tuple, "RadixNode"] = {}
self.page_idx = page_idx
self.tokens = tuple(tokens)
self.lock_ref = 0
class RadixCache:
"""Page-granular radix prefix index with exact token matching."""
def __init__(self, page_size: int): def __init__(self, page_size: int):
self._page_size = page_size self._page_size = page_size
self._root = RadixNode()
self._page_to_node: Dict[int, RadixNode] = {}
# Retained as an introspection-compatible map; matching never relies on
# this lossy value.
self._page_to_hash: Dict[int, int] = {} self._page_to_hash: Dict[int, int] = {}
self._hash_to_page: Dict[int, int] = {}
self._lock = threading.Lock() self._lock = threading.Lock()
def evict(self, idx: int): def evict(self, idx: int):
with self._lock: with self._lock:
h = self._page_to_hash.pop(idx, None) node = self._page_to_node.pop(idx, None)
if h is not None: self._page_to_hash.pop(idx, None)
self._hash_to_page.pop(h, None) if node is None:
return
node.page_idx = None
parent = node.parent
if parent is not None:
parent.children.pop(node.tokens, None)
def has_page(self, idx: int) -> bool: def has_page(self, idx: int) -> bool:
with self._lock: with self._lock:
return idx in self._page_to_hash return idx in self._page_to_node
def lookup(self, token_ids: List[int]) -> List[int]: def lookup(self, token_ids: List[int]) -> List[int]:
with self._lock: with self._lock:
full_pages = len(token_ids) // self._page_size full_pages = len(token_ids) // self._page_size
hits: List[int] = [] hits: List[int] = []
node = self._root
for i in range(full_pages): for i in range(full_pages):
h = page_hash(token_ids, i, self._page_size) start = i * self._page_size
p = self._hash_to_page.get(h) page_tokens = tuple(token_ids[start : start + self._page_size])
if p is None: child = node.children.get(page_tokens)
if child is None or child.page_idx is None:
break break
hits.append(p) hits.append(child.page_idx)
node = child
return hits return hits
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int): def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
with self._lock: with self._lock:
h = page_hash(token_ids, logical_page_idx, self._page_size) full_pages = len(token_ids) // self._page_size
old_h = self._page_to_hash.pop(page_idx, None) if logical_page_idx >= full_pages:
if old_h is not None: return
self._hash_to_page.pop(old_h, None) old = self._page_to_node.pop(page_idx, None)
self._page_to_hash[page_idx] = h self._page_to_hash.pop(page_idx, None)
self._hash_to_page[h] = page_idx if old is not None and old.parent is not None:
old.parent.children.pop(old.tokens, None)
node = self._root
for i in range(logical_page_idx + 1):
start = i * self._page_size
page_tokens = tuple(token_ids[start : start + self._page_size])
child = node.children.get(page_tokens)
if child is None:
child = RadixNode(node, page_tokens)
node.children[page_tokens] = child
node = child
if node.page_idx is not None and node.page_idx != page_idx:
replaced = node.page_idx
self._page_to_node.pop(replaced, None)
self._page_to_hash.pop(replaced, None)
node.page_idx = page_idx
self._page_to_node[page_idx] = node
self._page_to_hash[page_idx] = page_hash(
token_ids, logical_page_idx, self._page_size
)
def release(self, pages: List[int]) -> None:
with self._lock:
for page_idx in pages:
node = self._page_to_node.get(page_idx)
if node is not None and node.lock_ref:
node.lock_ref -= 1
class ReqToTokenPool:
"""Maps [req_idx, pos] -> physical token slot in KV storage.
Each row is one request; each column is a sequence position. The value
at [req_idx, pos] is the flat index into the KV storage buffers.
"""
def __init__(self, size: int, max_context_len: int, device: torch.device):
self.size = size
self.max_context_len = max_context_len
self.req_to_token = torch.zeros(
(size, max_context_len), dtype=torch.long, device=device
)
self.free_slots = list(range(size))
self._lock = threading.Lock()
def alloc(self, num_reqs: int) -> Optional[List[int]]:
with self._lock:
if num_reqs > len(self.free_slots):
return None
slots = self.free_slots[:num_reqs]
self.free_slots = self.free_slots[num_reqs:]
return slots
def free(self, req_indices: List[int]):
with self._lock:
self.free_slots.extend(req_indices)
def write(self, indices, values):
self.req_to_token[indices] = values
class KVStorage:
"""Token-level KV cache storage.
Buffers: [n_layers, size, n_kv_heads, head_dim]. Each token occupies
one slot indexed by ReqToTokenPool.
"""
def __init__(
self,
size: int,
n_layers: int,
n_kv_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.size = size
self.k_buffer = torch.empty(
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
)
self.v_buffer = torch.empty(
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
)
def get_key_buffer(self, layer_id: int) -> Tensor:
return self.k_buffer[layer_id]
def get_value_buffer(self, layer_id: int) -> Tensor:
return self.v_buffer[layer_id]
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
self.k_buffer[layer_id, loc] = k
self.v_buffer[layer_id, loc] = v
@dataclass
class KVCache:
"""Pure data struct passed to model for KV cache I/O.
The attention layer does raw buffer indexing — no methods, no abstraction.
Attributes:
k_buffer: [n_layers, size, n_kv_heads, head_dim]
v_buffer: [n_layers, size, n_kv_heads, head_dim]
req_to_token: [num_reqs, max_ctx_len] — index table
req_pool_indices: [batch_size] — row indices into req_to_token
seq_lens: [batch_size] — per-request total sequence lengths
out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
kv_indptr: [batch+1] int32 — prefix sum of seq_lens, precomputed once
per step so the attention backend avoids rebuilding it per layer.
"""
k_buffer: Tensor
v_buffer: Tensor
req_to_token: Tensor
req_pool_indices: Tensor
seq_lens: Tensor
out_cache_loc: Tensor
max_len: int = 0
kv_indptr: Optional[Tensor] = None
qo_indptr: Optional[Tensor] = None
class PagePool: class PagePool:
"""Orchestrates allocator (page management) and PrefixCache (content addressing).""" """Top-level KV cache manager.
def __init__(self, allocator: Allocator, prefix: PrefixCache): Combines KVStorage + ReqToTokenPool + Allocator + RadixCache.
self._alloc = allocator
self._prefix = prefix
self._alloc.on_evict = prefix.evict
@property Args:
def allocator(self) -> Allocator: n_layers: Number of transformer layers.
return self._alloc n_kv_heads: Number of KV attention heads.
head_dim: Dimension per head.
@property max_batch_size: Maximum concurrent requests.
def prefix(self) -> PrefixCache: max_seq_len: Maximum sequence length per request.
return self._prefix device, dtype: Tensor device and dtype.
page_size: Page size for paged mode (1 = token-level).
def alloc(self) -> int: n_tokens: Total token slots for paged mode. None = contiguous mode
return self._alloc.alloc() (pre-allocates max_batch_size * max_seq_len).
"""
def free(self, idx: int):
keep = self._prefix.has_page(idx)
self._alloc.free(idx, keep_cached=keep)
if not keep:
self._prefix.evict(idx)
def inc_ref(self, idx: int):
self._alloc.inc_ref(idx)
def lookup(self, token_ids: List[int]) -> List[int]:
hits = self._prefix.lookup(token_ids)
for p in hits:
self._alloc.touch(p)
return hits
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
self._prefix.record(page_idx, token_ids, logical_page_idx)
class TaskTable:
"""Maps task_ids to page tables and cached token counts."""
def __init__(self, page_size: int):
self._page_size = page_size
self._pages: Dict[str, List[int]] = {}
self._cached: Dict[str, int] = {}
self._lock = threading.Lock()
def set(self, task_id: str, page_table: List[int], cached: int):
with self._lock:
self._pages[task_id] = page_table
self._cached[task_id] = cached
def get(self, task_id: str) -> List[int]:
with self._lock:
return self._pages.get(task_id, [])
def get_cached(self, task_id: str) -> int:
with self._lock:
return self._cached.get(task_id, 0)
def pop(self, task_id: str) -> Tuple[List[int], int]:
with self._lock:
pages = self._pages.pop(task_id, [])
cached = self._cached.pop(task_id, 0)
return pages, cached
def get_ref(self, task_id: str) -> List[int]:
with self._lock:
return self._pages.setdefault(task_id, [])
def table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
with self._lock:
states = [self._pages.get(tid, []) for tid in task_ids]
max_pages = max((len(s) for s in states), default=0)
rows = [s + [-1] * (max_pages - len(s)) for s in states]
return torch.tensor(rows, dtype=torch.long, device=device)
class Storage:
"""KV-cache tensor storage with paged write/gather."""
def __init__( def __init__(
self, self,
n_layers: int, n_layers: int,
n_pages: int,
page_size: int,
n_kv_heads: int, n_kv_heads: int,
head_dim: int, head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.page_size = page_size
self.k_cache = torch.empty(
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
device=device,
dtype=dtype,
)
self.v_cache = torch.empty(
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
device=device,
dtype=dtype,
)
def write(
self,
layer_id: int,
page_table: Tensor,
start_pos: int,
k: Tensor,
v: Tensor,
):
seq_len = k.size(1)
if seq_len == 0:
return
page_size = self.page_size
written = 0
first_page = start_pos // page_size
last_page = (start_pos + seq_len - 1) // page_size
for pi in range(first_page, last_page + 1):
phys_pages = page_table[:, pi]
page_start = pi * page_size
write_start = max(page_start, start_pos)
write_end = min(page_start + page_size, start_pos + seq_len)
offset = write_start - page_start
chunk = write_end - write_start
valid = phys_pages >= 0
if not valid.all():
if valid.any():
valid_pages = phys_pages[valid]
self.k_cache[layer_id, valid_pages, offset : offset + chunk] = k[
valid, written : written + chunk
]
self.v_cache[layer_id, valid_pages, offset : offset + chunk] = v[
valid, written : written + chunk
]
written += chunk
continue
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
:, written : written + chunk
]
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
:, written : written + chunk
]
written += chunk
def gather(
self, layer_id: int, page_table: Tensor, total_len: int
) -> Tuple[Tensor, Tensor]:
safe = page_table.clamp(min=0)
k = self.k_cache[layer_id, safe]
v = self.v_cache[layer_id, safe]
k = k.flatten(1, 2)
v = v.flatten(1, 2)
if (page_table < 0).any():
invalid = (
(page_table < 0)
.unsqueeze(-1)
.expand(-1, -1, self.page_size)
.flatten(1, 2)
)
invalid = invalid[:, :, None, None].expand_as(k)
k = k.masked_fill(invalid, 0.0)
v = v.masked_fill(invalid, 0.0)
k = k[:, :total_len]
v = v[:, :total_len]
return k, v
class CacheView(ABC):
"""Abstract view passed to attention layers for KV-cache I/O."""
@abstractmethod
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
@abstractmethod
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
class KVCache(ABC):
"""Abstract KV-cache facade for scheduler/executor."""
@abstractmethod
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
@abstractmethod
def task_free(self, task_id: str): ...
@abstractmethod
def task_extend(self, task_id: str, pos: int) -> bool: ...
@abstractmethod
def bind_tasks(
self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> CacheView: ...
def task_cached(self, task_id: str) -> int:
return 0
def task_record_hashes(
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
): ...
class PageCacheView(CacheView):
"""Bundles Storage + page_table + total_len for attention layers."""
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
self._storage = storage
self._page_table = page_table
self._total_len = total_len
def write(self, layer_id: int, k: Tensor, v: Tensor):
start_pos = self._total_len - k.size(1)
self._storage.write(layer_id, self._page_table, start_pos, k, v)
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
return self._storage.gather(layer_id, self._page_table, self._total_len)
class PageCache(KVCache):
"""Paged KV-cache with prefix sharing."""
def __init__(
self,
n_layers: int,
n_pages: int,
page_size: int,
n_kv_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.page_size = page_size
self._pool = PagePool(Allocator(n_pages), PrefixCache(page_size))
self._table = TaskTable(page_size)
self._storage = Storage(
n_layers, n_pages, page_size, n_kv_heads, head_dim, device, dtype
)
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
hits = self._pool.lookup(prompt_ids)
cached = len(hits) * self.page_size
for p in hits:
self._pool.inc_ref(p)
remaining = len(prompt_ids) - cached
n_new = (
(remaining + self.page_size - 1) // self.page_size if remaining > 0 else 0
)
new_pages: List[int] = []
if n_new > 0:
for _ in range(n_new):
p = self._pool.alloc()
if p < 0:
for hp in hits:
self._pool.free(hp)
for np in new_pages:
self._pool.free(np)
return False
new_pages.append(p)
self._table.set(task_id, hits + new_pages, cached)
return True
def task_free(self, task_id: str):
page_table, _ = self._table.pop(task_id)
for idx in page_table:
self._pool.free(idx)
def task_extend(self, task_id: str, pos: int) -> bool:
page_table = self._table.get(task_id)
needed = (pos + 1 + self.page_size - 1) // self.page_size
while len(page_table) < needed:
p = self._pool.alloc()
if p < 0:
return False
page_table.append(p)
return True
def task_cached(self, task_id: str) -> int:
return self._table.get_cached(task_id)
def task_record_hashes(
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
):
page_table = self._table.get(task_id)
full_pages = len(prompt_ids) // self.page_size
for i in range(start_logical_page, full_pages):
self._pool.record(page_table[i], prompt_ids, i)
def bind_tasks(
self,
task_ids: List[str],
total_len: int,
device: torch.device,
write_positions: Optional[Tensor] = None,
) -> PageCacheView:
page_table = self._table.table_tensor(task_ids, device)
return PageCacheView(self._storage, page_table, total_len)
class ContiguousCacheView(CacheView):
"""Contiguous KV-cache view for attention layers."""
def __init__(
self,
cache: "ContiguousCache",
batch_indices: Tensor,
total_len: int = 0,
write_positions: Optional[Tensor] = None,
):
self._cache = cache
self._batch_indices = batch_indices
self._total_len = total_len
self._write_positions = write_positions
def write(self, layer_id: int, k: Tensor, v: Tensor):
seq_len = k.size(1)
indices = self._batch_indices
if self._write_positions is not None and seq_len == 1:
pos = self._write_positions
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
for s, p in zip(indices.tolist(), pos.tolist()):
cur = self._cache._slot_len.get(s, 0)
if p + 1 > cur:
self._cache._slot_len[s] = p + 1
else:
start_pos = self._total_len - seq_len
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
new_len = start_pos + seq_len
for s in indices.tolist():
cur = self._cache._slot_len.get(s, 0)
if new_len > cur:
self._cache._slot_len[s] = new_len
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
max_len = max(
self._cache._slot_len.get(int(s), 0) for s in self._batch_indices.tolist()
)
indices = self._batch_indices
k = self._cache.k[layer_id, indices, :max_len]
v = self._cache.v[layer_id, indices, :max_len]
return k, v
class ContiguousCache(KVCache):
"""Contiguous per-slot KV cache (default implementation)."""
def __init__(
self,
n_layers: int,
max_batch_size: int, max_batch_size: int,
max_seq_len: int, max_seq_len: int,
n_kv_heads: int,
head_dim: int,
device: torch.device, device: torch.device,
dtype: torch.dtype, dtype: torch.dtype,
page_size: int = 1,
n_tokens: Optional[int] = None,
): ):
self.page_size = page_size
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len self.max_seq_len = max_seq_len
self.k = torch.zeros( self.device = device
n_layers, self.dtype = dtype
max_batch_size, self.n_layers = n_layers
max_seq_len, self.n_kv_heads = n_kv_heads
n_kv_heads, self.head_dim = head_dim
head_dim,
device=device, self.contiguous = n_tokens is None
dtype=dtype, if self.contiguous:
self.n_tokens = max_batch_size * max_seq_len
else:
self.n_tokens = n_tokens
self._storage = KVStorage(
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
) )
self.v = torch.zeros( self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
n_layers,
max_batch_size, if self.contiguous:
max_seq_len, for i in range(max_batch_size):
n_kv_heads, self._req_pool.req_to_token[i] = torch.arange(
head_dim, i * max_seq_len, (i + 1) * max_seq_len, device=device
device=device,
dtype=dtype,
) )
self._slot_len: Dict[int, int] = {} self._alloc: Optional[Allocator] = None
self._task_slot: Dict[str, int] = {} self._prefix: Optional[RadixCache] = None
self._free_slots = list(range(max_batch_size)) else:
self._device = device n_pages = self.n_tokens // page_size
self._alloc = Allocator(n_pages)
self._prefix = RadixCache(page_size) if page_size > 1 else None
if self._prefix is not None:
self._alloc.on_evict = self._prefix.evict
self._task_req: Dict[str, int] = {}
self._task_len: Dict[int, int] = {}
self._task_cached: Dict[str, int] = {}
self._task_slots: Dict[str, List[int]] = {}
self._task_pages: Dict[str, List[int]] = {}
self._lock = threading.Lock()
# Steady-state decode validation state: the ordered task set and its
# Python seq_lens mirror. When the same set advances every sequence
# by exactly one token per step, bind_tasks updates the stable
# buffers in-place (+=1 / +=inc) instead of re-cumsumming. Any
# task-set change is a miss and rebuilds.
self._bind_sig: Optional[tuple] = None
self._bind_seq_lens: Optional[List[int]] = None
# ---- task lifecycle ----
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
if not self._free_slots: req_slots = self._req_pool.alloc(1)
if req_slots is None:
return False return False
slot = self._free_slots.pop(0) req_idx = req_slots[0]
self._task_slot[task_id] = slot self._task_req[task_id] = req_idx
self._slot_len[slot] = 0
if self.contiguous:
self._task_len[req_idx] = len(prompt_ids)
self._task_cached[task_id] = 0
return True
n_tokens_needed = len(prompt_ids)
cached = 0
if self._prefix is not None:
hits = self._prefix.lookup(prompt_ids)
cached = len(hits) * self.page_size
for p in hits:
self._alloc.inc_ref(p)
self._task_pages[task_id] = list(hits)
self._task_slots[task_id] = []
else:
self._task_pages[task_id] = []
self._task_slots[task_id] = []
remaining = n_tokens_needed - cached
if remaining > 0:
if self.page_size == 1:
slots = self._alloc_tokens(remaining)
if slots is None:
for p in self._task_pages[task_id]:
self._alloc.free(p)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
self._task_slots[task_id] = slots
else:
n_new_pages = (remaining + self.page_size - 1) // self.page_size
new_pages = []
for _ in range(n_new_pages):
p = self._alloc.alloc()
if p < 0:
for hp in self._task_pages[task_id]:
self._alloc.free(hp)
for np_ in new_pages:
self._alloc.free(np_)
self._req_pool.free([req_idx])
del self._task_req[task_id]
return False
new_pages.append(p)
self._task_pages[task_id].extend(new_pages)
self._write_req_to_token(task_id, prompt_ids, cached)
self._task_len[req_idx] = len(prompt_ids)
self._task_cached[task_id] = cached
return True return True
def task_free(self, task_id: str): def task_free(self, task_id: str):
slot = self._task_slot.pop(task_id, None) req_idx = self._task_req.pop(task_id, None)
if slot is not None: if req_idx is None:
self._slot_len.pop(slot, None) return
self._free_slots.append(slot) self._task_len.pop(req_idx, None)
self._task_cached.pop(task_id, None)
if not self.contiguous:
if self._prefix is not None:
for p in self._task_pages.get(task_id, []):
keep = self._prefix.has_page(p)
self._alloc.free(p, keep_cached=keep)
if not keep:
self._prefix.evict(p)
else:
for p in self._task_pages.get(task_id, []):
self._alloc.free(p)
self._task_pages.pop(task_id, None)
self._task_slots.pop(task_id, None)
self._req_pool.free([req_idx])
def task_extend(self, task_id: str, pos: int) -> bool: def task_extend(self, task_id: str, pos: int) -> bool:
return pos < self.max_seq_len req_idx = self._task_req.get(task_id)
if req_idx is None or pos >= self.max_seq_len:
return False
# Paged mode must also claim a physical slot for the new token;
# contiguous mode's block is pre-allocated so this is a no-op.
if not self.contiguous and not self._extend_slot(task_id, req_idx, pos):
return False
self._task_len[req_idx] = pos + 1
return True
def _extend_slot(self, task_id: str, req_idx: int, pos: int) -> bool:
"""Allocate the physical slot for one extended token (paged mode)."""
if self.page_size == 1:
slots = self._alloc_tokens(1)
if slots is None:
return False
self._task_slots.setdefault(task_id, []).extend(slots)
self._req_pool.req_to_token[req_idx, pos] = slots[0]
return True
page_idx = pos // self.page_size
existing = self._task_pages.get(task_id, [])
if page_idx >= len(existing):
p = self._alloc.alloc()
if p < 0:
return False
existing.append(p)
self._task_pages[task_id] = existing
page_offset = pos % self.page_size
page = existing[page_idx]
token_slot = page * self.page_size + page_offset
self._req_pool.req_to_token[req_idx, pos] = token_slot
return True
def task_cached(self, task_id: str) -> int: def task_cached(self, task_id: str) -> int:
slot = self._task_slot.get(task_id) return self._task_cached.get(task_id, 0)
if slot is None:
return 0 def task_record_hashes(
return self._slot_len.get(slot, 0) self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
):
if self._prefix is None or self.contiguous:
return
pages = self._task_pages.get(task_id, [])
full_pages = len(prompt_ids) // self.page_size
for i in range(start_logical_page, min(full_pages, len(pages))):
self._prefix.record(pages[i], prompt_ids, i)
def task_cacheable_ids(
self, task_id: str, prompt_ids: List[int], output_ids: List[int]
):
"""Return the sequence whose KV entries are already materialized.
The first sampled output is produced by prompt prefill, and the last
sampled output has not been decoded into KV yet. Therefore the cache
can safely retain the prompt plus every output except the last one.
"""
return list(prompt_ids) + list(output_ids[:-1])
# ---- bind for forward ----
def bind_tasks( def bind_tasks(
self, self,
task_ids: List[str], task_ids: List[str],
total_len: int, workspace: InferenceWorkspace,
device: torch.device, device: Optional[torch.device] = None,
write_positions: Optional[Tensor] = None, start_pos: Optional[int] = None,
) -> ContiguousCacheView: ) -> KVCache:
slots = [self._task_slot[tid] for tid in task_ids] if device is None:
batch_indices = torch.tensor(slots, dtype=torch.long, device=device) device = workspace.device
return ContiguousCacheView( req_indices = [self._task_req[tid] for tid in task_ids]
self, batch_indices, total_len, write_positions=write_positions # Per-request lengths come from the pool's own tracking (task_alloc
# sets len(prompt_ids); task_extend sets pos+1), so callers need not
# pass them.
seq_lens = [self._task_len[req_idx] for req_idx in req_indices]
b = len(task_ids)
sig = tuple(task_ids)
# Write into the caller's workspace buffers (fixed addresses, sized
# to max_batch/max_seq at init) — the sole owner of the per-step
# KV bind tensors.
rpi_buf = workspace.req_pool_indices
sl_buf = workspace.seq_lens
kvp_buf = workspace.kv_indptr
inc_buf = workspace.inc
ocl_buf = workspace.out_cache_loc
incremental = (
start_pos is None
and self._bind_sig is not None
and self._bind_sig == sig
and self._bind_seq_lens is not None
and len(self._bind_seq_lens) == b
and all(s == p + 1 for s, p in zip(seq_lens, self._bind_seq_lens))
) )
if incremental:
# Steady-state decode: advance the stable buffers in-place.
# Normal-mode buffers keep ``+=`` legal regardless of whether
# this runs inside ``torch.inference_mode()``.
sl_buf[:b] += 1
kvp_buf[: b + 1] += inc_buf[: b + 1]
req_pool_indices = rpi_buf[:b]
seq_lens_t = sl_buf[:b]
kv_indptr = kvp_buf[: b + 1]
else:
# Cold path: fill the stable buffers from fresh host tensors.
rpi_buf[:b].copy_(
torch.tensor(req_indices, dtype=torch.long, device=device)
)
sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device))
kvp_buf[: b + 1].zero_()
kvp_buf[1 : b + 1] = sl_buf[:b].cumsum(0).to(torch.int32)
req_pool_indices = rpi_buf[:b]
seq_lens_t = sl_buf[:b]
kv_indptr = kvp_buf[: b + 1]
self._bind_sig = sig
self._bind_seq_lens = list(seq_lens)
if start_pos is not None:
seq_len = seq_lens[0]
out_cache_loc = self._req_pool.req_to_token[
req_pool_indices, start_pos:seq_len
]
# Ragged query segmentation for the prefill kernel, computed once
# (was rebuilt per layer in CudaBackend.fwd_prefill).
q_len = seq_len - start_pos
workspace.qo_indptr[: b + 1].copy_(
torch.arange(b + 1, dtype=torch.int32, device=device) * q_len
)
qo_indptr = workspace.qo_indptr[: b + 1]
else:
write_pos = seq_lens_t - 1
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
ocl_buf[:b].copy_(loc)
out_cache_loc = ocl_buf[:b]
qo_indptr = None
return KVCache(
k_buffer=self._storage.k_buffer,
v_buffer=self._storage.v_buffer,
req_to_token=self._req_pool.req_to_token,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens_t,
out_cache_loc=out_cache_loc,
max_len=max(seq_lens),
kv_indptr=kv_indptr,
qo_indptr=qo_indptr,
)
# ---- internals ----
def _alloc_tokens(self, n: int) -> Optional[List[int]]:
if self.page_size != 1:
raise RuntimeError("_alloc_tokens is for page_size=1 only")
slots = []
for _ in range(n):
p = self._alloc.alloc()
if p < 0:
for s in slots:
self._alloc.free(s)
return None
slots.append(p)
return slots
def _write_req_to_token(self, task_id: str, prompt_ids: List[int], cached: int):
req_idx = self._task_req[task_id]
total = len(prompt_ids)
if self.contiguous:
return
if self.page_size == 1:
slots = self._task_slots.get(task_id, [])
all_slots = slots[: total - cached]
if all_slots:
self._req_pool.req_to_token[req_idx, cached:total] = torch.tensor(
all_slots, dtype=torch.long, device=self.device
)
else:
pages = self._task_pages.get(task_id, [])
for pos in range(cached, total):
page_idx = pos // self.page_size
page_offset = pos % self.page_size
if page_idx < len(pages):
token_slot = pages[page_idx] * self.page_size + page_offset
self._req_pool.req_to_token[req_idx, pos] = token_slot
+188 -74
View File
@@ -1,10 +1,13 @@
import logging import logging
from dataclasses import dataclass
from typing import List, Optional from typing import List, Optional
import torch import torch
from torch import Tensor
from astrai.inference.core.cache import KVCache from astrai.inference.core.cache import PagePool
from astrai.inference.core.task import Task from astrai.inference.core.task import Task
from astrai.inference.core.workspace import InferenceWorkspace
from astrai.inference.sample import sample from astrai.inference.sample import sample
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer
@@ -12,6 +15,42 @@ from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@dataclass
class SamplingBatchInfo:
"""Per-batch sampling parameters, cached across decode steps.
Sampling params are constant for a given ordered task set, so they are
built once (pinned-memory async H2D) and reused until the task set
changes. ``top_ks`` is int32 to match the native consumers.
"""
temperatures: Tensor # float32 [B]
top_ks: Tensor # int32 [B]
top_ps: Tensor # float32 [B]
freq_penalties: Tensor # float32 [B]
has_freq: bool # any frequency_penalty != 0 (avoids per-step GPU .any())
def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
pin = str(device).startswith("cuda")
freq_penalties = torch.tensor(
[t.frequency_penalty for t in tasks], dtype=torch.float32, pin_memory=pin
).to(device, non_blocking=True)
return SamplingBatchInfo(
temperatures=torch.tensor(
[t.temperature for t in tasks], dtype=torch.float32, pin_memory=pin
).to(device, non_blocking=True),
top_ks=torch.tensor(
[t.top_k for t in tasks], dtype=torch.int32, pin_memory=pin
).to(device, non_blocking=True),
top_ps=torch.tensor(
[t.top_p for t in tasks], dtype=torch.float32, pin_memory=pin
).to(device, non_blocking=True),
freq_penalties=freq_penalties,
has_freq=bool((freq_penalties != 0).any()),
)
class Executor: class Executor:
"""Model forward passes for prefill and decode phases.""" """Model forward passes for prefill and decode phases."""
@@ -19,7 +58,7 @@ class Executor:
self, self,
model: AutoModel, model: AutoModel,
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
kv_cache: KVCache, kv_cache: PagePool,
device: Optional[str] = None, device: Optional[str] = None,
dtype: Optional[torch.dtype] = None, dtype: Optional[torch.dtype] = None,
): ):
@@ -29,9 +68,82 @@ class Executor:
self.device = device or next(model.parameters()).device self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype self.dtype = dtype or next(model.parameters()).dtype
def execute_prefill(self, tasks: List[Task], prompt_len: int, start_pos: int = 0): # Per-step decode cache for the steady-state case where the same
# ordered task set decodes one token per step. Sampling params are
# constant across steps; position_ids grows by exactly 1. Single-slot:
# any task-set change is a cache miss.
self._decode_cache: Optional[tuple] = None
# Pre-allocated fixed-shape buffers for the decode hot path
# (input_ids, decode mask, KV bind metadata). Eagerly sized at init
# so the workspace is CUDA-graph-capture friendly — no allocation
# during capture.
self._workspace = InferenceWorkspace(
max_batch_size=kv_cache.max_batch_size,
max_seq_len=kv_cache.max_seq_len,
device=self.device,
dtype=self.dtype,
)
def _sample_logits(
self,
logits: Tensor,
tasks: List[Task],
return_logprobs: bool = False,
info: Optional[SamplingBatchInfo] = None,
):
info = info or _build_sampling_batch_info(tasks, self.device)
if info.has_freq:
history_lists = [
t.prompt_ids[-t.rep_window :] + t.output_ids for t in tasks
]
history_lens = [len(ids) for ids in history_lists]
max_len = max(history_lens, default=0)
padded_ids = torch.zeros(
len(tasks), max_len, dtype=torch.long, device=self.device
)
padded_mask = torch.zeros(
len(tasks), max_len, dtype=torch.bool, device=self.device
)
for i, ids in enumerate(history_lists):
length = len(ids)
padded_ids[i, :length] = torch.as_tensor(
ids, dtype=torch.long, device=self.device
)
padded_mask[i, :length] = True
else:
padded_ids = None
padded_mask = None
result = sample(
logits,
temperature=info.temperatures,
top_k=info.top_ks,
top_p=info.top_ps,
frequency_penalty=info.freq_penalties,
input_ids=padded_ids,
input_mask=padded_mask,
return_logprobs=return_logprobs,
)
if not return_logprobs:
return result.tolist()
tokens, logprobs = result
tokens_list = tokens.tolist()
logprobs_list = logprobs.tolist()
for task, logprob in zip(tasks, logprobs_list):
task.output_logprobs.append(float(logprob))
return list(zip(tokens_list, logprobs_list))
def execute_prefill(
self,
tasks: List[Task],
prompt_len: int,
start_pos: int = 0,
return_logprobs: bool = False,
):
if start_pos >= prompt_len: if start_pos >= prompt_len:
return return []
tasks = sorted(tasks, key=lambda t: t.task_id) tasks = sorted(tasks, key=lambda t: t.task_id)
batch_sz = len(tasks) batch_sz = len(tasks)
@@ -43,85 +155,87 @@ class Executor:
) )
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
position_ids = (
with torch.inference_mode(): torch.arange(start_pos, prompt_len, dtype=torch.long, device=self.device)
self.model(
input_ids,
position_ids=torch.arange(
start_pos, prompt_len, dtype=torch.long, device=self.device
)
.unsqueeze(0) .unsqueeze(0)
.expand(batch_sz, -1), .expand(batch_sz, -1)
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
) )
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
def execute_decode(self, tasks: List[Task]) -> List[int]: prompt_len, device=self.device
if not tasks:
return []
input_ids = torch.tensor(
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks],
dtype=torch.long,
device=self.device,
)
position_ids = torch.tensor(
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
)
total_len = position_ids.max().item() + 1
task_ids = [t.task_id for t in tasks]
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
freq_penalties = torch.tensor(
[t.frequency_penalty for t in tasks], device=self.device
)
history_lists = []
mask_lists = []
for t in tasks:
window = t.rep_window
prompt_part = t.prompt_ids[-window:]
ids = prompt_part + t.output_ids
history_lists.append(ids)
mask_lists.append([True] * len(ids))
max_len = max(len(h) for h in history_lists)
padded_ids = torch.zeros(
len(tasks), max_len, dtype=torch.long, device=self.device
)
padded_mask = torch.zeros(
len(tasks), max_len, dtype=torch.bool, device=self.device
)
for i, (h, m) in enumerate(zip(history_lists, mask_lists)):
padded_ids[i, : len(h)] = torch.tensor(
h, dtype=torch.long, device=self.device
)
padded_mask[i, : len(m)] = torch.tensor(
m, dtype=torch.bool, device=self.device
) )
with torch.inference_mode(): with torch.inference_mode():
outputs = self.model( outputs = self.model(
input_ids.unsqueeze(1), input_ids,
paged_cache=self.kv_cache.bind_tasks( input_mask=input_mask,
position_ids=position_ids,
kv_cache=self.kv_cache.bind_tasks(
task_ids, task_ids,
total_len, self._workspace,
self.device, start_pos=start_pos,
write_positions=position_ids, ),
)
logits = outputs["logits"][:, -1, :]
return tasks, self._sample_logits(logits, tasks, return_logprobs)
def execute_decode(
self, tasks: List[Task], return_logprobs: bool = False
) -> List[int]:
"""Decode next token for each task.
Args:
return_logprobs: When ``True``, also record (and return)
the log-probability of each sampled token under the
post-strategy sampling distribution. The logprob is
appended to ``task.output_logprobs`` and the return
list becomes ``List[Tuple[int, float]]``.
Returns:
``List[int]`` of sampled token IDs, or
``List[Tuple[int, float]]`` of ``(token_id, logprob)`` when
``return_logprobs`` is ``True``.
"""
if not tasks:
return []
input_ids = self._workspace.fill_input_ids(
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
).unsqueeze(1)
task_ids = [t.task_id for t in tasks]
sig = tuple(task_ids)
cur_positions = [t.next_pos for t in tasks]
cached = self._decode_cache
if (
cached is not None
and cached[0] == sig
and cur_positions == [p + 1 for p in cached[1]]
):
_, _, info, position_ids = cached
position_ids += 1
self._decode_cache = (sig, cur_positions, info, position_ids)
else:
info = _build_sampling_batch_info(tasks, self.device)
position_ids = torch.tensor(
cur_positions, dtype=torch.long, device=self.device
)
self._decode_cache = (sig, cur_positions, info, position_ids)
total_len = max(t.next_pos for t in tasks) + 1
input_mask = self._workspace.decode_mask(position_ids, total_len)
with torch.inference_mode():
outputs = self.model(
input_ids,
input_mask=input_mask,
kv_cache=self.kv_cache.bind_tasks(
task_ids,
self._workspace,
), ),
position_ids=position_ids.unsqueeze(1), position_ids=position_ids.unsqueeze(1),
) )
logits = outputs["logits"][:, -1, :] logits = outputs["logits"][:, -1, :]
return sample( return self._sample_logits(logits, tasks, return_logprobs, info=info)
logits,
temperature=temperatures,
top_k=top_ks,
top_p=top_ps,
frequency_penalty=freq_penalties,
input_ids=padded_ids,
input_mask=padded_mask,
).tolist()
+185 -58
View File
@@ -1,10 +1,11 @@
import logging import logging
import threading import threading
import uuid
from typing import Any, Dict, List, Optional, Tuple from typing import Any, Dict, List, Optional, Tuple
import torch import torch
from astrai.inference.core.cache import ContiguousCache, KVCache from astrai.inference.core.cache import PagePool
from astrai.inference.core.executor import Executor from astrai.inference.core.executor import Executor
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
@@ -22,45 +23,43 @@ class InferenceScheduler:
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
max_batch_size: int = 16, max_batch_size: int = 16,
max_seq_len: Optional[int] = None, max_seq_len: Optional[int] = None,
max_prompt_len: int = 2048,
device: Optional[str] = None, device: Optional[str] = None,
dtype: Optional[torch.dtype] = None, dtype: Optional[torch.dtype] = None,
cache: Optional[KVCache] = None, cache: Optional[PagePool] = None,
): ):
config = model.config config = model.config
if max_seq_len is not None: if max_seq_len is not None:
self.max_seq_len = max_seq_len self.max_seq_len = max_seq_len
elif config.max_len is not None: elif config.max_position_embeddings is not None:
self.max_seq_len = config.max_len self.max_seq_len = config.max_position_embeddings
else: else:
raise ValueError( raise ValueError(
"max_seq_len must be provided either as argument " "max_seq_len must be provided either as argument "
"or in model config (config.max_len)" "or in model config (config.max_position_embeddings)"
) )
self.device = device or next(model.parameters()).device self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype self.dtype = dtype or next(model.parameters()).dtype
head_dim = config.dim // config.n_heads head_dim = config.hidden_size // config.num_attention_heads
if cache is not None: if cache is not None:
self._cache = cache self._cache = cache
else: else:
self._cache = ContiguousCache( self._cache = PagePool(
config.n_layers, n_layers=config.num_hidden_layers,
max_batch_size, n_kv_heads=config.num_key_value_heads,
self.max_seq_len, head_dim=head_dim,
config.n_kv_heads, max_batch_size=max_batch_size,
head_dim, max_seq_len=self.max_seq_len,
self.device, device=self.device,
self.dtype, dtype=self.dtype,
) )
self._task_mgr = TaskManager( self._task_mgr = TaskManager(
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
max_seq_len=self.max_seq_len, max_seq_len=self.max_seq_len,
max_prompt_len=max_prompt_len,
) )
self._executor = Executor( self._executor = Executor(
@@ -84,6 +83,80 @@ class InferenceScheduler:
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
return self._task_mgr.get_stats() return self._task_mgr.get_stats()
def _step(
self, tasks: List[Task], return_logprobs: bool = False
) -> Tuple[List[Task], List[Task]]:
"""Advance every active task by one token (prefill + decode).
Single shared primitive for both the continuous-batching loop and
the synchronous ``run_batch`` path, so the two cannot drift.
Tasks must already be allocated in the KV cache. Tasks without output
are prefilled first and sample their first token from the final prompt
position. Tasks with output extend the cache by one position and decode
from their latest generated token.
Args:
tasks: Active tasks to advance by one token.
return_logprobs: Forwarded to ``execute_decode``; per-token
logprobs are recorded on each task's ``output_logprobs``.
Returns:
``(decoded, aborted)``: tasks that produced a new token (its ID
already appended to ``output_ids``) and tasks that hit the
sequence cap and were marked ``ABORTED``.
"""
cache = self._cache
to_prefill = [t for t in tasks if t.output_tokens == 0 and t.prompt_ids]
prefilled_ids = set()
produced: List[Task] = []
if to_prefill:
for t in to_prefill:
t.input_tokens = len(t.prompt_ids)
groups: Dict[Tuple[int, int], List[Task]] = {}
for t in to_prefill:
start_pos = min(cache.task_cached(t.task_id), len(t.prompt_ids) - 1)
groups.setdefault((len(t.prompt_ids), start_pos), []).append(t)
for (prompt_len, start_pos), group in groups.items():
prefilled, step_out = self._executor.execute_prefill(
group, prompt_len, start_pos, return_logprobs=return_logprobs
)
for t, out in zip(prefilled, step_out):
t.output_ids.append(out[0] if return_logprobs else out)
t.output_tokens += 1
prefilled_ids.add(t.task_id)
produced.append(t)
start_logical_page = start_pos // getattr(cache, "page_size", 64)
for t in group:
cache.task_record_hashes(
t.task_id, t.prompt_ids, start_logical_page
)
decoded: List[Task] = []
aborted: List[Task] = []
for t in tasks:
if t.task_id in prefilled_ids:
continue
if cache.task_extend(t.task_id, t.next_pos):
decoded.append(t)
else:
t.status = TaskStatus.ABORTED
aborted.append(t)
if decoded:
step_out = self._executor.execute_decode(
decoded, return_logprobs=return_logprobs
)
for t, out in zip(decoded, step_out):
t.output_ids.append(out[0] if return_logprobs else out)
t.output_tokens += 1
produced.append(t)
return produced, aborted
def _run_generation_loop(self): def _run_generation_loop(self):
stop_ids = self._task_mgr.tokenizer.stop_ids stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache cache = self._cache
@@ -91,6 +164,13 @@ class InferenceScheduler:
while not self._stop_event.is_set(): while not self._stop_event.is_set():
finished = self._task_mgr.remove_finished_tasks(stop_ids) finished = self._task_mgr.remove_finished_tasks(stop_ids)
for task in finished: for task in finished:
if task.status == TaskStatus.FINISHED:
cache.task_record_hashes(
task.task_id,
cache.task_cacheable_ids(
task.task_id, task.prompt_ids, task.output_ids
),
)
cache.task_free(task.task_id) cache.task_free(task.task_id)
active = self._task_mgr.get_active_tasks() active = self._task_mgr.get_active_tasks()
@@ -110,55 +190,17 @@ class InferenceScheduler:
self._task_mgr.wait_for_tasks(timeout=1.0) self._task_mgr.wait_for_tasks(timeout=1.0)
continue continue
to_prefill = [ active = self._task_mgr.get_active_tasks()
t
for t in self._task_mgr.get_active_tasks()
if t.output_tokens == 0
and cache.task_cached(t.task_id) < len(t.prompt_ids)
]
if to_prefill:
for t in to_prefill:
t.input_tokens = len(t.prompt_ids)
groups: Dict[Tuple[int, int], List[Task]] = {} decoded, aborted = self._step(active)
for t in to_prefill:
key = (
len(t.prompt_ids),
cache.task_cached(t.task_id),
)
groups.setdefault(key, []).append(t)
for (prompt_len, start_pos), group in groups.items(): for t in aborted:
self._executor.execute_prefill(group, prompt_len, start_pos)
start_logical_page = start_pos // getattr(
cache, "page_size", 64
)
for t in group:
cache.task_record_hashes(
t.task_id, t.prompt_ids, start_logical_page
)
decode_tasks = self._task_mgr.get_active_tasks()
valid: List[Task] = []
for t in sorted(decode_tasks, key=lambda t: t.task_id):
if cache.task_extend(t.task_id, t.next_pos):
valid.append(t)
else:
t.status = TaskStatus.ABORTED
self._task_mgr.invoke_callback(t.task_id, STOP) self._task_mgr.invoke_callback(t.task_id, STOP)
if valid: for t in decoded:
next_tokens = self._executor.execute_decode(valid)
for t, ntok in zip(valid, next_tokens):
t.output_ids.append(ntok)
t.output_tokens += 1
new_text = t.decode_new_token(self._task_mgr.tokenizer) new_text = t.decode_new_token(self._task_mgr.tokenizer)
if new_text: if new_text:
self._task_mgr.invoke_callback(t.task_id, new_text) self._task_mgr.invoke_callback(t.task_id, new_text)
for t in valid:
if t.is_finished(stop_ids): if t.is_finished(stop_ids):
remaining = t.flush_remaining(self._task_mgr.tokenizer) remaining = t.flush_remaining(self._task_mgr.tokenizer)
if remaining: if remaining:
@@ -194,6 +236,91 @@ class InferenceScheduler:
self._cache.task_free(task.task_id) self._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._cache.task_free(task.task_id)
self._task_mgr.clear_queues() self._task_mgr.clear_queues()
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
def run_batch(
self,
prompt_ids_list: List[List[int]],
*,
max_tokens: Optional[int] = None,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
return_logprobs: bool = False,
) -> List[List[int]]:
"""Synchronous batch generation without the scheduler thread.
Accepts already-tokenized prompts (no string round-trip) and runs
prefill + decode to completion on the calling thread. Designed for
RL rollout, where logprobs of the behaviour policy must be collected
alongside generated tokens.
Args:
prompt_ids_list: ``B`` prompts, each a list of token IDs.
max_tokens: Maximum tokens to generate per prompt. ``None``
uses ``self.max_seq_len - len(prompt_ids)``.
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
parameters (uniform across the batch).
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
Returns:
``List[List[int]]`` of generated token IDs per prompt, or —
when ``return_logprobs`` is ``True`` —
``List[Tuple[List[int], List[float]]]``.
"""
stop_ids = self._task_mgr.tokenizer.stop_ids
cache = self._cache
seq_cap = self.max_seq_len
tasks: List[Task] = []
for ids in prompt_ids_list:
if len(ids) >= seq_cap:
tasks.append(None)
continue
t_max = max_tokens
if t_max is None:
t_max = seq_cap - len(ids)
else:
t_max = min(t_max, seq_cap - len(ids))
task = Task(
task_id=f"batch_{uuid.uuid4().hex[:8]}",
prompt_ids=list(ids),
max_tokens=t_max,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
)
if not cache.task_alloc(task.task_id, task.prompt_ids):
tasks.append(None)
continue
task.input_tokens = len(task.prompt_ids)
tasks.append(task)
try:
live = [t for t in tasks if t is not None]
while live:
decoded, _ = self._step(live, return_logprobs=return_logprobs)
live = [t for t in decoded if not t.is_finished(stop_ids)]
finally:
for t in tasks:
if t is not None:
cache.task_free(t.task_id)
results: List[Any] = []
for t in tasks:
if t is None:
results.append(([], []) if return_logprobs else [])
elif return_logprobs:
results.append((list(t.output_ids), list(t.output_logprobs)))
else:
results.append(list(t.output_ids))
return results
+25 -37
View File
@@ -6,6 +6,8 @@ from collections import deque
from enum import Enum from enum import Enum
from typing import Any, Callable, Deque, Dict, List, Optional from typing import Any, Callable, Deque, Dict, List, Optional
from tokenizers.decoders import DecodeStream
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -14,37 +16,30 @@ STOP = object()
class StreamDecoder: class StreamDecoder:
"""Incremental decoder for byte-level BPE streaming. """Incremental decoder backed by the tokenizers library's DecodeStream.
Byte-level BPE may split a single Unicode character (e.g. em-dash, Delegates to the Rust-native streaming decoder which maintains an
smart quotes) across multiple tokens. Decoding such a token in O(1) bounded token buffer internally (via prefix drain), avoiding
isolation produces U+FFFD (replacement char). This decoder the O(n²) cost of re-decoding the full history on each step.
accumulates token IDs and only emits text once the trailing
characters are complete, buffering incomplete multi-byte sequences Multi-byte UTF-8 sequences split across token boundaries are
until the next token arrives. buffered until complete; ``push`` returns "" while the trailing
sequence is still incomplete.
""" """
__slots__ = ("_tokenizer", "_ids", "_emitted") __slots__ = ("_stream", "_tok")
def __init__(self, tokenizer: AutoTokenizer): def __init__(self, tokenizer: AutoTokenizer):
self._tokenizer = tokenizer self._tok = tokenizer._tokenizer
self._ids: List[int] = [] self._stream = DecodeStream(skip_special_tokens=True)
self._emitted: str = ""
def push(self, token_id: int) -> str: def push(self, token_id: int) -> str:
"""Append a token ID and return newly completed text. """Append a token ID and return newly completed text.
Returns "" while a multi-byte character is still incomplete. Returns "" while a multi-byte character is still incomplete.
""" """
self._ids.append(token_id) chunk = self._stream.step(self._tok, token_id)
full = self._tokenizer.decode(self._ids, skip_special_tokens=True) return chunk or ""
if full.endswith("\ufffd"):
return ""
if len(full) > len(self._emitted):
diff = full[len(self._emitted) :]
self._emitted = full
return diff
return ""
class TaskStatus(Enum): class TaskStatus(Enum):
@@ -81,6 +76,7 @@ class Task:
self.status = TaskStatus.PENDING self.status = TaskStatus.PENDING
self.output_ids: List[int] = [] self.output_ids: List[int] = []
self.output_logprobs: List[float] = []
self.input_tokens: int = 0 self.input_tokens: int = 0
self.output_tokens: int = 0 self.output_tokens: int = 0
self.arrival_time = time.time() self.arrival_time = time.time()
@@ -100,23 +96,17 @@ class Task:
def flush_remaining(self, tokenizer: AutoTokenizer) -> str: def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
"""Emit any text still buffered in the decoder. """Emit any text still buffered in the decoder.
Called when generation terminates (max_tokens reached, stop With the Rust-native DecodeStream, the stream is always in a
sequence, or external removal) to avoid dropping a final correct state — any completed text was already emitted by the
incomplete-looking fragment that is actually complete when last ``push``. A trailing incomplete multi-byte sequence has no
adjacent to the stop token. valid text to emit, so this is a no-op.
""" """
if self._decoder is None or not self.output_ids:
return ""
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
if len(full) > len(self._decoder._emitted):
diff = full[len(self._decoder._emitted) :]
self._decoder._emitted = full
return diff
return "" return ""
@property @property
def next_pos(self) -> int: def next_pos(self) -> int:
return self.input_tokens + len(self.output_ids) # The first output is sampled from prefill and enters KV on the next step.
return self.input_tokens + max(0, len(self.output_ids) - 1)
def is_finished(self, stop_ids: List[int]) -> bool: def is_finished(self, stop_ids: List[int]) -> bool:
if self.max_tokens is not None and self.output_tokens >= self.max_tokens: if self.max_tokens is not None and self.output_tokens >= self.max_tokens:
@@ -134,12 +124,10 @@ class TaskManager:
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
max_batch_size: int = 16, max_batch_size: int = 16,
max_seq_len: int = 8192, max_seq_len: int = 8192,
max_prompt_len: int = 512,
): ):
self.tokenizer = tokenizer self.tokenizer = tokenizer
self.max_batch_size = max_batch_size self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len self.max_seq_len = max_seq_len
self.max_prompt_len = max_prompt_len
self.waiting_queue: Deque[Task] = deque() self.waiting_queue: Deque[Task] = deque()
self.active_tasks: List[Task] = [] self.active_tasks: List[Task] = []
@@ -164,10 +152,10 @@ class TaskManager:
) -> str: ) -> str:
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}" task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
prompt_ids = self.tokenizer.encode(prompt) prompt_ids = self.tokenizer.encode(prompt)
if len(prompt_ids) > self.max_prompt_len: if len(prompt_ids) > self.max_seq_len:
prompt_ids = prompt_ids[-self.max_prompt_len :] prompt_ids = prompt_ids[-self.max_seq_len :]
if len(prompt_ids) >= self.max_seq_len: if len(prompt_ids) > self.max_seq_len:
if stream_callback: if stream_callback:
stream_callback(STOP) stream_callback(STOP)
return task_id return task_id
+111
View File
@@ -0,0 +1,111 @@
"""Pre-allocated buffers for the inference decode hot path.
Mirrors SGLang's pre-allocated input buffers (``input_buffers.py``): tensors
are sized once to the server's maximum dimensions and sliced to the live
batch each step, so the per-token decode loop never calls
``torch.empty``/``torch.zeros``/``torch.arange`` for the hot shapes. Fills
go through ``out=`` variants (``torch.ge``) which write into the stable
buffers instead of allocating fresh results.
All buffers are allocated eagerly at init (nothing is lazy), so the
workspace is CUDA-graph-capture friendly: the decode step reads/writes
fixed-address tensors with no allocation during capture.
"""
import torch
from torch import Tensor
class InferenceWorkspace:
"""Reusable fixed-shape per-step buffers for decode.
Families of buffers, all sized to ``max_batch_size`` / ``max_seq_len``
and sliced via views each step:
- ``decode_mask``: a ``[B, 1, total_len]`` validity mask, the RHS
``arange`` pre-computed so only a single ``torch.ge(out=)`` runs per
step.
- ``input_ids``: per-step token IDs filled from host (pinned, double-
buffered so an in-flight async H2D copy never races the next fill).
- KV-cache bind metadata (``req_pool_indices``, ``seq_lens``,
``kv_indptr``, ``inc``, ``out_cache_loc``), written by
``PagePool.bind_tasks`` when the Executor passes this workspace.
No re-allocation while the server's bounds are respected.
"""
def __init__(
self,
max_batch_size: int,
max_seq_len: int,
device: torch.device,
dtype: torch.dtype,
):
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.device = device
self.dtype = dtype
# ``position_ids[:, None, None] >= arange`` RHS, reused every step.
self.arange = torch.arange(max_seq_len, device=device)
# Decode validity mask: [max_batch, 1, max_seq_len] bool.
self.input_mask = torch.empty(
(max_batch_size, 1, max_seq_len), dtype=torch.bool, device=device
)
# Per-step token IDs. Values come from host Python lists every
# step, so the device buffer is pre-allocated (stable address for
# CUDA-graph capture) and filled via a host staging buffer. A
# double buffer keeps a copy in flight from being overwritten by
# the next fill.
self.input_ids = torch.empty((max_batch_size,), dtype=torch.long, device=device)
self._pin = [
torch.empty((max_batch_size,), dtype=torch.long),
torch.empty((max_batch_size,), dtype=torch.long),
]
self._pin_idx = 0
# KV-cache bind metadata (fixed shape, written by ``PagePool.bind_tasks``
# when the Executor passes this workspace). Stable addresses make the
# decode forward CUDA-graph capturable.
self.req_pool_indices = 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(
(max_batch_size + 1,), dtype=torch.int32, device=device
)
self.qo_indptr = torch.empty(
(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(
(max_batch_size, 1), dtype=torch.long, device=device
)
def fill_input_ids(self, ids: "list[int]") -> Tensor:
"""Write ``ids`` into the device buffer and return ``[B]``.
Host values are staged through the double buffer and copied into the
stable device buffer (``copy_`` without pinning is synchronous, so
the alternating buffers guard against an in-flight transfer).
"""
b = len(ids)
pin = self._pin[self._pin_idx]
self._pin_idx ^= 1
for i, v in enumerate(ids):
pin[i] = v
self.input_ids[:b].copy_(pin[:b])
return self.input_ids[:b]
def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor:
"""Return the ``[B, 1, total_len]`` validity mask for this step.
Written into the pre-allocated buffer via ``torch.ge(out=)`` — no
new tensor is allocated. ``position_ids`` is the current step's
``[B]`` positions; ``total_len`` must not exceed ``max_seq_len``.
"""
b = position_ids.size(0)
out = self.input_mask[:b, :, :total_len]
torch.ge(position_ids[:, None, None], self.arange[:total_len], out=out)
return out
+2 -5
View File
@@ -8,7 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
import torch import torch
import torch.nn as nn import torch.nn as nn
from astrai.inference.core.cache import KVCache from astrai.inference.core.cache import PagePool
from astrai.inference.core.scheduler import InferenceScheduler from astrai.inference.core.scheduler import InferenceScheduler
from astrai.inference.core.task import STOP from astrai.inference.core.task import STOP
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
@@ -111,9 +111,7 @@ class InferenceEngine:
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
max_batch_size: int = 1, max_batch_size: int = 1,
max_seq_len: Optional[int] = None, max_seq_len: Optional[int] = None,
max_prompt_len: int = 2048, cache: Optional[PagePool] = None,
page_size: int = 128,
cache: Optional[KVCache] = None,
): ):
self.model = model self.model = model
self.tokenizer = tokenizer self.tokenizer = tokenizer
@@ -122,7 +120,6 @@ class InferenceEngine:
tokenizer=self.tokenizer, tokenizer=self.tokenizer,
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
max_seq_len=max_seq_len, max_seq_len=max_seq_len,
max_prompt_len=max_prompt_len,
cache=cache, cache=cache,
) )
+81 -20
View File
@@ -276,7 +276,8 @@ class SamplingPipeline(BaseSamplingStrategy):
filter_value: float = -float("inf"), filter_value: float = -float("inf"),
input_ids: Optional[Tensor] = None, input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
) -> Tensor: return_logprobs: bool = False,
):
"""Apply strategies then sample (softmax + multinomial). """Apply strategies then sample (softmax + multinomial).
Short-circuits to ``argmax`` when temperature is exactly 0 Short-circuits to ``argmax`` when temperature is exactly 0
@@ -286,21 +287,41 @@ class SamplingPipeline(BaseSamplingStrategy):
logits: Raw logits ``[batch, vocab_size]``. logits: Raw logits ``[batch, vocab_size]``.
input_ids: Previously generated token IDs ``[batch, seq_len]``. input_ids: Previously generated token IDs ``[batch, seq_len]``.
input_mask: Boolean mask for ``input_ids`` padding. input_mask: Boolean mask for ``input_ids`` padding.
return_logprobs: If ``True``, return ``(tokens, logprobs)``
where ``logprobs[i]`` is the log-probability of
``tokens[i]`` under the (post-strategy) sampling
distribution.
Returns: Returns:
Sampled token IDs ``[batch]``. Sampled token IDs ``[batch]``, or — when ``return_logprobs``
is ``True`` — a ``(token_ids, chosen_logprobs)`` tuple.
""" """
for s in self.strategies: if self._is_greedy_pipeline():
if isinstance(s, TemperatureStrategy) and self._is_greedy(s.temperature): tokens = logits.argmax(dim=-1)
return logits.argmax(dim=-1) if not return_logprobs:
break return tokens
log_probs = torch.log_softmax(logits.float(), dim=-1)
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
return torch.multinomial( transformed = self.apply(logits, filter_value, input_ids, input_mask)
torch.softmax( tokens = torch.multinomial(
self.apply(logits, filter_value, input_ids, input_mask), dim=-1 torch.softmax(transformed, dim=-1), num_samples=1
),
num_samples=1,
).squeeze(-1) ).squeeze(-1)
if not return_logprobs:
return tokens
log_probs = torch.log_softmax(transformed.float(), dim=-1)
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
def _is_greedy_pipeline(self) -> bool:
"""True if the first strategy is greedy temperature (temp=0)."""
if not self.strategies:
return False
first = self.strategies[0]
return isinstance(first, TemperatureStrategy) and self._is_greedy(
first.temperature
)
@torch.inference_mode() @torch.inference_mode()
@@ -313,31 +334,71 @@ def sample(
input_ids: Optional[Tensor] = None, input_ids: Optional[Tensor] = None,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
filter_value: float = -float("inf"), filter_value: float = -float("inf"),
) -> Tensor: return_logprobs: bool = False,
):
"""Apply sampling strategies then sample (softmax + multinomial). """Apply sampling strategies then sample (softmax + multinomial).
Shortcut for ``SamplingPipeline(...).sample(logits)``. Shortcut for ``SamplingPipeline(...).sample(logits, return_logprobs=)``.
When **temperature** is exactly 0 (scalar or single-element tensor) When **temperature** is exactly 0 (scalar or single-element tensor)
the function short-circuits to ``argmax`` for deterministic decode. the function short-circuits to ``argmax`` for deterministic decode.
When **frequency_penalty** is 0 (the common decode case), the entire
frequency penalty computation — including the O(batch * vocab) count
tensor allocation — is skipped.
Args: Args:
logits: Raw logits ``[batch, vocab_size]``. logits: Raw logits ``[batch, vocab_size]``.
frequency_penalty: Penalty per occurrence for repeated tokens frequency_penalty: Penalty per occurrence for repeated tokens
(0.0 disables, range -2.0~2.0). (0.0 disables, range -2.0~2.0).
input_ids: Previously generated token IDs ``[batch, seq_len]``. input_ids: Previously generated token IDs ``[batch, seq_len]``.
input_mask: Boolean mask for ``input_ids`` padding. input_mask: Boolean mask for ``input_ids`` padding.
return_logprobs: If ``True``, also return the log-probability
of each sampled token under the (post-strategy) sampling
distribution — useful for RL rollout (PPO/GRPO importance
ratios).
Returns: Returns:
Sampled token IDs ``[batch]``. Sampled token IDs ``[batch]``, or — when ``return_logprobs`` is
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
``chosen_logprobs`` has shape ``[batch]``.
""" """
if SamplingPipeline._is_greedy(temperature): greedy = (
return logits.argmax(dim=-1) (
return SamplingPipeline( isinstance(temperature, Tensor)
[ and temperature.numel() == 1
and temperature.item() == 0
)
if isinstance(temperature, Tensor)
else temperature == 0
)
if greedy:
tokens = logits.argmax(dim=-1)
if not return_logprobs:
return tokens
log_probs = torch.log_softmax(logits.float(), dim=-1)
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
return tokens, chosen
has_freq = (
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
if isinstance(frequency_penalty, Tensor)
else frequency_penalty != 0
)
strategies: List[BaseSamplingStrategy] = [
TemperatureStrategy(temperature), TemperatureStrategy(temperature),
TopKStrategy(top_k), TopKStrategy(top_k),
TopPStrategy(top_p), TopPStrategy(top_p),
FrequencyPenaltyStrategy(frequency_penalty),
] ]
).sample(logits, filter_value, input_ids, input_mask) if has_freq:
strategies.append(FrequencyPenaltyStrategy(frequency_penalty))
return SamplingPipeline(strategies).sample(
logits,
filter_value=filter_value,
input_ids=input_ids,
input_mask=input_mask,
return_logprobs=return_logprobs,
)
+2 -1
View File
@@ -9,7 +9,7 @@ from astrai.model.components.lora import (
merge_lora, merge_lora,
save_lora, save_lora,
) )
from astrai.model.components.mlp import MLP from astrai.model.components.mlp import MLP, DeepSeekMoE
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
from astrai.model.encoder import EmbeddingEncoder from astrai.model.encoder import EmbeddingEncoder
from astrai.model.transformer import AutoRegressiveLM from astrai.model.transformer import AutoRegressiveLM
@@ -19,6 +19,7 @@ __all__ = [
"Linear", "Linear",
"RMSNorm", "RMSNorm",
"MLP", "MLP",
"DeepSeekMoE",
"GQA", "GQA",
"DecoderBlock", "DecoderBlock",
# Models # Models
+7 -6
View File
@@ -40,11 +40,12 @@ def _disable_random_init(enable: bool = True):
setattr(nn.init, n, fn) setattr(nn.init, n, fn)
class AutoModel(BaseFactory["AutoModel"], nn.Module): class ModelFactory(BaseFactory[nn.Module]):
""" """Pure factory for model dispatch, separated from nn.Module state."""
Autoregressive language model base class.
Provides model loading/saving, registration, and generation.
""" class AutoModel(nn.Module):
"""Model base class with loading/saving and generation."""
def __init__(self, config: BaseModelConfig): def __init__(self, config: BaseModelConfig):
super().__init__() super().__init__()
@@ -68,7 +69,7 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
config = ConfigFactory.load(raw) config = ConfigFactory.load(raw)
model_type = config.model_type or "autoregressive_lm" model_type = config.model_type or "autoregressive_lm"
actual_cls = AutoModel.get_component_class(model_type) actual_cls = ModelFactory.get_component_class(model_type)
with _disable_random_init(enable=disable_random_init): with _disable_random_init(enable=disable_random_init):
model = actual_cls(config) model = actual_cls(config)
+4 -4
View File
@@ -1,12 +1,12 @@
from astrai.model.components.attention import GQA, MLA, repeat_kv from astrai.extension.rotary_backend import apply_rotary_emb
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
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
from astrai.model.components.mlp import MLP from astrai.model.components.mlp import MLP, DeepSeekMoE
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
from astrai.model.components.rope import ( from astrai.model.components.rope import (
RotaryEmbedding, RotaryEmbedding,
apply_rotary_emb,
get_rotary_emb, get_rotary_emb,
) )
@@ -14,6 +14,7 @@ __all__ = [
"Linear", "Linear",
"RMSNorm", "RMSNorm",
"MLP", "MLP",
"DeepSeekMoE",
"Embedding", "Embedding",
"GQA", "GQA",
"MLA", "MLA",
@@ -21,5 +22,4 @@ __all__ = [
"RotaryEmbedding", "RotaryEmbedding",
"apply_rotary_emb", "apply_rotary_emb",
"get_rotary_emb", "get_rotary_emb",
"repeat_kv",
] ]
+9 -43
View File
@@ -5,22 +5,12 @@ 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.rotary_backend import apply_rotary_emb
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.inference.core.cache import CacheView from astrai.inference.core.cache import KVCache
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
from astrai.model.components.rope import apply_rotary_emb
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
bs, slen, n_heads, head_dim = x.shape
if n_rep == 1:
return x
return (
x[:, :, :, None, :]
.expand(bs, slen, n_heads, n_rep, head_dim)
.reshape(bs, slen, n_heads * n_rep, head_dim)
)
class AttnFactory(BaseFactory[nn.Module]): class AttnFactory(BaseFactory[nn.Module]):
@@ -75,10 +65,9 @@ class GQA(nn.Module):
x: Tensor, x: Tensor,
rotary_emb: Tensor, rotary_emb: Tensor,
attn_mask: Tensor = None, attn_mask: Tensor = None,
paged_cache: Optional[CacheView] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False,
) -> Tensor: ) -> Tensor:
is_causal = attn_mask is None
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)
v = self._split_heads(self.v_proj(x), self.n_kv_heads) v = self._split_heads(self.v_proj(x), self.n_kv_heads)
@@ -87,19 +76,7 @@ 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)
if paged_cache is not None: sdqa_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
paged_cache.write(self.layer_id, k, v)
k, v = paged_cache.gather(self.layer_id)
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
sdqa_out = (
F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
.permute(0, 2, 1, 3)
.contiguous()
.flatten(2)
)
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))
@@ -162,10 +139,10 @@ class MLA(nn.Module):
x: Tensor, x: Tensor,
rotary_emb: Tensor, rotary_emb: Tensor,
attn_mask: Tensor = None, attn_mask: Tensor = None,
paged_cache: Optional[CacheView] = None, kv_cache: Optional[KVCache] = None,
is_causal: bool = False,
) -> Tensor: ) -> Tensor:
bsz, seq_len, _ = x.size() bsz, seq_len, _ = x.size()
is_causal = attn_mask is None
q = self.q_proj(x) q = self.q_proj(x)
q = q.view(bsz, seq_len, self.n_heads, self.head_dim) q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
@@ -194,18 +171,7 @@ class MLA(nn.Module):
q = self.q_norm(q) q = self.q_norm(q)
k = self.k_norm(k) k = self.k_norm(k)
if paged_cache is not None: attn_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
paged_cache.write(self.layer_id, k, v)
k, v = paged_cache.gather(self.layer_id)
q = q.permute(0, 2, 1, 3)
k = k.permute(0, 2, 1, 3)
v = v.permute(0, 2, 1, 3)
attn_out = F.scaled_dot_product_attention(
q, k, v, attn_mask, is_causal=is_causal
)
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
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))
+47 -12
View File
@@ -1,39 +1,74 @@
from dataclasses import asdict from dataclasses import asdict
from typing import Optional from typing import Optional, TypedDict
import torch.nn as nn import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.inference.core.cache import CacheView from astrai.inference.core.cache import KVCache
from astrai.model.components.attention import AttnFactory from astrai.model.components.attention import AttnFactory
from astrai.model.components.mlp import FFNFactory from astrai.model.components.mlp import FFNFactory, RouterStats
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
class DecoderOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
router_stats: Optional[RouterStats]
class DecoderBlock(nn.Module): class DecoderBlock(nn.Module):
def __init__(self, config, layer_id: int): def __init__(self, config, layer_id: int):
super().__init__() super().__init__()
cfg = asdict(config) cfg = asdict(config)
cfg["down_init_std"] = 0.02 / (2 * config.n_layers) ** 0.5 cfg.update(
dim=config.hidden_size,
dim_ffn=config.intermediate_size,
n_layers=config.num_hidden_layers,
n_heads=config.num_attention_heads,
n_kv_heads=config.num_key_value_heads,
norm_eps=config.rms_norm_eps,
down_init_std=0.02 / (2 * config.num_hidden_layers) ** 0.5,
)
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id) self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
self.input_norm = RMSNorm(config.dim, config.norm_eps) self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.post_attention_norm = RMSNorm(config.dim, config.norm_eps) self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.mlp = FFNFactory.create(config.ffn_type, **cfg) ffn_type = self._resolve_ffn_type(config, layer_id)
self.mlp = FFNFactory.create(ffn_type, **cfg)
@staticmethod
def _resolve_ffn_type(config, layer_id: int) -> str:
if config.ffn_type != "moe":
return config.ffn_type
mlp_only = config.mlp_only_layers or []
if layer_id in mlp_only:
return "mlp"
if config.decoder_sparse_step > 1:
if (layer_id + 1) % config.decoder_sparse_step != 0:
return "mlp"
return "moe"
def forward( def forward(
self, self,
x: Tensor, x: Tensor,
rotary_emb: Tensor, rotary_emb: Tensor,
attention_mask: Optional[Tensor] = None, attention_mask: Optional[Tensor] = None,
paged_cache: Optional[CacheView] = None, kv_cache: Optional[KVCache] = None,
) -> Tensor: is_causal: bool = False,
) -> DecoderOutput:
attn_output = self.attention( attn_output = self.attention(
self.input_norm(x), self.input_norm(x),
rotary_emb, rotary_emb,
attention_mask, attention_mask,
paged_cache, kv_cache,
is_causal,
) )
x = attn_output + x x = attn_output + x
x = self.mlp(self.post_attention_norm(x)) + x normalized = self.post_attention_norm(x)
mlp_output = self.mlp(normalized)
x = mlp_output["hidden_states"] + x
return x return {
"hidden_states": x,
"aux_loss": mlp_output["aux_loss"],
"router_stats": mlp_output.get("router_stats"),
}
+8 -3
View File
@@ -1,11 +1,12 @@
import logging import logging
from dataclasses import asdict, dataclass from dataclasses import asdict
from pathlib import Path from pathlib import Path
from typing import Optional, Set from typing import Optional, Set
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from pydantic.dataclasses import dataclass
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
from astrai.serialization import ( from astrai.serialization import (
@@ -39,8 +40,12 @@ class LoRALinear(nn.Module):
self.r = r self.r = r
self.scaling = alpha / r self.scaling = alpha / r
self.lora_A = nn.Parameter(torch.randn(r, self.weight.shape[1]) / r) device = self.weight.device
self.lora_B = nn.Parameter(torch.zeros(self.weight.shape[0], r)) dtype = self.weight.dtype
lora_a = torch.randn(r, self.weight.shape[1], device=device, dtype=dtype) / r
lora_b = torch.zeros(self.weight.shape[0], r, device=device, dtype=dtype)
self.lora_A = nn.Parameter(lora_a)
self.lora_B = nn.Parameter(lora_b)
self._merged = False self._merged = False
def forward(self, x): def forward(self, x):
+99 -22
View File
@@ -1,3 +1,5 @@
from typing import Optional, TypedDict
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
@@ -11,6 +13,28 @@ class FFNFactory(BaseFactory[nn.Module]):
pass pass
class RouterStats(TypedDict):
"""Per-layer MoE routing statistics for training diagnostics.
Both tensors are detached monitoring data produced during forward.
"""
probs: Tensor
topk_indices: Tensor
class FFNOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
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):
@@ -19,10 +43,10 @@ class MLP(nn.Module):
self.gate = Linear(dim, dim_ffn) self.gate = Linear(dim, dim_ffn)
self.down = Linear(dim_ffn, dim, init_std=down_init_std) self.down = Linear(dim_ffn, dim, init_std=down_init_std)
def forward(self, x: Tensor) -> Tensor: def forward(self, x: Tensor) -> FFNOutput:
gated = self.up(x) * F.silu(self.gate(x)) gated = self.up(x) * F.silu(self.gate(x))
out = self.down(gated) out = self.down(gated)
return out return {"hidden_states": out, "aux_loss": None, "router_stats": None}
@FFNFactory.register("moe") @FFNFactory.register("moe")
@@ -36,6 +60,9 @@ class DeepSeekMoE(nn.Module):
n_activated_experts: int = 2, n_activated_experts: int = 2,
topk_method: str = "greedy", topk_method: str = "greedy",
n_layers: int = 1, n_layers: int = 1,
moe_intermediate_size: Optional[int] = None,
shared_expert_intermediate_size: Optional[int] = None,
norm_topk_prob: bool = True,
): ):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
@@ -43,6 +70,16 @@ class DeepSeekMoE(nn.Module):
self.n_shared_experts = n_shared_experts self.n_shared_experts = n_shared_experts
self.n_activated_experts = n_activated_experts self.n_activated_experts = n_activated_experts
self.topk_method = topk_method self.topk_method = topk_method
self.norm_topk_prob = norm_topk_prob
expert_dim_ffn = (
moe_intermediate_size if moe_intermediate_size is not None else dim_ffn
)
shared_dim_ffn = (
shared_expert_intermediate_size
if shared_expert_intermediate_size is not None
else dim_ffn
)
self.router = Linear(dim, n_routed_experts, bias=False) self.router = Linear(dim, n_routed_experts, bias=False)
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
@@ -50,51 +87,91 @@ class DeepSeekMoE(nn.Module):
self.shared_experts = nn.ModuleList( self.shared_experts = nn.ModuleList(
[ [
MLP(dim, dim_ffn, down_init_std=down_init_std) MLP(dim, shared_dim_ffn, down_init_std=down_init_std)
for _ in range(n_shared_experts) for _ in range(n_shared_experts)
] ]
) )
self.routed_experts = nn.ModuleList( self.routed_experts = nn.ModuleList(
[ [
MLP(dim, dim_ffn, down_init_std=down_init_std) MLP(dim, expert_dim_ffn, down_init_std=down_init_std)
for _ in range(n_routed_experts) for _ in range(n_routed_experts)
] ]
) )
def forward(self, x: Tensor) -> Tensor: def forward(self, x: Tensor) -> FFNOutput:
include_aux_loss = self.training and torch.is_grad_enabled()
bsz, seq_len, dim = x.shape bsz, seq_len, dim = x.shape
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_out = self._routed_forward(x_flat) routed_output = self._routed_forward(x_flat, include_aux_loss)
out = (shared_out + routed_out).view(bsz, seq_len, dim) out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim)
return out return {
"hidden_states": out,
"aux_loss": routed_output["aux_loss"],
"router_stats": routed_output["router_stats"],
}
def _shared_forward(self, x: Tensor) -> Tensor: def _shared_forward(self, x: Tensor) -> Tensor:
if self.n_shared_experts == 0: if self.n_shared_experts == 0:
return torch.zeros_like(x) return torch.zeros_like(x)
return sum(e(x) for e in self.shared_experts) / self.n_shared_experts return (
sum(e(x)["hidden_states"] for e in self.shared_experts)
/ self.n_shared_experts
)
def _routed_forward(self, x: Tensor) -> Tensor: def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
N, D = x.shape N, D = x.shape
K = self.n_activated_experts K = self.n_activated_experts
E = self.n_routed_experts
router_logits = self.router(x) router_logits = self.router(x)
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype) router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1) topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1, sorted=False)
if self.norm_topk_prob:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
output = torch.zeros(N, D, device=x.device, dtype=x.dtype) aux_loss = None
for expert_idx in range(self.n_routed_experts): router_stats = None
expert_mask = topk_indices == expert_idx if include_aux_loss:
token_idx, k_idx = expert_mask.nonzero(as_tuple=True) expert_load = F.one_hot(topk_indices, num_classes=E).float()
if token_idx.numel() == 0: expert_load = expert_load.mean(dim=(0, 1))
continue router_prob = router_probs.float().mean(dim=0)
expert_input = x[token_idx] aux_loss = E * (expert_load * router_prob).sum()
expert_output = self.routed_experts[expert_idx](expert_input) router_stats = {
weights = topk_weights[token_idx, k_idx].unsqueeze(-1) "probs": router_probs.detach(),
output.index_add_(0, token_idx, expert_output * weights) "topk_indices": topk_indices,
}
return output # Grouped dispatch: sort (token, slot) pairs by expert so each expert
# consumes one contiguous slice instead of a per-expert mask scan.
flat_experts = topk_indices.reshape(-1)
sorted_experts, order = torch.sort(flat_experts)
flat_tokens = x.repeat_interleave(K, dim=0)[order]
flat_weights = topk_weights.reshape(-1, 1)[order]
boundaries = torch.cumsum(
torch.bincount(sorted_experts, minlength=E), dim=0
).tolist()
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
start = 0
for expert_idx, end in enumerate(boundaries):
if end == start:
continue
expert_output = self.routed_experts[expert_idx](flat_tokens[start:end])[
"hidden_states"
]
output.index_add_(
0,
order[start:end] // K,
expert_output * flat_weights[start:end],
)
start = end
return {
"hidden_states": output,
"aux_loss": aux_loss,
"router_stats": router_stats,
}
+17 -15
View File
@@ -11,28 +11,23 @@ def get_rotary_emb(
base: float = 10000, base: float = 10000,
device: Optional[torch.device] = None, device: Optional[torch.device] = None,
) -> Tensor: ) -> Tensor:
"""Precompute cos/sin tables for rotary embedding.
Returns:
[max_len, dim/2, 2] (f32) — [cos, sin] pairs.
"""
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim) theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
t = torch.arange(0, max_len, dtype=torch.float64, device=device) t = torch.arange(0, max_len, dtype=torch.float64, device=device)
freqs = torch.outer(t, theta).float() freqs = torch.outer(t, theta).float()
cos = torch.cos(freqs) cos = torch.cos(freqs)
sin = torch.sin(freqs) sin = torch.sin(freqs)
return torch.complex(cos, sin) return torch.stack([cos, sin], dim=-1)
def ntk_base(base: float, dim: int, factor: float) -> float: def ntk_base(base: float, dim: int, factor: float) -> float:
return base * (factor ** (dim / (dim - 2))) return base * (factor ** (dim / (dim - 2)))
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
dtype = x.dtype
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
x_complex = torch.view_as_complex(x_)
freqs_cis = freqs_cis.unsqueeze(2)
x_rotated = x_complex * freqs_cis
x_out = torch.view_as_real(x_rotated).flatten(-2)
return x_out.to(dtype)
class RotaryEmbedding(nn.Module): class RotaryEmbedding(nn.Module):
def __init__( def __init__(
self, self,
@@ -56,16 +51,23 @@ class RotaryEmbedding(nn.Module):
self._set_rotary_buffer(self.max_len) self._set_rotary_buffer(self.max_len)
def _set_rotary_buffer(self, max_len: int): def _set_rotary_buffer(self, max_len: int):
rotary_emb = get_rotary_emb(self.dim, max_len, self.base) freqs_cis = get_rotary_emb(self.dim, max_len, self.base)
freqs_cis = torch.view_as_real(rotary_emb)
self.register_buffer("freqs_cis", freqs_cis, persistent=False) self.register_buffer("freqs_cis", freqs_cis, persistent=False)
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor: def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
"""Lookup cos/sin for the given positions.
Args:
x: [batch, seq_len, ...] — only batch and seq_len are used.
position_ids: [batch, seq_len] optional position indices.
Returns:
[batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
"""
if position_ids is None: if position_ids is None:
position_ids = ( position_ids = (
torch.arange(x.size(1), device=x.device) torch.arange(x.size(1), device=x.device)
.unsqueeze(0) .unsqueeze(0)
.expand(x.size(0), -1) .expand(x.size(0), -1)
) )
position_freq_cis = self.freqs_cis[position_ids].float() return self.freqs_cis[position_ids].float()
return torch.view_as_complex(position_freq_cis)
+17 -9
View File
@@ -5,7 +5,7 @@ import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.config.model_config import EncoderConfig from astrai.config.model_config import EncoderConfig
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel, ModelFactory
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding from astrai.model.components.embedding import Embedding
from astrai.model.components.norm import RMSNorm from astrai.model.components.norm import RMSNorm
@@ -13,25 +13,33 @@ from astrai.model.components.rope import RotaryEmbedding
from astrai.model.transformer import process_attention_mask from astrai.model.transformer import process_attention_mask
@AutoModel.register("embedding") @ModelFactory.register("embedding")
class EmbeddingEncoder(AutoModel): class EmbeddingEncoder(AutoModel):
def __init__(self, config: EncoderConfig): def __init__(self, config: EncoderConfig):
super().__init__(config) super().__init__(config)
self.config = config self.config = config
rope_dim = config.dim // config.n_heads rope_dim = config.hidden_size // config.num_attention_heads
rope_base = config.rope_theta if config.rope_theta is not None else 10000 rope_base = config.rope_theta if config.rope_theta is not None else 10000
self.rotary_embedding = RotaryEmbedding( self.rotary_embedding = RotaryEmbedding(
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling rope_dim,
config.max_position_embeddings,
rope_base,
rope_scaling=config.rope_scaling,
) )
self.embed_tokens = Embedding( self.embed_tokens = Embedding(
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha config.vocab_size,
config.hidden_size,
neftune_alpha=config.neftune_alpha,
) )
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)] [
DecoderBlock(config, layer_id)
for layer_id in range(config.num_hidden_layers)
]
) )
self.norm = RMSNorm(config.dim, config.norm_eps) self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.pooling_type = config.pooling_type or "mean" self.pooling_type = config.pooling_type or "mean"
self.normalize_embeddings = config.normalize_embeddings or False self.normalize_embeddings = config.normalize_embeddings or False
@@ -59,10 +67,10 @@ class EmbeddingEncoder(AutoModel):
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)
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=False) attn_mask = process_attention_mask(input_mask)
for layer in self.layers: for layer in self.layers:
x = layer(x, rotary_emb, attn_mask, paged_cache=None) x = layer(x, rotary_emb, attn_mask)["hidden_states"]
hidden_states = self.norm(x) hidden_states = self.norm(x)
+48 -39
View File
@@ -5,8 +5,8 @@ import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.config.model_config import AutoRegressiveLMConfig from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.inference.core.cache import CacheView from astrai.inference.core.cache import KVCache
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel, ModelFactory
from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding from astrai.model.components.embedding import Embedding
from astrai.model.components.linear import Linear from astrai.model.components.linear import Linear
@@ -15,35 +15,18 @@ from astrai.model.components.rope import RotaryEmbedding
def process_attention_mask( def process_attention_mask(
input_tensor: Tensor, input_mask: Optional[Tensor],
position_ids: Optional[Tensor],
input_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Optional[Tensor]: ) -> Optional[Tensor]:
if position_ids is None: if input_mask is None:
return None return None
if input_mask is not None and input_mask.dim() > 2: if input_mask.dim() == 2:
return input_mask[:, None, None, :]
if input_mask.dim() == 3:
return input_mask[:, None, :, :]
return input_mask return input_mask
device = input_tensor.device
B = input_tensor.size(0)
T = position_ids.max().item() + 1
if input_mask is None: @ModelFactory.register("autoregressive_lm")
if position_ids.min().item() == 0 and is_causal:
return None
attend = torch.ones(B, 1, T, dtype=torch.bool, device=device)
else:
attend = input_mask[:, :T].to(device=device, dtype=torch.bool).unsqueeze(1)
if is_causal:
causal = position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
attend = attend & causal
return attend.unsqueeze(1)
@AutoModel.register("autoregressive_lm")
class AutoRegressiveLM(AutoModel): class AutoRegressiveLM(AutoModel):
"""Autoregressive language model with paged KV cache.""" """Autoregressive language model with paged KV cache."""
@@ -53,24 +36,32 @@ class AutoRegressiveLM(AutoModel):
rope_dim = ( rope_dim = (
config.qk_rope_head_dim config.qk_rope_head_dim
if config.attn_type == "mla" if config.attn_type == "mla"
else config.dim // config.n_heads else config.hidden_size // config.num_attention_heads
) )
rope_base = config.rope_theta if config.rope_theta is not None else 10000 rope_base = config.rope_theta if config.rope_theta is not None else 10000
self.rotary_embedding = RotaryEmbedding( self.rotary_embedding = RotaryEmbedding(
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling rope_dim,
config.max_position_embeddings,
rope_base,
rope_scaling=config.rope_scaling,
) )
self.embed_tokens = Embedding( self.embed_tokens = Embedding(
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha config.vocab_size,
config.hidden_size,
neftune_alpha=config.neftune_alpha,
) )
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)] [
DecoderBlock(config, layer_id)
for layer_id in range(config.num_hidden_layers)
]
) )
self.norm = RMSNorm(config.dim, config.norm_eps) self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.lm_head = Linear(config.dim, config.vocab_size) self.lm_head = Linear(config.hidden_size, config.vocab_size)
if self.config.tie_weight is True: if self.config.tie_word_embeddings is True:
self.lm_head.weight = self.embed_tokens.weight self.lm_head.weight = self.embed_tokens.weight
self.apply(self._init_weights) self.apply(self._init_weights)
@@ -85,7 +76,7 @@ class AutoRegressiveLM(AutoModel):
state_dict = dict(state_dict) state_dict = dict(state_dict)
if self.config.tie_weight is True: if self.config.tie_word_embeddings is True:
# same tensor for embed and lm_head # same tensor for embed and lm_head
if embed_key in state_dict: if embed_key in state_dict:
state_dict[lm_head_key] = state_dict[embed_key] state_dict[lm_head_key] = state_dict[embed_key]
@@ -101,7 +92,7 @@ class AutoRegressiveLM(AutoModel):
destination=destination, prefix=prefix, keep_vars=keep_vars destination=destination, prefix=prefix, keep_vars=keep_vars
) )
if self.config.tie_weight is True: if self.config.tie_word_embeddings is True:
lm_head_key = prefix + "lm_head.weight" lm_head_key = prefix + "lm_head.weight"
if lm_head_key in state_dict: if lm_head_key in state_dict:
del state_dict[lm_head_key] del state_dict[lm_head_key]
@@ -112,19 +103,37 @@ class AutoRegressiveLM(AutoModel):
self, self,
input_ids: Tensor, input_ids: Tensor,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
paged_cache: Optional[CacheView] = None, kv_cache: Optional[KVCache] = None,
position_ids: Optional[Tensor] = None, position_ids: Optional[Tensor] = None,
) -> Dict[str, Tensor]: ) -> Dict[str, Tensor]:
assert input_ids.ndim == 2 assert input_ids.ndim == 2
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)
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=True) attn_mask = process_attention_mask(input_mask)
use_sdpa_causal_mask = attn_mask is None
aux_losses = []
router_stats_list = []
for layer in self.layers: for layer in self.layers:
x = layer(x, rotary_emb, attn_mask, paged_cache) layer_output = layer(
x,
rotary_emb,
attn_mask,
kv_cache,
use_sdpa_causal_mask,
)
x = layer_output["hidden_states"]
stats = layer_output.get("router_stats")
if stats is not None:
aux_losses.append(layer_output["aux_loss"])
router_stats_list.append(stats)
hidden_states = self.norm(x) hidden_states = self.norm(x)
logits = self.lm_head(hidden_states) logits = self.lm_head(hidden_states)
return {"logits": logits, "hidden_states": hidden_states} output = {"logits": logits, "hidden_states": hidden_states}
if aux_losses:
output["aux_loss"] = torch.stack(aux_losses).mean()
output["router_stats"] = router_stats_list
return output
+38
View File
@@ -0,0 +1,38 @@
"""Optimizer implementations and factory registration."""
from astrai.optim.composite import (
OptimizerFactory,
composite_state_dict,
composite_step,
composite_zero_grad,
refresh_param_groups,
)
from astrai.optim.mano_adamw import Mano, ManoAdamW
from astrai.optim.muon_adamw import MuonAdamW
from astrai.optim.nora_nadamw import (
NAdamW,
Nora,
NoraNAdamW,
OptimizerParameterGroups,
nora_direction,
nora_lr_scale,
partition_optimizer_parameters,
)
__all__ = [
"Mano",
"ManoAdamW",
"MuonAdamW",
"NAdamW",
"Nora",
"NoraNAdamW",
"OptimizerFactory",
"OptimizerParameterGroups",
"composite_state_dict",
"composite_step",
"composite_zero_grad",
"nora_direction",
"nora_lr_scale",
"partition_optimizer_parameters",
"refresh_param_groups",
]
+71
View File
@@ -0,0 +1,71 @@
"""Shared infrastructure for the optim package.
This module hosts two things:
* ``OptimizerFactory`` — the registry for built-in optimizers. Defining it
here (rather than in ``__init__.py``) lets each optimizer module import it
and register itself with a decorator, avoiding circular imports.
* Composite-optimizer helpers — ``step``/``zero_grad``/``state_dict``/
``param_groups`` delegation shared by every optimizer that routes different
parameter groups through distinct sub-optimizers.
"""
from typing import Any
import torch
from torch.optim import Optimizer
from astrai.factory import BaseFactory
class OptimizerFactory(BaseFactory[Optimizer]):
"""Factory for built-in training optimizers."""
def composite_step(
sub_optimizers: list[Optimizer],
closure=None,
) -> torch.Tensor | None:
"""Run ``step`` on every sub-optimizer, invoking the closure once.
The closure (if given) is executed inside ``torch.enable_grad`` exactly
once before any sub-optimizer steps, matching the contract of a single
``Optimizer.step``. Sub-optimizers receive ``None`` so they do not
re-execute it.
"""
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for sub in sub_optimizers:
sub.step()
return loss
def composite_zero_grad(
sub_optimizers: list[Optimizer],
set_to_none: bool = True,
) -> None:
for sub in sub_optimizers:
sub.zero_grad(set_to_none=set_to_none)
def composite_state_dict(
named_sub_optimizers: dict[str, Optimizer | None],
) -> dict[str, Any]:
"""Serialize sub-optimizers, preserving ``None`` slots."""
return {
name: sub.state_dict() if sub is not None else None
for name, sub in named_sub_optimizers.items()
}
def refresh_param_groups(
sub_optimizers: list[Optimizer],
) -> list[dict]:
"""Concatenate param_groups from every non-None sub-optimizer."""
groups: list[dict] = []
for sub in sub_optimizers:
if sub is not None:
groups.extend(sub.param_groups)
return groups
+214
View File
@@ -0,0 +1,214 @@
"""Mano manifold optimizer combined with AdamW.
Mano projects the momentum onto the tangent space of the Oblique manifold
(axis-wise tangent projection) and normalizes it, replacing the expensive
Newton-Schulz iteration in Muon with a cheaper manifold normalization.
Reference: https://arxiv.org/abs/2601.23000
"""
import math
import torch
from torch import nn, optim
from torch.optim import Optimizer
from astrai.optim.composite import (
OptimizerFactory,
composite_state_dict,
composite_step,
composite_zero_grad,
refresh_param_groups,
)
from astrai.optim.nora_nadamw import partition_optimizer_parameters
class Mano(Optimizer):
"""Manifold Normalized Optimizer for two-dimensional matrices.
Each step alternates the projection axis (dim 0 / dim 1) to restrike the
manifold along both rows and columns. The tangent momentum is computed
without normalizing the parameter itself (v2 simplification) and the
epsilon is added (not clamped) to the norm denominator.
"""
def __init__(
self,
params,
lr: float = 1e-3,
weight_decay: float = 0.1,
momentum: float = 0.95,
nesterov: bool = True,
eps: float = 1e-8,
):
if lr < 0:
raise ValueError(f"Invalid learning rate: {lr}")
if weight_decay < 0:
raise ValueError(f"Invalid weight decay: {weight_decay}")
if not 0 <= momentum <= 1:
raise ValueError(f"Invalid momentum: {momentum}")
if eps <= 0:
raise ValueError(f"Invalid epsilon: {eps}")
defaults = {
"lr": lr,
"weight_decay": weight_decay,
"momentum": momentum,
"nesterov": nesterov,
"eps": eps,
"steps": 0,
}
super().__init__(params, defaults)
for group in self.param_groups:
for param in group["params"]:
if param.ndim != 2:
raise ValueError(
f"Mano only supports 2D matrices, got shape {tuple(param.shape)}"
)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
lr = group["lr"]
weight_decay = group["weight_decay"]
momentum = group["momentum"]
nesterov = group["nesterov"]
eps = group["eps"]
dim = int(group["steps"] % 2)
for param in group["params"]:
if param.grad is None:
continue
if param.grad.is_sparse:
raise RuntimeError("Mano does not support sparse gradients")
grad = param.grad
state = self.state[param]
momentum_buffer = state.get("momentum_buffer")
if momentum_buffer is None:
momentum_buffer = torch.zeros_like(grad)
momentum_buffer.mul_(momentum).add_(grad)
update = (
grad.add(momentum_buffer, alpha=momentum)
if nesterov
else momentum_buffer
)
tangent = update - (
torch.sum(update * param.data, dim=dim, keepdim=True) * param.data
)
direction = tangent / (
torch.norm(tangent, p=2, dim=dim, keepdim=True) + eps
)
if weight_decay != 0:
param.mul_(1 - lr * weight_decay)
adjusted_lr = lr * 0.2 * math.sqrt(direction.shape[dim])
param.add_(direction, alpha=-adjusted_lr)
state["momentum_buffer"] = momentum_buffer
group["steps"] += 1
return loss
@OptimizerFactory.register("mano_adamw")
class ManoAdamW(Optimizer):
"""Mano for internal linear weights and AdamW for remaining parameters."""
optimizer_name = "mano_adamw"
def __init__(
self,
model: nn.Module,
lr: float = 3e-4,
weight_decay: float = 0.1,
momentum: float = 0.95,
nesterov: bool = True,
):
groups = partition_optimizer_parameters(model)
all_params = [
*groups.nora,
*groups.nadamw_decay,
*groups.nadamw_no_decay,
]
if not all_params:
raise ValueError(
"Cannot build an optimizer for a model with no trainable parameters"
)
super().__init__(all_params, {})
self.mano = (
Mano(
groups.nora,
lr=lr,
weight_decay=weight_decay,
momentum=momentum,
nesterov=nesterov,
)
if groups.nora
else None
)
adamw_groups = []
if groups.nadamw_decay:
adamw_groups.append(
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
)
if groups.nadamw_no_decay:
adamw_groups.append({"params": groups.nadamw_no_decay, "weight_decay": 0.0})
self.adamw = (
optim.AdamW(
adamw_groups,
lr=lr,
betas=(0.9, 0.95),
fused=True,
)
if adamw_groups
else None
)
self.param_groups = refresh_param_groups([self.mano, self.adamw])
@torch.no_grad()
def step(self, closure=None):
return composite_step(
[opt for opt in (self.mano, self.adamw) if opt is not None],
closure,
)
def zero_grad(self, set_to_none: bool = True):
composite_zero_grad(
[opt for opt in (self.mano, self.adamw) if opt is not None],
set_to_none,
)
def state_dict(self) -> dict:
return composite_state_dict({"mano": self.mano, "adamw": self.adamw})
def load_state_dict(self, state_dict: dict):
if "muon" in state_dict or "nora" in state_dict:
raise ValueError(
"Checkpoint uses a different optimizer; select the matching "
"--optimizer to resume it"
)
if "mano" not in state_dict or "adamw" not in state_dict:
raise ValueError(
"Checkpoint optimizer state is not compatible with mano_adamw"
)
saved_mano = state_dict["mano"]
saved_adamw = state_dict["adamw"]
if (self.mano is None) != (saved_mano is None):
raise ValueError("Checkpoint Mano parameter groups do not match the model")
if (self.adamw is None) != (saved_adamw is None):
raise ValueError("Checkpoint AdamW parameter groups do not match the model")
if self.mano is not None:
self.mano.load_state_dict(saved_mano)
if self.adamw is not None:
self.adamw.load_state_dict(saved_adamw)
self.param_groups = refresh_param_groups([self.mano, self.adamw])
+95
View File
@@ -0,0 +1,95 @@
"""Legacy Muon + AdamW combined optimizer."""
from typing import Any
import torch
from torch import Tensor, nn, optim
from astrai.optim.composite import (
OptimizerFactory,
composite_state_dict,
composite_step,
composite_zero_grad,
refresh_param_groups,
)
@OptimizerFactory.register("muon_adamw")
class MuonAdamW(optim.Optimizer):
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
optimizer_name = "muon_adamw"
def __init__(
self,
model: nn.Module,
lr: float = 3e-4,
weight_decay: float = 0.1,
momentum: float = 0.95,
nesterov: bool = True,
ns_steps: int = 5,
adjust_lr_fn: str = "match_rms_adamw",
):
defaults = {
"lr": lr,
"weight_decay": weight_decay,
"momentum": momentum,
"nesterov": nesterov,
"ns_steps": ns_steps,
"adjust_lr_fn": adjust_lr_fn,
}
params = [param for param in model.parameters() if param.requires_grad]
super().__init__(params, defaults)
matrix_params: list[Tensor] = []
other_params: list[Tensor] = []
for name, param in model.named_parameters():
if not param.requires_grad:
continue
if (
param.dim() >= 2
and "norm" not in name
and "bias" not in name
and "embed" not in name
and "lm_head" not in name
):
matrix_params.append(param)
else:
other_params.append(param)
self.muon = optim.Muon(
matrix_params,
lr=lr,
weight_decay=weight_decay,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adjust_lr_fn=adjust_lr_fn,
)
self.adamw = optim.AdamW(
[{"params": other_params, "weight_decay": 0.0}],
lr=lr,
betas=(0.9, 0.95),
fused=True,
)
self.param_groups = refresh_param_groups([self.muon, self.adamw])
@torch.no_grad()
def step(self, closure=None):
return composite_step([self.muon, self.adamw], closure)
def zero_grad(self, set_to_none: bool = True):
composite_zero_grad([self.muon, self.adamw], set_to_none)
def state_dict(self) -> dict[str, Any]:
return composite_state_dict({"muon": self.muon, "adamw": self.adamw})
def load_state_dict(self, state_dict: dict[str, Any]):
if "muon" not in state_dict or "adamw" not in state_dict:
raise ValueError(
"Checkpoint optimizer state is not compatible with muon_adamw"
)
self.muon.load_state_dict(state_dict["muon"])
self.adamw.load_state_dict(state_dict["adamw"])
self.param_groups = refresh_param_groups([self.muon, self.adamw])
+372
View File
@@ -0,0 +1,372 @@
"""Nora matrix optimizer combined with Nesterov AdamW."""
import math
from dataclasses import dataclass
from typing import Any
import torch
from torch import Tensor, nn
from torch.distributed.tensor import DTensor, Shard
from torch.optim import Optimizer
from astrai.model.components.embedding import Embedding
from astrai.model.components.linear import Linear
from astrai.model.components.lora import LoRALinear
from astrai.model.components.norm import RMSNorm
from astrai.optim.composite import (
OptimizerFactory,
composite_state_dict,
composite_step,
composite_zero_grad,
refresh_param_groups,
)
NORA_EPS = 1e-10
def _row_normalize(tensor: Tensor, eps: float) -> Tensor:
return tensor / tensor.norm(dim=-1, keepdim=True).clamp(min=eps)
def nora_direction(update: Tensor, param: Tensor, eps: float = NORA_EPS) -> Tensor:
"""Project an update onto each parameter row's tangent space and normalize."""
theta_hat = _row_normalize(param.to(torch.float32), eps)
update_fp32 = update.to(torch.float32)
radial = (update_fp32 * theta_hat).sum(dim=-1, keepdim=True) * theta_hat
direction = _row_normalize(update_fp32 - radial, eps)
return direction.to(update.dtype)
def nora_lr_scale(lr: float, shape: torch.Size) -> float:
"""Scale Nora's LR for tall ``[d_out, d_in]`` linear weights."""
return lr * math.sqrt(max(1.0, shape[-2] / shape[-1]))
def _validate_complete_rows(param: Tensor) -> None:
if not isinstance(param, DTensor):
return
last_dim = param.ndim - 1
for placement in param.placements:
if isinstance(placement, Shard) and placement.dim % param.ndim == last_dim:
raise ValueError(
"Nora requires complete parameter rows, but this DTensor is sharded "
"along its last dimension"
)
class Nora(Optimizer):
"""Normalized Orthogonal Row Alignment for two-dimensional matrices."""
def __init__(
self,
params,
lr: float = 5e-3,
weight_decay: float = 0.0,
momentum: float = 0.95,
beta: float = 0.95,
nesterov: bool = True,
eps: float = NORA_EPS,
):
if lr < 0:
raise ValueError(f"Invalid learning rate: {lr}")
if weight_decay < 0:
raise ValueError(f"Invalid weight decay: {weight_decay}")
if not 0 <= momentum <= 1:
raise ValueError(f"Invalid momentum: {momentum}")
if not 0 <= beta < 1:
raise ValueError(f"Invalid beta: {beta}")
if eps <= 0:
raise ValueError(f"Invalid epsilon: {eps}")
defaults = {
"lr": lr,
"weight_decay": weight_decay,
"momentum": momentum,
"beta": beta,
"nesterov": nesterov,
"eps": eps,
}
super().__init__(params, defaults)
for group in self.param_groups:
for param in group["params"]:
if param.ndim != 2:
raise ValueError(
f"Nora only supports 2D matrices, got shape {tuple(param.shape)}"
)
_validate_complete_rows(param)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
lr = group["lr"]
weight_decay = group["weight_decay"]
momentum = group["momentum"]
beta = group["beta"]
nesterov = group["nesterov"]
eps = group["eps"]
for param in group["params"]:
if param.grad is None:
continue
if param.grad.is_sparse:
raise RuntimeError("Nora does not support sparse gradients")
grad = param.grad
state = self.state[param]
momentum_buffer = state.get("momentum_buffer")
if momentum_buffer is None:
momentum_buffer = torch.zeros_like(grad)
momentum_buffer.lerp_(grad, 1 - beta)
update = (
grad.lerp(momentum_buffer, momentum)
if nesterov
else momentum_buffer
)
direction = nora_direction(update, param, eps)
if weight_decay != 0:
param.mul_(1 - lr * weight_decay)
param.add_(direction, alpha=-nora_lr_scale(lr, param.shape))
state["momentum_buffer"] = momentum_buffer
return loss
class NAdamW(Optimizer):
"""AdamW using the reference Nesterov first-moment update."""
def __init__(
self,
params,
lr: float = 3e-4,
betas: tuple[float, float] = (0.9, 0.999),
eps: float = 1e-8,
weight_decay: float = 0.1,
):
beta1, beta2 = betas
if lr < 0:
raise ValueError(f"Invalid learning rate: {lr}")
if not 0 <= beta1 < 1 or not 0 <= beta2 < 1:
raise ValueError(f"Invalid betas: {betas}")
if eps <= 0:
raise ValueError(f"Invalid epsilon: {eps}")
if weight_decay < 0:
raise ValueError(f"Invalid weight decay: {weight_decay}")
defaults = {
"lr": lr,
"betas": betas,
"eps": eps,
"weight_decay": weight_decay,
}
super().__init__(params, defaults)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
beta1, beta2 = group["betas"]
eps = group["eps"]
lr = group["lr"]
weight_decay = group["weight_decay"]
for param in group["params"]:
if param.grad is None:
continue
if param.grad.is_sparse:
raise RuntimeError("NAdamW does not support sparse gradients")
grad = param.grad
state = self.state[param]
if not state:
state["step"] = 0
state["m"] = torch.zeros_like(param)
state["v"] = torch.zeros_like(param)
state["step"] += 1
first_moment = state["m"]
second_moment = state["v"]
first_moment.mul_(beta1).add_(grad, alpha=1 - beta1)
second_moment.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
bias_correction1 = 1 - beta1 ** state["step"]
bias_correction2 = 1 - beta2 ** state["step"]
nesterov_moment = (
beta1 * first_moment + (1 - beta1) * grad
) / bias_correction1
corrected_second_moment = second_moment / bias_correction2
if weight_decay != 0:
param.mul_(1 - lr * weight_decay)
param.addcdiv_(
nesterov_moment,
corrected_second_moment.sqrt().add_(eps),
value=-lr,
)
return loss
@dataclass
class OptimizerParameterGroups:
nora: list[Tensor]
nadamw_decay: list[Tensor]
nadamw_no_decay: list[Tensor]
def partition_optimizer_parameters(model: nn.Module) -> OptimizerParameterGroups:
"""Partition trainable parameters by module role and parameter identity."""
nora_ids: set[int] = set()
no_decay_ids: set[int] = set()
for module_name, module in model.named_modules():
if isinstance(module, LoRALinear):
for param in module.parameters(recurse=False):
if param.requires_grad:
no_decay_ids.add(id(param))
continue
if isinstance(module, (Embedding, RMSNorm)):
for param in module.parameters(recurse=False):
if param.requires_grad:
no_decay_ids.add(id(param))
continue
if not isinstance(module, Linear):
continue
if module.bias is not None and module.bias.requires_grad:
no_decay_ids.add(id(module.bias))
if not module.weight.requires_grad:
continue
if module_name.rsplit(".", 1)[-1] == "lm_head":
no_decay_ids.add(id(module.weight))
elif module.weight.ndim == 2:
nora_ids.add(id(module.weight))
nora: list[Tensor] = []
nadamw_decay: list[Tensor] = []
nadamw_no_decay: list[Tensor] = []
seen: set[int] = set()
for param in model.parameters():
param_id = id(param)
if not param.requires_grad or param_id in seen:
continue
seen.add(param_id)
if param_id in no_decay_ids or param.ndim <= 1:
nadamw_no_decay.append(param)
elif param_id in nora_ids:
nora.append(param)
else:
nadamw_decay.append(param)
trainable_ids = {id(param) for param in model.parameters() if param.requires_grad}
grouped_ids = {id(param) for param in [*nora, *nadamw_decay, *nadamw_no_decay]}
if grouped_ids != trainable_ids:
missing = len(trainable_ids - grouped_ids)
extra = len(grouped_ids - trainable_ids)
raise RuntimeError(
f"Optimizer parameter partition is incomplete: missing={missing}, extra={extra}"
)
return OptimizerParameterGroups(nora, nadamw_decay, nadamw_no_decay)
@OptimizerFactory.register("nora_nadamw")
class NoraNAdamW(Optimizer):
"""Nora for internal linear weights and NAdamW for remaining parameters."""
optimizer_name = "nora_nadamw"
def __init__(
self,
model: nn.Module,
lr: float = 3e-4,
weight_decay: float = 0.1,
nora_lr: float = 5e-3,
nora_weight_decay: float = 0.0,
nora_beta: float = 0.95,
nora_momentum: float = 0.95,
):
groups = partition_optimizer_parameters(model)
all_params = [
*groups.nora,
*groups.nadamw_decay,
*groups.nadamw_no_decay,
]
if not all_params:
raise ValueError(
"Cannot build an optimizer for a model with no trainable parameters"
)
super().__init__(all_params, {})
self.nora = (
Nora(
groups.nora,
lr=nora_lr,
weight_decay=nora_weight_decay,
momentum=nora_momentum,
beta=nora_beta,
)
if groups.nora
else None
)
nadamw_groups = []
if groups.nadamw_decay:
nadamw_groups.append(
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
)
if groups.nadamw_no_decay:
nadamw_groups.append(
{"params": groups.nadamw_no_decay, "weight_decay": 0.0}
)
self.nadamw = NAdamW(nadamw_groups, lr=lr) if nadamw_groups else None
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
@torch.no_grad()
def step(self, closure=None):
return composite_step(
[opt for opt in (self.nora, self.nadamw) if opt is not None],
closure,
)
def zero_grad(self, set_to_none: bool = True):
composite_zero_grad(
[opt for opt in (self.nora, self.nadamw) if opt is not None],
set_to_none,
)
def state_dict(self) -> dict[str, Any]:
return composite_state_dict({"nora": self.nora, "nadamw": self.nadamw})
def load_state_dict(self, state_dict: dict[str, Any]):
if "muon" in state_dict or "adamw" in state_dict:
raise ValueError(
"Checkpoint uses muon_adamw state; select optimizer='muon_adamw' "
"to resume it"
)
if "nora" not in state_dict or "nadamw" not in state_dict:
raise ValueError(
"Checkpoint optimizer state is not compatible with nora_nadamw"
)
saved_nora = state_dict["nora"]
saved_nadamw = state_dict["nadamw"]
if (self.nora is None) != (saved_nora is None):
raise ValueError("Checkpoint Nora parameter groups do not match the model")
if (self.nadamw is None) != (saved_nadamw is None):
raise ValueError(
"Checkpoint NAdamW parameter groups do not match the model"
)
if self.nora is not None:
self.nora.load_state_dict(saved_nora)
if self.nadamw is not None:
self.nadamw.load_state_dict(saved_nadamw)
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
+4 -3
View File
@@ -7,8 +7,9 @@ from astrai.parallel.executor import (
FSDPExecutor, FSDPExecutor,
GradientState, GradientState,
NoneExecutor, NoneExecutor,
broadcast_state_dict,
create_ref_model,
) )
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
from astrai.parallel.setup import ( from astrai.parallel.setup import (
get_current_device, get_current_device,
get_rank, get_rank,
@@ -25,8 +26,6 @@ __all__ = [
"only_on_rank", "only_on_rank",
"setup_parallel", "setup_parallel",
"spawn_parallel_fn", "spawn_parallel_fn",
"RowParallelLinear",
"ColumnParallelLinear",
"ExecutorFactory", "ExecutorFactory",
"BaseExecutor", "BaseExecutor",
"GradientState", "GradientState",
@@ -35,4 +34,6 @@ __all__ = [
"NoneExecutor", "NoneExecutor",
"DDPExecutor", "DDPExecutor",
"FSDPExecutor", "FSDPExecutor",
"create_ref_model",
"broadcast_state_dict",
] ]
+193 -76
View File
@@ -4,17 +4,19 @@ import contextlib
import logging import logging
import os import os
from contextlib import contextmanager from contextlib import contextmanager
from typing import Optional, Tuple from typing import Any, Callable, Dict, Optional, Tuple
import torch import torch
import torch.distributed as dist import torch.distributed as dist
import torch.nn as nn import torch.nn as nn
from torch.distributed.fsdp import FullStateDictConfig, StateDictType from torch.distributed.fsdp import (
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP FSDPModule,
fully_shard,
)
from torch.distributed.tensor import DTensor
from torch.nn.parallel import DistributedDataParallel as DDP from torch.nn.parallel import DistributedDataParallel as DDP
from torch.optim import Optimizer from torch.optim import Optimizer
from torch.optim.lr_scheduler import LRScheduler from torch.optim.lr_scheduler import LRScheduler
from torch.utils.data import DataLoader
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.parallel.setup import get_rank, get_world_size from astrai.parallel.setup import get_rank, get_world_size
@@ -22,6 +24,82 @@ from astrai.parallel.setup import get_rank, get_world_size
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def broadcast_state_dict(
state_dict: Optional[Dict[str, torch.Tensor]],
src: int = 0,
) -> Optional[Dict[str, torch.Tensor]]:
"""Broadcast a state_dict from *src* rank to all ranks.
Tensors stay on their original device (GPU) for the broadcast.
All ranks must call this collectively.
On non-distributed runs, returns *state_dict* unchanged.
"""
if not dist.is_initialized() or dist.get_world_size() == 1:
return state_dict
rank = dist.get_rank()
# Broadcast metadata (keys, shapes, dtypes, device) so non-src ranks
# can allocate matching empty tensors on the correct device.
if rank == src:
device = next(iter(state_dict.values())).device
metadata = [
(k, tuple(v.shape), v.dtype, str(device)) for k, v in state_dict.items()
]
else:
metadata = None
metadata_list = [metadata]
dist.broadcast_object_list(metadata_list, src=src)
metadata = metadata_list[0]
# Non-src ranks allocate empty tensors with the broadcasted metadata.
if rank != src:
state_dict = {
k: torch.empty(s, dtype=d, device=torch.device(dev))
for k, s, d, dev in metadata
}
# Broadcast each tensor in-place.
for tensor in state_dict.values():
dist.broadcast(tensor, src=src)
return state_dict
def create_ref_model(
model_fn: Callable[[], nn.Module],
executor: Optional["BaseExecutor"] = None,
model: Optional[nn.Module] = None,
state_dict: Optional[Dict[str, torch.Tensor]] = None,
device: Optional[str] = None,
) -> Optional[nn.Module]:
"""Create a frozen reference model from executor or state dict.
In distributed mode (FSDP), ``unwrap_model`` returns ``None`` on
non-rank-0. The state_dict is broadcast from rank-0 to all ranks
so every rank gets a complete copy.
"""
if state_dict is None and executor is not None and model is not None:
state_dict = executor.unwrap_model(model)
# FSDP's unwrap_model returns None on non-rank-0. Broadcast from
# rank-0 so every rank receives a complete state_dict.
if executor is not None and executor.use_distributed:
state_dict = broadcast_state_dict(state_dict)
if state_dict is None:
return None
ref_model = model_fn()
ref_model.load_state_dict(state_dict)
ref_model.requires_grad_(False)
ref_model.eval()
if device is not None:
ref_model = ref_model.to(device=device)
return ref_model
class GradientState: class GradientState:
def __init__(self, grad_accum_steps: int = 1): def __init__(self, grad_accum_steps: int = 1):
self.num_steps = max(grad_accum_steps, 1) self.num_steps = max(grad_accum_steps, 1)
@@ -86,19 +164,28 @@ class BaseExecutor:
def prepare( def prepare(
self, self,
model: nn.Module, model_fn: Callable[[], nn.Module],
optimizer: Optional[Optimizer] = None, optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
dataloader: Optional[DataLoader] = None, scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
scheduler: Optional[LRScheduler] = None, before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
) -> Tuple[ after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
nn.Module, Optional[Optimizer], Optional[DataLoader], Optional[LRScheduler] ) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
]: model = model_fn()
if before_wrap is not None:
model = before_wrap(model)
model = self._prepare_model(model) model = self._prepare_model(model)
if optimizer is not None: if after_wrap is not None:
model = after_wrap(model)
optimizer = None
scheduler = None
if optimizer_fn is not None:
optimizer = optimizer_fn(model)
if scheduler_fn is not None:
scheduler = scheduler_fn(optimizer)
optimizer = AccumOptimizer(optimizer, self.gradient_state) optimizer = AccumOptimizer(optimizer, self.gradient_state)
if scheduler is not None: if scheduler is not None:
scheduler = AccumScheduler(scheduler, self.gradient_state) scheduler = AccumScheduler(scheduler, self.gradient_state)
return model, optimizer, dataloader, scheduler return model, optimizer, scheduler
def _prepare_model(self, model: nn.Module) -> nn.Module: def _prepare_model(self, model: nn.Module) -> nn.Module:
return model return model
@@ -148,14 +235,7 @@ class BaseExecutor:
def grad_accum_steps(self) -> int: def grad_accum_steps(self) -> int:
return self.gradient_state.num_steps return self.gradient_state.num_steps
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float: def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
if max_norm is None:
total_norm = torch.norm(
torch.stack(
[p.grad.norm(2) for p in model.parameters() if p.grad is not None]
)
)
return total_norm.item()
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
if isinstance(total_norm, torch.Tensor): if isinstance(total_norm, torch.Tensor):
return total_norm.item() return total_norm.item()
@@ -234,78 +314,115 @@ class DDPExecutor(BaseExecutor):
@ExecutorFactory.register("fsdp") @ExecutorFactory.register("fsdp")
class FSDPExecutor(BaseExecutor): class FSDPExecutor(BaseExecutor):
"""FSDP executor using `torch.distributed.fsdp.fully_shard` (per-module API).
Wraps each child module individually via ``fully_shard``.
Skips the root model because ``ABC + Generic[T]`` in the MRO makes
``fully_shard``'s dynamic ``__class__`` assignment fail at the CPython level.
Original ``Parameter`` objects are preserved (as DTensors) — no
``FlatParameter``, no ``use_orig_params=True`` hack.
"""
def __init__( def __init__(
self, self,
grad_accum_steps: int = 1, grad_accum_steps: int = 1,
process_group=None, mesh: Optional[Any] = None,
sharding_strategy=None, mp_policy: Optional[Any] = None,
cpu_offload=None, reshard_after_forward: bool = False,
auto_wrap_policy=None,
backward_prefetch=None,
mixed_precision=None,
ignored_modules=None,
param_init_fn=None,
sync_module_states: bool = False,
forward_prefetch: bool = False,
limit_all_gathers: bool = True,
ignored_states=None,
device_mesh=None,
): ):
super().__init__(grad_accum_steps=grad_accum_steps) super().__init__(grad_accum_steps=grad_accum_steps)
self._fsdp_kwargs = { self._mesh = mesh
k: v self._mp_policy = mp_policy
for k, v in dict( self._reshard_after_forward = reshard_after_forward
process_group=process_group,
sharding_strategy=sharding_strategy,
cpu_offload=cpu_offload,
auto_wrap_policy=auto_wrap_policy,
backward_prefetch=backward_prefetch,
mixed_precision=mixed_precision,
ignored_modules=ignored_modules,
param_init_fn=param_init_fn,
sync_module_states=sync_module_states,
forward_prefetch=forward_prefetch,
limit_all_gathers=limit_all_gathers,
use_orig_params=True,
ignored_states=ignored_states,
device_mesh=device_mesh,
).items()
if v is not None
}
self._original_model: Optional[nn.Module] = None
def _prepare_model(self, model: nn.Module) -> nn.Module: def _prepare_model(self, model: nn.Module) -> nn.Module:
if not self.use_distributed: if not self.use_distributed:
logger.warning("FSDP backend selected but world_size=1, model not wrapped") logger.warning("FSDP backend selected but world_size=1, model not wrapped")
return model return model
self._original_model = model
device_id = torch.device("cuda", get_rank()) kwargs = dict(
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs) mesh=self._mesh,
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size()) mp_policy=self._mp_policy,
reshard_after_forward=self._reshard_after_forward,
)
kwargs = {k: v for k, v in kwargs.items() if v is not None}
for child in model.children():
if isinstance(child, nn.ModuleList):
for sub in child:
fully_shard(sub, **kwargs)
else:
fully_shard(child, **kwargs)
logger.info(
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
len(list(model.children())),
)
return model return model
@contextmanager
def _no_sync(self, model: nn.Module): def _no_sync(self, model: nn.Module):
if isinstance(model, FSDP): fsdp_modules = [m for m in model.modules() if isinstance(m, FSDPModule)]
return model.no_sync() if fsdp_modules:
return contextlib.nullcontext() for m in fsdp_modules:
m.set_requires_gradient_sync(False, recurse=True)
try:
yield
finally:
for m in fsdp_modules:
m.set_requires_gradient_sync(True, recurse=True)
else:
yield
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float: def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
if max_norm is None: if not self.use_distributed:
return super().clip_grad_norm(model, max_norm) return super().clip_grad_norm(model, max_norm)
if isinstance(model, FSDP) and self.use_distributed:
total_norm = model.clip_grad_norm_(max_norm) # FSDP params are DTensors (sharded across ranks).
if isinstance(total_norm, torch.Tensor): # torch.nn.utils.clip_grad_norm_ computes LOCAL norm per rank,
# so we must all-reduce to get the global norm before clipping.
local_norm = torch.nn.utils.get_total_norm(
[p.grad for p in model.parameters() if p.grad is not None],
)
if isinstance(local_norm, DTensor):
local_norm = local_norm.to_local()
total_norm_sq = local_norm**2
dist.all_reduce(total_norm_sq, op=dist.ReduceOp.SUM)
total_norm = total_norm_sq.sqrt()
clip_coef = max_norm / (total_norm + 1e-6)
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
for p in model.parameters():
if p.grad is not None:
p.grad.mul_(clip_coef_clamped)
return total_norm.item() return total_norm.item()
return total_norm
return super().clip_grad_norm(model, max_norm)
def unwrap_model(self, model: nn.Module): def unwrap_model(self, model: nn.Module):
if isinstance(model, FSDP) and self.use_distributed: if not self.use_distributed:
with FSDP.state_dict_type(
model,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
return model.state_dict() return model.state_dict()
return model.state_dict() # unshard() and full_tensor() are collective ops — all ranks must
# participate. Non-rank-0 ranks still call them but discard results.
for module in model.modules():
if isinstance(module, FSDPModule):
module.unshard()
state_dict = model.state_dict()
result = {}
for k, v in state_dict.items():
if isinstance(v, DTensor):
full = v.full_tensor()
if get_rank() == 0:
result[k] = full
elif get_rank() == 0:
result[k] = v
for module in model.modules():
if isinstance(module, FSDPModule):
module.reshard()
if get_rank() != 0:
return None
return result
-115
View File
@@ -1,115 +0,0 @@
from typing import Dict
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
class ParallelModel(nn.Module):
def __init__(self, process_group: dist.ProcessGroup):
super().__init__()
self.process_group = process_group
self.rank = dist.get_rank(self.process_group)
self.world_size = dist.get_world_size(self.process_group)
class RowParallelLinear(ParallelModel):
def __init__(
self,
process_group: dist.ProcessGroup,
in_features: int,
out_features: int,
bias: bool = True,
reduce_results: bool = True,
):
super().__init__(process_group)
self.in_features = in_features
self.out_features = out_features
self.in_features_per_rank = in_features // self.world_size
self.reduce_results = reduce_results
if in_features % self.world_size != 0:
raise ValueError(
f"in_features must be divisible by world_size. Got {in_features} and {self.world_size}"
)
self.weight = nn.Parameter(torch.empty(out_features, self.in_features_per_rank))
self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None
def forward(self, input: Tensor) -> Tensor:
output = F.linear(input, self.weight)
if self.reduce_results:
dist.all_reduce(output, op=dist.ReduceOp.SUM, group=self.process_group)
if self.bias is not None:
output += self.bias
return output
def load_state_dict(self, state_dict: Dict[str, Tensor]):
full_weight = state_dict.get("weight")
full_bias = state_dict.get("bias")
start_idx = self.rank * self.in_features_per_rank
end_idx = start_idx + self.in_features_per_rank
weight_slice = full_weight[:, start_idx:end_idx]
self.weight.data.copy_(weight_slice)
if self.bias is not None:
self.bias.data.copy_(full_bias)
class ColumnParallelLinear(ParallelModel):
def __init__(
self,
process_group: dist.ProcessGroup,
in_features: int,
out_features: int,
bias: bool = True,
gather_results: bool = True,
):
super().__init__(process_group)
self.in_features = in_features
self.out_features = out_features
self.out_features_per_rank = out_features // self.world_size
self.gather_results = gather_results
if out_features % self.world_size != 0:
raise ValueError(
f"out_features must be divisible by world_size. Got {out_features} and {self.world_size}"
)
self.weight = nn.Parameter(
torch.empty(self.out_features_per_rank, self.in_features)
)
self.bias = (
nn.Parameter(torch.zeros(self.out_features_per_rank)) if bias else None
)
def forward(self, input: Tensor) -> Tensor:
output = F.linear(input, self.weight, self.bias)
if self.gather_results:
output_list = [torch.empty_like(output) for _ in range(self.world_size)]
dist.all_gather(output_list, output, group=self.process_group)
output = torch.cat(output_list, dim=-1)
return output
def load_state_dict(self, state_dict: Dict[str, Tensor]):
full_weight = state_dict.get("weight")
full_bias = state_dict.get("bias")
start_idx = self.rank * self.out_features_per_rank
end_idx = start_idx + self.out_features_per_rank
weight_slice = full_weight[start_idx:end_idx, :]
self.weight.data.copy_(weight_slice)
if self.bias is not None:
bias_slice = full_bias[start_idx:end_idx]
self.bias.data.copy_(bias_slice)
+44 -2
View File
@@ -1,5 +1,8 @@
import logging
import os import os
import signal
import socket import socket
import threading
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from functools import wraps from functools import wraps
@@ -9,6 +12,10 @@ import torch
import torch.distributed as dist import torch.distributed as dist
import torch.multiprocessing as mp import torch.multiprocessing as mp
from astrai.signal_handler import install_early_signal_handlers
logger = logging.getLogger(__name__)
def find_free_port() -> str: def find_free_port() -> str:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
@@ -115,6 +122,7 @@ def _run_single_rank(
func: Callable, func: Callable,
kwargs: dict, kwargs: dict,
): ):
install_early_signal_handlers()
with setup_parallel( with setup_parallel(
rank=rank, rank=rank,
world_size=world_size, world_size=world_size,
@@ -155,6 +163,7 @@ class TorchrunStrategy(LaunchStrategy):
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set.""" """External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
def launch(self, func: Callable, **kwargs): def launch(self, func: Callable, **kwargs):
install_early_signal_handlers()
rank = int(os.environ["RANK"]) rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"]) world_size = int(os.environ["WORLD_SIZE"])
local_rank = int(os.environ.get("LOCAL_RANK", rank)) local_rank = int(os.environ.get("LOCAL_RANK", rank))
@@ -188,6 +197,7 @@ class LocalStrategy(LaunchStrategy):
_run_single_rank(0, *args) _run_single_rank(0, *args)
return return
install_early_signal_handlers()
ctx = mp.start_processes( ctx = mp.start_processes(
_run_single_rank, _run_single_rank,
args=args, args=args,
@@ -195,14 +205,46 @@ class LocalStrategy(LaunchStrategy):
start_method=self.start_method, start_method=self.start_method,
join=False, join=False,
) )
parent_stop = threading.Event()
original_handlers = {}
def _parent_handler(signum, frame):
sig = signal.Signals(signum)
logger.warning(
"Parent (pid=%d) received %s, forwarding to children...",
os.getpid(),
sig.name,
)
parent_stop.set()
for p in ctx.processes:
if p.is_alive():
p.terminate()
for sig in (signal.SIGTERM, signal.SIGINT):
prev = signal.signal(sig, _parent_handler)
if prev not in (signal.SIG_DFL, signal.SIG_IGN, None, _parent_handler):
original_handlers[sig] = prev
try: try:
while not ctx.join(): while not ctx.join() and not parent_stop.is_set():
pass pass
except BaseException: except BaseException:
logger.warning(
"Parent received unexpected exception, terminating children..."
)
for p in ctx.processes: for p in ctx.processes:
if p.is_alive():
p.terminate() p.terminate()
ctx.join()
raise raise
finally:
for sig, handler in original_handlers.items():
signal.signal(sig, handler)
for p in ctx.processes:
p.join()
ctx.join()
def _detect_launcher() -> str: def _detect_launcher() -> str:
+211 -6
View File
@@ -94,6 +94,97 @@ class SectionRenderer:
return all_ids, loss_mask return all_ids, loss_mask
def process_sections_batch(
self,
items: list[dict],
sections: list,
config,
tokenizer,
*,
is_top_level=False,
filter_text=True,
):
"""Render and tokenize a group of records with batched Rust tokenization."""
has_template = any(s.get("template") for s in sections)
is_text_config = not has_template and all(
s["action"] == "train" for s in sections
)
plans: list[list[tuple[str, str, bool]]] = []
for item in items:
plan: list[tuple[str, str, bool]] = []
first_section = True
for sec in sections:
field = sec["field"]
action = sec["action"]
use_template = sec.get("template", False)
add_special = sec.get(
"add_special_tokens", not use_template and first_section
)
if use_template:
messages = item.get(field)
if not isinstance(messages, list) or not messages:
continue
for msg in messages:
role = msg.get("role", "")
rendered = tokenizer.apply_chat_template(
[msg], tokenize=False, add_generation_prompt=False
)
plan.append(
(rendered, _resolve_action(action, role, config), False)
)
else:
text = str(item.get(field, ""))
if not text.strip():
continue
if is_text_config and filter_text:
pp = config.preprocessing
if pp.min_chars > 0 and len(text) < pp.min_chars:
continue
if len(text) > pp.max_chars:
continue
plan.append((text, action, add_special))
first_section = False
plans.append(plan)
encoded: dict[tuple[int, int], list[int]] = {}
for add_special in (False, True):
refs = [
(item_idx, unit_idx, text)
for item_idx, plan in enumerate(plans)
for unit_idx, (text, _, add) in enumerate(plan)
if add == add_special
]
if not refs:
continue
ids_batch = tokenizer.encode(
[text for _, _, text in refs], add_special_tokens=add_special
)
for (item_idx, unit_idx, _), ids in zip(refs, ids_batch):
encoded[(item_idx, unit_idx)] = ids
outputs = []
max_len = config.preprocessing.max_seq_len
for item_idx, plan in enumerate(plans):
all_ids = []
loss_mask = []
if is_top_level and has_template and tokenizer.bos_token_id is not None:
all_ids.append(tokenizer.bos_token_id)
loss_mask.append(0)
for unit_idx, (_, action, _) in enumerate(plan):
ids = encoded[(item_idx, unit_idx)]
all_ids.extend(ids)
loss_mask.extend([1 if action == "train" else 0] * len(ids))
all_ids = all_ids[:max_len]
loss_mask = loss_mask[: len(all_ids)]
if not all_ids or (is_top_level and has_template and len(all_ids) <= 1):
outputs.append((None, None))
else:
outputs.append((all_ids, loss_mask))
return outputs
def process_list_field(self, item: dict, sections: list, config, tokenizer): def process_list_field(self, item: dict, sections: list, config, tokenizer):
"""Tokenize a list-valued field, preserving per-element boundaries. """Tokenize a list-valued field, preserving per-element boundaries.
@@ -147,6 +238,42 @@ class SectionRenderer:
return None, None return None, None
return per_item_ids, per_item_masks return per_item_ids, per_item_masks
def process_list_field_batch(self, items, sections, config, tokenizer):
per_item_ids = [[] for _ in items]
per_item_masks = [[] for _ in items]
for sec in sections:
wrappers = []
owners = []
field = sec["field"]
for item_idx, item in enumerate(items):
values = item.get(field)
if not isinstance(values, list):
continue
for val in values:
if sec.get("template", False) and not isinstance(val, list):
continue
wrappers.append({field: val if isinstance(val, list) else str(val)})
owners.append(item_idx)
rendered = self.process_sections_batch(
wrappers,
[sec],
config,
tokenizer,
is_top_level=False,
filter_text=False,
)
for owner, (ids, mask) in zip(owners, rendered):
if ids:
per_item_ids[owner].append(ids)
per_item_masks[owner].append(mask)
return [
(ids, masks) if ids else (None, None)
for ids, masks in zip(per_item_ids, per_item_masks)
]
@staticmethod @staticmethod
def is_value_section(sections: list) -> bool: def is_value_section(sections: list) -> bool:
return len(sections) == 1 and sections[0].get("action") == "value" return len(sections) == 1 and sections[0].get("action") == "value"
@@ -214,6 +341,9 @@ class BaseMaskBuilder(ABC):
@abstractmethod @abstractmethod
def build(self, item: dict, config, tokenizer) -> Optional[dict]: ... def build(self, item: dict, config, tokenizer) -> Optional[dict]: ...
def build_batch(self, items: list[dict], config, tokenizer) -> list[Optional[dict]]:
return [self.build(item, config, tokenizer) for item in items]
class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]): class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
pass pass
@@ -248,6 +378,27 @@ class SingleOutputMaskBuilder(BaseMaskBuilder):
result["loss_mask"] = mask result["loss_mask"] = mask
return result return result
def build_batch(self, items, config, tokenizer):
sections = config.input.sections
if not sections:
return [None] * len(items)
rendered = self.renderer.process_sections_batch(
items, sections, config, tokenizer, is_top_level=True
)
results = []
for item, (ids, mask) in zip(items, rendered):
if ids is None:
results.append(None)
continue
result = {
"sequence": ids,
"domain": _extract_domain(item, config.output.domain_key),
}
if not all(m == 1 for m in mask):
result["loss_mask"] = mask
results.append(result)
return results
@MaskBuilderFactory.register("multi") @MaskBuilderFactory.register("multi")
class MultiOutputMaskBuilder(BaseMaskBuilder): class MultiOutputMaskBuilder(BaseMaskBuilder):
@@ -265,7 +416,11 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
return None return None
result: dict = {} result: dict = {}
any_output = False required_outputs = {
output_key
for output_key, spec in sources_spec.items()
if spec.get("sections")
}
for output_key, spec in sources_spec.items(): for output_key, spec in sources_spec.items():
sections = spec.get("sections", []) sections = spec.get("sections", [])
@@ -277,7 +432,6 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
if ids is None: if ids is None:
continue continue
result[output_key] = ids result[output_key] = ids
any_output = True
continue continue
list_field = spec.get("list_field", False) list_field = spec.get("list_field", False)
@@ -293,7 +447,6 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
result[output_key] = ids result[output_key] = ids
if mask is not None: if mask is not None:
result[mask_key] = mask result[mask_key] = mask
any_output = True
continue continue
ids, mask = self.renderer.process_sections( ids, mask = self.renderer.process_sections(
@@ -309,14 +462,60 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
elif "mask_key" in spec: elif "mask_key" in spec:
result[mask_key] = mask result[mask_key] = mask
any_output = True if not required_outputs or not required_outputs.issubset(result):
if not any_output:
return None return None
result["domain"] = _extract_domain(item, config.output.domain_key) result["domain"] = _extract_domain(item, config.output.domain_key)
return result return result
def build_batch(self, items, config, tokenizer):
sources_spec = getattr(config.input, "sources", None)
if not sources_spec:
return [None] * len(items)
results = [{} for _ in items]
required_outputs = {
output_key
for output_key, spec in sources_spec.items()
if spec.get("sections")
}
for output_key, spec in sources_spec.items():
sections = spec.get("sections", [])
if not sections:
continue
if self.renderer.is_value_section(sections):
for item, result in zip(items, results):
value = self.renderer.extract_raw_value(item, sections)
if value is not None:
result[output_key] = value
continue
mask_key = spec.get("mask_key", f"{output_key}_mask")
if spec.get("list_field", False):
rendered = self.renderer.process_list_field_batch(
items, sections, config, tokenizer
)
else:
rendered = self.renderer.process_sections_batch(
items, sections, config, tokenizer, is_top_level=True
)
for result, (ids, mask) in zip(results, rendered):
if ids is None:
continue
result[output_key] = ids
if spec.get("list_field", False) or not all(m == 1 for m in mask):
result[mask_key] = mask
elif "mask_key" in spec:
result[mask_key] = mask
return [
({**result, "domain": _extract_domain(item, config.output.domain_key)})
if required_outputs and required_outputs.issubset(result)
else None
for item, result in zip(items, results)
]
@MaskBuilderFactory.register("sectioned") @MaskBuilderFactory.register("sectioned")
class SectionedMaskBuilder(BaseMaskBuilder): class SectionedMaskBuilder(BaseMaskBuilder):
@@ -335,3 +534,9 @@ class SectionedMaskBuilder(BaseMaskBuilder):
if sources_spec: if sources_spec:
return self._multi.build(item, config, tokenizer) return self._multi.build(item, config, tokenizer)
return self._single.build(item, config, tokenizer) return self._single.build(item, config, tokenizer)
def build_batch(self, items, config, tokenizer):
sources_spec = getattr(config.input, "sources", None)
if sources_spec:
return self._multi.build_batch(items, config, tokenizer)
return self._single.build_batch(items, config, tokenizer)
+40 -11
View File
@@ -1,7 +1,7 @@
"""Config-driven JSONL preprocessing pipeline. """Config-driven JSONL preprocessing pipeline.
Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id sharding and flush to ``.bin`` storage. Packing, position-id
generation and storage writing are each delegated to pluggable strategies, generation and storage writing are each delegated to pluggable strategies,
dispatched by configuration keys. dispatched by configuration keys.
@@ -23,7 +23,6 @@ import tqdm
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.core import ( from astrai.preprocessing.core import (
build_preprocessing_components, build_preprocessing_components,
iter_raw_records,
primary_ids, primary_ids,
) )
from astrai.preprocessing.packing import PackingStrategyFactory from astrai.preprocessing.packing import PackingStrategyFactory
@@ -81,6 +80,9 @@ class Pipeline:
def transform(self, item: dict) -> Optional[dict]: def transform(self, item: dict) -> Optional[dict]:
return self.mask_builder.build(item, self.config, self.tokenizer) return self.mask_builder.build(item, self.config, self.tokenizer)
def transform_batch(self, items: list[dict]) -> list[Optional[dict]]:
return self.mask_builder.build_batch(items, self.config, self.tokenizer)
def run(self): def run(self):
domains: dict = defaultdict(lambda: defaultdict(list)) domains: dict = defaultdict(lambda: defaultdict(list))
total_tokens = 0 total_tokens = 0
@@ -89,19 +91,31 @@ class Pipeline:
pp = self.config.preprocessing pp = self.config.preprocessing
for item in tqdm.tqdm( progress = tqdm.tqdm(desc="Tokenizing", unit="docs", mininterval=0.5)
self._iter_items(), desc="Tokenizing", unit="docs", mininterval=0.5 stop = False
): for items in self._iter_batches(pp.batch_size):
if pp.max_items and count >= pp.max_items: progress.update(len(items))
break
try: try:
result = self.transform(item) results = self.transform_batch(items)
except Exception: except Exception:
logger.warning( logger.warning(
"Failed to process item #%d, skipping", count + 1, exc_info=True "Failed to process batch, retrying records individually",
exc_info=True,
) )
continue results = []
for item in items:
try:
results.append(self.transform(item))
except Exception:
logger.warning(
"Failed to process item, skipping", exc_info=True
)
results.append(None)
for result in results:
if pp.max_items and count >= pp.max_items:
stop = True
break
if result is None: if result is None:
continue continue
@@ -122,6 +136,10 @@ class Pipeline:
self._flush(domains, shard_idx) self._flush(domains, shard_idx)
domains.clear() domains.clear()
total_tokens = 0 total_tokens = 0
if stop:
break
progress.close()
if total_tokens > 0: if total_tokens > 0:
self._flush(domains, shard_idx) self._flush(domains, shard_idx)
@@ -150,6 +168,17 @@ class Pipeline:
continue continue
yield json.loads(line) yield json.loads(line)
def _iter_batches(self, batch_size: int):
batch_size = max(1, batch_size)
batch = []
for item in self._iter_items():
batch.append(item)
if len(batch) >= batch_size:
yield batch
batch = []
if batch:
yield batch
def _flush(self, domains, shard_idx): def _flush(self, domains, shard_idx):
for domain, keys in domains.items(): for domain, keys in domains.items():
idx = shard_idx[domain] idx = shard_idx[domain]
+2 -21
View File
@@ -1,7 +1,7 @@
"""Storage writer strategies for pipeline output. """Storage writer strategies for pipeline output.
The :class:`StoreWriter` abstraction decouples the pipeline from the The :class:`StoreWriter` abstraction decouples the pipeline from the
concrete storage format (bin / h5). The pipeline builds a ``{key: concrete storage format (bin). The pipeline builds a ``{key:
List[Tensor]}`` dict and delegates the write to the writer selected List[Tensor]}`` dict and delegates the write to the writer selected
by ``output.storage_format``. by ``output.storage_format``.
""" """
@@ -15,7 +15,7 @@ from typing import Dict, List
import torch import torch
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.serialization import save_bin, save_h5 from astrai.serialization import save_bin
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -54,22 +54,3 @@ class BinWriter(StoreWriter):
exc_info=True, exc_info=True,
) )
raise raise
@StoreWriterFactory.register("h5")
class H5Writer(StoreWriter):
def save(self, output_dir, domain, shard_idx, tensors):
chunk_dir = os.path.join(output_dir, domain)
file_path = os.path.join(chunk_dir, f"data_{shard_idx:04d}.h5")
try:
save_h5(chunk_dir, f"data_{shard_idx:04d}", tensors)
except Exception:
if os.path.exists(file_path):
os.remove(file_path)
logger.error(
"Failed to write shard %s/data_%04d.h5, cleaned up partial output",
domain,
shard_idx,
exc_info=True,
)
raise
-4
View File
@@ -20,9 +20,7 @@ from astrai.serialization.checkpoint import (
from astrai.serialization.dataset import ( from astrai.serialization.dataset import (
load_bin, load_bin,
load_bin_offsets, load_bin_offsets,
load_h5,
save_bin, save_bin,
save_h5,
) )
__all__ = [ __all__ = [
@@ -39,7 +37,5 @@ __all__ = [
"save_torch", "save_torch",
"load_bin", "load_bin",
"load_bin_offsets", "load_bin_offsets",
"load_h5",
"save_bin", "save_bin",
"save_h5",
] ]
+5 -46
View File
@@ -1,55 +1,14 @@
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary).""" """Dataset storage serialization helpers (memory-mapped binary)."""
import json import json
import os import os
from pathlib import Path
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
import h5py
import numpy as np import numpy as np
import torch import torch
from torch import Tensor from torch import Tensor
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
os.makedirs(file_path, exist_ok=True)
full_file_path = os.path.join(file_path, f"{file_name}.h5")
with h5py.File(full_file_path, "w") as f:
for key, tensors in tensor_group.items():
grp = f.create_group(key)
for idx, tensor in enumerate(tensors):
arr = tensor.cpu().numpy()
grp.create_dataset(f"data_{idx}", data=arr)
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
tensor_group: Dict[str, List[Tensor]] = {}
root_path = Path(file_path)
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
h5_files = [root_path]
else:
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
for h5_file in h5_files:
with h5py.File(h5_file, "r") as f:
for key in f.keys():
grp = f[key]
dsets = []
for dset_name in grp.keys():
dset = grp[dset_name]
tensor = torch.from_numpy(dset[:])
if share_memory:
tensor = tensor.share_memory_()
dsets.append(tensor)
if tensor_group.get(key) is None:
tensor_group[key] = []
tensor_group[key].extend(dsets)
return tensor_group
def save_bin( def save_bin(
file_path: str, file_path: str,
tensor_group: Dict[str, List[Tensor]], tensor_group: Dict[str, List[Tensor]],
@@ -65,7 +24,7 @@ def save_bin(
offsets, preserving backward compatibility. offsets, preserving backward compatibility.
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
not supported in bin format use H5 for those. not supported in bin format use JSONL for those.
""" """
os.makedirs(file_path, exist_ok=True) os.makedirs(file_path, exist_ok=True)
record_keys = set(record_keys or []) record_keys = set(record_keys or [])
@@ -74,7 +33,7 @@ def save_bin(
if tensors and isinstance(tensors[0], list): if tensors and isinstance(tensors[0], list):
raise ValueError( raise ValueError(
f"Nested key '{key}' (List[List[Tensor]]) is not supported " f"Nested key '{key}' (List[List[Tensor]]) is not supported "
f"in bin format. Use H5 or JSONL storage instead." f"in bin format. Use JSONL storage instead."
) )
cat = torch.cat(tensors, dim=0) cat = torch.cat(tensors, dim=0)
entry: Dict[str, Any] = { entry: Dict[str, Any] = {
@@ -100,7 +59,7 @@ def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
arr = np.memmap( arr = np.memmap(
os.path.join(file_path, f"{key}.bin"), os.path.join(file_path, f"{key}.bin"),
dtype=info["dtype"], dtype=info["dtype"],
mode="r", mode="c",
shape=tuple(info["shape"]), shape=tuple(info["shape"]),
) )
segments[key] = [torch.from_numpy(arr)] segments[key] = [torch.from_numpy(arr)]
@@ -112,7 +71,7 @@ def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
Returns an empty dict when no key has offsets (legacy bin files), Returns an empty dict when no key has offsets (legacy bin files),
in which case record-mode access falls back to per-record segment in which case record-mode access falls back to per-record segment
indexing (H5/JSONL layout). indexing (JSONL layout).
""" """
with open(os.path.join(file_path, "meta.json"), "r") as f: with open(os.path.join(file_path, "meta.json"), "r") as f:
meta = json.load(f) meta = json.load(f)
+53
View File
@@ -0,0 +1,53 @@
import logging
import os
import signal
import threading
logger = logging.getLogger(__name__)
_early_stop = threading.Event()
_active_context = None
def _early_handler(signum: int, frame):
sig = signal.Signals(signum)
logger.warning(
"Received %s (pid=%d), requesting graceful training stop...",
sig.name,
os.getpid(),
)
_early_stop.set()
if _active_context is not None:
_active_context.request_stop()
def install_early_signal_handlers():
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
_unblock_signals()
def _unblock_signals():
try:
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
blocked = {signal.SIGTERM, signal.SIGINT} & mask
if blocked:
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
except (AttributeError, OSError):
pass
def register_signal_handlers(context):
global _active_context
_active_context = context
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _early_handler)
if _early_stop.is_set():
context.request_stop()
logger.warning("Signal was received during initialization, stopping...")
def unregister_signal_handlers():
global _active_context
_active_context = None
_early_stop.clear()
+3 -1
View File
@@ -1,8 +1,10 @@
from astrai.tokenize.chat_template import ChatTemplate, MessageType from astrai.tokenize.chat_template import ChatTemplate, MessageType
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer, Message, Messages
__all__ = [ __all__ = [
"AutoTokenizer", "AutoTokenizer",
"ChatTemplate", "ChatTemplate",
"MessageType", "MessageType",
"Message",
"Messages",
] ]
+18 -3
View File
@@ -38,12 +38,27 @@ class ChatTemplate:
The compiled :class:`~jinja2.Template` holds a dynamically-generated The compiled :class:`~jinja2.Template` holds a dynamically-generated
``root`` render function whose ``__module__`` is ``None``; under ``root`` render function whose ``__module__`` is ``None``; under
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based ``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
multiprocessing. By deferring compilation to first access, the multiprocessing. :meth:`__getstate__` drops the cached template so
default pickle protocol serialises only ``template_str``; each that pickle serialises only ``template_str``; each worker rebuilds
worker rebuilds the cache on first render. the cache on first render.
""" """
return Template(self.template_str) return Template(self.template_str)
def __getstate__(self) -> Dict[str, Any]:
"""Exclude the cached Jinja2 template from pickling.
``Template.root_render_func`` is a dynamically generated closure
that cannot be pickled by reference. Dropping ``_compiled`` here
lets :class:`cached_property` rebuild it on first access after
unpickle.
"""
state = self.__dict__.copy()
state.pop("_compiled", None)
return state
def __setstate__(self, state: Dict[str, Any]) -> None:
self.__dict__.update(state)
@classmethod @classmethod
def from_string( def from_string(
cls, cls,
+53 -35
View File
@@ -10,12 +10,16 @@ from tokenizers import Tokenizer
from astrai.tokenize.chat_template import ChatTemplate from astrai.tokenize.chat_template import ChatTemplate
Message = Dict[str, str]
"""Single chat message with ``role`` and ``content`` keys."""
Messages = List[Message]
"""Single conversation — a list of messages."""
class AutoTokenizer: class AutoTokenizer:
"""Base tokenizer class with automatic loading support""" """Base tokenizer class with automatic loading support"""
TOKENIZER_CLASSES = {} # Registry for auto-loading
def __init__( def __init__(
self, self,
path: Optional[Union[str, Path]] = None, path: Optional[Union[str, Path]] = None,
@@ -102,17 +106,6 @@ class AutoTokenizer:
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f: with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
json.dump(config, f, ensure_ascii=False, indent=2) json.dump(config, f, ensure_ascii=False, indent=2)
@classmethod
def register_tokenizer(cls, name: str, tokenizer_class: type):
"""
Register a new tokenizer class.
Args:
name: Name to register the tokenizer class under
tokenizer_class: The tokenizer class to register
"""
cls.TOKENIZER_CLASSES[name] = tokenizer_class
def encode( def encode(
self, self,
tokens: Union[str, List[str]], tokens: Union[str, List[str]],
@@ -120,7 +113,16 @@ class AutoTokenizer:
is_pretokenized: bool = False, is_pretokenized: bool = False,
add_special_tokens: bool = True, add_special_tokens: bool = True,
) -> List: ) -> List:
"""Encode text to tokens or token IDs.""" """Encode text to token IDs.
Accepts both single strings and batches:
- ``encode("hello")`` ``[123, 456]``
- ``encode(["hello", "world"])`` ``[[123, 456], [789]]``
Batches are tokenised in parallel via the Rust backend's
``encode_batch`` (uses all available CPU cores).
"""
if self._tokenizer is None: if self._tokenizer is None:
raise RuntimeError( raise RuntimeError(
"Tokenizer not initialized. Load or create a tokenizer first." "Tokenizer not initialized. Load or create a tokenizer first."
@@ -133,15 +135,13 @@ class AutoTokenizer:
add_special_tokens=add_special_tokens, add_special_tokens=add_special_tokens,
) )
return encoded.ids if out_ids else encoded.tokens return encoded.ids if out_ids else encoded.tokens
else:
encoded_list = self._tokenizer.encode_batch( encoded_list = self._tokenizer.encode_batch(
tokens, tokens,
is_pretokenized=is_pretokenized, is_pretokenized=is_pretokenized,
add_special_tokens=add_special_tokens, add_special_tokens=add_special_tokens,
) )
return [ return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
]
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str: def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
"""Decode token IDs to text.""" """Decode token IDs to text."""
@@ -227,45 +227,63 @@ class AutoTokenizer:
def apply_chat_template( def apply_chat_template(
self, self,
messages: List[Dict[str, str]], messages: Union[Messages, List[Messages]],
system_prompt: Optional[str] = None, system_prompt: Optional[str] = None,
tokenize: bool = True, tokenize: bool = True,
add_generation_prompt: bool = True, add_generation_prompt: bool = True,
**kwargs, **kwargs,
) -> Union[str, List[int]]: ) -> Union[str, List[int], List[str], List[List[int]]]:
""" """Apply the chat template and optionally tokenize.
Apply the chat template to messages and optionally tokenize the result.
Accepts both single conversations and batches:
- ``apply_chat_template([msg1, msg2])`` ``"..."`` or ``[ids]``
- ``apply_chat_template([[msg1, msg2], [msg3]])`` ``["..", ".."]``
or ``[[ids], [ids]]``
Batches render each conversation list and tokenise all at once via
:meth:`encode` (``List[str]`` Rust parallel ``encode_batch``).
Args: Args:
messages: List of message dicts with 'role' and 'content'. messages: Single conversation (``Messages``) or batch of
system_prompt: Optional system prompt string (auto-converted to first message). conversations (``BatchMessages``).
system_prompt: Optional system prompt prepended (single mode only).
tokenize: Whether to return token IDs (True) or raw string (False). tokenize: Whether to return token IDs (True) or raw string (False).
add_generation_prompt: Whether to add the generation prompt (default: True). add_generation_prompt: Whether to add the generation prompt.
**kwargs: Additional variables to pass to the template. **kwargs: Additional template variables.
Returns: Returns:
Either the rendered string or list of token IDs. Single mode: ``str`` or ``List[int]``.
Batch mode: ``List[str]`` or ``List[List[int]]``.
Raises:
RuntimeError: If chat template is not set.
""" """
if self._chat_template is None: if self._chat_template is None:
raise RuntimeError( raise RuntimeError(
"Chat template not set. Use set_chat_template() to set a template first." "Chat template not set. Use set_chat_template() to set a template first."
) )
# Auto-convert system_prompt to first message if provided is_batch = bool(messages) and isinstance(messages[0], list)
if is_batch:
rendered = [
self._chat_template.render(
messages=msgs,
add_generation_prompt=add_generation_prompt,
**kwargs,
)
for msgs in messages
]
if tokenize:
return self.encode(rendered) # List[str] → batch encode
return rendered
# Single conversation
if system_prompt: if system_prompt:
messages = [{"role": "system", "content": system_prompt}] + list(messages) messages = [{"role": "system", "content": system_prompt}] + list(messages)
# Render the template
rendered = self._chat_template.render( rendered = self._chat_template.render(
messages=messages, messages=messages,
add_generation_prompt=add_generation_prompt, add_generation_prompt=add_generation_prompt,
**kwargs, **kwargs,
) )
if tokenize: if tokenize:
return self.encode(rendered) return self.encode(rendered)
return rendered return rendered
+72
View File
@@ -22,6 +22,51 @@ def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, fl
return total_sq.sqrt().item() return total_sq.sqrt().item()
class GradSNRTracker:
"""Track gradient signal-to-noise ratio via EMA of first/second moments.
SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2)
The tracker accumulates per-parameter EMA moments across optimizer steps.
Call ``update`` after backward (before ``optimizer.step``) and read
``snr`` to get the aggregate SNR across all parameters.
"""
def __init__(self, beta: float = 0.999, eps: float = 1e-8):
self.beta = beta
self.eps = eps
self._first: Dict[int, torch.Tensor] = {}
self._second: Dict[int, torch.Tensor] = {}
@torch.no_grad()
def update(self, model: nn.Module) -> None:
beta = self.beta
for param in model.parameters():
if param.grad is None:
continue
pid = id(param)
g = param.grad.detach()
if pid not in self._first:
self._first[pid] = g.clone()
self._second[pid] = g.pow(2).clone()
else:
self._first[pid].mul_(beta).add_(g, alpha=1 - beta)
self._second[pid].mul_(beta).addcmul_(g, g, value=1 - beta)
@property
def snr(self) -> float:
if not self._first:
return 0.0
total_signal = 0.0
total_noise = 0.0
for m, v in zip(self._first.values(), self._second.values()):
signal = m.pow(2).sum().item()
noise = (v - m.pow(2)).clamp(min=0).sum().item()
total_signal += signal
total_noise += noise
return total_signal / (total_noise + self.eps)
def ctx_get_loss(ctx): def ctx_get_loss(ctx):
return ctx.loss return ctx.loss
@@ -36,3 +81,30 @@ def ctx_get_val_loss(ctx):
def ctx_get_grad_norm(ctx): def ctx_get_grad_norm(ctx):
return ctx.grad_norm return ctx.grad_norm
def ctx_get_grad_snr(ctx):
tracker = getattr(ctx, "grad_snr_tracker", None)
if tracker is None:
return None
return tracker.snr
def ctx_get_moe_aux_loss(ctx):
return ctx.strategy._moe_metrics.get("aux_loss")
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")
+421
View File
@@ -0,0 +1,421 @@
"""Online rollout runner for RL training.
Provides:
- :class:`RawRollout` generation output container (no reward yet)
- :class:`RolloutResult` a :class:`RawRollout` with rewards attached
- :class:`BaseRewardModel` pluggable reward interface
- :class:`RolloutGenerator` KV-cache-backed generation of grouped
responses + decoding (no reward); delegates the generation loop to
:class:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`
so rollout and the production inference server share one code path
- :class:`RolloutRunner` orchestrates generation + scoring with a
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
so callers do not need to rely on object identity to detect refreshes.
"""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple
import torch
from torch import Tensor
from astrai.inference.core.scheduler import InferenceScheduler
@dataclass(kw_only=True)
class RawRollout:
"""Generation output before reward scoring.
Produced by :class:`RolloutGenerator`; consumed by :class:`RolloutRunner`
to assemble a :class:`RolloutResult` once rewards are attached.
Fields are designed to cover all common RL algorithms:
GRPO, PPO, Online DPO, Rejection Sampling, etc.
Fields:
prompts: Tokenized prompts, shape ``[B, P_len]``.
prompt_mask: Boolean mask for real prompt tokens, shape ``[B, P_len]``.
responses: Generated response token IDs, shape ``[B, G, R_max]``.
response_mask: Boolean mask for real (non-pad) response tokens,
shape ``[B, G, R_max]``.
logprobs_old: Per-token log-probs under the behaviour policy,
shape ``[B, G, R_max]``.
prompt_texts: Decoded prompt strings (for reward models that
need text).
response_texts: Decoded response strings, shape ``[B, G]``
(for reward models).
"""
prompts: Tensor
prompt_mask: Tensor
responses: Tensor
response_mask: Tensor
logprobs_old: Tensor
prompt_texts: List[str] = field(default_factory=list)
response_texts: List[List[str]] = field(default_factory=list)
@dataclass(kw_only=True)
class RolloutResult(RawRollout):
"""A :class:`RawRollout` with reward scoring attached.
Produced by :class:`RolloutRunner` once the :class:`BaseRewardModel`
has scored the decoded responses.
Fields:
rewards: Reward per response, shape ``[B, G]``.
"""
rewards: Tensor
class BaseRewardModel(ABC):
"""Pluggable reward model interface.
Subclasses should implement ``score()`` to return a ``[B, G]`` float
tensor of rewards. Implementations can be:
* A loaded reward model (e.g. ArmoRM, Skywork-Reward)
* An external API call
* A rule-based function (format, length, keyword matching)
"""
@abstractmethod
def score(self, prompts: List[str], responses: List[List[str]]) -> Tensor:
"""Score each generated response.
Args:
prompts: Raw prompt strings, length ``B``.
responses: Generated response strings, shape ``[B, G]``.
Returns:
Float tensor of shape ``[B, G]``.
"""
...
_PAD = 0
class RolloutGenerator:
"""Pure generation + decoding for a group of responses per prompt.
Delegates the prefill/decode loop to
:meth:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`,
which uses a real KV cache (no O() recompute). Has no dependency
on any reward model; can be reused in isolation for offline
generation, qualitative sampling, or eval pipelines.
"""
def __init__(
self,
scheduler: InferenceScheduler,
tokenizer,
max_tokens: int = 1024,
group_size: int = 8,
temperature: float = 1.0,
top_k: int = 0,
top_p: float = 1.0,
frequency_penalty: float = 0.0,
rep_window: int = 64,
):
self.scheduler = scheduler
self.tokenizer = tokenizer
self.max_tokens = max_tokens
self.group_size = group_size
self.temperature = temperature
self.top_k = top_k
self.top_p = top_p
self.frequency_penalty = frequency_penalty
self.rep_window = rep_window
@torch.no_grad()
def generate(self, batch: Dict) -> RawRollout:
"""Expand prompts by ``group_size`` and generate one response each.
Accepted batch formats (per sample, repeated B times):
- **messages**: ``{"messages": [{"role": "user", "content": "..."}, ...]}``
- **instruction + input + output**: ``{"instruction": "...",
"input": "...", "output": "..."}`` mapped to ``system`` /
``user`` / ``assistant`` messages; ``input`` and ``output``
are optional and skipped when empty.
Both are rendered through the tokenizer's chat template with
``add_generation_prompt=True`` so rollout prompts match the
format the policy was SFT-trained on.
"""
model = self.scheduler._executor.model
was_training = model.training
model.eval()
try:
return self._generate_eval(batch)
finally:
model.train(was_training)
def _generate_eval(self, batch: Dict) -> RawRollout:
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
B = len(prompt_texts)
G = self.group_size
# Re-expand flat list to G copies per prompt for run_batch.
expanded_prompt_ids: List[List[int]] = []
for ids in flat_prompt_ids:
expanded_prompt_ids.extend([list(ids)] * G)
results = self.scheduler.run_batch(
expanded_prompt_ids,
max_tokens=self.max_tokens,
temperature=self.temperature,
top_k=self.top_k,
top_p=self.top_p,
frequency_penalty=self.frequency_penalty,
rep_window=self.rep_window,
return_logprobs=True,
)
if len(results) != B * G:
raise RuntimeError(
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
)
for token_ids, logprobs in results:
if len(token_ids) != len(logprobs):
raise RuntimeError(
"Rollout scheduler returned misaligned token IDs and logprobs"
)
# Each element is (token_ids, logprobs); pad to max length.
max_len = 0
for token_ids, _lp in results:
max_len = max(max_len, len(token_ids))
max_len = max(max_len, 1)
device = self.scheduler.device
P_len = max(len(ids) for ids in flat_prompt_ids)
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
prompt_mask = torch.zeros(B, P_len, dtype=torch.bool, device=device)
for i, ids in enumerate(flat_prompt_ids):
prompts_tensor[i, -len(ids) :] = torch.tensor(
ids, dtype=torch.long, device=device
)
prompt_mask[i, -len(ids) :] = True
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
logprobs_old = torch.zeros((B, G, max_len), dtype=torch.float, device=device)
flat_idx = 0
response_texts: List[List[str]] = [[] for _ in range(B)]
for i in range(B):
for g in range(G):
token_ids, lps = results[flat_idx]
flat_idx += 1
n = len(token_ids)
if n:
responses[i, g, :n] = torch.tensor(
token_ids, dtype=torch.long, device=device
)
response_mask[i, g, :n] = True
logprobs_old[i, g, :n] = torch.tensor(
lps, dtype=torch.float, device=device
)
response_texts[i].append(
self.tokenizer.decode(token_ids, skip_special_tokens=True)
)
return RawRollout(
prompts=prompts_tensor,
prompt_mask=prompt_mask,
responses=responses,
response_mask=response_mask,
logprobs_old=logprobs_old,
prompt_texts=prompt_texts,
response_texts=response_texts,
)
def _prepare_prompts(self, batch: Dict) -> Tuple[List[str], List[List[int]]]:
"""Render batch prompts to ``(texts, token_id_lists)``.
Returns two parallel lists of length B (number of prompts in
the batch). Dispatches by batch keys:
- ``"messages"``: treated as a pre-built message list per sample.
- ``"instruction"`` (optionally ``"input"`` and ``"output"``): mapped
to ``system`` / ``user`` / ``assistant`` messages respectively.
Both paths go through the tokenizer's chat template with
``add_generation_prompt=True``.
"""
if "messages" in batch:
messages_list = batch["messages"]
elif "instruction" in batch:
instructions = batch["instruction"]
B = len(instructions)
inputs = batch.get("input") or [""] * B
outputs = batch.get("output") or [""] * B
messages_list = [
self._instruction_to_messages(i, u, o)
for i, u, o in zip(instructions, inputs, outputs)
]
else:
raise ValueError(
"Rollout batch must contain either 'messages' or "
"'instruction' (optionally 'input'/'output'); got keys: "
f"{list(batch.keys())}"
)
try:
prompt_texts = self.tokenizer.apply_chat_template(
messages_list, tokenize=False, add_generation_prompt=True
)
if (
not isinstance(prompt_texts, list)
or len(prompt_texts) != len(messages_list)
or not all(isinstance(text, str) for text in prompt_texts)
):
raise TypeError("Tokenizer does not support batched chat templates")
flat_prompt_ids = self.tokenizer.encode(prompt_texts)
if len(flat_prompt_ids) != len(messages_list) or not all(
isinstance(ids, list) for ids in flat_prompt_ids
):
raise TypeError("Tokenizer does not support batched encoding")
except (TypeError, IndexError, KeyError):
# Keep compatibility with lightweight tokenizer adapters that only
# implement the single-conversation template API.
prompt_texts = []
flat_prompt_ids = []
for messages in messages_list:
text = self.tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
ids = self.tokenizer.apply_chat_template(
messages, tokenize=True, add_generation_prompt=True
)
prompt_texts.append(text)
flat_prompt_ids.append(list(ids))
return prompt_texts, flat_prompt_ids
@staticmethod
def _instruction_to_messages(
instruction: str, inp: str = "", output: str = ""
) -> List[Dict[str, str]]:
"""Map instruction/input/output to chat messages.
Role mapping follows the convention used throughout the
preprocessing pipeline: ``instruction`` system, ``input``
user, ``output`` assistant. Empty fields are skipped so a
bare instruction produces a ``[system]`` list and the chat
template's ``add_generation_prompt`` adds the assistant header
for sampling.
"""
messages: List[Dict[str, str]] = []
if instruction:
messages.append({"role": "system", "content": instruction})
if inp:
messages.append({"role": "user", "content": inp})
if output:
messages.append({"role": "assistant", "content": output})
return messages
class RolloutRunner:
"""Produces :class:`RolloutResult` from a prompt batch.
Composes a :class:`RolloutGenerator` (generation + decoding) with a
:class:`BaseRewardModel` (scoring). Maintains an internal cache so
the same batch prompt can be replayed for multiple gradient steps.
A new rollout is triggered every ``rollout_interval`` calls to
:meth:`step` (or after :meth:`clear_cache`).
The ``__call__`` contract returns a ``(RolloutResult, is_fresh)``
tuple callers must use the boolean to detect a refreshed rollout
rather than relying on object identity.
Usage::
generator = RolloutGenerator(policy, tokenizer, pipeline, ...)
runner = RolloutRunner(generator, reward_model, rollout_interval=512)
result, is_fresh = runner(prompt_batch)
if is_fresh:
... # e.g. sync behaviour policy
"""
def __init__(
self,
generator: RolloutGenerator,
reward_model: BaseRewardModel,
rollout_interval: int = 512,
):
self.generator = generator
self.reward_model = reward_model
self.rollout_interval = rollout_interval
self._cache: Optional[RolloutResult] = None
self._cache_key = None
self._steps_since_rollout: int = 0
def step(self):
"""Advance the internal counter (call once per optimizer step)."""
self._steps_since_rollout += 1
def clear_cache(self):
"""Force next call to re-run rollout."""
self._cache = None
self._cache_key = None
@staticmethod
def _batch_key(batch: Dict):
"""Build a stable key for the prompt fields accepted by the generator."""
def freeze(value):
if isinstance(value, dict):
return tuple(sorted((key, freeze(val)) for key, val in value.items()))
if isinstance(value, (list, tuple)):
return tuple(freeze(item) for item in value)
return value
fields = ("messages", "instruction", "input", "output")
return tuple(
(field, freeze(batch[field])) for field in fields if field in batch
)
def _score(self, raw: RawRollout) -> RolloutResult:
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts)
if not isinstance(rewards, Tensor):
rewards = torch.as_tensor(rewards, dtype=torch.float32)
expected_shape = raw.responses.shape[:2]
if rewards.shape != expected_shape:
raise ValueError(
f"Reward model returned shape {tuple(rewards.shape)}, "
f"expected {tuple(expected_shape)}"
)
if not torch.isfinite(rewards).all():
raise ValueError("Reward model returned non-finite values")
device = raw.prompts.device
return RolloutResult(
prompts=raw.prompts,
prompt_mask=raw.prompt_mask,
responses=raw.responses,
response_mask=raw.response_mask,
rewards=rewards.to(device=device),
logprobs_old=raw.logprobs_old,
prompt_texts=raw.prompt_texts,
response_texts=raw.response_texts,
)
def __call__(self, batch: Dict[str, Tensor]) -> Tuple[RolloutResult, bool]:
"""Return ``(cached or fresh) RolloutResult`` plus an ``is_fresh`` flag.
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
or when the cache is empty.
"""
cache_key = self._batch_key(batch)
if (
self._cache is None
or cache_key != self._cache_key
or self._steps_since_rollout >= self.rollout_interval
):
raw = self.generator.generate(batch)
self._cache = self._score(raw)
self._cache_key = cache_key
self._steps_since_rollout = 0
return self._cache, True
return self._cache, False
+352 -47
View File
@@ -1,7 +1,7 @@
"""Training strategy implementations with factory pattern.""" """Training strategy implementations with factory pattern."""
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Callable, Dict, Union from typing import Callable, Dict, List, Optional, TypedDict, Union
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -9,17 +9,20 @@ import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
from astrai.model.components.mlp import RouterStats
from astrai.parallel.executor import broadcast_state_dict
from astrai.trainer.rollout import RolloutResult
def create_ref_model( class LossOutput(TypedDict):
model_fn: Callable[[], nn.Module], state_dict: Dict[str, Tensor] loss: Tensor
) -> nn.Module: metrics: Dict[str, float]
"""Create a frozen reference model from model_fn + full state dict."""
ref_model = model_fn()
ref_model.load_state_dict(state_dict) class LogprobsOutput(TypedDict):
ref_model.requires_grad_(False) logprobs: Tensor
ref_model.eval() aux_loss: Optional[Tensor]
return ref_model router_stats: Optional[List[RouterStats]]
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]: def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
@@ -28,17 +31,19 @@ def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
def get_logprobs( def get_logprobs(
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]], model: nn.Module,
input_ids: Tensor, input_ids: Tensor,
mask: Tensor, attn_mask: Tensor,
loss_mask: Tensor,
reduction: str, reduction: str,
) -> Tensor: ) -> LogprobsOutput:
"""Compute token-wise log probabilities from model outputs. """Compute token-wise log probabilities from model outputs.
Args: Args:
model: The language model model: The language model
input_ids: Input token IDs of shape [batch_size, seq_len] input_ids: Input token IDs of shape [batch_size, seq_len]
mask: Attention mask of shape [batch_size, seq_len] attn_mask: Attention mask passed to the model (may include causal).
loss_mask: Per-token mask for loss reduction.
reduction: How to reduce over sequence dimension ("mean", "sum", "none") reduction: How to reduce over sequence dimension ("mean", "sum", "none")
Returns: Returns:
@@ -51,9 +56,13 @@ def get_logprobs(
) )
shifted_input_ids = input_ids[:, 1:] shifted_input_ids = input_ids[:, 1:]
shifted_mask = mask[:, 1:] shifted_loss_mask = loss_mask[:, 1:]
logits = model(input_ids[:, :-1], mask[:, :-1])["logits"] outputs = model(
input_ids[:, :-1],
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
)
logits = outputs["logits"]
log_probs = torch.log_softmax(logits.float(), dim=-1) log_probs = torch.log_softmax(logits.float(), dim=-1)
token_logprobs = torch.gather( token_logprobs = torch.gather(
@@ -61,13 +70,18 @@ def get_logprobs(
).squeeze(-1) ).squeeze(-1)
if reduction == "mean": if reduction == "mean":
return (token_logprobs * shifted_mask).sum(dim=-1) / shifted_mask.sum( logprobs = (token_logprobs * shifted_loss_mask).sum(
dim=-1 dim=-1
).clamp(min=1.0) ) / shifted_loss_mask.sum(dim=-1).clamp(min=1.0)
elif reduction == "sum": elif reduction == "sum":
return (token_logprobs * shifted_mask).sum(dim=-1) logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
else: else:
return token_logprobs * shifted_mask logprobs = token_logprobs * shifted_loss_mask
return {
"logprobs": logprobs,
"aux_loss": outputs.get("aux_loss"),
"router_stats": outputs.get("router_stats"),
}
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor: def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
@@ -86,8 +100,78 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
return (same_doc & causal).unsqueeze(1) return (same_doc & causal).unsqueeze(1)
def _collect_moe_diagnostics(
router_stats_list: List[RouterStats],
) -> Dict[str, float]:
"""Collect MoE routing diagnostic metrics from per-layer router stats.
Args:
router_stats_list: One :class:`RouterStats` dict per MoE layer with
keys ``probs`` (N, E) and ``topk_indices`` (N, K), both detached.
Returns:
Dict with keys: router_entropy, dead_expert_fraction,
load_imbalance_mean, load_imbalance_max. Values are averaged
across layers.
"""
layer_entropies: List[Tensor] = []
layer_dead_fractions: List[Tensor] = []
layer_imbalance_means: List[Tensor] = []
layer_imbalance_maxs: List[Tensor] = []
for stats in router_stats_list:
probs = stats["probs"].float()
topk_indices = stats["topk_indices"]
num_experts = probs.shape[-1]
if num_experts == 0:
continue
probs = probs.reshape(-1, num_experts)
if probs.numel() == 0:
continue
# Router entropy
entropy = -(probs * torch.log(probs.clamp_min(1e-8))).sum(dim=-1).mean()
# Load from the actual dispatch: one-hot sum of top-k assignments.
expert_counts = F.one_hot(topk_indices, num_experts).sum(dim=(0, 1)).float()
ideal_load = expert_counts.mean() # N*K / E
load_ratios = expert_counts / max(float(ideal_load), 1.0)
imbalance_mean = (load_ratios - 1.0).abs().mean()
imbalance_max = load_ratios.max()
dead_fraction = (expert_counts == 0).float().mean()
layer_entropies.append(entropy)
layer_dead_fractions.append(dead_fraction)
layer_imbalance_means.append(imbalance_mean)
layer_imbalance_maxs.append(imbalance_max)
if not layer_entropies:
return {}
return {
"router_entropy": float(torch.stack(layer_entropies).mean().cpu().item()),
"dead_expert_fraction": float(
torch.stack(layer_dead_fractions).mean().cpu().item()
),
"load_imbalance_mean": float(
torch.stack(layer_imbalance_means).mean().cpu().item()
),
"load_imbalance_max": float(
torch.stack(layer_imbalance_maxs).mean().cpu().item()
),
}
class BaseStrategy(ABC): class BaseStrategy(ABC):
"""Abstract base class for training strategies.""" """Abstract base class for training strategies.
When a :class:`~astrai.trainer.rollout.RolloutRunner` is injected via
:meth:`set_rollout_runner`, the strategy transparently switches to
online mode: each ``__call__`` produces a :class:`RolloutResult`,
converts it to a training batch via :meth:`prepare_from_rollout`, and
then computes the loss. Without a runner the strategy runs in
offline mode and consumes the batch directly.
"""
def __init__( def __init__(
self, self,
@@ -98,7 +182,10 @@ class BaseStrategy(ABC):
self.model = model self.model = model
self.device = device self.device = device
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_metrics: Dict[str, float] = {}
self.extra_kwargs = kwargs self.extra_kwargs = kwargs
self._rollout_runner = None
@abstractmethod @abstractmethod
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
@@ -112,9 +199,96 @@ class BaseStrategy(ABC):
""" """
raise NotImplementedError raise NotImplementedError
def __call__(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
"""Allow calling strategy directly as a callable.""" return self._normalize_output(self.compute_loss(batch))
return self.compute_loss(batch)
def _loss_output(
self,
task_loss: Tensor,
metrics: Dict[str, Tensor],
aux_loss: Optional[Tensor] = None,
router_stats: Optional[List[RouterStats]] = None,
) -> LossOutput:
total_loss = task_loss
if aux_loss is not None:
weighted_aux_loss = self.moe_aux_loss_coef * aux_loss
total_loss = total_loss + weighted_aux_loss
metrics["moe_aux_loss"] = aux_loss
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
self._refresh_moe_diagnostics(aux_loss, router_stats)
metrics["loss"] = total_loss
return {
"loss": total_loss,
"metrics": {name: value.detach().item() for name, value in metrics.items()},
}
@staticmethod
def _normalize_output(output: Union[LossOutput, Tensor]) -> LossOutput:
if isinstance(output, dict):
return output
return {"loss": output, "metrics": {"loss": output.detach().item()}}
def supports_online(self) -> bool:
"""Whether this strategy can operate with a rollout runner.
Base implementation returns ``False``; strategies that implement
:meth:`prepare_from_rollout` should override to return ``True``.
"""
return False
def set_rollout_runner(self, runner):
"""Inject a :class:`RolloutRunner` to enable online rollout mode."""
self._rollout_runner = runner
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
"""Map a :class:`RolloutResult` to the batch layout expected by
:meth:`compute_loss`.
Strategies that return ``True`` from :meth:`supports_online` must
override this. Default raises :class:`NotImplementedError`.
"""
raise NotImplementedError(
f"{type(self).__name__} does not support online rollout"
)
def _on_rollout_refresh(self):
"""Hook fired when a fresh rollout result is produced.
Override to refresh stale state (e.g. syncing the behaviour
policy). Default is a no-op.
"""
pass
def _refresh_moe_diagnostics(
self,
aux_loss: Tensor,
router_stats: Optional[List[RouterStats]] = None,
) -> None:
"""Collect MoE routing diagnostics from the latest forward pass.
Populates ``self._moe_metrics`` with router entropy, dead expert
fraction, load imbalance, and aux_loss. Called from
:meth:`_loss_output` when an MoE aux loss is present.
"""
self._moe_metrics = _collect_moe_diagnostics(router_stats or [])
self._moe_metrics["aux_loss"] = float(aux_loss.detach().cpu().item())
def on_optimizer_step(self):
"""Advance online rollout state after a successful optimizer step."""
if self._rollout_runner is not None:
self._rollout_runner.step()
def __call__(self, batch: Dict[str, Tensor]) -> LossOutput:
"""Run offline or online forward depending on runner injection."""
if self._rollout_runner is None:
return self.compute_loss_output(batch)
result, is_fresh = self._rollout_runner(batch)
if is_fresh:
self._on_rollout_refresh()
train_batch = self.prepare_from_rollout(result)
return self.compute_loss_output(train_batch)
class StrategyFactory(BaseFactory["BaseStrategy"]): class StrategyFactory(BaseFactory["BaseStrategy"]):
@@ -141,6 +315,7 @@ class SEQStrategy(BaseStrategy):
"""Standard next-token prediction training strategy. """Standard next-token prediction training strategy.
Computes cross-entropy loss for next token prediction. Computes cross-entropy loss for next token prediction.
Optionally adds MoE load balancing auxiliary loss.
""" """
def __init__( def __init__(
@@ -154,9 +329,13 @@ class SEQStrategy(BaseStrategy):
self.label_smoothing = label_smoothing self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: 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:
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"]
logits = self.model(input_ids=input_ids)["logits"] outputs = self.model(input_ids=input_ids)
logits = outputs["logits"]
loss = F.cross_entropy( loss = F.cross_entropy(
input=logits.flatten(0, 1).float(), input=logits.flatten(0, 1).float(),
@@ -164,7 +343,12 @@ class SEQStrategy(BaseStrategy):
label_smoothing=self.label_smoothing, label_smoothing=self.label_smoothing,
) )
return loss return self._loss_output(
loss,
{"task_loss": loss},
outputs.get("aux_loss"),
outputs.get("router_stats"),
)
@StrategyFactory.register("sft") @StrategyFactory.register("sft")
@@ -172,6 +356,7 @@ class SFTStrategy(BaseStrategy):
"""Supervised Fine-tuning strategy with loss masking. """Supervised Fine-tuning strategy with loss masking.
Applies cross-entropy loss only to tokens where loss_mask is True. Applies cross-entropy loss only to tokens where loss_mask is True.
Optionally adds MoE load balancing auxiliary loss.
""" """
def __init__( def __init__(
@@ -185,6 +370,9 @@ class SFTStrategy(BaseStrategy):
self.label_smoothing = label_smoothing self.label_smoothing = label_smoothing
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: 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:
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 = (
batch["input_ids"], batch["input_ids"],
@@ -196,9 +384,10 @@ class SFTStrategy(BaseStrategy):
ignore_index = -100 ignore_index = -100
input_mask = make_doc_boundary_mask(position_ids) input_mask = make_doc_boundary_mask(position_ids)
target_ids = target_ids.masked_fill(~loss_mask, ignore_index) target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
logits = self.model( outputs = self.model(
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
)["logits"] )
logits = outputs["logits"]
loss = F.cross_entropy( loss = F.cross_entropy(
input=logits.flatten(0, 1).float(), input=logits.flatten(0, 1).float(),
@@ -207,7 +396,12 @@ class SFTStrategy(BaseStrategy):
label_smoothing=self.label_smoothing, label_smoothing=self.label_smoothing,
) )
return loss return self._loss_output(
loss,
{"task_loss": loss},
outputs.get("aux_loss"),
outputs.get("router_stats"),
)
@StrategyFactory.register("dpo") @StrategyFactory.register("dpo")
@@ -233,19 +427,43 @@ class DPOStrategy(BaseStrategy):
self.reduction = reduction self.reduction = reduction
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: 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:
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"]
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"] chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0) concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
concat_mask = torch.cat([chosen_mask, rejected_mask], dim=0) concat_loss_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
log_pi = get_logprobs(self.model, concat_ids, concat_mask, self.reduction) # Build full attention mask: key-padding + causal
key_pad = concat_ids.bool()[:, None, None, :] # [B*2, 1, 1, S]
S = key_pad.shape[-1]
causal = torch.tril(
torch.ones(S, S, dtype=torch.bool, device=concat_ids.device)
)[None, None, :, :] # [1, 1, S, S]
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
policy_output = get_logprobs(
self.model,
concat_ids,
full_mask,
concat_loss_mask,
self.reduction,
)
log_pi = policy_output["logprobs"]
aux_loss = policy_output["aux_loss"]
with torch.no_grad(): with torch.no_grad():
log_ref = get_logprobs( ref_output = get_logprobs(
self.ref_model, concat_ids, concat_mask, self.reduction self.ref_model,
concat_ids,
full_mask,
concat_loss_mask,
self.reduction,
) )
log_ref = ref_output["logprobs"]
log_pi_chosen = log_pi[: chosen_ids.shape[0]] log_pi_chosen = log_pi[: chosen_ids.shape[0]]
log_pi_rejected = log_pi[chosen_ids.shape[0] :] log_pi_rejected = log_pi[chosen_ids.shape[0] :]
@@ -258,7 +476,35 @@ class DPOStrategy(BaseStrategy):
ratio_diff = pi_log_ratio - ref_log_ratio ratio_diff = pi_log_ratio - ref_log_ratio
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean() dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
return dpo_loss return self._loss_output(
dpo_loss,
{"dpo_loss": dpo_loss},
aux_loss,
policy_output.get("router_stats"),
)
def supports_online(self) -> bool:
return True
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
"""Pick best/worst response per prompt by reward as chosen/rejected."""
rewards = result.rewards
responses = result.responses
masks = result.response_mask
best = rewards.argmax(dim=-1)
worst = rewards.argmin(dim=-1)
B = responses.shape[0]
idx = torch.arange(B, device=responses.device)
chosen = responses[idx, best]
chosen_mask = masks[idx, best].float()
rejected = responses[idx, worst]
rejected_mask = masks[idx, worst].float()
return {
"chosen": chosen,
"chosen_mask": chosen_mask,
"rejected": rejected,
"rejected_mask": rejected_mask,
}
@StrategyFactory.register("grpo") @StrategyFactory.register("grpo")
@@ -301,9 +547,16 @@ class GRPOStrategy(BaseStrategy):
def sync_old_model(self): def sync_old_model(self):
"""Copy current policy weights to old model.""" """Copy current policy weights to old model."""
self.old_model.load_state_dict(self.executor.unwrap_model(self.model)) state_dict = self.executor.unwrap_model(self.model)
if self.executor.use_distributed:
state_dict = broadcast_state_dict(state_dict)
if state_dict is not None:
self.old_model.load_state_dict(state_dict)
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: 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:
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
prompts = batch["prompts"] prompts = batch["prompts"]
responses = batch["responses"] responses = batch["responses"]
@@ -314,6 +567,12 @@ class GRPOStrategy(BaseStrategy):
responses_flat = responses.view(-1, response_len) responses_flat = responses.view(-1, response_len)
masks_flat = masks.view(-1, response_len) masks_flat = masks.view(-1, response_len)
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1) prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
prompt_mask = batch.get("prompt_mask")
if prompt_mask is None:
prompt_mask = prompts.ne(0)
prompt_mask_expanded = (
prompt_mask.unsqueeze(1).expand(-1, group_size, -1).flatten(0, 1)
)
prompt_len = prompt_expanded.size(1) prompt_len = prompt_expanded.size(1)
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1) full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
@@ -321,21 +580,40 @@ class GRPOStrategy(BaseStrategy):
# response tokens. get_logprobs shifts the mask by one position, so # response tokens. get_logprobs shifts the mask by one position, so
# the first response token's logprob (predicted from the last prompt # the first response token's logprob (predicted from the last prompt
# token) is correctly included. # token) is correctly included.
full_masks = torch.cat([torch.zeros_like(prompt_expanded), masks_flat], dim=-1) full_masks = torch.cat(
[torch.zeros_like(prompt_expanded, dtype=torch.bool), masks_flat], dim=-1
)
# Build full attention mask: key-padding + causal
key_pad = torch.cat([prompt_mask_expanded, masks_flat.bool()], dim=-1)[
:, None, None, :
]
S = key_pad.shape[-1]
causal = torch.tril(
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device)
)[None, None, :, :]
attn_mask = key_pad & causal
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len). # get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
# Response token logprobs occupy the last ``response_len`` positions # Response token logprobs occupy the last ``response_len`` positions
# (the first response token is predicted from the last prompt token). # (the first response token is predicted from the last prompt token).
token_log_probs_policy = get_logprobs( policy_output = get_logprobs(
self.model, full_sequences, full_masks, "none" self.model, full_sequences, attn_mask, full_masks, "none"
)[:, prompt_len - 1 :] )
token_log_probs_policy = policy_output["logprobs"]
aux_loss = policy_output["aux_loss"]
token_log_probs_policy = token_log_probs_policy[:, prompt_len - 1 :]
with torch.no_grad(): with torch.no_grad():
token_log_probs_old = get_logprobs( old_output = get_logprobs(
self.old_model, full_sequences, full_masks, "none" self.old_model, full_sequences, attn_mask, full_masks, "none"
)[:, prompt_len - 1 :] )
token_log_probs_ref = get_logprobs( token_log_probs_old = old_output["logprobs"]
self.ref_model, full_sequences, full_masks, "none" token_log_probs_old = token_log_probs_old[:, prompt_len - 1 :]
)[:, prompt_len - 1 :] ref_output = get_logprobs(
self.ref_model, full_sequences, attn_mask, full_masks, "none"
)
token_log_probs_ref = ref_output["logprobs"]
token_log_probs_ref = token_log_probs_ref[:, prompt_len - 1 :]
# Reshape to [B, G, response_len] # Reshape to [B, G, response_len]
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1) token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
@@ -368,6 +646,33 @@ class GRPOStrategy(BaseStrategy):
kl_per_token = r - torch.log(r + eps) - 1.0 kl_per_token = r - torch.log(r + eps) - 1.0
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
total_loss = policy_loss + kl_penalty task_loss = policy_loss + kl_penalty
return self._loss_output(
task_loss,
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
aux_loss,
policy_output.get("router_stats"),
)
return total_loss def supports_online(self) -> bool:
return True
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
return {
"prompts": result.prompts,
"prompt_mask": result.prompt_mask,
"responses": result.responses,
"masks": result.response_mask,
"rewards": result.rewards,
}
def _on_rollout_refresh(self):
"""Sync the behaviour policy whenever a fresh rollout arrives."""
self.sync_old_model()
# Factory aliases: online variants use the same strategy class; the
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
# online mode, so no separate subclass is needed.
StrategyFactory.register("online_grpo")(GRPOStrategy)
StrategyFactory.register("online_dpo")(DPOStrategy)
+41 -11
View File
@@ -17,9 +17,15 @@ 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_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_router_entropy,
ctx_get_val_loss, ctx_get_val_loss,
) )
from astrai.trainer.train_context import TrainContext from astrai.trainer.train_context import TrainContext
@@ -235,7 +241,7 @@ class ProgressBarCallback(TrainCallback):
class MetricCallback(TrainCallback): class MetricCallback(TrainCallback):
def __init__( def __init__(
self, self,
log_dir: str, ckpt_dir: str,
save_interval: int, save_interval: int,
metrics: List[str] = None, metrics: List[str] = None,
val_step: int = 0, val_step: int = 0,
@@ -246,8 +252,7 @@ class MetricCallback(TrainCallback):
self.val_step = val_step self.val_step = val_step
self._next_val_step = 0 self._next_val_step = 0
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs" self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint"
self.log_dir.mkdir(parents=True, exist_ok=True)
self.log_cache = [] self.log_cache = []
@@ -256,14 +261,37 @@ class MetricCallback(TrainCallback):
"lr": ctx_get_lr, "lr": ctx_get_lr,
"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,
"moe_aux_loss": ctx_get_moe_aux_loss,
"router_entropy": ctx_get_router_entropy,
"dead_expert_fraction": ctx_get_dead_expert_fraction,
"load_imbalance_mean": ctx_get_load_imbalance_mean,
"load_imbalance_max": ctx_get_load_imbalance_max,
} }
def _metrics(self, context: TrainContext, names): def _metrics(self, context: TrainContext, names):
return { metrics = dict(context.metrics)
m: self._metric_funcs[m](context) for name in names:
for m in names metric_fn = self._metric_funcs.get(name)
if self._metric_funcs[m](context) is not None if metric_fn is None:
} continue
value = metric_fn(context)
if value is not None:
metrics[name] = value
selected = set(context.metrics) | set(names)
selected.discard("*")
result = {name: metrics[name] for name in selected if name in metrics}
if context.world_size > 1 and dist.is_initialized() and result:
metric_names = sorted(result)
values = torch.tensor(
[result[name] for name in metric_names],
dtype=torch.float32,
device=get_current_device(),
)
dist.all_reduce(values, op=dist.ReduceOp.SUM)
values /= context.world_size
result.update(zip(metric_names, values.tolist()))
return result
@only_on_rank(0) @only_on_rank(0)
def _append(self, event_type: str, context: TrainContext, **extra): def _append(self, event_type: str, context: TrainContext, **extra):
@@ -285,8 +313,8 @@ class MetricCallback(TrainCallback):
with torch.no_grad(): with torch.no_grad():
for batch in context.val_dataloader: for batch in context.val_dataloader:
loss = context.strategy(batch) loss_output = context.strategy(batch)
total_loss += loss.item() total_loss += loss_output["loss"].item()
num_batches += 1 num_batches += 1
if context.world_size > 1 and dist.is_initialized(): if context.world_size > 1 and dist.is_initialized():
@@ -306,13 +334,15 @@ class MetricCallback(TrainCallback):
@only_on_rank(0) @only_on_rank(0)
def _flush(self, epoch, step): def _flush(self, epoch, step):
log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl" log_file = self.ckpt_dir / f"epoch_{epoch}_step_{step}" / "metric.jsonl"
log_file.parent.mkdir(parents=True, exist_ok=True) log_file.parent.mkdir(parents=True, exist_ok=True)
with open(log_file, "w") as f: with open(log_file, "w") as f:
for log in self.log_cache: for log in self.log_cache:
f.write(json.dumps(log) + "\n") f.write(json.dumps(log) + "\n")
def on_optimizer_step(self, context): def on_optimizer_step(self, context):
context.grad_snr_tracker.update(context.model)
if ( if (
context.val_dataloader is not None context.val_dataloader is not None
and self.val_step > 0 and self.val_step > 0
+137 -61
View File
@@ -1,3 +1,5 @@
import logging
import threading
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Optional, Self from typing import Any, Dict, Optional, Self
@@ -8,12 +10,18 @@ from torch.utils.data import DataLoader, random_split
from astrai.config.train_config import TrainConfig from astrai.config.train_config import TrainConfig
from astrai.dataset import RDSampler from astrai.dataset import RDSampler
from astrai.inference.core.scheduler import InferenceScheduler
from astrai.model.components.lora import inject_lora from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory 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, load_json
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model from astrai.tokenize import AutoTokenizer
from astrai.trainer.metric_util import GradSNRTracker
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
logger = logging.getLogger(__name__)
@dataclass @dataclass
@@ -27,11 +35,12 @@ class TrainContext:
config: TrainConfig = field(default=None) config: TrainConfig = field(default=None)
model_config: dict = field(default_factory=dict) model_config: dict = field(default_factory=dict)
executor: BaseExecutor = field(default=None) executor: BaseExecutor = field(default=None)
epoch: int = field(default=0) epoch: int = field(default=0)
consumed_samples: int = field(default=0) consumed_samples: int = field(default=0)
loss: float = field(default=0.0) loss: float = field(default=0.0)
metrics: Dict[str, float] = field(default_factory=dict)
grad_norm: Optional[float] = field(default=None) grad_norm: Optional[float] = field(default=None)
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
val_dataloader: Optional[DataLoader] = field(default=None) val_dataloader: Optional[DataLoader] = field(default=None)
val_loss: Optional[float] = field(default=None) val_loss: Optional[float] = field(default=None)
@@ -39,6 +48,15 @@ class TrainContext:
rank: int = field(default=0) rank: int = field(default=0)
kwargs: Dict[str, Any] = field(default_factory=dict) kwargs: Dict[str, Any] = field(default_factory=dict)
_stop_event: threading.Event = field(default_factory=threading.Event)
@property
def stop_requested(self) -> bool:
return self._stop_event.is_set()
def request_stop(self) -> None:
self._stop_event.set()
@property @property
def optimizer_step(self) -> int: def optimizer_step(self) -> int:
return self.consumed_samples // ( return self.consumed_samples // (
@@ -72,61 +90,72 @@ class TrainContextBuilder:
**cfg.executor_kwargs, **cfg.executor_kwargs,
) )
model = cfg.model_fn()
model = model.to(device=device)
model_config = {} model_config = {}
if self._param_path: if self._param_path:
config_path = Path(self._param_path) / "config.json" config_path = Path(self._param_path) / "config.json"
if config_path.exists(): if config_path.exists():
model_config = load_json(config_path) model_config = load_json(config_path)
if not model_config and hasattr(model, "config"): preloaded_state_dict = None
model_config = model.config.to_dict() preloaded_epoch = cfg.start_epoch
preloaded_consumed = cfg.start_samples * get_world_size()
preloaded_checkpoint = None
if self._param_path:
checkpoint = Checkpoint.load_any(self._param_path)
if checkpoint is not None:
preloaded_state_dict = checkpoint.state_dict
if checkpoint.config:
model_config = checkpoint.config
if self._resume:
preloaded_epoch = checkpoint.epoch
per_step = (
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
)
preloaded_consumed = (
checkpoint.consumed_samples // per_step
) * per_step
preloaded_checkpoint = checkpoint
if not model_config and hasattr(cfg.model_fn(), "config"):
model_config = cfg.model_fn().config.to_dict()
def _before_wrap(m):
m = m.to(device=device)
if cfg.lora is not None:
inject_lora(
m,
r=cfg.lora.r,
alpha=cfg.lora.alpha,
target_modules=set(cfg.lora.target_modules),
)
if preloaded_state_dict is not None:
m.load_state_dict(preloaded_state_dict, strict=False)
return m
def _after_wrap(m):
if cfg.compile_mode is not None:
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
m = torch.compile(m, mode=cfg.compile_mode)
return m
context = TrainContext( context = TrainContext(
model=model,
world_size=get_world_size(), world_size=get_world_size(),
rank=get_rank(), rank=get_rank(),
config=cfg, config=cfg,
model_config=model_config, model_config=model_config,
executor=executor, executor=executor,
epoch=preloaded_epoch,
consumed_samples=preloaded_consumed,
checkpoint=preloaded_checkpoint,
) )
if self._param_path: context.model, context.optimizer, context.scheduler = executor.prepare(
checkpoint = Checkpoint.load_any(self._param_path) cfg.model_fn,
if checkpoint is not None: cfg.optimizer_fn,
model.load_state_dict(checkpoint.state_dict, strict=False) cfg.scheduler_fn,
if checkpoint.config: before_wrap=_before_wrap,
context.model_config = checkpoint.config after_wrap=_after_wrap,
if self._resume:
context.epoch = checkpoint.epoch or cfg.start_epoch
if checkpoint.consumed_samples > 0:
per_step = (
cfg.batch_per_device
* context.world_size
* cfg.grad_accum_steps
) )
context.consumed_samples = (
checkpoint.consumed_samples // per_step
) * per_step
else:
context.consumed_samples = (
cfg.start_samples * context.world_size
)
context.checkpoint = checkpoint
if cfg.lora is not None:
inject_lora(
model,
r=cfg.lora.r,
alpha=cfg.lora.alpha,
target_modules=set(cfg.lora.target_modules),
)
context.optimizer = cfg.optimizer_fn(model)
context.scheduler = cfg.scheduler_fn(context.optimizer)
train_dataset = cfg.dataset train_dataset = cfg.dataset
val_dataset = cfg.val_dataset val_dataset = cfg.val_dataset
@@ -141,6 +170,15 @@ class TrainContextBuilder:
) )
sampler_offset = context.consumed_samples // context.world_size sampler_offset = context.consumed_samples // context.world_size
if self._resume and sampler_offset > 0:
offset = context.world_size - 1
num_samples_per_replica = (
len(train_dataset) + offset
) // context.world_size
if num_samples_per_replica > 0:
context.epoch = sampler_offset // num_samples_per_replica
sampler = RDSampler( sampler = RDSampler(
data_source=train_dataset, data_source=train_dataset,
start_epoch=context.epoch, start_epoch=context.epoch,
@@ -175,15 +213,6 @@ class TrainContextBuilder:
collate_fn=cfg.collate_fn, collate_fn=cfg.collate_fn,
) )
context.model, context.optimizer, context.dataloader, context.scheduler = (
executor.prepare(
model,
context.optimizer,
context.dataloader,
context.scheduler,
)
)
if context.checkpoint and context.checkpoint.extra: if context.checkpoint and context.checkpoint.extra:
extra = context.checkpoint.extra extra = context.checkpoint.extra
for name in ("optimizer", "scheduler"): for name in ("optimizer", "scheduler"):
@@ -193,18 +222,25 @@ class TrainContextBuilder:
obj.load_state_dict(extra[name]) obj.load_state_dict(extra[name])
strategy_kwargs = dict(cfg.extra_kwargs) strategy_kwargs = dict(cfg.extra_kwargs)
strategy_kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
if cfg.strategy in ("dpo", "grpo"): needs_ref = cfg.strategy in (
ref_model = create_ref_model( "dpo",
cfg.model_fn, executor.unwrap_model(context.model) "grpo",
).to(device=device) "online_grpo",
strategy_kwargs["ref_model"] = ref_model "online_dpo",
)
needs_old = cfg.strategy in ("grpo", "online_grpo")
if cfg.strategy == "grpo": if needs_ref:
old_model = create_ref_model( strategy_kwargs["ref_model"] = create_ref_model(
cfg.model_fn, executor.unwrap_model(context.model) cfg.model_fn, executor=executor, model=context.model, device=device
).to(device=device) )
strategy_kwargs["old_model"] = old_model
if needs_old:
strategy_kwargs["old_model"] = create_ref_model(
cfg.model_fn, executor=executor, model=context.model, device=device
)
context.strategy = StrategyFactory.create( context.strategy = StrategyFactory.create(
cfg.strategy, cfg.strategy,
@@ -214,4 +250,44 @@ class TrainContextBuilder:
**strategy_kwargs, **strategy_kwargs,
) )
# Enable online rollout when the train_type is an ``online_*`` variant.
is_online = cfg.strategy.startswith("online_")
if is_online:
if not context.strategy.supports_online():
raise ValueError(
f"Strategy '{cfg.strategy}' does not support online rollout"
)
if cfg.reward_model_fn is None:
raise ValueError("reward_model_fn is required for online RL strategies")
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
reward_model = cfg.reward_model_fn()
group_size = strategy_kwargs.get("group_size", 1)
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
max_seq_len = getattr(context.model.config, "max_position_embeddings", None)
scheduler = InferenceScheduler(
model=context.model,
tokenizer=tokenizer,
max_batch_size=rollout_batch_size,
max_seq_len=max_seq_len,
)
generator = RolloutGenerator(
scheduler=scheduler,
tokenizer=tokenizer,
max_tokens=cfg.rollout_max_tokens,
group_size=group_size,
temperature=cfg.rollout_temperature,
top_k=cfg.rollout_top_k,
top_p=cfg.rollout_top_p,
)
runner = RolloutRunner(
generator=generator,
reward_model=reward_model,
rollout_interval=cfg.rollout_interval,
)
context.strategy.set_rollout_runner(runner)
return context return context
+26 -4
View File
@@ -1,8 +1,14 @@
import logging import logging
from typing import List, Optional from typing import List, Optional
import torch.distributed as dist
from astrai.config import TrainConfig from astrai.config import TrainConfig
from astrai.parallel.setup import spawn_parallel_fn from astrai.parallel.setup import spawn_parallel_fn
from astrai.signal_handler import (
register_signal_handlers,
unregister_signal_handlers,
)
from astrai.trainer.train_callback import ( from astrai.trainer.train_callback import (
CallbackFactory, CallbackFactory,
TrainCallback, TrainCallback,
@@ -36,7 +42,7 @@ class Trainer:
), ),
CallbackFactory.create( CallbackFactory.create(
"metric", "metric",
log_dir=cfg.log_dir, ckpt_dir=cfg.ckpt_dir,
save_interval=cfg.ckpt_interval, save_interval=cfg.ckpt_interval,
metrics=cfg.metrics, metrics=cfg.metrics,
val_step=cfg.val_step, val_step=cfg.val_step,
@@ -58,6 +64,7 @@ class Trainer:
.with_param_path(param_path, resume=resume) .with_param_path(param_path, resume=resume)
.build() .build()
) )
register_signal_handlers(context)
executor = context.executor executor = context.executor
self._call_callbacks("on_train_begin", context) self._call_callbacks("on_train_begin", context)
@@ -65,15 +72,20 @@ class Trainer:
context.model.train() context.model.train()
for epoch in range(context.epoch, context.config.n_epoch): for epoch in range(context.epoch, context.config.n_epoch):
if context.stop_requested:
break
context.epoch = epoch context.epoch = epoch
self._call_callbacks("on_epoch_begin", context) self._call_callbacks("on_epoch_begin", context)
for batch in context.dataloader: for batch in context.dataloader:
if context.stop_requested:
break
with executor.accumulate(context.model): with executor.accumulate(context.model):
self._call_callbacks("on_batch_begin", context) self._call_callbacks("on_batch_begin", context)
loss = context.strategy(batch) loss_output = context.strategy(batch)
context.loss = loss.item() context.loss = loss_output["loss"].item()
stand_loss = loss / executor.grad_accum_steps context.metrics = loss_output["metrics"]
stand_loss = loss_output["loss"] / executor.grad_accum_steps
executor.backward(stand_loss) executor.backward(stand_loss)
context.consumed_samples += ( context.consumed_samples += (
context.config.batch_per_device * context.world_size context.config.batch_per_device * context.world_size
@@ -83,6 +95,7 @@ class Trainer:
if executor.sync_gradients: if executor.sync_gradients:
self._call_callbacks("on_optimizer_step", context) self._call_callbacks("on_optimizer_step", context)
context.optimizer.step() context.optimizer.step()
context.strategy.on_optimizer_step()
context.optimizer.zero_grad() context.optimizer.zero_grad()
if context.scheduler: if context.scheduler:
@@ -90,12 +103,21 @@ class Trainer:
self._call_callbacks("on_epoch_end", context) self._call_callbacks("on_epoch_end", context)
if context.stop_requested:
logger.warning(
"Training interrupted by signal, saving emergency checkpoint..."
)
self._call_callbacks("on_error", context)
except Exception as e: except Exception as e:
logger.error("Training failed: %s", str(e), exc_info=True) logger.error("Training failed: %s", str(e), exc_info=True)
self._call_callbacks("on_error", context) self._call_callbacks("on_error", context)
raise raise
finally: finally:
self._call_callbacks("on_train_end", context) self._call_callbacks("on_train_end", context)
if executor.use_distributed and dist.is_initialized():
dist.barrier()
unregister_signal_handlers()
def train(self, param_path: Optional[str] = None, resume: bool = False): def train(self, param_path: Optional[str] = None, resume: bool = False):
cfg = self.train_config cfg = self.train_config
+74
View File
@@ -0,0 +1,74 @@
cmake_minimum_required(VERSION 3.18)
project(astrai_kernels LANGUAGES CUDA CXX)
set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
set(CMAKE_CUDA_STANDARD 17)
find_package(CUDAToolkit REQUIRED)
if(NOT DEFINED TORCH_HOME)
set(TORCH_HOME "$ENV{TORCH_HOME}")
endif()
if(NOT TORCH_HOME)
message(FATAL_ERROR "TORCH_HOME must point at the torch install dir (site-packages/torch)")
endif()
if(NOT DEFINED PYTHON_INCLUDE_DIR)
set(PYTHON_INCLUDE_DIR "/usr/include/python${PYTHON_VERSION_MAJOR}.${PYTHON_VERSION_MINOR}")
endif()
if(NOT DEFINED ASTRAI_CUDA_ARCH)
if(DEFINED ENV{ASTRAI_CUDA_ARCH})
set(ASTRAI_CUDA_ARCH "$ENV{ASTRAI_CUDA_ARCH}")
else()
set(ASTRAI_CUDA_ARCH 80)
endif()
endif()
set(TORCH_LIB_DIR "${TORCH_HOME}/lib")
set(CUDA_LIB_DIR "/usr/local/cuda/lib64")
set(CXX_FLAGS -O3 -funroll-loops)
set(NVCC_FLAGS -O3
--expt-relaxed-constexpr
--use_fast_math
"--ptxas-options=-O3,-v"
--extra-device-vectorization
--threads=16)
set(TORCH_LIBS
"${TORCH_LIB_DIR}/libtorch_python.so"
"${TORCH_LIB_DIR}/libtorch_cuda.so"
"${TORCH_LIB_DIR}/libc10_cuda.so"
"${TORCH_LIB_DIR}/libtorch_cpu.so"
"${TORCH_LIB_DIR}/libtorch.so"
"${TORCH_LIB_DIR}/libc10.so"
CUDA::cudart)
set(CMAKE_CUDA_ARCHITECTURES "${ASTRAI_CUDA_ARCH}")
set(KERNELS attn_decode attn_prefill attn_paged_decode attn_paged_prefill rotary_emb)
foreach(name ${KERNELS})
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${name}.cu")
target_compile_definitions(${name} PRIVATE TORCH_EXTENSION_NAME=${name})
target_include_directories(${name} PRIVATE
"${TORCH_HOME}/include"
"${TORCH_HOME}/include/torch/csrc/api/include"
"${PYTHON_INCLUDE_DIR}")
target_link_libraries(${name} PRIVATE ${TORCH_LIBS})
target_link_options(${name} PRIVATE "-Wl,-rpath,${TORCH_LIB_DIR}")
target_compile_options(${name} PRIVATE
$<$<COMPILE_LANGUAGE:CXX>:${CXX_FLAGS}>
$<$<COMPILE_LANGUAGE:CUDA>:${NVCC_FLAGS}>)
set_target_properties(${name} PROPERTIES
PREFIX ""
SUFFIX ".${PY_SOABI}.so"
LIBRARY_OUTPUT_DIRECTORY "${CMAKE_CURRENT_SOURCE_DIR}/../astrai/extension/lib")
endforeach()
-48
View File
@@ -1,48 +0,0 @@
from pathlib import Path
def _arch_flags() -> list[str]:
import torch
if torch.cuda.is_available():
cap = torch.cuda.get_device_capability()
else:
cap = (8, 0)
ver = f"{cap[0]}{cap[1]}"
flags = [f"-gencode=arch=compute_{ver},code=sm_{ver}"]
# tensor-core mma path (mma.sync.m16n8k16.bf16) requires sm_80+; decide the
# kernel dispatch at build time via this define rather than at runtime.
if cap[0] < 8:
flags.append("-DASTRAI_NO_MMA")
return flags
_kernels_dir = Path("csrc/kernels")
REGISTRY: dict[str, dict] = {}
CXX_FLAGS = ["-O3", "-funroll-loops"]
NVCC_FLAGS = [
"-O3",
"--expt-relaxed-constexpr",
"--use_fast_math",
"--ptxas-options=-O3,-v",
"--extra-device-vectorization",
"--threads=8",
]
def register(name: str, sources: list[str] | None = None, **kwargs):
if sources is None:
sources = [str(_kernels_dir / f"{name}.cu")]
REGISTRY[name] = {
"sources": sources,
"cxx_flags": [*CXX_FLAGS],
"nvcc_flags": [*NVCC_FLAGS, *_arch_flags()],
"extra_link_args": kwargs.pop("extra_link_args", []),
**kwargs,
}
register("attn_decode")
register("attn_prefill")
register("attn_paged_decode")
+36 -38
View File
@@ -1,13 +1,27 @@
#pragma once #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> template<typename T, typename AT = float>
struct AttentionParams { struct AttentionParams {
// ---- shared across all paths ----
int batch; int batch;
int q_head; int q_head;
int kv_head; int kv_head;
int q_len;
int kv_len;
int head_dim; int head_dim;
int use_mask; int use_mask;
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
@@ -16,53 +30,37 @@ struct AttentionParams {
// Q strides (element offsets for each dim — layout-agnostic) // Q strides (element offsets for each dim — layout-agnostic)
int q_stride_b, q_stride_h, q_stride_l, q_stride_d; int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
// KV strides (K and V share the same layout — only base pointers differ)
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
// Mask: 2D [batch, kv_len] (mask_q_stride=0) or 3D [batch, q_len, kv_len] // Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
int mask_b_stride; // = kv_len (both 2D and 3D) // or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len int mask_b_stride; // batch stride
int mask_h_stride; // head stride (0 = broadcast across heads)
const T* __restrict__ q; int mask_q_stride; // q stride (0 = all q rows share)
const T* __restrict__ k;
const T* __restrict__ v;
const bool* __restrict__ mask; const bool* __restrict__ mask;
const T* __restrict__ q;
T* __restrict__ o; T* __restrict__ o;
AT* __restrict__ o_part; AT* __restrict__ o_part;
AT* __restrict__ ml_part; AT* __restrict__ ml_part;
};
template<typename T, typename AT = float> // ---- contiguous K/V mode ----
struct PagedAttentionParams {
int batch;
int q_head;
int kv_head;
int q_len; int q_len;
int kv_len; int kv_len;
int head_dim; int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
int use_mask; const T* __restrict__ k;
int causal_offset; const T* __restrict__ v;
float scale;
int num_splits; // ---- paged (SGLang flat pool) mode ----
int page_size;
int max_pages;
// Q strides (layout-agnostic)
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
// Mask strides (2D or 3D)
int mask_b_stride;
int mask_q_stride;
const T* __restrict__ q;
const T* __restrict__ k_cache; const T* __restrict__ k_cache;
const T* __restrict__ v_cache; const T* __restrict__ v_cache;
const bool* __restrict__ mask;
const int64_t* __restrict__ page_table;
T* __restrict__ o; // Indexing
AT* __restrict__ o_part; const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len]
AT* __restrict__ ml_part; const int64_t* __restrict__ req_pool_indices; // [batch]
const int* __restrict__ kv_indptr; // [batch+1]
const int* __restrict__ qo_indptr; // [batch+1] or nullptr (decode)
int max_context_len; // req_to_token stride (dim 1)
int max_seq_len; // max per-request seq_len (host-side, for split computation)
int total_q; // total Q tokens across all requests (host-side, for grid)
int max_q_len; // max per-request q_len (host-side, for prefill grid)
}; };
+9 -50
View File
@@ -1,51 +1,6 @@
#include "attn_decode_split_kv.cuh" #include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "attn_entry_utils.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_decode_split_kv_mma.cuh"
#endif
// Scalar fallback: one warp per query head, split-KV across grid.z.
static void launch_scalar_decode(AttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
alloc_split_partials(p);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, group_size), smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#ifndef ASTRAI_NO_MMA
// MMA head-packing requires G <= 16 (BR=16 rows). sm_80+ tensor-core
// + cp.async wins even at G=1 (decode is memory-bound, not compute-bound).
// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256
// (double-buffer would be 32 KB, near the 48 KB static cap — keep single
// to preserve occupancy).
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_mma_decode(AttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
template <int HEAD_DIM>
static void dispatch_decode(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16) {
launch_mma_decode<HEAD_DIM, 32>(p);
return;
}
#endif
launch_scalar_decode(p);
}
torch::Tensor attn_decode( torch::Tensor attn_decode(
torch::Tensor q, torch::Tensor q,
torch::Tensor k, torch::Tensor k,
@@ -55,17 +10,21 @@ torch::Tensor attn_decode(
double scale, double scale,
int64_t layout int64_t layout
) { ) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
AttentionParams<bf16> p; AttentionParams<bf16> p;
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p); attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1"); TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32"); TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
// O matches Q's original layout
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == 1) ? O.transpose(1, 2) : O; auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr(); p.o = (bf16*)O_view.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p); alloc_split_partials(p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O; return O;
} }
@@ -77,6 +36,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
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,
py::arg("layout") = 0, py::arg("layout") = (int64_t)BHLD,
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)"); "GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
} }
+50 -37
View File
@@ -2,16 +2,16 @@
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include <float.h> #include <float.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_kv_source.cuh"
using bf16 = __nv_bfloat16; #include "attn_warp_utils.cuh"
constexpr int DC_CHUNK = 64; constexpr int DC_CHUNK = 64;
__device__ inline float warp_reduce_sum(float val) { // Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified
for (int offset = 16; offset > 0; offset >>= 1) // across contiguous and paged (SGLang flat-pool) K/V via the KV template
val += __shfl_xor_sync(0xFFFFFFFF, val, offset); // parameter. For decode the query is the last token, so its valid range
return val; // [0, seq_len) IS the causal range; KV::decode_attend_len expresses that
} // bound per addressing mode (contig clips to causal_offset, paged = seq_len).
template <int HEAD_DIM, typename KV, bool IsCausal, bool HasMask>
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) { __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head; int batch = blockIdx.x / p.kv_head;
int kv_head = blockIdx.x % p.kv_head; int kv_head = blockIdx.x % p.kv_head;
@@ -21,63 +21,76 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int lane = threadIdx.x; int lane = threadIdx.x;
int hd_per_thread = p.head_dim / 32; int hd_per_thread = p.head_dim / 32;
const int seq_len = KV::kv_len(p, batch);
const KVContext kctx = KV::template make_ctx<HEAD_DIM>(p, batch, kv_head);
// Q: [batch, q_head, q_len=1, head_dim] — stride-based // Q: [batch, q_head, q_len=1, head_dim] — stride-based
float q_reg[8]; float q_reg[8];
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h int q_off = KV::q_decode_base(p, batch, q_head)
+ lane * hd_per_thread * p.q_stride_d; + lane * hd_per_thread * p.q_stride_d;
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
int mask_base = batch * p.mask_b_stride;
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f}; float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
extern __shared__ __align__(16) bf16 k_smem[]; extern __shared__ __align__(16) bf16 k_smem[];
// Split-KV: each split processes a contiguous subset of chunks // Split-KV: each split processes a contiguous subset of chunks
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; int chunks_total = (seq_len + DC_CHUNK - 1) / DC_CHUNK;
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits; int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
int ch_begin = split * chunks_per_split; int ch_begin = split * chunks_per_split;
int ch_end = min(chunks_total, ch_begin + chunks_per_split); int ch_end = min(chunks_total, ch_begin + chunks_per_split);
for (int ci = ch_begin; ci < ch_end; ci++) { for (int ci = ch_begin; ci < ch_end; ci++) {
int chunk_start = ci * DC_CHUNK; int chunk_start = ci * DC_CHUNK;
int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start); int this_chunk = min(DC_CHUNK, seq_len - chunk_start);
// Load K into shared memory (gather from strided global) // Load K into shared memory (addressing via KV policy; paged guards
// empty slots with zero-fill).
int total = this_chunk * p.head_dim; int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) { for (int i = threadIdx.y * 32 + lane; i < total;
i += blockDim.x * blockDim.y) {
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 kv_idx = chunk_start + s; int kc = chunk_start + s;
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d; KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true);
k_smem[i] = p.k[g_off]; k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
} }
__syncthreads(); __syncthreads();
for (int s = 0; s < this_chunk; s++) { for (int s = 0; s < this_chunk; s++) {
float partial = 0.0f; float partial = 0.0f;
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++)
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]); partial += q_reg[i] * __bfloat162float(
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = warp_reduce_sum(partial) * p.scale; partial = warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s; int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx]) if constexpr (HasMask) {
if (!p.mask[mask_base + kv_idx])
partial = -FLT_MAX; partial = -FLT_MAX;
if (p.causal_offset >= 0 && kv_idx > p.causal_offset) }
if constexpr (IsCausal) {
if (kv_idx >= KV::decode_attend_len(p, batch))
partial = -FLT_MAX; partial = -FLT_MAX;
}
float new_m = fmaxf(m, partial); float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m); float alpha = __expf(m - new_m);
float beta = expf(partial - new_m); float beta = __expf(partial - new_m);
d = d * alpha + beta; d = d * alpha + beta;
// V: stride-based read // V read via KV policy; when masked (beta == 0) or the slot is
int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d; // empty the term vanishes, so no extra branches are needed.
for (int i = 0; i < hd_per_thread; i++) for (int i = 0; i < hd_per_thread; i++) {
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta; KVAddr a = KV::kv_addr(p, kctx, kv_idx, lane * hd_per_thread + i, true);
float vv = a.valid
? __bfloat162float(*reinterpret_cast<const bf16*>(a.v))
: 0.0f;
acc_reg[i] = fmaf(acc_reg[i], alpha, vv * beta);
}
m = new_m; m = new_m;
} }
__syncthreads(); __syncthreads();
@@ -85,7 +98,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
// ---- write UN-normalised partials for this split ---- // ---- write UN-normalised partials for this split ----
size_t bh = (size_t)batch * p.q_head + q_head; size_t bh = (size_t)batch * p.q_head + q_head;
size_t slot = bh * p.num_splits + split; size_t slot = bh * MAX_SPLITS + split;
int d0 = lane * hd_per_thread; int d0 = lane * hd_per_thread;
for (int i = 0; i < hd_per_thread; i++) { for (int i = 0; i < hd_per_thread; i++) {
int dd = d0 + i; int dd = d0 + i;
@@ -97,9 +110,10 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
} }
} }
// Reduce split-K partials into the final bf16 output. One block per (batch, // Split-combine: merges the per-split partials (o_part/ml_part) into the
// q_head); each thread folds across all splits with a single-pass // final normalised O. KV selects the O addressing (contig batch stride vs
// online-rescale reduction (expf + FMA counts halved vs 3-pass original). // paged row stride).
template <typename KV>
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) { __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int bh = blockIdx.x; int bh = blockIdx.x;
int d = threadIdx.x; int d = threadIdx.x;
@@ -108,7 +122,7 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int batch = bh / p.q_head; int batch = bh / p.q_head;
int q_head = bh % p.q_head; int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits; size_t split_base = (size_t)bh * MAX_SPLITS;
const float* mlp = p.ml_part + split_base * 2; const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim; const float* op = p.o_part + split_base * p.head_dim;
@@ -120,13 +134,12 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
float nm = fmaxf(m, mi); float nm = fmaxf(m, mi);
float corr = __expf(m - nm); float corr = __expf(m - nm);
float e = __expf(mi - nm); float e = __expf(mi - nm);
acc = acc * corr + op[s * p.head_dim + d] * e; acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
l = l * corr + li * e; l = fmaf(l, corr, li * e);
m = nm; m = nm;
} }
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
// Stride-based output write (q_len=1 for decode, so stride_l not needed) int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_stride_d;
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
p.o[o_off] = __float2bfloat16(acc * inv); p.o[o_off] = __float2bfloat16(acc * inv);
} }
+97 -97
View File
@@ -2,160 +2,160 @@
#include <cfloat> #include <cfloat>
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_kv_source.cuh"
#include "attn_mma_utils.cuh" #include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16; // Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing. // parameter. Decode has q_len == 1, so we pack G = q_head/kv_head query
// heads into the M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs
// into a single GEMM that reuses each loaded K/V tile across all G heads.
// //
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core // KV = ContigKV (dense tensors) or PagedKV (flat pool + req_to_token).
// work on its own. But GQA gives us G = q_head / kv_head query heads that all // IsCausal and HasMask are compile-time bools — no runtime branch in the
// share one kv_head. We pack those G heads into the M=16 rows of // inner compute loop.
// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that //
// reuses each loaded K/V tile across all G heads (K/V load is the decode // Traits = KernelTraits<HEAD_DIM, BC=16, WARPS=1, STAGES=2>.
// bottleneck, so the reuse is the win, not the flops). The KV sequence is template <typename Traits, typename KV, bool IsCausal, bool HasMask>
// partitioned across gridDim.z blocks so that a decode with only
// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head,
// split) block computes an UN-normalised partial (Oacc, m, l) over its KV
// slice; the combine kernel below reduces across splits. Fixes the "grid too
// small" bottleneck (0.04 waves/SM → many blocks) for long-context,
// small-batch decode.
template <int HEAD_DIM, int BC, int STAGES = 2>
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) { __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x; const int lane = threadIdx.x;
const int gid = lane >> 2; const int gid = lane >> 2;
const int tid4 = lane & 3; const int tid4 = lane & 3;
const int kv_head = blockIdx.x; const int pass = blockIdx.x / p.kv_head;
const int kv_head = blockIdx.x % p.kv_head;
const int batch = blockIdx.y; const int batch = blockIdx.y;
const int split = blockIdx.z; const int split = blockIdx.z;
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G;
// Double-buffered shared memory for K/V (no sQ needed — Q goes direct constexpr int MAX_G = 16;
// from global to registers). const int G_total = p.q_head / p.kv_head;
__shared__ __align__(16) bf16 sK[STAGES * BC * LD]; const int g_begin = pass * MAX_G;
__shared__ __align__(16) bf16 sV[STAGES * BC * LD]; const int G = min(MAX_G, G_total - g_begin);
const int q_head0 = kv_head * G_total + g_begin;
// ---- Load Q directly from global into mma A-operand registers ---- // Per-request seq_len (paged reads kv_indptr; contig uses p.kv_len).
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h; const int seq_len = KV::kv_len(p, batch);
const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head);
// Double-buffered shared memory for K/V (no sQ needed)
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
// Load Q directly from global into mma A-operand registers.
const int q_base = KV::q_decode_base(p, batch, q_head0);
const int qra = gid; const int qra = gid;
const int qrb = gid + 8; const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G; const bool va = qra < G, vb = qrb < G;
unsigned Qa[KD][4]; unsigned Qa[Traits::KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d, load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa); qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4]; float Oacc[Traits::DN8][4];
#pragma unroll #pragma unroll
for (int j = 0; j < DN8; j++) for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base — [batch, kv_head, kv_len, head_dim] const int tiles_total = (seq_len + Traits::BC - 1) / Traits::BC;
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
const int tiles_total = (p.kv_len + BC - 1) / BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits; const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split; const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split); const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// ---- Load tile lambda: predicated cp.async, unified full/partial ---- // ---- Load tile lambda: predicated cp.async (addressing via KV policy) ----
auto load_tile = [&](int ti, int buf) { auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * BC * LD; bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * BC * LD; bf16* dV = sV + buf * Traits::BC * Traits::LD;
#pragma unroll #pragma unroll
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) { for (int i = lane * Traits::VEC; i < Traits::TOTAL;
int r = i / HEAD_DIM, d = i % HEAD_DIM; i += Traits::NUM_THREADS * Traits::VEC) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < p.kv_len; bool valid = kc < seq_len;
int off = r * LD + swiz_col(d, r, SWIZ_MASK); KVAddr a = KV::kv_addr(p, kctx, kc, d, valid);
// KV stride-based: contiguous within head_dim (stride_d == 1 typically) int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d; cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dK[off], &p.k[g_off], valid); cp_async_16_pred(&dV[off], a.v, a.valid);
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
} }
cp_async_commit(); cp_async_commit();
}; };
// ---- Prologue: issue first tile load ---- // ---- Multi-stage cp.async pipeline ----
if (ti_begin < ti_end) { // Prologue loads STAGES tiles; each loop iteration waits only for the
load_tile(ti_begin, 0); // oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
} // tile loads stay in flight and overlap with the current tile's compute.
constexpr int STAGES = Traits::STAGES;
const int ntiles = ti_end - ti_begin;
for (int ti = ti_begin; ti < ti_end; ti++) { auto process_tile = [&](int it, int buf) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0; const bf16* bK = sK + buf * Traits::BC * Traits::LD;
int buf = (ti - ti_begin) & BUF_MASK; const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = (ti_begin + it) * Traits::BC;
// Wait for current tile, then issue next tile's prefetch (overlaps float Sacc[Traits::NC8][4];
// with this tile's compute). Single syncwarp covers both hazards. mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
// When STAGES==1, no prefetch — load happens at end of prior iter.
cp_async_wait_group<0>();
__syncwarp();
if constexpr (STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
const bf16* bK = sK + buf * BC * LD;
const bf16* bV = sV + buf * BC * LD;
int kv0 = ti * BC;
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant // Decode: q_len=1, so qrow0=qrow1=0. Paged treats [0, seq_len) as
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len; // the causal range (query is the last token); contig clips to the
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc, // causal_offset bound. Dead code eliminated when IsCausal == false.
int maxc = IsCausal ? KV::decode_attend_len(p, batch) : seq_len;
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
0, 0, 0, 0,
p.mask_b_stride, 0, p.mask_b_stride, 0, 0,
batch, batch, 0,
p.mask, has_mask, p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane); Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
__syncwarp(); };
if constexpr (STAGES == 1) { if (ntiles >= STAGES) {
if (ti + 1 < ti_end) #pragma unroll
load_tile(ti + 1, 0); for (int i = 0; i < STAGES; i++)
load_tile(ti_begin + i, i);
for (int it = 0; it < ntiles; it++) {
cp_async_wait_group<STAGES - 1>();
__syncwarp();
process_tile(it, it & (STAGES - 1));
__syncwarp();
if (it + STAGES < ntiles)
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
} }
} else {
// Fewer tiles than stages: load all, wait for all, process.
for (int i = 0; i < ntiles; i++)
load_tile(ti_begin + i, i);
cp_async_wait_group<0>();
__syncwarp();
for (int it = 0; it < ntiles; it++)
process_tile(it, it);
} }
// ---- write UN-normalised partials for this split ---- // ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t { auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h; size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split; return bh * MAX_SPLITS + split;
}; };
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4; int d = dn8 * 8 + 2 * tid4;
int r0 = gid, r1 = gid + 8; int r0 = gid, r1 = gid + 8;
if (r0 < G) { if (r0 < G) {
int h = q_head0 + r0; int h = q_head0 + r0;
float* op = p.o_part + split_slot(h) * HEAD_DIM; float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][0]; op[d] = Oacc[dn8][0];
op[d + 1] = Oacc[dn8][1]; op[d + 1] = Oacc[dn8][1];
} }
if (r1 < G) { if (r1 < G) {
int h = q_head0 + r1; int h = q_head0 + r1;
float* op = p.o_part + split_slot(h) * HEAD_DIM; float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
op[d] = Oacc[dn8][2]; op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3]; op[d + 1] = Oacc[dn8][3];
} }
+205
View File
@@ -0,0 +1,205 @@
#pragma once
// Shared attention dispatchers — used by both production .cu and test .cu.
// No torch dependency; pure CUDA.
//
// The paged and contiguous kernels are unified by the KVSource policy
// (ContigKV / PagedKV from attn_kv_source.cuh), so each launcher struct
// 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
// covered by KV::host_q_len / KV::host_kv_len.
#include <cuda_runtime.h>
#include <algorithm>
#include "attn_warp_utils.cuh"
#include "attn_kv_source.cuh"
#include "attn_prefill_split_q.cuh"
#include "attn_decode_split_kv.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_prefill_split_q_mma.cuh"
#include "attn_decode_split_kv_mma.cuh"
#endif
// 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,
// avoiding excessive loop/prologue overhead when tiles are small.
//
// Target total grid blocks (`TARGET_BLOCKS`) rather than scaling splits by SM
// count. Decode blocks are single-warp (32 threads) and a SM hosts ~11 of
// them, so the old `2*sm/base` cap badly undersplit at large batch (B=16 got
// 3 splits, optimal ~8). Measured (L20, grid search): bandwidth saturates
// near 256-512 total blocks; 512 minimizes worst-case latency across the
// B x kv grid; more is pure oversplit overhead.
constexpr int DECODE_TARGET_BLOCKS = 512;
inline int compute_num_splits(int base_blocks, int tiles_total,
int min_tiles_per_split = 1) {
int n = (DECODE_TARGET_BLOCKS + base_blocks - 1) / base_blocks;
int max_by_work = tiles_total / min_tiles_per_split;
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
}
// Dispatch IsCausal × HasMask — eliminates the duplicated 4-way if/else
// ladder that appeared in each dispatch_* function. FN must be a function
// template <int HEAD_DIM, bool IsCausal, bool HasMask>; HEAD_DIM is forwarded
// as the first template argument so callers only spell it once.
//
// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launcher<KV>::template launch, HEAD_DIM, p, stream);
#define DISPATCH_CAUSAL_MASK(is_causal, has_mask, FN, HEAD_DIM, ...) \
do { \
if (is_causal) { \
if (has_mask) FN<HEAD_DIM, true, true>(__VA_ARGS__); \
else FN<HEAD_DIM, true, false>(__VA_ARGS__); \
} else { \
if (has_mask) FN<HEAD_DIM, false, true>(__VA_ARGS__); \
else FN<HEAD_DIM, false, false>(__VA_ARGS__); \
} \
} while (0)
// ======================================================================
// Prefill launchers (KV selects ContigKV or PagedKV addressing)
// ======================================================================
#ifndef ASTRAI_NO_MMA
template <typename KV>
struct PrefillLauncherMMA {
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
constexpr int WARPS = 4;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
int q_len = KV::host_q_len(p);
dim3 grid((q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS),
p.q_head, p.batch);
dim3 block(Traits::NUM_THREADS);
attn_prefill_split_q_mma_kernel<Traits, KV, IsCausal, HasMask>
<<<grid, block, 0, stream>>>(p);
}
};
#endif
template <typename KV>
struct PrefillLauncherScalar {
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
constexpr int G = 8, ROWS = 32, P_BC = 32;
int q_len = KV::host_q_len(p);
dim3 grid((q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
dim3 block(G, ROWS);
attn_prefill_split_q_kernel_t<HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>
<<<grid, block, 0, stream>>>(p);
}
};
template <int HEAD_DIM>
static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherMMA<ContigKV>::template launch,
HEAD_DIM, p, stream);
#else
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherScalar<ContigKV>::template launch,
HEAD_DIM, p, stream);
#endif
}
template <int HEAD_DIM>
static inline void dispatch_paged_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
#ifndef ASTRAI_NO_MMA
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherMMA<PagedKV>::template launch,
HEAD_DIM, p, stream);
#else
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
PrefillLauncherScalar<PagedKV>::template launch,
HEAD_DIM, p, stream);
#endif
}
// ======================================================================
// Decode launchers (KV selects ContigKV or PagedKV addressing)
// ======================================================================
#ifndef ASTRAI_NO_MMA
// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM).
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
// the 176-byte spill that STAGES=1+BC=32 suffered.
template <typename KV>
struct DecodeLauncherMMA {
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
int G = p.q_head / p.kv_head;
constexpr int MAX_G = 16;
int num_passes = (G + MAX_G - 1) / MAX_G;
constexpr int BC = 16;
int kv_len = KV::host_kv_len(p);
int tiles_total = (kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2);
constexpr int STAGES = 2;
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
attn_decode_split_kv_mma_kernel<Traits, KV, IsCausal, HasMask>
<<<grid, 32, 0, stream>>>(p);
}
};
#endif
template <typename KV>
struct DecodeLauncherScalar {
template <int HEAD_DIM, bool IsCausal, bool HasMask>
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
int kv_len = KV::host_kv_len(p);
int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, g);
attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask>
<<<grid, block, smem, stream>>>(p);
}
};
template <int HEAD_DIM>
static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
int group_size = p.q_head / p.kv_head;
#ifndef ASTRAI_NO_MMA
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherMMA<ContigKV>::template launch,
HEAD_DIM, p, group_size, stream);
#else
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherScalar<ContigKV>::template launch,
HEAD_DIM, p, group_size, stream);
#endif
attn_decode_combine_kernel<ContigKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
}
template <int HEAD_DIM>
static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
bool is_causal = (p.causal_offset >= 0);
bool has_mask = (p.use_mask && p.mask);
int group_size = p.q_head / p.kv_head;
#ifndef ASTRAI_NO_MMA
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherMMA<PagedKV>::template launch,
HEAD_DIM, p, group_size, stream);
#else
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
DecodeLauncherScalar<PagedKV>::template launch,
HEAD_DIM, p, group_size, stream);
#endif
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
}
+179 -45
View File
@@ -1,36 +1,35 @@
#pragma once #pragma once
#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 "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16; using bf16 = __nv_bfloat16;
inline int compute_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
return std::max(1, std::min(n, std::min(tiles_total, 32)));
}
// 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, arg) // Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
// Expands to: fn<32>(arg); fn<64>(arg); etc. // Expands to: fn<32>(args...); fn<64>(args...); etc.
#define DISPATCH_HEAD_DIM(hd, fn, arg) \ #define DISPATCH_HEAD_DIM(hd, fn, ...) \
switch (hd) { \ switch (hd) { \
case 32: fn<32>(arg); break; \ case 32: fn<32>(__VA_ARGS__); break; \
case 64: fn<64>(arg); break; \ case 64: fn<64>(__VA_ARGS__); break; \
case 128: fn<128>(arg); break; \ case 128: fn<128>(__VA_ARGS__); break; \
case 256: fn<256>(arg); break; \ case 256: fn<256>(__VA_ARGS__); break; \
default: \ default: \
TORCH_CHECK(false, "unsupported head_dim ", hd, \ TORCH_CHECK(false, "unsupported head_dim ", hd, \
" (supported: 32, 64, 128, 256)"); \ " (supported: 32, 64, 128, 256)"); \
} }
// 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
// skips them. Allocators are therefore left uninitialized (torch::empty); the
// per-call memset (torch::zeros / torch::full) was pure overhead.
template<typename P> template<typename P>
inline void alloc_split_partials(P& p) { inline void alloc_split_partials(P& p) {
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt); auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt); auto ml_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
p.o_part = (float*)o_part.data_ptr(); p.o_part = (float*)o_part.data_ptr();
p.ml_part = (float*)ml_part.data_ptr(); p.ml_part = (float*)ml_part.data_ptr();
} }
@@ -38,7 +37,7 @@ inline void alloc_split_partials(P& p) {
// ---- Shared Q-dims + strides extraction ---- // ---- Shared Q-dims + strides extraction ----
template <typename P> template <typename P>
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) { inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
if (layout == 1) q = q.transpose(1, 2); if (layout == BLHD) q = q.transpose(1, 2);
p.batch = (int)q.size(0); p.batch = (int)q.size(0);
p.q_head = (int)q.size(1); p.q_head = (int)q.size(1);
p.q_len = (int)q.size(2); p.q_len = (int)q.size(2);
@@ -50,6 +49,9 @@ inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
} }
// ---- Shared mask packing ---- // ---- Shared mask packing ----
// Accepts 2D [batch, kv_len], 3D [batch, q_len, kv_len],
// or 4D [batch, n_heads, q_len, kv_len].
// Head/q dimensions with size 1 broadcast (stride set to 0).
template <typename P> template <typename P>
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) { inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
if (p.use_mask) { if (p.use_mask) {
@@ -60,18 +62,26 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch"); TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
if (m.dim() == 2) { if (m.dim() == 2) {
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_q_stride = 0;
} else if (m.dim() == 3) { } else if (m.dim() == 3) {
TORCH_CHECK(m.size(1) == p.q_len, "mask q_len mismatch"); TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
p.mask_b_stride = (int)m.stride(0); p.mask_b_stride = (int)m.stride(0);
p.mask_q_stride = (int)m.stride(1); p.mask_h_stride = 0;
p.mask_q_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
} else if (m.dim() == 4) {
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
} else { } else {
TORCH_CHECK(false, "mask must be 2D [batch, kv_len] or 3D [batch, q_len, kv_len]"); TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
} }
p.mask = m.data_ptr<bool>(); p.mask = m.data_ptr<bool>();
} else { } else {
p.mask = nullptr; p.mask = nullptr;
p.mask_b_stride = 0; p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_q_stride = 0;
} }
} }
@@ -99,7 +109,7 @@ inline void attn_pack_params(
extract_q_dims_and_strides(q, layout, p); extract_q_dims_and_strides(q, layout, p);
if (layout == 1) k = k.transpose(1, 2), v = v.transpose(1, 2); if (layout == BLHD) k = k.transpose(1, 2), v = v.transpose(1, 2);
p.kv_head = (int)k.size(1); p.kv_head = (int)k.size(1);
p.kv_len = (int)k.size(2); p.kv_len = (int)k.size(2);
@@ -124,54 +134,178 @@ inline void attn_pack_params(
pack_mask(mask, p); pack_mask(mask, p);
} }
// ---- attn_pack_paged_params ---- // ---- attn_pack_paged_decode_params ----
// SGLang-style: flat KV pool + req_to_token indexing + variable
// seq_lens via kv_indptr. Q is [batch, q_head, head_dim] (q_len=1 per req).
template<typename T> template<typename T>
inline void attn_pack_paged_params( inline void attn_pack_paged_decode_params(
torch::Tensor q, torch::Tensor q,
torch::Tensor page_table,
torch::Tensor k_cache, torch::Tensor k_cache,
torch::Tensor v_cache, torch::Tensor v_cache,
int64_t page_size, torch::Tensor req_to_token,
int64_t kv_len, torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
int64_t max_seq_len,
c10::optional<torch::Tensor> mask, c10::optional<torch::Tensor> mask,
int64_t causal_offset, int64_t causal_offset,
double scale, double scale,
int64_t layout, AttentionParams<T>& p
PagedAttentionParams<T>& p
) { ) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
TORCH_CHECK(q.is_cuda() && page_table.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() && kv_indptr.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(page_table.dtype() == torch::kLong, "page_table must be int64"); TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes"); TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
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.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
TORCH_CHECK(q.dim() == 3, "q must be 3D [batch, q_head, head_dim]");
extract_q_dims_and_strides(q, layout, p); p.batch = (int)q.size(0);
p.q_head = (int)q.size(1);
p.kv_head = (int)k_cache.size(2); p.head_dim = (int)q.size(2);
p.kv_len = (int)kv_len; p.kv_head = (int)k_cache.size(1);
p.page_size = (int)page_size; TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
p.max_pages = (int)page_table.size(1);
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32"); TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
TORCH_CHECK(k_cache.size(1) == page_size, TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
"k_cache dim 1 must equal page_size, got ",
k_cache.size(1), " vs ", page_size); p.q_stride_l = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1);
p.q_stride_d = (int)q.stride(2);
p.k_cache = (const T*)k_cache.data_ptr();
p.v_cache = (const T*)v_cache.data_ptr();
p.q = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int64_t>();
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
p.kv_indptr = kv_indptr.data_ptr<int>();
p.qo_indptr = nullptr;
p.max_context_len = (int)req_to_token.size(1);
p.max_seq_len = (int)max_seq_len;
p.total_q = p.batch; // decode: 1 Q token per request
p.max_q_len = 1;
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);
p.page_table = page_table.data_ptr<int64_t>(); if (p.use_mask) {
auto m = mask.value();
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = 0;
p.mask = m.data_ptr<bool>();
} else {
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}
p.o = nullptr;
p.o_part = nullptr;
p.ml_part = nullptr;
}
// ---- attn_pack_paged_prefill_params ----
// SGLang-style: flat KV pool + req_to_token + ragged batch via qo_indptr.
// Q is [total_q, q_head, head_dim] (flattened across all requests).
template<typename T>
inline void attn_pack_paged_prefill_params(
torch::Tensor q,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
torch::Tensor qo_indptr,
c10::optional<torch::Tensor> mask,
int64_t max_q_len,
int64_t causal_offset,
double scale,
AttentionParams<T>& p
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
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(kv_indptr.is_cuda() && qo_indptr.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q 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(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
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(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(q.dim() == 3, "q must be 3D [total_q, q_head, head_dim]");
p.q_head = (int)q.size(1);
p.head_dim = (int)q.size(2);
p.kv_head = (int)k_cache.size(1);
p.batch = (int)req_pool_indices.size(0);
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
TORCH_CHECK(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]");
p.q_stride_l = (int)q.stride(0);
p.q_stride_h = (int)q.stride(1);
p.q_stride_d = (int)q.stride(2);
p.k_cache = (const T*)k_cache.data_ptr(); p.k_cache = (const T*)k_cache.data_ptr();
p.v_cache = (const T*)v_cache.data_ptr(); p.v_cache = (const T*)v_cache.data_ptr();
p.q = (const T*)q.data_ptr(); p.q = (const T*)q.data_ptr();
p.req_to_token = req_to_token.data_ptr<int64_t>();
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
p.kv_indptr = kv_indptr.data_ptr<int>();
p.qo_indptr = qo_indptr.data_ptr<int>();
p.max_context_len = (int)req_to_token.size(1);
p.total_q = (int)q.size(0); // prefill: flattened Q across all requests
p.max_q_len = (int)max_q_len;
// max_seq_len is unused by the prefill path (decode uses it for split
// computation); fill with max_q_len only to keep the POD struct defined.
p.max_seq_len = p.max_q_len;
p.causal_offset = (int)causal_offset;
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
if (p.use_mask) {
auto m = mask.value();
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
if (m.dim() == 2) {
TORCH_CHECK(m.size(1) <= p.max_context_len, "mask kv_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = 0;
p.mask_q_stride = 0;
} else if (m.dim() == 4) {
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_head, "mask head mismatch");
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.max_q_len, "mask q_len mismatch");
TORCH_CHECK(m.size(3) <= p.max_context_len, "mask kv_len mismatch");
p.mask_b_stride = (int)m.stride(0);
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
} else {
TORCH_CHECK(false, "mask must be 2D or 4D");
}
p.mask = m.data_ptr<bool>();
} else {
p.mask = nullptr;
p.mask_b_stride = 0;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
p.o = nullptr; p.o = nullptr;
p.o_part = nullptr; p.o_part = nullptr;
p.ml_part = nullptr; p.ml_part = nullptr;
pack_mask(mask, p);
} }
+159
View File
@@ -0,0 +1,159 @@
#pragma once
#include <cuda_bf16.h>
#include "attn_common.h"
// ============================================================================
// KVSource policies — the single dimension along which the paged and
// non-paged attention kernels differ. Each kernel is templated on one of
// these (ContigKV / PagedKV) and stays fully generic: the policy owns every
// place where "where does K/V live" and "what is this request's seq_len"
// are answered. All methods are __host__ __device__ so the same policy
// serves both the device kernels (addressing, seq_len) and the host-side
// launchers (grid / split computation).
//
// ContigKV: K/V are dense [batch, kv_head, kv_len, head_dim] tensors.
// Params fields used: k, v, kv_stride_*, kv_len, q_len,
// q_stride_b, causal_offset.
// PagedKV: K/V live in a flat pool [size, kv_head, head_dim] indexed via
// req_to_token. Params fields used: k_cache, v_cache,
// req_to_token, req_pool_indices, kv_indptr, qo_indptr,
// max_context_len, q_stride_l.
//
// Addressing state that is constant across a whole kernel invocation for one
// (batch, kv_head) pair is captured once by make_ctx<HEAD_DIM>() and passed
// to kv_addr, so the load loops never redo the hoistable base computation
// (e.g. the req_pool_indices global read) element-by-element.
// ============================================================================
// Every policy method is static + callable from both host and device code.
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
using bf16 = __nv_bfloat16;
// Hoisted per-(batch, kv_head) addressing context.
struct KVContext {
int kv_base; // contig: batch*kv_stride_b + kv_head*kv_stride_h
int64_t req_idx; // paged: req_pool_indices[batch]
int64_t rtt_stride; // paged: max_context_len
int64_t pool_stride; // paged: kv_head * HEAD_DIM
int64_t head_off; // paged: kv_head * HEAD_DIM
};
// Per-element K/V global addresses for one (kc, d) position of a K/V tile.
// The pointers are ALWAYS the computed addresses (never nullptr) — callers
// gate on `valid` (cp.async src_size=0, or a guarded scalar deref). `valid`
// starts as "within the request's seq_len"; the paged policy further degrades
// it when req_to_token maps the position to a negative slot (empty padding).
// This matches the original hand-rolled load loops, where the address was
// always formed and the predicate decided whether anything was read.
struct KVAddr {
const void* k;
const void* v;
bool valid;
};
// ---- Contiguous K/V ----
struct ContigKV {
static constexpr bool kPaged = false;
// host-side length hooks (grid + split computation in the launchers)
HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams<bf16>& p) {
return p.q_len;
}
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
return p.kv_len;
}
// prefill: element offset of the request's Q rows (kernel adds qrow*q_stride_l)
HOST_DEV_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_b + q_head * p.q_stride_h;
}
// decode: same offset (q_len == 1, so there is no row stride component)
HOST_DEV_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_b + q_head * p.q_stride_h;
}
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
return p.kv_len;
}
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
return p.q_len;
}
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
return p.causal_offset;
}
// decode: exclusive bound of the single query's attend range
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1);
}
template <int HEAD_DIM>
HOST_DEV_FORCEINLINE KVContext make_ctx(
const AttentionParams<bf16>& p, int batch, int kv_head) {
KVContext c = {};
c.kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
return c;
}
HOST_DEV_FORCEINLINE KVAddr kv_addr(
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
const int g_off = c.kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
return {&p.k[g_off], &p.v[g_off], valid};
}
};
// ---- Paged (SGLang-style flat pool) K/V ----
struct PagedKV {
static constexpr bool kPaged = true;
HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams<bf16>& p) {
return p.max_q_len;
}
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
return p.max_seq_len;
}
// prefill: Q rows start at qo_indptr[batch] (ragged batch base)
HOST_DEV_FORCEINLINE int q_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return p.qo_indptr[batch] * p.q_stride_l + q_head * p.q_stride_h;
}
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
HOST_DEV_FORCEINLINE int q_decode_base(
const AttentionParams<bf16>& p, int batch, int q_head) {
return batch * p.q_stride_l + q_head * p.q_stride_h;
}
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
return p.kv_indptr[batch + 1] - p.kv_indptr[batch];
}
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
return p.qo_indptr[batch + 1] - p.qo_indptr[batch];
}
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
return kv_len(p, batch) - q_len(p, batch);
}
// decode: the query is the last token, so [0, seq_len) IS its causal range
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
return kv_len(p, batch);
}
template <int HEAD_DIM>
HOST_DEV_FORCEINLINE KVContext make_ctx(
const AttentionParams<bf16>& p, int batch, int kv_head) {
KVContext c = {};
c.req_idx = p.req_pool_indices[batch];
c.rtt_stride = (int64_t)p.max_context_len;
c.pool_stride = (int64_t)p.kv_head * HEAD_DIM;
c.head_off = (int64_t)kv_head * HEAD_DIM;
return c;
}
HOST_DEV_FORCEINLINE KVAddr kv_addr(
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
const int64_t slot = valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : 0;
const bool ok = valid && (slot >= 0);
const int64_t gmem_off = slot * c.pool_stride + c.head_off + d;
return {&p.k_cache[gmem_off], &p.v_cache[gmem_off], ok};
}
};
+84 -80
View File
@@ -3,10 +3,47 @@
#include <cuda_fp16.h> #include <cuda_fp16.h>
#include <cuda_runtime.h> #include <cuda_runtime.h>
// Shared MMA utilities for tensor-core GQA kernels. // Predicated cp.async (4-operand form) requires CUDA 11.2+.
// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing. // bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA).
#if CUDART_VERSION < 11020
#error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)."
#endif
// ============================================================================
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
//
// Bundles all dimension-dependent constants so device functions only need a
// single Traits template parameter rather than scattered <KD, NC8, KT2, ...>.
// ============================================================================
template <int HEAD_DIM_, int BC_, int WARPS_, int STAGES_>
struct KernelTraits {
static constexpr int HEAD_DIM = HEAD_DIM_;
static constexpr int BC = BC_; // K/V tile size along seq dim
static constexpr int WARPS = WARPS_; // warps per block
static constexpr int STAGES = STAGES_; // double-buffer stages (1 or 2)
static constexpr int BR = 16; // Q rows per warp (mma M=16)
// Derived: mma.sync.m16n8k16 tile counts
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides
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 DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
static constexpr int LD = HEAD_DIM; // smem leading dim
// XOR swizzle chunk bits for ldmatrix bank-conflict avoidance.
// mask = log2(LD/8) bits, clamped to stay within LD.
static constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
static constexpr int NUM_THREADS = WARPS * 32;
static constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
static constexpr int TOTAL = BC * HEAD_DIM; // total elements per tile
};
// ---- PTX wrappers ----
using bf16 = __nv_bfloat16;
// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32
__device__ __forceinline__ void mma16816(float* d, const unsigned* a, __device__ __forceinline__ void mma16816(float* d, const unsigned* a,
const unsigned* b, const float* c) { const unsigned* b, const float* c) {
asm volatile( asm volatile(
@@ -37,9 +74,7 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
} }
// ldmatrix: cooperatively load mma fragments from smem (one instruction per // ldmatrix: cooperatively load mma fragments from smem (one instruction per
// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the // 16x16 / 16x8 tile) with the exact register layout mma expects.
// scalar per-thread fragment packing, cutting shared-load instructions and bank
// conflicts. Each lane supplies the shared address of one 8-wide row.
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) { __device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p); unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];" asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
@@ -60,29 +95,19 @@ __device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
} }
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity. // XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
// Eliminates ldmatrix bank conflicts without LD padding: consecutive rows
// land in distinct bank groups. swiz_col(d, r, mask) = ((d>>3)^(r&mask))<<3 | (d&7).
// mask must cover log2(HEAD_DIM/8) chunk bits but stay within LD: use 7 for
// HEAD_DIM>=64 (8+ chunks), 3 for HEAD_DIM=32 (4 chunks). Default 7 keeps
// existing HEAD_DIM>=64 call sites working unchanged.
__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: copy 16 bytes (8 bf16) from global to shared memory directly.
// bypassing registers. Eliminates shared-store bank conflicts and cuts
// load-loop instruction count in half (1 cp.async vs 1 LDG + 1 STS).
// Requires sm_80+.
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) { __device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr); unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;" asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
:: "r"(smem_addr), "l"(gmem_ptr)); :: "r"(smem_addr), "l"(gmem_ptr));
} }
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill the // Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// destination (src-size operand = 0 → no bytes read from src, so an // src_size=0 → no bytes read from src, so out-of-bounds src address is safe.
// out-of-bounds src address is never dereferenced). Lets full and partial
// tiles share one uniform async load path — no scalar fallback branch.
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr, __device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr, const void* gmem_ptr,
bool pred) { bool pred) {
@@ -100,9 +125,6 @@ __device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;"); asm volatile("cp.async.wait_all;");
} }
// Wait until at most N commit groups are still in flight. Used for
// double-buffered pipelining: wait_group<1> lets the next tile's cp.async
// continue while ensuring the current tile's data is ready.
template <int N> template <int N>
__device__ __forceinline__ void cp_async_wait_group() { __device__ __forceinline__ void cp_async_wait_group() {
asm volatile("cp.async.wait_group %0;" :: "n"(N)); asm volatile("cp.async.wait_group %0;" :: "n"(N));
@@ -139,78 +161,65 @@ __device__ inline void load_q_mma_frags(
} }
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
// Shared MMA compute functions — used by both decode and prefill MMA kernels.
// Extracted because S=Q@K^T, online softmax, and P@V are structurally identical
// between the two kernels; only the per-row causal/mask bounds differ.
// ---------------------------------------------------------------------------
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the // S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
// caller to avoid bf16 precision loss). // caller to avoid bf16 precision loss).
// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as // Traits provides KD, NC8, LD, and SWIZ_MASK.
// runtime ints lets the compiler fold them while keeping the signature clean. // ---------------------------------------------------------------------------
template <int KD, int NC8> template <typename Traits>
__device__ inline void mma_compute_scores( __device__ inline void mma_compute_scores(
const unsigned Qa[KD][4], const unsigned Qa[Traits::KD][4],
const bf16* __restrict__ sK, const bf16* __restrict__ sK,
int LD,
int SWIZ_MASK,
int lane, int lane,
float Sacc[NC8][4]) float Sacc[Traits::NC8][4])
{ {
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) { for (int n8 = 0; n8 < Traits::NC8; n8++) {
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f; Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
int krow_l = n8 * 8 + (lane & 7); int krow_l = n8 * 8 + (lane & 7);
int kcol_h = (lane & 8) ? 8 : 0; int kcol_h = (lane & 8) ? 8 : 0;
#pragma unroll #pragma unroll
for (int kt = 0; kt < KD; kt++) { for (int kt = 0; kt < Traits::KD; kt++) {
unsigned b[2]; unsigned b[2];
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]); ldmatrix_x2(b, &sK[krow_l * Traits::LD
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]); mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
} }
} }
} }
// ---------------------------------------------------------------------------
// Online softmax + Oacc rescale for one K/V tile. // Online softmax + Oacc rescale for one K/V tile.
// maxc0/maxc1: per-row KV column bounds (prefill: per-query-row causal limits; //
// decode: same value for both rows since q_len==1). // HasMask is a compile-time template bool: when false, the mask branch is
// qrow0/qrow1: query row indices (for 3D mask indexing; decode passes 0). // entirely dead-code-eliminated from the inner unrolled loop.
// mask_b_stride/mask_q_stride: mask layout (2D: mask_q_stride=0; 3D: =kv_len). // ---------------------------------------------------------------------------
// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm), template <typename Traits, bool HasMask>
// rescales Oacc by exp(m_old - nm), and updates m/l — all in place.
template <int NC8, int DN8>
__device__ inline void mma_softmax_tile( __device__ inline void mma_softmax_tile(
int kv0, int kv0,
int maxc0, int maxc0, int maxc1,
int maxc1, int qrow0, int qrow1,
int qrow0, int mask_b_stride, int mask_h_stride, int mask_q_stride,
int qrow1, int mask_batch, int mask_head,
int mask_b_stride,
int mask_q_stride,
int mask_batch,
const bool* __restrict__ mask, const bool* __restrict__ mask,
bool has_mask, float Sacc[Traits::NC8][4],
float Sacc[NC8][4], float Oacc[Traits::DN8][4],
float Oacc[DN8][4],
float& m0, float& m1, float& m0, float& m1,
float& l0, float& l1, float& l0, float& l1,
int lane) int lane)
{ {
int tid4 = lane & 3; int tid4 = lane & 3;
// Mask out-of-bounds / masked columns: set -FLT_MAX so expf → 0 downstream
// without per-element sentinel checks. Compute tile-local row maxima.
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX; float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride; int mask_base0 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow0 * mask_q_stride;
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride; int mask_base1 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow1 * mask_q_stride;
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) { for (int n8 = 0; n8 < Traits::NC8; n8++) {
int cc = kv0 + n8 * 8 + 2 * tid4; int cc = kv0 + n8 * 8 + 2 * tid4;
int c1 = cc + 1; int c1 = cc + 1;
bool b0 = (cc >= maxc0) || (has_mask && !mask[mask_base0 + cc]); bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base0 + c1]); bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base1 + cc]); bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base1 + c1]); bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
float s0 = b0 ? -FLT_MAX : Sacc[n8][0]; float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
float s1 = b1 ? -FLT_MAX : Sacc[n8][1]; float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
float s2 = b2 ? -FLT_MAX : Sacc[n8][2]; float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
@@ -220,29 +229,20 @@ __device__ inline void mma_softmax_tile(
rmax0 = fmaxf(rmax0, fmaxf(s0, s1)); rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
rmax1 = fmaxf(rmax1, fmaxf(s2, s3)); rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
} }
// Warp-reduce row maxima across the 4-lane thread group (xor 1, xor 2).
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1)); rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2)); rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1)); rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2)); rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
// nm = max(running max m, tile-local max rmax) — updated running maximum.
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1); float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
// corr rescales Oacc and l by exp(m_old - nm). When all-masked (m == nm ==
// -FLT_MAX), exp(0) = 1 — correct, no guard needed.
float corr0 = __expf(m0 - nm0); float corr0 = __expf(m0 - nm0);
float corr1 = __expf(m1 - nm1); float corr1 = __expf(m1 - nm1);
// pn guards only the all-masked-row edge: if nm == -FLT_MAX, exp(S - nm)
// gives 1 not 0 for masked entries. Two scalar masks replace 4*NC8
// per-element comparisons.
float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f; float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f;
float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f; float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f;
// P = exp(S - nm) for each element. Masked entries (Sacc = -FLT_MAX) give
// exp(-inf) ≈ 0 naturally; pn zero-fills the all-masked-row edge.
float rsum0 = 0.0f, rsum1 = 0.0f; float rsum0 = 0.0f, rsum1 = 0.0f;
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) { for (int n8 = 0; n8 < Traits::NC8; n8++) {
float p0 = pn0 * __expf(Sacc[n8][0] - nm0); float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
float p1 = pn0 * __expf(Sacc[n8][1] - nm0); float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
float p2 = pn1 * __expf(Sacc[n8][2] - nm1); float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
@@ -261,22 +261,25 @@ __device__ inline void mma_softmax_tile(
m0 = nm0; m1 = nm1; m0 = nm0; m1 = nm1;
#pragma unroll #pragma unroll
for (int j = 0; j < DN8; j++) { for (int j = 0; j < Traits::DN8; j++) {
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0; Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1; Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
} }
} }
// ---------------------------------------------------------------------------
// O += P @ V (Sacc must contain P = attention weights after softmax). // O += P @ V (Sacc must contain P = attention weights after softmax).
template <int DN8, int KT2> // Traits provides DN8, KT2, LD, and SWIZ_MASK.
// ---------------------------------------------------------------------------
template <typename Traits>
__device__ inline void mma_pv_accumulate( __device__ inline void mma_pv_accumulate(
float Sacc[][4], float Sacc[][4],
const bf16* __restrict__ sV, const bf16* __restrict__ sV,
int LD, int SWIZ_MASK, int lane, int lane,
float Oacc[DN8][4]) float Oacc[Traits::DN8][4])
{ {
#pragma unroll #pragma unroll
for (int kt2 = 0; kt2 < KT2; kt2++) { for (int kt2 = 0; kt2 < Traits::KT2; kt2++) {
unsigned Pa[4]; unsigned Pa[4];
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]); Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]); Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
@@ -284,9 +287,10 @@ __device__ inline void mma_pv_accumulate(
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]); Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
int vrow_l = kt2 * 16 + (lane & 15); int vrow_l = kt2 * 16 + (lane & 15);
#pragma unroll #pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) { for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
unsigned b[2]; unsigned b[2];
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]); ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]); mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
} }
} }
+23 -59
View File
@@ -1,82 +1,46 @@
#include "attn_paged_decode_split_kv.cuh" #include "attn_dispatchers.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_paged_decode_split_kv_mma.cuh"
#endif
#include "attn_entry_utils.cuh" #include "attn_entry_utils.cuh"
static void launch_paged_scalar_decode(PagedAttentionParams<bf16>& p) {
int group_size = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
alloc_split_partials(p);
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits);
dim3 block = dim3(32, group_size);
paged_attn_decode_split_kv_kernel<<<grid, block, smem>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#ifndef ASTRAI_NO_MMA
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
alloc_split_partials(p);
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
template <int HEAD_DIM>
static void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
int G = p.q_head / p.kv_head;
if (G >= 1 && G <= 16 && p.page_size >= 32) {
launch_paged_mma_decode<HEAD_DIM, 32>(p);
return;
}
#endif
launch_paged_scalar_decode(p);
}
torch::Tensor attn_paged_decode( torch::Tensor attn_paged_decode(
torch::Tensor q, torch::Tensor q,
torch::Tensor page_table,
torch::Tensor k_cache, torch::Tensor k_cache,
torch::Tensor v_cache, torch::Tensor v_cache,
int64_t page_size, torch::Tensor req_to_token,
int64_t kv_len, torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
int64_t max_seq_len,
c10::optional<torch::Tensor> mask, c10::optional<torch::Tensor> mask,
int64_t causal_offset, int64_t causal_offset,
double scale, double scale
int64_t layout
) { ) {
PagedAttentionParams<bf16> p; const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
attn_pack_paged_params(q, page_table, k_cache, v_cache, auto stream = at::cuda::getCurrentCUDAStream();
page_size, kv_len, mask, causal_offset, scale, layout, p);
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); AttentionParams<bf16> p;
auto O_view = (layout == 1) ? O.transpose(1, 2) : O; attn_pack_paged_decode_params(q, k_cache, v_cache,
p.o = (bf16*)O_view.data_ptr(); req_to_token, req_pool_indices, kv_indptr,
max_seq_len, mask, causal_offset, scale, p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p); auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
p.o = (bf16*)O.data_ptr();
alloc_split_partials(p);
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O; return O;
} }
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_paged_decode", &attn_paged_decode, m.def("attn_paged_decode", &attn_paged_decode,
py::arg("q"), py::arg("q"),
py::arg("page_table"),
py::arg("k_cache"), py::arg("k_cache"),
py::arg("v_cache"), py::arg("v_cache"),
py::arg("page_size"), py::arg("req_to_token"),
py::arg("kv_len"), py::arg("req_pool_indices"),
py::arg("kv_indptr"),
py::arg("max_seq_len"),
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,
py::arg("layout") = 0, "SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
"Paged GQA decode — split-KV with direct page-table access.");
} }
-147
View File
@@ -1,147 +0,0 @@
#pragma once
#include <cuda_bf16.h>
#include <float.h>
#include "attn_common.h"
using bf16 = __nv_bfloat16;
constexpr int PDC_CHUNK = 64;
__device__ inline float paged_warp_reduce_sum(float val) {
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
// Split-KV scalar decode: one warp per query head, grid.z partitions KV.
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
int batch = blockIdx.x / p.kv_head;
int kv_head = blockIdx.x % p.kv_head;
int split = blockIdx.z;
int group_size = blockDim.y;
int q_head = kv_head * group_size + threadIdx.y;
int lane = threadIdx.x;
int hd_per_thread = p.head_dim / 32;
// Q: stride-based [batch, q_head, q_len=1, head_dim]
float q_reg[8];
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
+ lane * hd_per_thread * p.q_stride_d;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
extern __shared__ __align__(16) bf16 k_smem[];
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
int ch_begin = split * chunks_per_split;
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
const int mask_base = batch * p.mask_b_stride;
for (int ci = ch_begin; ci < ch_end; ci++) {
int chunk_start = ci * PDC_CHUNK;
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
int s = i / p.head_dim;
int d_dim = i % p.head_dim;
int pos = chunk_start + s;
int logical_page = pos / p.page_size;
int page_offset = pos % p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page];
if (phys_page >= 0) {
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
+ (int64_t)page_offset * p.kv_head * p.head_dim
+ (int64_t)kv_head * p.head_dim
+ d_dim;
k_smem[i] = p.k_cache[off];
} else {
k_smem[i] = __float2bfloat16(0.0f);
}
}
__syncthreads();
for (int s = 0; s < this_chunk; s++) {
float partial = 0.0f;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
partial = paged_warp_reduce_sum(partial) * p.scale;
int kv_idx = chunk_start + s;
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
partial = -FLT_MAX;
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
partial = -FLT_MAX;
float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m);
float beta = expf(partial - new_m);
d = d * alpha + beta;
int pos = chunk_start + s;
int logical_page = pos / p.page_size;
int page_offset = pos % p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page];
if (phys_page >= 0) {
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
+ (int64_t)page_offset * p.kv_head * p.head_dim
+ (int64_t)kv_head * p.head_dim;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
} else {
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta;
}
m = new_m;
}
__syncthreads();
}
size_t bh = (size_t)batch * p.q_head + q_head;
size_t slot = bh * p.num_splits + split;
int d0 = lane * hd_per_thread;
#pragma unroll
for (int i = 0; i < hd_per_thread; i++)
p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i];
if (lane == 0) {
p.ml_part[slot * 2] = m;
p.ml_part[slot * 2 + 1] = d;
}
}
__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
int bh = blockIdx.x;
int d = threadIdx.x;
if (d >= p.head_dim) return;
int batch = bh / p.q_head;
int q_head = bh % p.q_head;
size_t split_base = (size_t)bh * p.num_splits;
const float* mlp = p.ml_part + split_base * 2;
const float* op = p.o_part + split_base * p.head_dim;
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
for (int s = 0; s < p.num_splits; s++) {
float mi = mlp[s * 2];
if (mi <= -FLT_MAX) continue;
float li = mlp[s * 2 + 1];
float nm = fmaxf(m, mi);
float corr = __expf(m - nm);
float e = __expf(mi - nm);
acc = acc * corr + op[s * p.head_dim + d] * e;
l = l * corr + li * e;
m = nm;
}
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
p.o[o_off] = __float2bfloat16(acc * inv);
}
@@ -1,170 +0,0 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
#include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16;
// Paged split-KV tensor-core decode via GQA head-packing.
// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V
// directly from the page pool through a page table, eliminating the gather
// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so
// the page-table lookup happens once per tile for cp.async.
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
constexpr int KD = HEAD_DIM / 16;
constexpr int NC8 = BC / 8;
constexpr int KT2 = BC / 16;
constexpr int DN8 = HEAD_DIM / 8;
constexpr int LD = HEAD_DIM;
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
constexpr int VEC = 8;
constexpr int TOTAL = BC * HEAD_DIM;
const int lane = threadIdx.x;
const int gid = lane >> 2;
const int tid4 = lane & 3;
const int kv_head = blockIdx.x;
const int batch = blockIdx.y;
const int split = blockIdx.z;
const int G = p.q_head / p.kv_head;
const int q_head0 = kv_head * G;
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
// ---- Load Q directly from global into mma A-operand registers ----
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
const int qra = gid;
const int qrb = gid + 8;
const bool va = qra < G, vb = qrb < G;
unsigned Qa[KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4];
#pragma unroll
for (int j = 0; j < DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
const int tiles_total = (p.kv_len + BC - 1) / BC;
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
const int ti_begin = split * tiles_per_split;
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
const int has_mask = p.use_mask && p.mask;
// Paged strides (constant for the block)
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM;
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM;
const int64_t head_off = (int64_t)kv_head * HEAD_DIM;
// ---- Load tile lambda: predicated cp.async, paged addressing ----
auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC;
bf16* dK = sK + buf * BC * LD;
bf16* dV = sV + buf * BC * LD;
int logical_page = kv0 / p.page_size;
int phys_page = p.page_table[batch * p.max_pages + logical_page];
bool page_valid = (phys_page >= 0);
#pragma unroll
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
int r = i / HEAD_DIM, d = i % HEAD_DIM;
int kc = kv0 + r;
bool valid = (kc < p.kv_len) && page_valid;
int page_off = kc % p.page_size;
int64_t gmem_base = (int64_t)phys_page * page_stride
+ (int64_t)page_off * pos_stride
+ head_off;
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
}
cp_async_commit();
};
// ---- Prologue: issue first tile load ----
if (ti_begin < ti_end) {
load_tile(ti_begin, 0);
}
for (int ti = ti_begin; ti < ti_end; ti++) {
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
int buf = (ti - ti_begin) & BUF_MASK;
cp_async_wait_group<0>();
__syncwarp();
if constexpr (STAGES > 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
}
const bf16* bK = sK + buf * BC * LD;
const bf16* bV = sV + buf * BC * LD;
int kv0 = ti * BC;
float Sacc[NC8][4];
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
#pragma unroll
for (int n8 = 0; n8 < NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
0, 0,
p.mask_b_stride, 0,
batch,
p.mask, has_mask,
Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
__syncwarp();
if constexpr (STAGES == 1) {
if (ti + 1 < ti_end)
load_tile(ti + 1, 0);
}
}
// ---- write UN-normalised partials for this split ----
auto split_slot = [&](int h) -> size_t {
size_t bh = (size_t)batch * p.q_head + h;
return bh * p.num_splits + split;
};
#pragma unroll
for (int dn8 = 0; dn8 < DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4;
int r0 = gid, r1 = gid + 8;
if (r0 < G) {
int h = q_head0 + r0;
float* op = p.o_part + split_slot(h) * HEAD_DIM;
op[d] = Oacc[dn8][0];
op[d + 1] = Oacc[dn8][1];
}
if (r1 < G) {
int h = q_head0 + r1;
float* op = p.o_part + split_slot(h) * HEAD_DIM;
op[d] = Oacc[dn8][2];
op[d + 1] = Oacc[dn8][3];
}
}
if (tid4 == 0) {
int r0 = gid, r1 = gid + 8;
if (r0 < G) {
int h = q_head0 + r0;
float* mp = p.ml_part + split_slot(h) * 2;
mp[0] = m0; mp[1] = l0;
}
if (r1 < G) {
int h = q_head0 + r1;
float* mp = p.ml_part + split_slot(h) * 2;
mp[0] = m1; mp[1] = l1;
}
}
}
+48
View File
@@ -0,0 +1,48 @@
#include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh"
torch::Tensor attn_paged_prefill(
torch::Tensor q,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor req_to_token,
torch::Tensor req_pool_indices,
torch::Tensor kv_indptr,
torch::Tensor qo_indptr,
c10::optional<torch::Tensor> mask,
int64_t max_q_len,
int64_t causal_offset,
double scale
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
AttentionParams<bf16> p;
attn_pack_paged_prefill_params(q, k_cache, v_cache,
req_to_token, req_pool_indices,
kv_indptr, qo_indptr, mask,
max_q_len, causal_offset, scale, p);
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
p.o = (bf16*)O.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("attn_paged_prefill", &attn_paged_prefill,
py::arg("q"),
py::arg("k_cache"),
py::arg("v_cache"),
py::arg("req_to_token"),
py::arg("req_pool_indices"),
py::arg("kv_indptr"),
py::arg("qo_indptr"),
py::arg("mask") = py::none(),
py::arg("max_q_len"),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
"SGLang-style paged prefill: flat KV pool + ragged batch.");
}
+8 -33
View File
@@ -1,35 +1,6 @@
#include "attn_prefill_split_q.cuh" #include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh" #include "attn_entry_utils.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_prefill_split_q_mma.cuh"
#endif
template <int HEAD_DIM>
static void dispatch_prefill(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
constexpr int WARPS = 4, BR = 16;
// KV tile: bigger tiles amortize the per-tile cp.async wait + barrier +
// loop overhead over more tensor-core work (this kernel is latency-bound,
// not compute/bandwidth-bound), so BC=32 wins ~6-8% over BC=16 for
// D<=128. D=256 stays at 16: BC=32 double-buffered would need 64KB smem,
// over the 48KB static cap.
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
dim3 block(WARPS * 32, 1, 1);
// Static shared memory — double-buffered K/V only (no sQ: Q goes direct
// to registers). 2*BC*LD bf16 each for sK and sV → 4*BC*HEAD_DIM*2 bytes.
// Occupancy is smem-capped: D=64→3 blocks/SM (16KB), D=128→1 (32KB),
// D=256→1 (32KB, BC=16).
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
#else
constexpr int G = 8, ROWS = 32, P_BC = 32;
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
dim3 block(G, ROWS, 1);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
#endif
}
torch::Tensor attn_prefill( torch::Tensor attn_prefill(
torch::Tensor q, torch::Tensor q,
torch::Tensor k, torch::Tensor k,
@@ -39,15 +10,19 @@ torch::Tensor attn_prefill(
double scale, double scale,
int64_t layout int64_t layout
) { ) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
AttentionParams<bf16> p; AttentionParams<bf16> p;
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p); attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16"); TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
auto O_view = (layout == 1) ? O.transpose(1, 2) : O; auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr(); p.o = (bf16*)O_view.data_ptr();
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p); DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O; return O;
} }
@@ -59,6 +34,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
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,
py::arg("layout") = 0, py::arg("layout") = (int64_t)BHLD,
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)"); "GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
} }
+37 -35
View File
@@ -2,16 +2,15 @@
#include <cfloat> #include <cfloat>
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_kv_source.cuh"
using bf16 = __nv_bfloat16; using bf16 = __nv_bfloat16;
// v9: group-split register blocking. G threads cooperate on one query row, // v9: group-split register blocking. G threads cooperate on one query row,
// each owning HEAD_DIM/G dims of qreg[]/acc[]. Small per-thread footprint keeps // each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
// occupancy high; the S dot product is reduced across the G-lane group with a // compile-time bools — the compiler eliminates dead branches.
// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce. // Unified across contiguous and paged (SGLang flat-pool) K/V via KV.
// Online (per-kv) softmax — cheap because acc[] is only HEAD_DIM/G long. // Templated on <HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>.
// Templated on <HEAD_DIM, G, ROWS, P_BC>. Block = (G, ROWS). G power-of-two,
// G*ROWS a multiple of 32 with groups warp-aligned.
template <int G> template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) { __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
@@ -21,8 +20,7 @@ __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
return v; return v;
} }
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to // load 8 contiguous bf16 from (16-byte aligned) smem as one float4
// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads.
__device__ __forceinline__ void ld8(const bf16* p, float* o) { __device__ __forceinline__ void ld8(const bf16* p, float* o) {
float4 raw = *reinterpret_cast<const float4*>(p); float4 raw = *reinterpret_cast<const float4*>(p);
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw); const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
@@ -34,7 +32,7 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) {
} }
} }
template <int HEAD_DIM, int G, int ROWS, int P_BC> template <int HEAD_DIM, 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;
@@ -45,19 +43,24 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
int row = threadIdx.y; // 0..ROWS-1 int row = threadIdx.y; // 0..ROWS-1
int q_row = q_tile * ROWS + row; int q_row = q_tile * ROWS + row;
int kv_head = q_head / (p.q_head / p.kv_head); // Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
const int seq_len = KV::kv_len(p, batch);
const int q_len = KV::q_len(p, batch);
const int causal_off = KV::causal_offset(p, batch);
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);
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM]; __shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
__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);
float qreg[DPT]; float qreg[DPT];
if (q_row < p.q_len) { if (q_row < q_len) {
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h int q_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
#pragma unroll #pragma unroll
for (int i = 0; i < DPT; i++) for (int i = 0; i < DPT; i++)
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * p.scale; qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
} }
float m = -FLT_MAX, l = 0.0f; float m = -FLT_MAX, l = 0.0f;
@@ -66,42 +69,41 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
for (int i = 0; i < DPT; i++) for (int i = 0; i < DPT; i++)
acc[i] = 0.0f; acc[i] = 0.0f;
// KV: stride-based base int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; int tiles = (seq_len + P_BC - 1) / P_BC;
int mask_batch_base = batch * p.mask_b_stride;
int tiles = (p.kv_len + P_BC - 1) / P_BC;
int tt = G * ROWS; int tt = G * ROWS;
int lid = row * G + gpos; int lid = row * G + gpos;
// per-group shuffle mask: only the G lanes of this row's group participate,
// so causal masking (differing loop bounds across rows in a warp) is safe.
int lane_in_warp = lid & 31; int lane_in_warp = lid & 31;
unsigned gmask = (G == 32) ? 0xFFFFFFFFu unsigned gmask = (G == 32) ? 0xFFFFFFFFu
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1))); : (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
for (int ti = 0; ti < tiles; ti++) { for (int ti = 0; ti < tiles; ti++) {
int kv0 = ti * P_BC; int kv0 = ti * P_BC;
int tlen = min(P_BC, p.kv_len - kv0); int tlen = min(P_BC, seq_len - kv0);
// Load K/V into shared memory from strided global // Load K/V into shared memory (addressing via KV policy; paged
// guards empty slots with zero-fill).
for (int i = lid; i < tlen * HEAD_DIM; i += tt) { for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
int s = i / HEAD_DIM; int s = i / HEAD_DIM;
int d_dim = i % HEAD_DIM; int d_dim = i % HEAD_DIM;
int kv_idx = kv0 + s; int kc = kv0 + s;
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d; KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true);
sK[i] = p.k[g_off]; sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
sV[i] = p.v[g_off]; sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
} }
__syncthreads(); __syncthreads();
int lim = tlen; int lim = tlen;
if (p.causal_offset >= 0 && q_row < p.q_len) { if constexpr (IsCausal) {
int ep = q_row + p.causal_offset + 1; if (q_row < q_len) {
int ep = causal_off + q_row + 1;
if (kv0 >= ep) if (kv0 >= ep)
lim = 0; lim = 0;
else if (kv0 + tlen > ep) else if (kv0 + tlen > ep)
lim = ep - kv0; lim = ep - kv0;
} }
}
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride; int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
for (int s = 0; s < lim; s++) { for (int s = 0; s < lim; s++) {
@@ -115,11 +117,13 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
for (int j = 0; j < 8; j++) for (int j = 0; j < 8; j++)
part = fmaf(qreg[i + j], k8[j], part); part = fmaf(qreg[i + j], k8[j], part);
} }
float dot = group_reduce_sum<G>(part, gmask); float dot = group_reduce_sum<G>(part, gmask) * p.scale;
int kv_idx = kv0 + s; int kv_idx = kv0 + s;
if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx]) if constexpr (HasMask) {
if (!p.mask[mask_row_base + kv_idx])
dot = -FLT_MAX; dot = -FLT_MAX;
}
float nm = fmaxf(m, dot); float nm = fmaxf(m, dot);
float al = __expf(m - nm); float al = __expf(m - nm);
@@ -140,11 +144,9 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
__syncthreads(); __syncthreads();
} }
if (q_row < p.q_len) { if (q_row < q_len) {
// O: stride-based write int o_off = q_base + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f;
#pragma unroll #pragma unroll
for (int i = 0; i < DPT; i++) for (int i = 0; i < DPT; i++)
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl); p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
+81 -123
View File
@@ -2,126 +2,89 @@
#include <cfloat> #include <cfloat>
#include <cuda_bf16.h> #include <cuda_bf16.h>
#include "attn_common.h" #include "attn_common.h"
#include "attn_kv_source.cuh"
#include "attn_mma_utils.cuh" #include "attn_mma_utils.cuh"
using bf16 = __nv_bfloat16; // Tensor-core prefill flash attention (raw mma.sync PTX), unified across
// contiguous and paged (SGLang flat-pool) K/V via the KV template parameter.
// Tensor-core prefill flash attention (raw mma.sync PTX).
// 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). Q fragments are loaded once // cores via mma.sync.m16n8k16 (f32 accumulate).
// straight from global into the mma A-operand layout (no smem staging) and
// kept resident in registers across the tile loop. S, O, and the online-softmax
// stats (m, l) also live in registers.
// Shared memory is statically sized via template parameters — no dynamic
// allocation. The mma fragment layout is used directly: the S accumulator
// (f32) maps element-for-element onto the P matrix_a (bf16) operand, so
// softmax needs no shuffle repack; row reductions fold across the 4-lane
// thread group. Templated on <HEAD_DIM, WARPS, BC> with BC a multiple of 16.
// //
// Software pipeline: K/V are double-buffered and loaded via cp.async one tile // KV = ContigKV (dense [batch, kv_head, kv_len, head_dim]) or PagedKV
// ahead, so the next tile streams from global memory while the current tile's // (flat pool + req_to_token, ragged batches via qo_indptr/kv_indptr).
// tensor-core math runs — hiding load latency (long_scoreboard). A single // IsCausal and HasMask are compile-time bools — the compiler eliminates all
// __syncthreads per tile both publishes the freshly loaded tile cross-warp and // dead branches in the inner compute loop (FA2-style).
// (because it runs before the next prefetch) guards the buffer being refilled,
// so no second barrier is needed. Predicated cp.async (cp_async_16_pred)
// zero-fills rows past kv_len, unifying full and partial tiles on one path.
// BC=32 (D<=128) amortizes the per-tile wait+barrier+loop overhead over more
// tensor-core work — this kernel is latency-bound (low occupancy from high
// register pressure), so fewer, larger tiles beat many tiny ones.
// //
// Optimizations: load Q fragments directly from global in mma A-operand layout // Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
// (no sQ staging, no prologue barriers); post-multiply scale in float after template <typename Traits, typename KV, bool IsCausal, bool HasMask>
// S=Q@K^T to avoid bf16 precision loss; packed bf16x2 output stores;
// causal tile skipping (block-level prefetch bound + warp-level compute skip);
// XOR swizzle (swiz_col) → eliminates ldmatrix bank conflicts without LD
// padding (LD=HEAD_DIM).
template <int HEAD_DIM, int WARPS, int BC>
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) { __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
constexpr int BR = 16;
constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles
constexpr int NC8 = BC / 8; // S n-tiles (N=8 each)
constexpr int KT2 = BC / 16; // P k-tiles (K=16 each)
constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each)
constexpr int LD = HEAD_DIM; // XOR swizzle (swiz_col) handles bank conflicts
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); // chunk bits, stay within LD
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 → rows gid, gid+8 const int gid = lane >> 2; // 0..7
const int tid4 = lane & 3; // 0..3 const int tid4 = lane & 3; // 0..3
const int nthreads = WARPS * 32;
const int q_head = blockIdx.y; const int q_head = blockIdx.y;
const int batch = blockIdx.z; const int batch = blockIdx.z;
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 int qrow0 = (blockIdx.x * WARPS + warp) * BR; const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
// ---- Static shared memory: double-buffered K/V ---- // Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
// K/V are double-buffered (STAGES=2): the next tile's cp.async load runs const int seq_len = KV::kv_len(p, batch);
// while the current tile's tensor-core math executes, hiding global-load const int q_len = KV::q_len(p, batch);
// latency (FA2-style software pipeline). No dynamic smem / carveout opt-in. const int causal_off = KV::causal_offset(p, batch);
constexpr int STAGES = 2; const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head);
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
__shared__ __align__(16) bf16 sV[STAGES * BC * LD]; // Static shared memory: double-buffered K/V (no sQ — Q goes direct
// to registers in mma A-operand layout).
__shared__ __align__(16) bf16 sK[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.
// stride_row = p.q_stride_l for prefill (multi-q rows across q_len). const int q_base = KV::q_base(p, batch, q_head);
// See attn_mma_utils.cuh for the shared template.
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
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 < p.q_len, vb = qrb < p.q_len; const bool va = qra < q_len, vb = qrb < q_len;
unsigned Qa[KD][4]; unsigned Qa[Traits::KD][4];
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_l, p.q_stride_d, load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
qra, qrb, va, vb, tid4, Qa); qra, qrb, va, vb, tid4, Qa);
float Oacc[DN8][4]; float Oacc[Traits::DN8][4];
#pragma unroll #pragma unroll
for (int j = 0; j < DN8; j++) for (int j = 0; j < Traits::DN8; j++)
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f; Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f; float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
// KV: stride-based base const int tiles = (seq_len + Traits::BC - 1) / Traits::BC;
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h; const int qr0 = qrow0 + gid;
const int tiles = (p.kv_len + BC - 1) / BC; const int qr1 = qrow0 + gid + 8;
const int qr0 = qrow0 + gid; // row for c0/c1
const int qr1 = qrow0 + gid + 8; // row for c2/c3
// Causal tile-skip bounds (no-op when causal_offset < 0) // Causal tile-skip bounds (dead code when IsCausal == false)
const int use_skip = (p.causal_offset >= 0) ? 1 : 0; const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
const int block_max_kv = const int block_max_kv =
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset; blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
const int has_mask = p.use_mask && p.mask; + causal_off;
// Last active tile: block-level causal bound (all warps in the block share
// the K/V load, so the prefetch range is the block max, not per-warp).
int t_end = tiles - 1; int t_end = tiles - 1;
if (use_skip) { if constexpr (IsCausal) {
int bt = block_max_kv / BC; int bt = block_max_kv / Traits::BC;
if (bt < t_end) t_end = bt; if (bt < t_end) t_end = bt;
} }
constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes) // ---- Load tile lambda: predicated cp.async (addressing via KV policy) ----
constexpr int TOTAL = BC * HEAD_DIM;
// ---- Load tile lambda: predicated cp.async ----
// Issue cp.async loads for tile `ti` into shared buffer `buf`. Predicated
// loads zero-fill rows past kv_len, so partial tiles need no scalar path.
auto load_tile = [&](int ti, int buf) { auto load_tile = [&](int ti, int buf) {
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
bf16* dK = sK + buf * BC * LD; bf16* dK = sK + buf * Traits::BC * Traits::LD;
bf16* dV = sV + buf * BC * LD; bf16* dV = sV + buf * Traits::BC * Traits::LD;
#pragma unroll #pragma unroll
for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) { for (int i = threadIdx.x * Traits::VEC; i < Traits::TOTAL;
int r = i / HEAD_DIM, d = i % HEAD_DIM; i += Traits::NUM_THREADS * Traits::VEC) {
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
int kc = kv0 + r; int kc = kv0 + r;
bool valid = kc < p.kv_len; bool valid = kc < seq_len;
int off = r * LD + swiz_col(d, r, SWIZ_MASK); KVAddr a = KV::kv_addr(p, kctx, kc, d, valid);
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d; int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], &p.k[g_off], valid); cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], &p.v[g_off], valid); cp_async_16_pred(&dV[off], a.v, a.valid);
} }
cp_async_commit(); cp_async_commit();
}; };
@@ -132,65 +95,60 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
for (int ti = 0; ti <= t_end; ti++) { for (int ti = 0; ti <= t_end; ti++) {
int buf = ti & 1; int buf = ti & 1;
// Wait for the current tile's async copies, then a single barrier: it // Wait for current tile, then publish cross-warp + guard buffer reuse.
// both publishes this tile's data cross-warp AND guarantees the prior
// compute on the buffer we are about to refill has finished. Issuing
// the next tile's load *after* this barrier lets one barrier cover both
// hazards (vs two), while the load still overlaps this tile's math.
cp_async_wait_group<0>(); 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);
const bf16* bK = sK + buf * BC * LD; const bf16* bK = sK + buf * Traits::BC * Traits::LD;
const bf16* bV = sV + buf * BC * LD; const bf16* bV = sV + buf * Traits::BC * Traits::LD;
int kv0 = ti * BC; int kv0 = ti * Traits::BC;
// Warp-level causal skip // Warp-level causal skip (dead branch eliminated when IsCausal == false)
if (!use_skip || kv0 <= max_kv) { if (!IsCausal || kv0 <= max_kv) {
// S = Q @ K^T + scale + online softmax + O += P @ V float Sacc[Traits::NC8][4];
float Sacc[NC8][4]; mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
// post-multiply scale in float (no bf16 precision loss from pre-scaling Q) // Post-multiply scale in float (no bf16 precision loss)
#pragma unroll #pragma unroll
for (int n8 = 0; n8 < NC8; n8++) for (int n8 = 0; n8 < Traits::NC8; n8++)
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale, Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale; Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
int maxc0 = (p.causal_offset >= 0) ? min(p.kv_len, qr0 + p.causal_offset + 1) int maxc0 = IsCausal ? min(seq_len, causal_off + qr0 + 1)
: p.kv_len; : seq_len;
int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1) int maxc1 = IsCausal ? min(seq_len, causal_off + qr1 + 1)
: p.kv_len; : seq_len;
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1, mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
qr0, qr1, qr0, qr1,
p.mask_b_stride, p.mask_q_stride, p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
batch, batch, q_head,
p.mask, has_mask, p.mask,
Sacc, Oacc, m0, m1, l0, l1, lane); Sacc, Oacc, m0, m1, l0, l1, lane);
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc); mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
} // if active (warp-level causal skip) }
} }
// ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair, // ---- write output: packed bf16x2 stores ----
// halves store count and removes the uncoalesced scalar-store penalty)
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;
// O: stride-based write const int o_base = KV::q_base(p, batch, q_head);
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h; #pragma unroll
#pragma unroll for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
for (int dn8 = 0; dn8 < DN8; dn8++) {
int d = dn8 * 8 + 2 * tid4; int d = dn8 * 8 + 2 * tid4;
if (qr0 < p.q_len) { if (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*>(&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v; *reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
} }
if (qr1 < p.q_len) { if (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*>(&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v; *reinterpret_cast<__nv_bfloat162*>(
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
} }
} }
} }
+13
View File
@@ -0,0 +1,13 @@
#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;
}
+98
View File
@@ -0,0 +1,98 @@
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_bf16.h>
__global__ void rotary_emb_kernel(
const __nv_bfloat16* __restrict__ x,
const float* __restrict__ freqs_cis,
__nv_bfloat16* __restrict__ out,
int batch,
int seq_len,
int n_heads,
int head_dim
) {
const int half_dim = head_dim >> 1;
const int total = batch * seq_len * n_heads * half_dim;
for (int idx = blockIdx.x * blockDim.x + threadIdx.x;
idx < total;
idx += gridDim.x * blockDim.x) {
int pair = idx % half_dim;
int tmp = idx / half_dim;
int head = tmp % n_heads;
tmp /= n_heads;
int seq = tmp % seq_len;
int b = tmp / seq_len;
int x_offset = ((b * seq_len + seq) * n_heads + head) * head_dim + (pair << 1);
int cs_offset = ((b * seq_len + seq) * half_dim + pair) * 2;
__nv_bfloat162 x_pair = *reinterpret_cast<const __nv_bfloat162*>(x + x_offset);
float x_even = __bfloat162float(__low2bfloat16(x_pair));
float x_odd = __bfloat162float(__high2bfloat16(x_pair));
float c = freqs_cis[cs_offset];
float s = freqs_cis[cs_offset + 1];
float out_even = x_even * c - x_odd * s;
float out_odd = x_even * s + x_odd * c;
__nv_bfloat162 out_pair = __floats2bfloat162_rn(out_even, out_odd);
*reinterpret_cast<__nv_bfloat162*>(out + x_offset) = out_pair;
}
}
torch::Tensor rotary_emb(
torch::Tensor x,
torch::Tensor freqs_cis
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(x));
auto stream = at::cuda::getCurrentCUDAStream();
TORCH_CHECK(x.is_cuda(), "x 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.dim() == 4, "x must be 4D [batch, seq_len, n_heads, head_dim]");
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.is_contiguous(), "freqs_cis must be contiguous");
TORCH_CHECK(freqs_cis.scalar_type() == torch::kFloat32, "freqs_cis must be f32");
int batch = x.size(0);
int seq_len = x.size(1);
int n_heads = x.size(2);
int head_dim = x.size(3);
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.size(1) == seq_len, "freqs_cis seq_len 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]");
auto out = torch::empty_like(x);
int half_dim = head_dim / 2;
int total = batch * seq_len * n_heads * half_dim;
int block = 256;
int grid = std::min((total + block - 1) / block, 1024);
rotary_emb_kernel<<<grid, block, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
freqs_cis.data_ptr<float>(),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
batch, seq_len, n_heads, head_dim
);
C10_CUDA_CHECK(cudaGetLastError());
return out;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("rotary_emb", &rotary_emb,
py::arg("x"),
py::arg("freqs_cis"),
"Fused rotary embedding (bf16 x, f32 freqs_cis [b,s,d/2,2], bf16 out)"
);
}
-201
View File
@@ -1,201 +0,0 @@
/*
Pure-C test:
nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
csrc/tests/attn_decode_test.cu -o test && ./test
*/
#include "test_utils.cuh"
#include "../kernels/attn_decode_split_kv.cuh"
#ifndef ASTRAI_NO_MMA
#include "../kernels/attn_decode_split_kv_mma.cuh"
#endif
// Split-K scratch (torch-free): the production launcher allocates these from
// torch; here we pass pre-allocated device buffers so the bench loop doesn't
// pay a cudaMalloc per iteration. Size for the maximum split count (32).
struct DecodeScratch {
float* o_part = nullptr;
float* ml_part = nullptr;
};
// Launch the production decode path (tensor-core head-packing MMA on sm_80+,
// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu.
#ifndef ASTRAI_NO_MMA
static bool decode_use_mma(const AttentionParams<bf16>& p) {
int G = p.q_head / p.kv_head;
return !p.use_mask && G > 1 && G <= 16;
}
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
int tiles_total = (p.kv_len + BC - 1) / BC;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
p.o_part = sc.o_part;
p.ml_part = sc.ml_part;
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
#endif
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
int gs = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
p.o_part = sc.o_part;
p.ml_part = sc.ml_part;
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM>
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
#ifndef ASTRAI_NO_MMA
if (decode_use_mma(p)) { launch_mma_decode<HEAD_DIM, 32>(p, sc); return; }
#endif
launch_scalar_decode(p, sc);
}
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
dispatch_by_head_dim(p.head_dim, [&]<int D>() { dispatch_decode_t<D>(p, sc); });
}
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
static void bench() {
const int cfgs[][5] = {
{1, 32, 4, 512, 128}, // B, Hq, Hk, kv_len, D
{1, 32, 4, 1024, 128},
{1, 32, 4, 2048, 128},
{1, 32, 4, 4096, 128},
{16, 32, 4, 2048, 128},
{32, 32, 4, 1024, 128},
};
const int WARMUP = 10, ITERS = 100;
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
print_bench_header();
for (int ci = 0; ci < 6; ci++) {
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
int sl = cfgs[ci][3], D = cfgs[ci][4];
size_t nQ = (size_t)B * Hq * D;
size_t nKV = (size_t)B * Hk * sl * D;
bf16 *dQ, *dK, *dV, *dO;
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
delete[] tmp;
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.use_mask = 0; p.causal_offset = -1;
p.scale = 1.0f / sqrtf((float)D);
set_default_strides(p);
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
auto launch = [&]() { dispatch_decode(p, sc); };
double flops = 4.0 * B * Hq * (double)sl * D;
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
B, Hq, Hk, 1, sl, D, 0);
print_bench_row(cfg, r);
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
cudaFree(sc.o_part); cudaFree(sc.ml_part);
}
}
int main() {
const int configs[][5] = {
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
{1, 32, 4, 512, 128},
{1, 32, 4, 1024, 128},
};
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
for (int ci = 0; ci < n_cfgs; ci++) {
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
bool* hMask=new bool[B*sl];
for (int i=0;i<B*sl;i++) hMask[i]=true;
bf16 *dQ,*dK,*dV,*dO,*tmp;
bool* dMask;
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
cudaMalloc(&dMask,B*sl);
tmp=new bf16[max(nQ,nKV)];
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
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.use_mask=0; p.causal_offset=-1;
p.scale=1.0f/sqrtf((float)D);
set_default_strides(p);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
// Split-K scratch (max 32 splits), sized for the production MMA path.
DecodeScratch sc;
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
double t0=now_ms();
dispatch_decode(p, sc);
cudaDeviceSynchronize();
double kms=now_ms()-t0;
cudaError_t err=cudaGetLastError();
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
bf16* hOut=new bf16[nQ];
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, -1);
float max_err=0;
for (size_t i=0;i<nQ;i++){
float d=fabsf(bf2f(hOut[i])-ref[i]);
if(d>max_err) max_err=d;
}
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
cudaFree(sc.o_part);cudaFree(sc.ml_part);
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
}
printf("All tests passed!\n");
bench();
return 0;
}
-332
View File
@@ -1,332 +0,0 @@
// Compile:
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
// -o /tmp/test_paged && /tmp/test_paged
#include <cstring>
#include "test_utils.cuh"
#include "../kernels/attn_paged_decode_split_kv.cuh"
#ifndef ASTRAI_NO_MMA
#include "../kernels/attn_paged_decode_split_kv_mma.cuh"
#endif
// Copy contiguous K/V from page pool (reference gather)
static void gather_kv_cpu(
const bf16* h_k_pool, const bf16* h_v_pool,
const int64_t* h_pt, int B, int Hkv, int kv_len,
int page_size, int head_dim,
bf16* h_k, bf16* h_v)
{
int max_pages = (kv_len + page_size - 1) / page_size;
size_t page_stride = (size_t)page_size * Hkv * head_dim;
for (int b = 0; b < B; b++) {
for (int pos = 0; pos < kv_len; pos++) {
int log_pg = pos / page_size;
int pg_off = pos % page_size;
int phys = (int)h_pt[b * max_pages + log_pg];
for (int h = 0; h < Hkv; h++) {
size_t src_base = (size_t)phys * page_stride
+ (size_t)pg_off * Hkv * head_dim
+ h * head_dim;
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + (size_t)pos * head_dim;
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
}
}
}
}
template <int HEAD_DIM>
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
#ifndef ASTRAI_NO_MMA
int G_check = p.q_head / p.kv_head;
bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
if (use_mma) {
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
int tiles_total = (p.kv_len + 32 - 1) / 32;
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
} else
#endif
{
int group_sz = p.q_head / p.kv_head;
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
paged_attn_decode_split_kv_kernel<<<
dim3(p.batch * p.kv_head, 1, p.num_splits),
dim3(32, group_sz), smem>>>(p);
}
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
}
template <int HEAD_DIM>
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) {
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM);
fflush(stdout);
int max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages;
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
size_t sz_o = sz_q;
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
int max_splits = 32;
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);
bf16 *d_q, *d_o_paged, *d_o_ref;
bf16 *d_k_pool, *d_v_pool;
int64_t* d_pt;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q);
cudaMalloc(&d_o_paged, sz_o);
cudaMalloc(&d_o_ref, sz_o);
cudaMalloc(&d_k_pool, sz_kv);
cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_pt, sz_pt);
cudaMalloc(&d_op, sz_op);
cudaMalloc(&d_ml, sz_ml);
srand(seed);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
h_q[i] = __float2bfloat16(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
for (int pg = 0; pg < n_phys_pages; pg++) {
for (int off = 0; off < page_size; off++) {
for (int h = 0; h < Hkv; h++) {
for (int d = 0; d < HEAD_DIM; d++) {
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM + h * HEAD_DIM + d;
h_k_pool[idx] = __float2bfloat16(v);
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
}
}
}
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_pt = (int64_t*)malloc(sz_pt);
int next_pg = 0;
for (int b = 0; b < B; b++)
for (int p = 0; p < max_pages; p++)
h_pt[b * max_pages + p] = next_pg++;
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_cont[i]);
h_v_f[i] = bf2f(h_v_cont[i]);
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, -1);
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
PagedAttentionParams<bf16, float> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
p.use_mask = 0; p.causal_offset = -1;
set_default_paged_strides(p);
p.num_splits = 1; p.scale = scale_val;
p.page_size = page_size; p.max_pages = max_pages;
p.page_table = d_pt;
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
p.o_part = d_op; p.ml_part = d_ml;
launch_paged_decode<HEAD_DIM>(p);
cudaDeviceSynchronize();
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
float max_err = 0.0f;
int bad_idx = -1;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
if (e > max_err) { max_err = e; bad_idx = i; }
}
bool pass = max_err < 0.02f;
if (pass) {
printf("PASS (max_abs_err=%.4e)\n", max_err);
} else {
int b = bad_idx / (Hq * HEAD_DIM);
int h = (bad_idx / HEAD_DIM) % Hq;
int d = bad_idx % HEAD_DIM;
printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
max_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
printf(" ref[0..7]:");
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
printf(" %.4f", h_o_ref[i]);
printf("\n got[0..7]:");
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
printf(" %.4f", h_o_paged[i]);
printf("\n");
}
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
free(h_k_cont); free(h_v_cont);
free(h_q_f); free(h_k_f); free(h_v_f);
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_o_ref);
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
cudaFree(d_op); cudaFree(d_ml);
return pass ? 0 : 1;
}
struct TestCase {
int head_dim;
int B, Hq, Hkv, kv_len, page_size, seed;
};
static const TestCase TESTS[] = {
{128, 1, 1, 1, 8, 128, 1},
{128, 1, 4, 4, 128, 128, 2},
{128, 2, 4, 4, 256, 128, 3},
{128, 1, 4, 1, 64, 64, 4},
{128, 1, 8, 2, 64, 128, 5},
{128, 2, 16, 4, 128, 128, 6},
{64, 1, 4, 2, 32, 128, 7},
{256, 1, 2, 1, 16, 128, 8},
{32, 1, 4, 2, 32, 64, 9},
{128, 3, 8, 2, 256, 128, 10},
{128, 2, 32, 8, 512, 128, 11},
#ifndef ASTRAI_NO_MMA
{128, 1, 16, 2, 256, 128, 12},
{128, 2, 32, 4, 512, 128, 13},
#endif
};
static int dispatch_test(const TestCase& tc) {
bool matched = false;
int r = 0;
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
matched = true;
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
});
return matched ? r : 1;
}
// Warmed-up, CUDA-event timed sweep over paged decode configs.
// Bytes = K + V read through page table (B*Hk*kv*D each), bf16.
template <int HEAD_DIM>
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
int max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages;
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
int max_splits = 32;
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);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t* d_pt;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_pt, sz_pt);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_pt = (int64_t*)malloc(sz_pt);
int next_pg = 0;
for (int b = 0; b < B; b++)
for (int p = 0; p < max_pages; p++)
h_pt[b * max_pages + p] = next_pg++;
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
free(h_pt);
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
PagedAttentionParams<bf16, float> pa;
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
pa.use_mask = 0; pa.causal_offset = -1;
set_default_paged_strides(pa);
pa.num_splits = 1; pa.scale = scale_val;
pa.page_size = page_size; pa.max_pages = max_pages;
pa.page_table = d_pt;
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
pa.o_part = d_op; pa.ml_part = d_ml;
const int WARMUP = 10, ITERS = 100;
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(pa); };
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
print_bench_row(cfg, r);
free(tmp);
cudaFree(d_q); cudaFree(d_o);
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
cudaFree(d_op); cudaFree(d_ml);
}
static void bench() {
printf("\n===== PAGED DECODE BENCH =====\n");
print_bench_header();
bench_config<128>(1, 32, 4, 512, 128);
bench_config<128>(1, 32, 4, 1024, 128);
bench_config<128>(1, 32, 4, 2048, 128);
bench_config<128>(1, 32, 4, 4096, 128);
bench_config<128>(16, 32, 4, 2048, 128);
bench_config<128>(32, 32, 4, 1024, 128);
}
int main() {
int n = sizeof(TESTS) / sizeof(TESTS[0]);
int fail = 0;
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
for (int i = 0; i < n; i++) {
fail += dispatch_test(TESTS[i]);
if (fail) break;
}
if (fail) {
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
return fail;
}
printf("\nAll %d tests passed!\n", n);
bench();
return 0;
}
+945
View File
@@ -0,0 +1,945 @@
// Compile:
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
// --extra-device-vectorization -Xcompiler -fopenmp \
// csrc/tests/attn_paged_test.cu \
// -o /tmp/test_paged && /tmp/test_paged
#include <cstring>
#include <vector>
#include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh"
// ---- CPU reference: paged decode with variable seq_lens ----
// Q: [B, Hq, D], K/V pool: [pool_size, Hkv, D]
// 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.
static void cpu_paged_decode_ref(
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* kv_indptr, const bool* mask, int mask_b_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len,
float* O)
{
float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hkv;
for (int b = 0; b < B; b++) {
int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int64_t req_idx = req_pool_indices[b];
#pragma omp parallel for schedule(dynamic)
for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep;
float mv = -INFINITY, sv = 0.0f;
float accum[256] = {0.0f};
for (int kj = 0; kj < seq_len; kj++) {
if (mask && !mask[b * mask_b_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f;
for (int d = 0; d < D; d++)
dot += Q[(b * Hq + h) * D + d] *
K_pool[slot * Hkv * D + kv_h * D + d];
dot *= scale;
float nm = fmaxf(mv, dot);
float a = expf(mv - nm);
float be = expf(dot - nm);
sv = sv * a + be;
for (int d = 0; d < D; d++)
accum[d] = accum[d] * a +
V_pool[slot * Hkv * D + kv_h * D + d] * be;
mv = nm;
}
float inv = 1.0f / sv;
for (int d = 0; d < D; d++)
O[(b * Hq + h) * D + d] = accum[d] * inv;
}
}
}
// ---- CPU reference: paged prefill with ragged batch ----
// Q: [total_q, Hq, D], K/V pool: [pool_size, Hkv, D]
// req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B]
// kv_indptr: [B+1], qo_indptr: [B+1].
// mask: [B, max_q_len, max_seq_len] bool (True=keep, q-local + kv-local
// positions) or NULL. Used only when causal==0 to apply an arbitrary
// attention mask on top of the (unused) causal logic.
static void cpu_paged_prefill_ref(
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* kv_indptr, const int* qo_indptr,
const bool* mask, int mask_q_stride, int mask_kv_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len, int causal,
float* O)
{
float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hkv;
for (int b = 0; b < B; b++) {
int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int q_len = qo_indptr[b + 1] - qo_indptr[b];
int causal_off = seq_len - q_len;
int64_t req_idx = req_pool_indices[b];
#pragma omp parallel for collapse(2) schedule(dynamic)
for (int h = 0; h < Hq; h++) {
for (int qi = 0; qi < q_len; qi++) {
int kv_h = h / n_rep;
float mv = -INFINITY, sv = 0.0f;
float accum[256] = {0.0f};
int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len;
for (int kj = 0; kj < lim; kj++) {
if (mask && !mask[b * mask_q_stride * mask_kv_stride
+ qi * mask_kv_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f;
for (int d = 0; d < D; d++)
dot += Q[(qo_indptr[b] + qi) * Hq * D + h * D + d] *
K_pool[slot * Hkv * D + kv_h * D + d];
dot *= scale;
float nm = fmaxf(mv, dot);
float a = expf(mv - nm);
float be = expf(dot - nm);
sv = sv * a + be;
for (int d = 0; d < D; d++)
accum[d] = accum[d] * a +
V_pool[slot * Hkv * D + kv_h * D + d] * be;
mv = nm;
}
float inv = 1.0f / sv;
for (int d = 0; d < D; d++)
O[(qo_indptr[b] + qi) * Hq * D + h * D + d] = accum[d] * inv;
}
}
}
}
// ---- paged validation table (kernel vs CPU ref, abs error only) ----
inline void print_paged_header() {
printf("%-58s | %11s | %6s\n",
"config", "max_err", "result");
printf("----------------------------------------------------------------"
"--------------------------------\n");
}
inline void print_paged_row(const char* cfg, float max_err, bool pass) {
printf("%-58s | %11.3e | %s\n",
cfg, max_err, pass ? "PASS" : "FAIL");
}
// ======================================================================
// DECODE TEST
// ======================================================================
template <int HEAD_DIM>
static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
int causal, int seed) {
// Variable seq_lens per request
srand(seed);
std::vector<int> seq_lens(B);
for (int b = 0; b < B; b++)
seq_lens[b] = 8 + rand() % (max_seq - 8);
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
int max_ctx = max_sl + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
char cfg[80];
snprintf(cfg, sizeof(cfg), "DECODE B=%d Hq=%d Hkv=%d D=%d max_sl=%d causal=%d",
B, Hq, Hkv, HEAD_DIM, max_sl, causal);
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_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
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_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_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)
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
// req_pool_indices: pick B random request rows
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
// kv_indptr: prefix sum of seq_lens
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_lens[b];
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
// CPU reference
float* h_q_f = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_paged_decode_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi,
nullptr, 0,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
// Kernel launch
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_q_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
print_paged_row(cfg, max_err, pass);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_q_f); free(h_k_f); free(h_v_f);
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_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_op); cudaFree(d_ml);
return pass ? 0 : 1;
}
// ======================================================================
// DECODE WITH MASK TEST (regression: 2D mask on mixed seq_lens)
// ======================================================================
template <int HEAD_DIM>
static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
int seed) {
srand(seed);
std::vector<int> seq_lens(B);
for (int b = 0; b < B; b++)
seq_lens[b] = 8 + rand() % (max_seq - 8);
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
int max_ctx = max_sl + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
char cfg[80];
snprintf(cfg, sizeof(cfg), "DECODE-MASK B=%d Hq=%d Hkv=%d D=%d max_sl=%d",
B, Hq, Hkv, HEAD_DIM, max_sl);
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_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
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_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi;
bool *d_mask;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi);
cudaMalloc(&d_mask, sz_mask);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_lens[b];
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
// Mask: keep first half of each request's kv range, drop the rest —
// exercises the HasMask path with per-request seq_len.
bool* h_mask = (bool*)malloc(sz_mask);
for (int b = 0; b < B; b++)
for (int k = 0; k < max_sl; k++)
h_mask[b * max_sl + k] = (k < seq_lens[b]) && (k % 2 == 0);
cudaMemcpy(d_mask, h_mask, sz_mask, cudaMemcpyHostToDevice);
float* h_q_f = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_paged_decode_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi,
h_mask, max_sl,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
p.causal_offset = -1; p.use_mask = 1;
p.mask = d_mask; p.mask_b_stride = max_sl;
p.mask_h_stride = 0; p.mask_q_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
print_paged_row(cfg, max_err, pass);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f);
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_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_mask);
cudaFree(d_op); cudaFree(d_ml);
return pass ? 0 : 1;
}
// ======================================================================
// PREFILL TEST
// ======================================================================
template <int HEAD_DIM>
static int run_prefill_test(int B, int Hq, int Hkv,
std::vector<int>& q_lens,
std::vector<int>& kv_lens,
int causal, int seed) {
int total_q = 0;
int max_sl = 0;
for (int b = 0; b < B; b++) {
total_q += q_lens[b];
max_sl = max(max_sl, kv_lens[b]);
}
int max_ctx = max_sl + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
char cfg[80];
snprintf(cfg, sizeof(cfg), "PREFILL B=%d Hq=%d Hkv=%d D=%d max_sl=%d causal=%d",
B, Hq, Hkv, HEAD_DIM, max_sl, causal);
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_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (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;
int64_t *d_rtt, *d_rpi;
int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
srand(seed);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + kv_lens[b];
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
int* h_qoi = (int*)malloc(sz_qoi);
h_qoi[0] = 0;
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_lens[b];
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
// CPU reference
float* h_q_f = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(total_q * Hq * HEAD_DIM, sizeof(float));
cpu_paged_prefill_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi, h_qoi,
nullptr, 0, 0,
B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref);
// Kernel launch
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
int max_ql = 0;
for (int b = 0; b < B; b++) max_ql = max(max_ql, q_lens[b]);
p.max_q_len = max_ql;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_q_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
print_paged_row(cfg, max_err, pass);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_qoi); free(h_q_f); free(h_k_f); free(h_v_f);
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_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
return pass ? 0 : 1;
}
// ======================================================================
// PREFILL WITH MASK TEST (regression: 4D causal mask on single request)
// ======================================================================
template <int HEAD_DIM>
static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
srand(seed);
int B = 1;
int total_q = q_len;
int seq_len = q_len; // pure prefill: kv_len == q_len
int max_ctx = seq_len + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
char cfg[80];
snprintf(cfg, sizeof(cfg), "PREFILL-MASK Hq=%d Hkv=%d D=%d q_len=%d",
Hq, Hkv, HEAD_DIM, q_len);
fflush(stdout);
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_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (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);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi, *d_qoi;
bool *d_mask;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
cudaMalloc(&d_mask, sz_mask);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
h_rpi[0] = 0;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0; h_kvi[1] = seq_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
int* h_qoi = (int*)malloc(sz_qoi);
h_qoi[0] = 0; h_qoi[1] = q_len;
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
// 4D causal mask [B, 1, q_len, q_len], True=keep.
bool* h_mask = (bool*)malloc(sz_mask);
for (int qi = 0; qi < q_len; qi++)
for (int kj = 0; kj < q_len; kj++)
h_mask[qi * q_len + kj] = (kj <= qi);
cudaMemcpy(d_mask, h_mask, sz_mask, cudaMemcpyHostToDevice);
float* h_q_f = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(total_q * Hq * HEAD_DIM, sizeof(float));
// CPU ref with causal=0 so it consults the mask (not the causal flag).
cpu_paged_prefill_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi, h_qoi,
h_mask, q_len, q_len,
B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref);
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = q_len;
p.max_q_len = q_len;
p.causal_offset = -1; p.use_mask = 1;
p.mask = d_mask; p.mask_b_stride = q_len * q_len;
p.mask_h_stride = 0; p.mask_q_stride = q_len;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
print_paged_row(cfg, max_err, pass);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_qoi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f);
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_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
cudaFree(d_mask);
return pass ? 0 : 1;
}
// ======================================================================
// BENCH
// ======================================================================
template <int HEAD_DIM>
static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
int max_ctx = seq_len + 16;
int pool_size = B * max_ctx;
int num_reqs = B;
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_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
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_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = seq_len;
p.causal_offset = 0; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); });
};
// Decode: q_len=1, query is the last token → attends to all [0, seq_len).
// FLOPs = 2 * (QK^T + PV) = 4 * B * Hq * seq_len * D.
double flops = 4.0 * B * Hq * (double)seq_len * HEAD_DIM;
BenchResult r = bench_kernel(launch, 3, 10, flops);
char cfg[64];
snprintf(cfg, sizeof(cfg), "DEC B=%2d Hq=%2d Hk=%d kv=%4d D=%3d",
B, Hq, Hkv, seq_len, HEAD_DIM);
print_bench_row(cfg, r);
free(tmp); free(h_rtt); free(h_rpi); free(h_kvi);
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_op); cudaFree(d_ml);
}
template <int HEAD_DIM>
static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int causal) {
int total_q = B * q_len;
int max_ctx = kv_len + 16;
int pool_size = B * max_ctx;
int num_reqs = B;
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_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (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;
int64_t *d_rtt, *d_rpi;
int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + kv_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
int* h_qoi = (int*)malloc(sz_qoi);
h_qoi[0] = 0;
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len;
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = kv_len;
p.total_q = total_q; p.max_q_len = q_len;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); });
};
// FLOPs = 2 * (QK^T + PV) = 4 * effective_qk_pairs * Hq * D.
// Non-causal: effective = q_len * kv_len.
// Causal: Q row qi attends to [0, causal_off + qi + 1) where
// causal_off = kv_len - q_len. Total KV accesses per request:
// sum_{qi=0}^{q_len-1} (kv_len - q_len + qi + 1)
// = q_len * (kv_len - q_len) + q_len * (q_len + 1) / 2.
double eff_kv;
if (causal) {
eff_kv = (double)q_len * (kv_len - q_len)
+ (double)q_len * (q_len + 1) / 2.0;
} else {
eff_kv = (double)q_len * kv_len;
}
double flops = 4.0 * B * Hq * eff_kv * HEAD_DIM;
BenchResult r = bench_kernel(launch, 3, 10, flops);
char cfg[80];
snprintf(cfg, sizeof(cfg), "PRE B=%d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d c=%d",
B, Hq, Hkv, q_len, kv_len, HEAD_DIM, causal);
print_bench_row(cfg, r);
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_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
}
int main() {
int fail = 0;
// ===== DECODE TESTS =====
printf("=== Paged Decode Tests ===\n");
print_paged_header();
fail += run_decode_test<128>(1, 32, 4, 512, 0, 1);
fail += run_decode_test<128>(1, 32, 4, 1024, 0, 2);
fail += run_decode_test<128>(4, 32, 4, 512, 0, 3);
fail += run_decode_test<128>(8, 32, 4, 1024, 0, 4);
fail += run_decode_test<128>(4, 32, 8, 2048, 0, 5);
fail += run_decode_test<128>(1, 16, 1, 256, 0, 6);
fail += run_decode_test<128>(2, 8, 2, 512, 1, 7);
fail += run_decode_test<64>(1, 4, 2, 256, 0, 8);
fail += run_decode_test<256>(1, 2, 1, 256, 0, 9);
fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10);
fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11);
// Decode with 2D mask (regression: mixed seq_lens + HasMask)
fail += run_decode_mask_test<128>(2, 8, 2, 256, 30);
fail += run_decode_mask_test<128>(4, 32, 4, 512, 31);
fail += run_decode_mask_test<64>(2, 4, 2, 128, 32);
if (fail) { printf("\nFAILED decode tests\n"); return fail; }
// ===== PREFILL TESTS =====
printf("\n=== Paged Prefill Tests ===\n");
print_paged_header();
// Single request, pure prefill (q_len == kv_len)
{
std::vector<int> ql = {512};
std::vector<int> kl = {512};
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 20);
}
{
std::vector<int> ql = {1024};
std::vector<int> kl = {1024};
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 21);
}
{
std::vector<int> ql = {2048};
std::vector<int> kl = {2048};
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 22);
}
// Ragged batch: different q_lens and kv_lens
{
std::vector<int> ql = {128, 256, 64};
std::vector<int> kl = {128, 256, 64};
fail += run_prefill_test<128>(3, 32, 4, ql, kl, 1, 23);
}
{
std::vector<int> ql = {64, 128, 256, 32};
std::vector<int> kl = {64, 128, 256, 32};
fail += run_prefill_test<128>(4, 32, 4, ql, kl, 1, 24);
}
// Extend: kv_len > q_len (append to existing cache)
{
std::vector<int> ql = {64, 128};
std::vector<int> kl = {256, 512};
fail += run_prefill_test<128>(2, 32, 4, ql, kl, 1, 25);
}
// Non-causal
{
std::vector<int> ql = {256, 128};
std::vector<int> kl = {256, 128};
fail += run_prefill_test<128>(2, 32, 4, ql, kl, 0, 26);
}
// Single token (q_len=1 per request, like decode but via prefill path)
{
std::vector<int> ql = {1, 1, 1, 1};
std::vector<int> kl = {128, 256, 64, 512};
fail += run_prefill_test<128>(4, 32, 4, ql, kl, 1, 27);
}
// D=64
{
std::vector<int> ql = {128, 64};
std::vector<int> kl = {128, 64};
fail += run_prefill_test<64>(2, 4, 2, ql, kl, 1, 28);
}
// D=256
{
std::vector<int> ql = {128, 64};
std::vector<int> kl = {128, 64};
fail += run_prefill_test<256>(2, 2, 1, ql, kl, 1, 29);
}
// Prefill with 4D causal mask (regression: single-request mask path)
fail += run_prefill_mask_test<128>(32, 4, 512, 40);
fail += run_prefill_mask_test<128>(32, 4, 1024, 41);
fail += run_prefill_mask_test<64>(4, 2, 256, 42);
if (fail) { printf("\nFAILED prefill tests\n"); return fail; }
printf("\nAll tests passed!\n");
// ===== BENCH =====
printf("\n===== PAGED DECODE BENCH =====\n");
print_bench_header();
bench_decode<128>(1, 32, 4, 512);
bench_decode<128>(1, 32, 4, 1024);
bench_decode<128>(1, 32, 4, 2048);
bench_decode<128>(1, 32, 4, 4096);
bench_decode<128>(4, 32, 4, 2048);
bench_decode<128>(16, 32, 4, 2048);
bench_decode<128>(32, 32, 4, 1024);
printf("\n===== PAGED PREFILL BENCH =====\n");
print_bench_header();
bench_prefill<128>(1, 32, 4, 512, 512, 0);
bench_prefill<128>(1, 32, 4, 1024, 1024, 0);
bench_prefill<128>(1, 32, 4, 2048, 2048, 0);
bench_prefill<128>(1, 32, 4, 2048, 2048, 1);
bench_prefill<128>(4, 32, 4, 2048, 2048, 1);
bench_prefill<128>(1, 32, 4, 4096, 4096, 1);
return 0;
}
-178
View File
@@ -1,178 +0,0 @@
/*
Pure-C test:
nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
csrc/tests/attn_prefill_test.cu -o test && ./test
*/
#include "test_utils.cuh"
#include "../kernels/attn_prefill_split_q.cuh"
#ifndef ASTRAI_NO_MMA
#include "../kernels/attn_prefill_split_q_mma.cuh"
#endif
// Launch the production prefill path (tensor-core MMA on sm_80+, else the
// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu.
template <int HEAD_DIM>
static void launch_prefill(AttentionParams<bf16>& p) {
#ifndef ASTRAI_NO_MMA
constexpr int WARPS = 4, BR = 16;
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
dim3 block(WARPS * 32, 1, 1);
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
#else
constexpr int G = 8, ROWS = 32, P_BC = 32;
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
dim3 block(G, ROWS, 1);
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
#endif
}
static void dispatch_prefill(AttentionParams<bf16>& p) {
switch (p.head_dim) {
case 64: launch_prefill<64>(p); break;
case 128: launch_prefill<128>(p); break;
default: printf("bench: unsupported D=%d\n", p.head_dim);
}
}
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls:
// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal).
static void bench() {
const int cfgs[][7] = {
{1,32,4,512,512,128,0},
{1,32,4,1024,1024,128,0},
{1,32,4,2048,2048,128,0},
{1,32,4,2048,2048,128,1},
{4,32,4,2048,2048,128,1},
{1,32,4,4096,4096,128,1},
};
int n = sizeof(cfgs)/sizeof(cfgs[0]);
const int WARMUP = 10, ITERS = 50;
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
printf("%-46s | %10s | %10s | %10s\n",
"config", "latency", "bandwidth", "throughput");
printf("---------------------------------------------------------------"
"----------------------------\n");
for (int ci = 0; ci < n; ci++) {
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
bf16 *dQ,*dK,*dV,*dO,*tmp;
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
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.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p);
p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
for (int i=0;i<WARMUP;i++) dispatch_prefill(p);
cudaDeviceSynchronize();
cudaError_t err=cudaGetLastError();
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
cudaEventRecord(s);
for (int i=0;i<ITERS;i++) dispatch_prefill(p);
cudaEventRecord(e); cudaEventSynchronize(e);
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
double flops = 4.0*B*Hq*(double)ql*kl*D;
if (causal) flops *= 0.5;
double tflops = flops/(ms*1e-3)/1e12;
// HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16.
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
double gbps = bytes/(ms*1e-3)/1e9;
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
B,Hq,Hk,ql,kl,D,causal);
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
cfg, ms, gbps, tflops);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
}
}
int main() {
const int configs[][7] = {
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal
};
int n_configs = sizeof(configs) / sizeof(configs[0]);
for (int ci = 0; ci < n_configs; ci++) {
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
int causal=configs[ci][6];
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
B,Hq,Hk,ql,kl,D,causal);
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
bf16 *dQ,*dK,*dV,*dO,*tmp;
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
tmp=new bf16[max(nQ,nKV)];
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
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.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p);
p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
double t0=now_ms();
dispatch_prefill(p);
cudaDeviceSynchronize();
double kms=now_ms()-t0;
cudaError_t err=cudaGetLastError();
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
bf16* hOut=new bf16[nQ];
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
float max_err=0;
for (size_t i=0;i<nQ;i++) {
float d=fabsf(bf2f(hOut[i])-ref[i]);
if(d>max_err) max_err=d;
}
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
}
printf("All tests passed!\n");
bench();
return 0;
}
+346
View File
@@ -0,0 +1,346 @@
/*
Pure-C test — uses shared dispatcher. Combines the decode (split-KV) and
prefill (split-Q) correctness checks + benchmarks into one binary.
nvcc -I csrc -arch=sm_89 -O3 \
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
-Xcompiler -fopenmp csrc/tests/attn_test.cu -o test && ./test
*/
#include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh"
// Split-K scratch (torch-free)
struct DecodeScratch {
float* o_part = nullptr;
float* ml_part = nullptr;
};
static void setup_scratch(AttentionParams<bf16>& p, DecodeScratch& sc) {
int max_splits = 32;
cudaMalloc(&sc.o_part, (size_t)p.batch * p.q_head * max_splits * p.head_dim * sizeof(float));
cudaMalloc(&sc.ml_part, (size_t)p.batch * p.q_head * max_splits * 2 * sizeof(float));
}
static void free_scratch(DecodeScratch& sc) {
cudaFree(sc.o_part); cudaFree(sc.ml_part);
}
// ======================================================================
// DECODE
// ======================================================================
static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
int gs = Hq / Hk;
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
bool* hMask=new bool[B*sl];
for (int i=0;i<B*sl;i++) hMask[i]=true;
bf16 *dQ,*dK,*dV,*dO,*tmp;
bool* dMask;
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
cudaMalloc(&dMask,B*sl);
tmp=new bf16[max(nQ,nKV)];
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
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.use_mask=0; p.causal_offset=causal?0:-1;
p.scale=1.0f/sqrtf((float)D);
set_default_strides(p);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
DecodeScratch sc;
setup_scratch(p, sc);
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
double t0=now_ms();
dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p, 0); });
cudaDeviceSynchronize();
(void)t0;
cudaError_t err=cudaGetLastError();
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
bf16* hOut=new bf16[nQ];
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
float max_abs_err=0, max_rel_err=0;
for (size_t i=0;i<nQ;i++){
float err=fabsf(bf2f(hOut[i])-ref[i]);
if(err>max_abs_err) max_abs_err=err;
float rel=err/fmaxf(fabsf(ref[i]), 1e-4f);
if(rel>max_rel_err) max_rel_err=rel;
}
const float atol=0.01f, rtol=0.01f;
bool pass=true;
for (size_t i=0;i<nQ;i++){
float err=fabsf(bf2f(hOut[i])-ref[i]);
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
}
char cfg[64];
snprintf(cfg, sizeof(cfg), "B=%2d Hq=%2d Hk=%d seq=%4d D=%3d causal=%d",
B, Hq, Hk, sl, D, causal);
print_test_row(cfg, max_abs_err, max_rel_err, pass);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
free_scratch(sc);
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
return pass ? 0 : 1;
}
static void bench_decode() {
const int cfgs[][5] = {
{1, 32, 4, 512, 128},
{1, 32, 4, 1024, 128},
{1, 32, 4, 2048, 128},
{1, 32, 4, 4096, 128},
{16, 32, 4, 2048, 128},
{32, 32, 4, 1024, 128},
};
const int WARMUP = 3, ITERS = 10;
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
print_bench_header();
for (int ci = 0; ci < 6; ci++) {
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
int sl = cfgs[ci][3], D = cfgs[ci][4];
size_t nQ = (size_t)B * Hq * D;
size_t nKV = (size_t)B * Hk * sl * D;
bf16 *dQ, *dK, *dV, *dO;
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
delete[] tmp;
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.use_mask = 0; p.causal_offset = -1;
p.scale = 1.0f / sqrtf((float)D);
set_default_strides(p);
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
DecodeScratch sc;
setup_scratch(p, sc);
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p, 0); }); };
double flops = 4.0 * B * Hq * (double)sl * D;
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops);
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
B, Hq, Hk, 1, sl, D, 0);
print_bench_row(cfg, r);
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
free_scratch(sc);
}
}
// ======================================================================
// PREFILL
// ======================================================================
static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
bf16 *dQ,*dK,*dV,*dO,*tmp;
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
tmp=new bf16[max(nQ,nKV)];
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
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.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p);
p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
double t0=now_ms();
dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p, 0); });
cudaDeviceSynchronize();
(void)t0;
cudaError_t err=cudaGetLastError();
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
bf16* hOut=new bf16[nQ];
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
float* ref=new float[nQ];
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
float max_abs_err=0, max_rel_err=0;
for (size_t i=0;i<nQ;i++) {
float err=fabsf(bf2f(hOut[i])-ref[i]);
if(err>max_abs_err) max_abs_err=err;
float rel=err/fmaxf(fabsf(ref[i]), 1e-4f);
if(rel>max_rel_err) max_rel_err=rel;
}
const float atol=0.01f, rtol=0.01f;
bool pass=true;
for (size_t i=0;i<nQ;i++) {
float err=fabsf(bf2f(hOut[i])-ref[i]);
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
}
char cfg[64];
snprintf(cfg, sizeof(cfg), "B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
B, Hq, Hk, ql, kl, D, causal);
print_test_row(cfg, max_abs_err, max_rel_err, pass);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
return pass ? 0 : 1;
}
static void bench_prefill() {
const int cfgs[][7] = {
{1,32,4,512,512,128,0},
{1,32,4,1024,1024,128,0},
{1,32,4,2048,2048,128,0},
{1,32,4,2048,2048,128,1},
{4,32,4,2048,2048,128,1},
{1,32,4,4096,4096,128,1},
};
int n = sizeof(cfgs)/sizeof(cfgs[0]);
const int WARMUP = 3, ITERS = 10;
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
print_bench_header();
for (int ci = 0; ci < n; ci++) {
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
bf16 *dQ,*dK,*dV,*dO,*tmp;
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
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.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p);
p.scale=1.0f/sqrtf((float)D);
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p, 0); }); };
for (int i=0;i<WARMUP;i++) launch();
cudaDeviceSynchronize();
cudaError_t err=cudaGetLastError();
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
cudaEventRecord(s);
for (int i=0;i<ITERS;i++) launch();
cudaEventRecord(e); cudaEventSynchronize(e);
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
double flops = 4.0*B*Hq*(double)ql*kl*D;
if (causal) flops *= 0.5;
double tflops = flops/(ms*1e-3)/1e12;
BenchResult r{ms, tflops};
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
B,Hq,Hk,ql,kl,D,causal);
print_bench_row(cfg, r);
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
}
}
// ======================================================================
// MAIN
// ======================================================================
int main() {
int fail = 0;
// ---- DECODE ----
{
const int configs[][6] = {
{1, 2, 1, 64, 32, 0},
{1, 32, 4, 512, 128, 0},
{1, 32, 4, 1024, 128, 0},
{1, 32, 4, 512, 128, 1},
};
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
printf("=== DECODE TESTS ===\n");
print_test_header();
for (int ci = 0; ci < n_cfgs; ci++) {
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5];
fail += run_decode_test(B, Hq, Hk, sl, D, causal);
if (fail) break;
}
if (fail) { printf("FAILED decode tests\n"); return fail; }
bench_decode();
}
// ---- PREFILL ----
{
const int configs[][7] = {
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
{1,32,4,512,512,128,0}, // standard
{1,32,4,128,256,128,0}, // medium
{1,4,2,256,256,128,1}, // causal
};
int n_configs = sizeof(configs) / sizeof(configs[0]);
printf("\n=== PREFILL TESTS ===\n");
print_test_header();
for (int ci = 0; ci < n_configs; ci++) {
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
int causal=configs[ci][6];
fail += run_prefill_test(B, Hq, Hk, ql, kl, D, causal);
if (fail) break;
}
if (fail) { printf("FAILED prefill tests\n"); return fail; }
bench_prefill();
}
printf("\nAll tests passed!\n");
return 0;
}
+26 -20
View File
@@ -18,16 +18,6 @@ inline double now_ms() {
return duration_cast<milliseconds>(steady_clock::now().time_since_epoch()).count(); return duration_cast<milliseconds>(steady_clock::now().time_since_epoch()).count();
} }
inline int compute_num_splits(int base_blocks, int tiles_total) {
int sm_count = 0;
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
if (n > tiles_total) n = tiles_total;
if (n > 32) n = 32;
if (n < 1) n = 1;
return n;
}
#define CUDA_CHECK(call) \ #define CUDA_CHECK(call) \
do { \ do { \
cudaError_t _e = (call); \ cudaError_t _e = (call); \
@@ -39,19 +29,18 @@ inline int compute_num_splits(int base_blocks, int tiles_total) {
struct BenchResult { struct BenchResult {
float ms; float ms;
double gbps;
double tflops; double tflops;
}; };
template <typename Fn> template <typename Fn>
BenchResult bench_kernel(Fn launch, int warmup, int iters, BenchResult bench_kernel(Fn launch, int warmup, int iters,
double flops, double bytes) { double flops) {
for (int i = 0; i < warmup; i++) launch(); for (int i = 0; i < warmup; i++) launch();
cudaDeviceSynchronize(); cudaDeviceSynchronize();
cudaError_t err = cudaGetLastError(); cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) { if (err != cudaSuccess) {
printf("CUDA error before bench: %s\n", cudaGetErrorString(err)); printf("CUDA error before bench: %s\n", cudaGetErrorString(err));
return {0, 0, 0}; return {0, 0};
} }
cudaEvent_t s, e; cudaEvent_t s, e;
@@ -62,19 +51,33 @@ BenchResult bench_kernel(Fn launch, int warmup, int iters,
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters; float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters;
cudaEventDestroy(s); cudaEventDestroy(e); cudaEventDestroy(s); cudaEventDestroy(e);
return {ms, bytes / (ms * 1e-3) / 1e9, flops / (ms * 1e-3) / 1e12}; return {ms, flops / (ms * 1e-3) / 1e12};
} }
inline void print_bench_header() { inline void print_bench_header() {
printf("%-46s | %10s | %10s | %10s\n", printf("%-46s | %10s | %10s\n",
"config", "latency", "bandwidth", "throughput"); "config", "latency", "TFLOP/s");
printf("---------------------------------------------------------------" printf("---------------------------------------------------------------"
"----------------------------\n"); "----------------------------\n");
} }
inline void print_bench_row(const char* cfg, const BenchResult& r) { inline void print_bench_row(const char* cfg, const BenchResult& r) {
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n", printf("%-46s | %7.4f ms | %6.2f\n",
cfg, r.ms, r.gbps, r.tflops); cfg, r.ms, r.tflops);
}
// ---- validation table (kernel vs CPU reference) ----
inline void print_test_header() {
printf("%-46s | %11s | %11s | %6s\n",
"config", "max_abs_err", "max_rel_err", "result");
printf("----------------------------------------------------------------"
"----------------------------\n");
}
inline void print_test_row(const char* cfg, float max_abs_err,
float max_rel_err, bool pass) {
printf("%-46s | %11.3e | %11.3e | %s\n",
cfg, max_abs_err, max_rel_err, pass ? "PASS" : "FAIL");
} }
template <int... Ds> template <int... Ds>
@@ -113,10 +116,11 @@ inline void set_default_strides(P& p) {
p.kv_stride_l = p.head_dim; p.kv_stride_l = p.head_dim;
p.kv_stride_d = 1; p.kv_stride_d = 1;
p.mask_b_stride = p.kv_len; p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_q_stride = 0;
} }
// Set default Q strides for contiguous b h l d layout on PagedAttentionParams. // Set default Q strides for a paged decode params struct.
template<typename P> template<typename P>
inline void set_default_paged_strides(P& p) { inline void set_default_paged_strides(P& p) {
p.q_stride_b = p.q_head * p.q_len * p.head_dim; p.q_stride_b = p.q_head * p.q_len * p.head_dim;
@@ -124,6 +128,7 @@ inline void set_default_paged_strides(P& p) {
p.q_stride_l = p.head_dim; p.q_stride_l = p.head_dim;
p.q_stride_d = 1; p.q_stride_d = 1;
p.mask_b_stride = p.kv_len; p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0;
p.mask_q_stride = 0; p.mask_q_stride = 0;
} }
@@ -143,9 +148,10 @@ static void cpu_attention_ref(
float scale = 1.0f / sqrtf((float)D); float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hk; int n_rep = Hq / Hk;
for (int b = 0; b < B; b++) { for (int b = 0; b < B; b++) {
#pragma omp parallel for collapse(2) schedule(dynamic)
for (int h = 0; h < Hq; h++) { for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep;
for (int qi = 0; qi < q_len; qi++) { for (int qi = 0; qi < q_len; qi++) {
int kv_h = h / n_rep;
float mv = -INFINITY, sv = 0.0f; float mv = -INFINITY, sv = 0.0f;
float accum[256] = {0.0f}; float accum[256] = {0.0f};
int lim = kv_len; int lim = kv_len;
+4
View File
@@ -3,6 +3,8 @@ services:
build: build:
context: . context: .
dockerfile: Dockerfile dockerfile: Dockerfile
args:
CUDA_TAG: ${CUDA_TAG:-cu128}
user: "${UID:-1000}:${GID:-1000}" user: "${UID:-1000}:${GID:-1000}"
ports: ports:
- "8000:8000" - "8000:8000"
@@ -29,6 +31,8 @@ services:
build: build:
context: . context: .
dockerfile: Dockerfile dockerfile: Dockerfile
args:
CUDA_TAG: ${CUDA_TAG:-cu128}
user: "${UID:-1000}:${GID:-1000}" user: "${UID:-1000}:${GID:-1000}"
ports: ports:
- "8000:8000" - "8000:8000"
@@ -1,9 +1,9 @@
<div align="center"> <div align="center">
<img src="../images/logo.png" width="auto" alt="Logo"> <img src="./images/logo.png" width="auto" alt="Logo">
<div> <div>
<a href="../../README.md">English</a> • <a href="../README.md">English</a> •
<a href="#chinese">中文</a> <a href="#chinese">中文</a>
</div> </div>
@@ -14,7 +14,7 @@
<div align="center"> <div align="center">
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python"> <img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license"> <img src="https://img.shields.io/badge/license-Apache--2.0-blue.svg" alt="license">
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release"> <img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars"> <img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks"> <img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
@@ -23,7 +23,7 @@
<br> <br>
<div align="center"> <div align="center">
<a href="../../README.md">English</a> • <a href="../README.md">English</a> •
<a href="#chinese">中文</a> • <a href="#chinese">中文</a> •
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> • <a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> • <a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
@@ -33,7 +33,7 @@
## 📖 目录 ## 📖 目录
- [特性](#特性) - [项目概览](#项目概览)
- [快速上手](#快速上手) - [快速上手](#快速上手)
- [演示](#演示) - [演示](#演示)
- [文档](#文档) - [文档](#文档)
@@ -46,15 +46,19 @@
<a id="chinese"></a> <a id="chinese"></a>
## 中文 ## 中文
### 特性 ### 项目概览
- 🚀 **高性能**: 训练与推理双向优化,高效并行 AstrAI 是一个覆盖模型构建、训练、评测与部署的端到端 Transformer 框架。项目以精简的 PyTorch 代码实现完整模型生命周期,包括声明式数据预处理、分布式训练、连续批处理推理,以及兼容 OpenAI 和 Anthropic 的服务接口
- 🔧 **灵活**: 支持 seq/sft/dpo/grpo 多种训练方式,可定制模型架构。
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。 | 领域 | 能力 |
- 📦 **轻量**: 依赖少,部署简单。 |---|---|
- 🔬 **研究友好**: 模块化设计,便于实验新想法。 | **模型** | 自回归语言模型与嵌入模型,支持 GQA、MLA、MoE、RoPE,以及可扩展的 Attention/FFN 组件 |
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。 | **训练** | 预训练(`seq`)、监督微调(`sft`)、DPO 和 GRPO,支持梯度累积、检查点、DDP 与 FSDP |
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。 | **数据** | 声明式 JSON 预处理、可配置掩码与样本打包、二进制/JSONL 存储和流式数据集 |
| **推理** | 连续批处理、分页 KV Cache、Radix 前缀缓存、流式生成,以及 Torch/CUDA/FlashAttention 后端 |
| **服务** | 基于 FastAPI 的 OpenAI 与 Anthropic 聊天补全协议,支持 SSE 流式输出和工具调用 |
| **评测** | Perplexity、MMLU、HumanEval、IFEval、IFD 和 ROUGE 评测工具 |
| **扩展** | 基于工厂与注册表扩展模型、数据集、训练策略、回调、内核和协议组件 |
### 快速上手 ### 快速上手
@@ -62,6 +66,8 @@
**1. 安装** **1. 安装**
AstrAI 需要 Python 3.12+,并精确固定 PyTorch 版本为 `2.11.0`。训练、`scripts/tools/generate.py`、生成式评估和生成演示需要 CUDA;CPU 支持仅适用于提供明确 CPU 设备路径的组件,例如 HTTP 服务和直接打分评估。
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
@@ -138,7 +144,7 @@ curl http://localhost:8000/v1/chat/completions \
# 下载模型权重(运行演示前必需) # 下载模型权重(运行演示前必需)
python scripts/demo/download.py # model → params/ python scripts/demo/download.py # model → params/
# 交互式流式聊天(多轮对话,保持历史记录 # 单轮交互式流式提示循环(不保留对话历史
python scripts/demo/stream_chat.py python scripts/demo/stream_chat.py
# 在 >> 后输入消息,输入 !exit 退出 # 在 >> 后输入消息,输入 !exit 退出
@@ -189,7 +195,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker ComposeGPU,默认) # Docker ComposeGPU,默认)
docker compose up -d docker compose up -d
# Docker Compose(仅 CPU # Docker Compose CPU 服务配置(不支持仅限 CUDA 的生成脚本和演示
docker compose --profile cpu up -d docker compose --profile cpu up -d
``` ```
@@ -219,22 +225,27 @@ curl -X POST http://localhost:8000/v1/messages \
curl http://localhost:8000/health curl http://localhost:8000/health
``` ```
SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)。 SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference.md)。
### 文档 ### 文档
| 文档 | 说明 | | 文档 | 说明 |
|------|------| |------|------|
| [CLI 参考](./params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) | | [快速上手](./get-started.md) | 安装与快速入门 |
| [架构文档](./architecture.md) | 系统架构、类图与设计模式 | | [CLI 参考](./guides/params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
| [训练文档](./training.md) | 训练循环、策略与公式 | | [数据预处理](./guides/preprocessing.md) | 声明式 JSON 驱动数据预处理 |
| [推理文档](./inference.md) | KVCache、连续批处理、采样与 HTTP API | | [训练文档](./guides/training.md) | 训练循环、策略与公式 |
| [数据流程](./dataflow.md) | 数据管道、存储后端与数据集架构 | | [推理文档](./guides/inference.md) | KVCache、连续批处理、采样与 HTTP API |
| [数据预处理](./preprocessing.md) | 声明式 JSON 驱动数据预处理 | | [评估文档](./guides/evaluation.md) | HumanEval、MMLU、PPL、ROUGE、IFD、IFEval |
| [分布式训练](./guides/distributed.md) | 多卡 DDP / FSDP 训练 |
| [架构文档](./developer/architecture.md) | 系统架构、类图与设计模式 |
| [数据流程](./developer/dataflow.md) | 数据管道、存储后端与数据集架构 |
| [内部实现](./developer/internals.md) | 训练原理:损失公式、回调生命周期、KV Cache |
| [CUDA 内核](./developer/cuda_kernels.md) | 自定义 CUDA 注意力内核与基准测试 |
### 贡献 ### 贡献
我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。 我们欢迎贡献!请参阅[贡献指南](../CONTRIBUTING.md)了解详情。
1. Fork 本仓库。 1. Fork 本仓库。
2. 创建功能分支。 2. 创建功能分支。
@@ -251,7 +262,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)
### 许可证 ### 许可证
本项目采用 [GPL-3.0 许可证](../../LICENSE)。 本项目采用 [Apache-2.0 许可证](../LICENSE)。
--- ---
File diff suppressed because it is too large Load Diff
+186
View File
@@ -0,0 +1,186 @@
# 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.
## Overview
| Kernel | File | Description |
|--------|------|-------------|
| `attn_decode` | `attn_decode.cu` | GQA decode attention (split-KV) |
| `attn_prefill` | `attn_prefill.cu` | GQA prefill attention (split-Q) |
| `attn_paged_decode` | `attn_paged_decode.cu` | Paged KV cache decode attention |
| `attn_paged_prefill` | `attn_paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) |
| `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist:
| Variant | File | Optimization |
|---------|------|--------------|
| Split-KV MMA decode | `attn_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+) |
> The paged and non-paged paths are ONE kernel templated on a `KVSource`
> policy (`ContigKV` / `PagedKV` in `attn_kv_source.cuh`); there are no
> separate `attn_paged_*.cuh` files anymore.
### Rotary Embedding Kernel
The `rotary_emb` kernel (`csrc/kernels/rotary_emb.cu`) fuses cos/sin lookup and rotation into a single kernel:
- One thread per (head, dim-pair), vectorized `__nv_bfloat162` load/store
- f32 cos/sin input, bf16 compute and output
- 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)
- 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).
## Build System
### Auto-detection
Kernels are built when **both** of these conditions are met:
1. `nvcc` is available on `PATH`
2. `torch.cuda.is_available()` returns `True`
Unless `CSRC_KERNELS=false` is set explicitly.
### Manual build
```bash
# During install
CSRC_KERNELS=true pip install -e . --no-build-isolation
# Rebuild after editing .cu/.cuh files
CSRC_KERNELS=true python setup.py build_ext --inplace
# Output: astrai/extension/lib/*.so
# Or invoke CMake directly
cmake -S csrc -B build/cmake \
-DTORCH_HOME=<site-packages>/torch \
-DPYTHON_INCLUDE_DIR=<python include> \
-DPY_SOABI=cpython-312-x86_64-linux-gnu
cmake --build build/cmake -j 16
```
### Architecture flags
`setup.py` passes the GPU compute capability to CMake via `ASTRAI_CUDA_ARCH` (default `89`, i.e. sm_89 / L20):
- **sm_80+** (Ampere and later): enables tensor-core MMA path (`mma.sync.m16n8k16.bf16`)
- **Below sm_80**: adds `-DASTRAI_NO_MMA` to disable the MMA path at compile time
### Build configuration
`csrc/CMakeLists.txt` defines the CUDA extension build:
```
NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
--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`.
## Attention Backend
`astrai/extension/attention_backend.py` provides the backend abstraction:
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (default)
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`)
Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
```python
from astrai.extension import attn_backend, ATTN_BACKEND
with attn_backend(ATTN_BACKEND.CUDA):
engine.generate("hello")
```
`CudaBackend` falls back to `TorchNativeBackend` when a kernel is not available.
### Rotary Backend
`astrai/extension/rotary_backend.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)
- **Torch fallback**: complex multiply (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd) or when kernel unavailable
No context-manager switching needed — the dispatch is automatic per call.
## 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/rotary_ops.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `rotary_backend.py`.
Interface (all functions):
```
is_causal: True = causal mask; False = non-causal
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)`.
## Standalone Testing
Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment. Example:
```bash
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
--ptxas-options=-O3,-v --extra-device-vectorization \
-Xcompiler -fopenmp csrc/tests/attn_test.cu -o /tmp/test && /tmp/test
```
Test files:
- `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks)
- `attn_paged_test.cu` — paged decode/prefill kernels
## Benchmarks
Hardware: NVIDIA L20 (sm_89, 46 GB), CUDA 12.8, driver 570.86.
Reproduce (decode + prefill in `attn_test.cu`, paged in `attn_paged_test.cu`):
```bash
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
--ptxas-options=-O3,-v --extra-device-vectorization \
-Xcompiler -fopenmp csrc/tests/attn_test.cu -o /tmp/test && /tmp/test
```
## Known Optimization Targets
- **Decode D=256**: spill eliminated (BC=16 + STAGES=2), but still 248 regs — further tiling could help.
- **Prefill single-batch**: bandwidth low (22 GB/s at q=kv=2048) — compute-bound at ~94 TFLOP/s (near L20 bf16 ceiling ~193 TFLOP/s for non-causal).
- **Decode single-batch**: bandwidth low (113 GB/s at kv=512, 13% of 864 GB/s theoretical) — small kv underutilizes SMs despite split-KV; scales to 757 GB/s (88%) at B=16+.
## File Layout
```
csrc/
├── CMakeLists.txt # CMake build: 5 kernel targets, torch/pybind11 linking
├── kernels/
│ ├── attn_common.h # Unified attention params (contig + paged modes)
│ ├── attn_decode.cu # Basic decode kernel (registered)
│ ├── attn_prefill.cu # Basic prefill kernel (registered)
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
│ ├── attn_paged_prefill.cu # Paged prefill kernel (registered)
│ ├── rotary_emb.cu # Fused rotary embedding kernel (registered)
│ ├── attn_decode_split_kv.cuh # Split-KV variant (contig + paged via KVSource)
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant (contig + paged)
│ ├── attn_prefill_split_q.cuh # Split-Q variant (contig + paged via KVSource)
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant (contig + paged)
│ ├── attn_kv_source.cuh # KVSource policies (ContigKV / PagedKV)
│ ├── attn_dispatchers.cuh # Kernel dispatch macros + KV-templated launchers
│ ├── attn_entry_utils.cuh # Entry point helpers
│ ├── attn_mma_utils.cuh # MMA utilities
│ └── attn_warp_utils.cuh # Warp-level utilities
└── tests/
├── test_utils.cuh # Shared test utilities
├── attn_test.cu # Decode + prefill kernels
└── attn_paged_test.cu # Paged decode/prefill kernels
```
Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.
> Document Update Time: 2026-07-31

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